"""OAuth + JWT auth routes (mounted at /api/auth and /api/me).""" import os from datetime import datetime, timedelta, timezone from typing import Optional from authlib.integrations.base_client import OAuthError from fastapi import APIRouter, Depends, HTTPException, Request, Response, status from fastapi.responses import RedirectResponse from pydantic import BaseModel from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.auth import ( REFRESH_COOKIE, clear_auth_cookies, create_access_token, email_domain_allowed, find_or_create_google_user, get_current_user, hash_refresh_token, issue_refresh_token, oauth, set_auth_cookies, ) from app.database import get_db from app.models import ( ActivityLog, DiscoveryConversation, DiscoveryProfile, RefreshToken, User, ) from sqlalchemy import func router = APIRouter(prefix="/api", tags=["auth"]) # -- response shapes ---------------------------------------------------------- class UserOut(BaseModel): id: str email: str display_name: str avatar_url: str | None = None role: str @classmethod def from_orm_user(cls, user: User) -> "UserOut": return cls( id=user.id, email=user.email, display_name=user.display_name, avatar_url=user.avatar_url, role=user.role, ) class RefreshIn(BaseModel): # Optional: browser clients send the refresh token via httpOnly cookie and # omit the body entirely; API clients may still post it explicitly. refresh_token: Optional[str] = None class AccessOut(BaseModel): access_token: str token_type: str = "bearer" class DisplayNamePatch(BaseModel): display_name: str class SessionOut(BaseModel): id: str device: str | None created_at: datetime expires_at: datetime # -- OAuth flow --------------------------------------------------------------- @router.get("/auth/login") async def login(request: Request): """Kick off the OAuth dance: 302 to Google's consent screen.""" redirect_uri = os.getenv("OAUTH_REDIRECT_URI") or str( request.url_for("auth_callback") ) return await oauth.google.authorize_redirect(request, redirect_uri) @router.get("/auth/callback", name="auth_callback") async def auth_callback(request: Request, db: AsyncSession = Depends(get_db)): """Google redirects here with `code`; exchange it, mint our tokens.""" try: token = await oauth.google.authorize_access_token(request) except OAuthError as exc: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail=f"OAuth error: {exc.error or exc}", ) from exc user_info = token.get("userinfo") if not user_info or not user_info.get("email") or not user_info.get("sub"): raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail="Google did not return email/sub", ) if user_info.get("email_verified") is False: raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="Email not verified by Google", ) if not email_domain_allowed(user_info["email"]): raise HTTPException( status_code=status.HTTP_403_FORBIDDEN, detail="Email domain not allowed", ) user = await find_or_create_google_user(db, user_info) access = create_access_token(user.id, user.email) refresh = await issue_refresh_token( db, user, request.headers.get("user-agent") ) # The browser drove this redirect flow, so hand the session back as # httpOnly cookies and bounce into the app rather than dumping JSON. redirect_to = os.getenv("POST_LOGIN_REDIRECT", "/static/discovery.html") response = RedirectResponse(url=redirect_to, status_code=status.HTTP_303_SEE_OTHER) set_auth_cookies(response, access, refresh) return response @router.post("/auth/refresh", response_model=AccessOut) async def refresh_access_token( request: Request, response: Response, body: Optional[RefreshIn] = None, db: AsyncSession = Depends(get_db), ): raw_refresh = (body.refresh_token if body else None) or request.cookies.get( REFRESH_COOKIE ) if not raw_refresh: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="No refresh token provided", ) stmt = select(RefreshToken).where( RefreshToken.token_hash == hash_refresh_token(raw_refresh) ) row = (await db.execute(stmt)).scalar_one_or_none() now = datetime.now(timezone.utc) if row is None or row.revoked_at is not None or _expired(row.expires_at, now): raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="Refresh token invalid or expired", ) user = await db.get(User, row.user_id) if user is None: raise HTTPException( status_code=status.HTTP_401_UNAUTHORIZED, detail="User no longer exists", ) access = create_access_token(user.id, user.email) # Refresh the access cookie in place; the refresh cookie is untouched. set_auth_cookies(response, access) return AccessOut(access_token=access) def _expired(expires_at: datetime, now: datetime) -> bool: # SQLite stores naive datetimes; compare in UTC. if expires_at.tzinfo is None: expires_at = expires_at.replace(tzinfo=timezone.utc) return expires_at < now @router.post("/auth/logout") async def logout( request: Request, response: Response, body: Optional[RefreshIn] = None, db: AsyncSession = Depends(get_db), ): """Revoke a single refresh token and clear the session cookies. Idempotent — an unknown/missing token still returns 200 with cookies cleared.""" raw_refresh = (body.refresh_token if body else None) or request.cookies.get( REFRESH_COOKIE ) if raw_refresh: stmt = select(RefreshToken).where( RefreshToken.token_hash == hash_refresh_token(raw_refresh) ) row = (await db.execute(stmt)).scalar_one_or_none() if row is not None and row.revoked_at is None: row.revoked_at = datetime.now(timezone.utc) await db.commit() clear_auth_cookies(response) return {"status": "ok"} # -- profile / sessions ------------------------------------------------------- @router.get("/auth/me", response_model=UserOut) @router.get("/me", response_model=UserOut) async def get_me(user: User = Depends(get_current_user)): return UserOut.from_orm_user(user) @router.patch("/me", response_model=UserOut) async def patch_me( body: DisplayNamePatch, user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): new_name = body.display_name.strip() if not new_name: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, detail="display_name cannot be empty", ) user.display_name = new_name await db.commit() await db.refresh(user) return UserOut.from_orm_user(user) class MeStats(BaseModel): conversations: int profiles: int locked_profiles: int activity_last_30d: int last_login_at: datetime | None @router.get("/me/stats", response_model=MeStats) async def my_stats( user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): convo_count = (await db.execute( select(func.count()) .select_from(DiscoveryConversation) .where(DiscoveryConversation.user_id == user.id) )).scalar_one() profile_count = (await db.execute( select(func.count()) .select_from(DiscoveryProfile) .where(DiscoveryProfile.user_id == user.id) )).scalar_one() locked_count = (await db.execute( select(func.count()) .select_from(DiscoveryProfile) .where(DiscoveryProfile.user_id == user.id) .where(DiscoveryProfile.locked.is_(True)) )).scalar_one() since = datetime.now(timezone.utc) - timedelta(days=30) activity_count = (await db.execute( select(func.count()) .select_from(ActivityLog) .where(ActivityLog.user_id == user.id) .where(ActivityLog.created_at >= since) )).scalar_one() return MeStats( conversations=int(convo_count or 0), profiles=int(profile_count or 0), locked_profiles=int(locked_count or 0), activity_last_30d=int(activity_count or 0), last_login_at=user.last_login_at, ) @router.get("/me/sessions", response_model=list[SessionOut]) async def list_sessions( user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): stmt = ( select(RefreshToken) .where(RefreshToken.user_id == user.id) .where(RefreshToken.revoked_at.is_(None)) .order_by(RefreshToken.created_at.desc()) ) rows = (await db.execute(stmt)).scalars().all() return [ SessionOut( id=r.id, device=r.device, created_at=r.created_at, expires_at=r.expires_at, ) for r in rows ] @router.delete("/me/sessions/{session_id}") async def revoke_session( session_id: str, user: User = Depends(get_current_user), db: AsyncSession = Depends(get_db), ): row = await db.get(RefreshToken, session_id) if row is None or row.user_id != user.id: raise HTTPException( status_code=status.HTTP_404_NOT_FOUND, detail="Session not found", ) if row.revoked_at is None: row.revoked_at = datetime.now(timezone.utc) await db.commit() return {"status": "revoked"}