eb18994b7d
- Migrate 660+ datetime.utcnow() across 153 files to datetime.now(UTC) - Migrate 30+ datetime.now() without UTC to datetime.now(UTC) - Convert all 170 DateTime columns to DateTime(timezone=True) - Add migrate_datetime_to_timestamptz() in universal_migration with SET LOCAL timezone='UTC' safety - Remove 70+ .replace(tzinfo=None) workarounds - Fix utcfromtimestamp → fromtimestamp(..., tz=UTC) - Fix fromtimestamp() without tz= (system_logs, backup_service, referral_diagnostics) - Fix fromisoformat/isoparse to ensure aware output (platega, yookassa, wata, miniapp, nalogo) - Fix strptime() to add .replace(tzinfo=UTC) (backup_service, referral_diagnostics) - Fix datetime.combine() to include tzinfo=UTC (remnawave_sync, traffic_monitoring) - Fix datetime.max/datetime.min sentinels with .replace(tzinfo=UTC) - Rename panel_datetime_to_naive_utc → panel_datetime_to_utc - Remove DTZ003 from ruff ignore list
166 lines
5.5 KiB
Python
166 lines
5.5 KiB
Python
"""OAuth 2.0 authentication routes for cabinet."""
|
|
|
|
from datetime import UTC, datetime
|
|
|
|
import structlog
|
|
from fastapi import APIRouter, Depends, HTTPException, status
|
|
from pydantic import BaseModel, Field
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.config import settings
|
|
from app.database.crud.user import (
|
|
create_user_by_oauth,
|
|
get_user_by_email,
|
|
get_user_by_oauth_provider,
|
|
set_user_oauth_provider_id,
|
|
)
|
|
from app.database.models import User
|
|
|
|
from ..auth.oauth_providers import (
|
|
OAuthUserInfo,
|
|
generate_oauth_state,
|
|
get_provider,
|
|
validate_oauth_state,
|
|
)
|
|
from ..dependencies import get_cabinet_db
|
|
from ..schemas.auth import AuthResponse
|
|
from .auth import _create_auth_response, _store_refresh_token
|
|
|
|
|
|
logger = structlog.get_logger(__name__)
|
|
|
|
router = APIRouter(prefix='/auth/oauth', tags=['Cabinet OAuth'])
|
|
|
|
|
|
async def _finalize_oauth_login(db: AsyncSession, user: User, provider: str) -> AuthResponse:
|
|
"""Update last login, create tokens, store refresh token."""
|
|
user.cabinet_last_login = datetime.now(UTC)
|
|
await db.commit()
|
|
auth_response = _create_auth_response(user)
|
|
await _store_refresh_token(db, user.id, auth_response.refresh_token, device_info=f'oauth:{provider}')
|
|
return auth_response
|
|
|
|
|
|
# --- Schemas ---
|
|
|
|
|
|
class OAuthProviderInfo(BaseModel):
|
|
name: str
|
|
display_name: str
|
|
|
|
|
|
class OAuthProvidersResponse(BaseModel):
|
|
providers: list[OAuthProviderInfo]
|
|
|
|
|
|
class OAuthAuthorizeResponse(BaseModel):
|
|
authorize_url: str
|
|
state: str
|
|
|
|
|
|
class OAuthCallbackRequest(BaseModel):
|
|
code: str = Field(..., description='Authorization code from provider')
|
|
state: str = Field(..., description='CSRF state token')
|
|
|
|
|
|
# --- Endpoints ---
|
|
|
|
|
|
@router.get('/providers', response_model=OAuthProvidersResponse)
|
|
async def get_oauth_providers():
|
|
"""Get list of enabled OAuth providers."""
|
|
providers_config = settings.get_oauth_providers_config()
|
|
providers = [
|
|
OAuthProviderInfo(name=name, display_name=cfg['display_name'])
|
|
for name, cfg in providers_config.items()
|
|
if cfg['enabled']
|
|
]
|
|
return OAuthProvidersResponse(providers=providers)
|
|
|
|
|
|
@router.get('/{provider}/authorize', response_model=OAuthAuthorizeResponse)
|
|
async def get_oauth_authorize_url(provider: str):
|
|
"""Get authorization URL for an OAuth provider."""
|
|
oauth_provider = get_provider(provider)
|
|
if not oauth_provider:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail=f'OAuth provider "{provider}" is not enabled',
|
|
)
|
|
|
|
state = await generate_oauth_state(provider)
|
|
authorize_url = oauth_provider.get_authorization_url(state)
|
|
|
|
return OAuthAuthorizeResponse(authorize_url=authorize_url, state=state)
|
|
|
|
|
|
@router.post('/{provider}/callback', response_model=AuthResponse)
|
|
async def oauth_callback(
|
|
provider: str,
|
|
request: OAuthCallbackRequest,
|
|
db: AsyncSession = Depends(get_cabinet_db),
|
|
):
|
|
"""Handle OAuth callback: exchange code, find/create user, return JWT."""
|
|
# 1. Validate CSRF state
|
|
if not await validate_oauth_state(request.state, provider):
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail='Invalid or expired OAuth state',
|
|
)
|
|
|
|
# 2. Get provider instance
|
|
oauth_provider = get_provider(provider)
|
|
if not oauth_provider:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail=f'OAuth provider "{provider}" is not enabled',
|
|
)
|
|
|
|
# 3. Exchange code for tokens
|
|
try:
|
|
token_data = await oauth_provider.exchange_code(request.code)
|
|
except Exception as exc:
|
|
logger.error('OAuth code exchange failed for', provider=provider, exc=exc)
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail='Failed to exchange authorization code',
|
|
) from exc
|
|
|
|
# 4. Fetch user info from provider
|
|
try:
|
|
user_info: OAuthUserInfo = await oauth_provider.get_user_info(token_data)
|
|
except Exception as exc:
|
|
logger.error('OAuth user info fetch failed for', provider=provider, exc=exc)
|
|
raise HTTPException(
|
|
status_code=status.HTTP_400_BAD_REQUEST,
|
|
detail='Failed to fetch user information from provider',
|
|
) from exc
|
|
|
|
# 5. Find user by provider ID
|
|
user = await get_user_by_oauth_provider(db, provider, user_info.provider_id)
|
|
if user:
|
|
logger.info('OAuth login via for existing user', provider=provider, user_id=user.id)
|
|
return await _finalize_oauth_login(db, user, provider)
|
|
|
|
# 6. Find user by email (if verified) and link provider
|
|
if user_info.email and user_info.email_verified:
|
|
user = await get_user_by_email(db, user_info.email)
|
|
if user:
|
|
await set_user_oauth_provider_id(db, user, provider, user_info.provider_id)
|
|
logger.info('OAuth login via linked to existing email user', provider=provider, user_id=user.id)
|
|
return await _finalize_oauth_login(db, user, provider)
|
|
|
|
# 7. Create new user
|
|
user = await create_user_by_oauth(
|
|
db=db,
|
|
provider=provider,
|
|
provider_id=user_info.provider_id,
|
|
email=user_info.email if user_info.email_verified else None,
|
|
email_verified=user_info.email_verified,
|
|
first_name=user_info.first_name,
|
|
last_name=user_info.last_name,
|
|
username=user_info.username,
|
|
)
|
|
logger.info('OAuth new user created via with id', provider=provider, user_id=user.id)
|
|
return await _finalize_oauth_login(db, user, provider)
|