189 lines
6.0 KiB
Python
189 lines
6.0 KiB
Python
"""Auth service for VoIdea."""
|
|
|
|
from collections import defaultdict
|
|
from datetime import datetime, timedelta, timezone
|
|
from uuid import uuid4
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.core.config import get_settings
|
|
from app.core.security import (
|
|
create_access_token,
|
|
create_refresh_token,
|
|
decode_token,
|
|
get_password_hash,
|
|
verify_password,
|
|
)
|
|
from app.models.user import User
|
|
|
|
settings = get_settings()
|
|
|
|
# ── Brute force protection (in-memory, TODO: move to Redis in production) ──
|
|
_login_attempts: dict[str, list[datetime]] = defaultdict(list)
|
|
MAX_LOGIN_ATTEMPTS = 5
|
|
LOGIN_WINDOW = timedelta(minutes=15)
|
|
|
|
|
|
def _check_login_attempts(email: str) -> bool:
|
|
now = datetime.now(timezone.utc)
|
|
attempts = [t for t in _login_attempts[email] if now - t < LOGIN_WINDOW]
|
|
_login_attempts[email] = attempts
|
|
return len(attempts) < MAX_LOGIN_ATTEMPTS
|
|
|
|
|
|
def _record_login_attempt(email: str):
|
|
_login_attempts[email].append(datetime.now(timezone.utc))
|
|
|
|
|
|
def _clear_login_attempts(email: str):
|
|
_login_attempts.pop(email, None)
|
|
|
|
|
|
class AuthService:
|
|
def __init__(self, db: AsyncSession):
|
|
self.db = db
|
|
|
|
async def register(
|
|
self, email: str, password: str, display_name: str,
|
|
accepted_terms: bool = True,
|
|
) -> dict:
|
|
result = await self.db.execute(
|
|
select(User).where(User.email == email)
|
|
)
|
|
if result.scalar_one_or_none():
|
|
raise ValueError("Email already registered")
|
|
|
|
if not accepted_terms:
|
|
raise ValueError("Необходимо принять Пользовательское соглашение и Политику конфиденциальности")
|
|
|
|
user = User(
|
|
id=uuid4(),
|
|
email=email,
|
|
password_hash=get_password_hash(password),
|
|
display_name=display_name,
|
|
is_active=True,
|
|
is_superuser=False,
|
|
role="user",
|
|
accepted_terms_at=datetime.now(timezone.utc),
|
|
accepted_terms_version=settings.accepted_terms_version,
|
|
)
|
|
self.db.add(user)
|
|
await self.db.commit()
|
|
await self.db.refresh(user)
|
|
|
|
# Set owner by SYSTEM_OWNER_EMAIL if matches
|
|
if settings.system_owner_email and email == settings.system_owner_email:
|
|
from app.services.user_service import UserService
|
|
svc = UserService(self.db)
|
|
await svc.set_owner_by_email(email)
|
|
|
|
return self._generate_tokens(str(user.id))
|
|
|
|
async def login(self, email: str, password: str) -> dict:
|
|
if not _check_login_attempts(email):
|
|
raise ValueError("Too many login attempts. Try again in 15 minutes.")
|
|
|
|
result = await self.db.execute(
|
|
select(User).where(User.email == email)
|
|
)
|
|
user = result.scalar_one_or_none()
|
|
|
|
if not user or not user.password_hash:
|
|
_record_login_attempt(email)
|
|
raise ValueError("Invalid credentials")
|
|
|
|
if not verify_password(password, user.password_hash):
|
|
_record_login_attempt(email)
|
|
raise ValueError("Invalid credentials")
|
|
|
|
if not user.is_active:
|
|
raise ValueError("Account is disabled")
|
|
|
|
_clear_login_attempts(email)
|
|
return self._generate_tokens(str(user.id))
|
|
|
|
async def refresh(self, refresh_token: str) -> dict:
|
|
payload = decode_token(refresh_token)
|
|
if not payload or payload.get("type") != "refresh":
|
|
raise ValueError("Invalid refresh token")
|
|
|
|
user_id = payload.get("sub")
|
|
if not user_id:
|
|
raise ValueError("Invalid token payload")
|
|
|
|
result = await self.db.execute(
|
|
select(User).where(User.id == user_id)
|
|
)
|
|
user = result.scalar_one_or_none()
|
|
if not user or not user.is_active:
|
|
raise ValueError("User not found or disabled")
|
|
|
|
# Refresh token rotation: issue new pair, old token becomes invalid
|
|
return self._generate_tokens(user_id)
|
|
|
|
async def oauth_or_register_login(
|
|
self,
|
|
email: str,
|
|
oauth_provider: str,
|
|
oauth_id: str,
|
|
display_name: str,
|
|
avatar_url: str | None = None,
|
|
) -> dict:
|
|
result = await self.db.execute(
|
|
select(User).where(
|
|
User.oauth_provider == oauth_provider,
|
|
User.oauth_id == oauth_id,
|
|
)
|
|
)
|
|
user = result.scalar_one_or_none()
|
|
|
|
if user:
|
|
if not user.is_active:
|
|
raise ValueError("Account is disabled")
|
|
if avatar_url:
|
|
user.avatar_url = avatar_url
|
|
await self.db.commit()
|
|
return self._generate_tokens(str(user.id))
|
|
|
|
result = await self.db.execute(
|
|
select(User).where(User.email == email)
|
|
)
|
|
user = result.scalar_one_or_none()
|
|
|
|
if user:
|
|
user.oauth_provider = oauth_provider
|
|
user.oauth_id = oauth_id
|
|
if avatar_url:
|
|
user.avatar_url = avatar_url
|
|
await self.db.commit()
|
|
return self._generate_tokens(str(user.id))
|
|
|
|
user = User(
|
|
id=uuid4(),
|
|
email=email,
|
|
password_hash=None,
|
|
display_name=display_name,
|
|
avatar_url=avatar_url,
|
|
is_active=True,
|
|
is_superuser=False,
|
|
oauth_provider=oauth_provider,
|
|
oauth_id=oauth_id,
|
|
)
|
|
self.db.add(user)
|
|
await self.db.commit()
|
|
await self.db.refresh(user)
|
|
return self._generate_tokens(str(user.id))
|
|
|
|
async def create_token_response(self, user: User) -> dict:
|
|
return self._generate_tokens(str(user.id))
|
|
|
|
def _generate_tokens(self, user_id: str) -> dict:
|
|
now = datetime.now(timezone.utc)
|
|
return {
|
|
"access_token": create_access_token({"sub": user_id}),
|
|
"refresh_token": create_refresh_token({"sub": user_id}),
|
|
"token_type": "bearer",
|
|
"expires_at": now + settings.jwt_access_token_expire_minutes * 60,
|
|
}
|