Benard John
feat: implement full-stack architecture with database models, authentication services, and web dashboard components
37b5223 | """ | |
| Ukweli — Authentication & Authorization | |
| Implements the multi-tier model from OAG_Kenya_Auth_Architecture.md: | |
| - Public: IP-fingerprinted, anonymous | |
| - Registered: JWT (email + OTP) | |
| - API: API key + bcrypt secret | |
| - Enterprise: mTLS / IP whitelist (reserved) | |
| The `require_auth(min_tier=...)` dependency is the single entry point | |
| used by every protected endpoint. It: | |
| 1. Inspects the request for an API key, a JWT, or neither. | |
| 2. Resolves the caller's tier and identity. | |
| 3. Increments the per-tier daily rate counter (Redis). | |
| 4. Returns an `AuthContext` with rate-limit info. | |
| Endpoints: | |
| POST /auth/register -> issue OTP | |
| POST /auth/verify -> verify OTP, issue JWT pair | |
| POST /auth/refresh -> rotate refresh token | |
| POST /auth/logout -> revoke current session | |
| GET /auth/me -> current user + rate info | |
| """ | |
| from __future__ import annotations | |
| import logging | |
| import secrets | |
| from datetime import datetime, timezone | |
| import bcrypt | |
| import jwt | |
| from fastapi import APIRouter, Cookie, Depends, HTTPException, Request, Response, status | |
| from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer | |
| from sqlalchemy import select | |
| from sqlalchemy.ext.asyncio import AsyncSession | |
| from app.config import get_settings | |
| from app.db.session import get_db_session | |
| from app.models.auth import ApiKey, AuthAuditLog, Session, User | |
| from app.models.enums import ( | |
| APIKeyStatus, | |
| AuthEventType, | |
| AuthTier, | |
| OTPPurpose, | |
| UserStatus, | |
| ) | |
| from app.models.schemas import ( | |
| AuthContext, | |
| RateLimitInfo, | |
| RegisterRequest, | |
| RegisterResponse, | |
| TokenResponse, | |
| UserOut, | |
| VerifyRequest, | |
| ) | |
| from app.services.auth.email_service import render_otp_email, send_email | |
| from app.services.auth.jwt_service import ( | |
| create_access_token, | |
| create_refresh_token, | |
| decode_token, | |
| hash_refresh_token, | |
| refresh_token_expiry, | |
| ) | |
| from app.services.auth.otp_service import issue_otp, verify_otp | |
| from app.services.auth.rate_limiter import ( | |
| RateInfo, | |
| check_and_increment, | |
| client_fingerprint, | |
| ) | |
| logger = logging.getLogger("ukweli.auth") | |
| router = APIRouter(prefix="/auth", tags=["Auth"]) | |
| # `auto_error=False` so unauthenticated requests on protected endpoints | |
| # can still fall through to the public tier when `min_tier=public`. | |
| bearer_scheme = HTTPBearer(auto_error=False) | |
| # Tier ordering (lowest -> highest). Higher tiers inherit lower-tier | |
| # capabilities, so `require_auth("registered")` accepts API-key users too. | |
| _TIER_ORDER = { | |
| AuthTier.PUBLIC.value: 0, | |
| AuthTier.REGISTERED.value: 1, | |
| AuthTier.API.value: 2, | |
| AuthTier.ENTERPRISE.value: 3, | |
| } | |
| REFRESH_COOKIE_NAME = "ukweli_refresh_token" | |
| # --------------------------------------------------------------------------- | |
| # API key helpers (used by require_auth + api_keys router) | |
| # --------------------------------------------------------------------------- | |
| def hash_api_secret(raw_secret: str) -> bytes: | |
| """bcrypt-hash an API key secret. The raw secret is never stored.""" | |
| return bcrypt.hashpw(raw_secret.encode("utf-8"), bcrypt.gensalt()) | |
| def generate_api_key_pair() -> tuple[str, str]: | |
| """ | |
| Generate a (key_id, secret) pair. | |
| Format: | |
| key_id -> ukweli_live_<22 url-safe chars> | |
| secret -> sk_<43 url-safe chars> | |
| """ | |
| key_id = f"ukweli_live_{secrets.token_urlsafe(16)}" | |
| secret = f"sk_{secrets.token_urlsafe(32)}" | |
| return key_id, secret | |
| def verify_api_secret(raw_secret: str, secret_hash: bytes) -> bool: | |
| try: | |
| return bcrypt.checkpw(raw_secret.encode("utf-8"), secret_hash) | |
| except (ValueError, TypeError): | |
| return False | |
| # --------------------------------------------------------------------------- | |
| # Audit-log helper | |
| # --------------------------------------------------------------------------- | |
| async def _audit( | |
| db: AsyncSession, | |
| *, | |
| event_type: AuthEventType, | |
| success: bool, | |
| user_id=None, | |
| api_key_id=None, | |
| fingerprint: str | None = None, | |
| ip: str | None = None, | |
| user_agent: str | None = None, | |
| error: str | None = None, | |
| metadata: dict | None = None, | |
| ) -> None: | |
| """Append an immutable audit-log row. Best-effort, never raises.""" | |
| try: | |
| entry = AuthAuditLog( | |
| event_type=event_type.value, | |
| user_id=user_id, | |
| api_key_id=api_key_id, | |
| fingerprint=fingerprint, | |
| ip_address=ip, | |
| user_agent=user_agent, | |
| success=success, | |
| error_message=error, | |
| metadata_extra=metadata or {}, | |
| ) | |
| db.add(entry) | |
| await db.flush() | |
| except Exception as exc: # pragma: no cover - audit is best-effort | |
| logger.warning("Audit log write failed: %s", exc) | |
| def _client_ip(request: Request) -> str | None: | |
| fwd = request.headers.get("X-Forwarded-For") | |
| if fwd: | |
| return fwd.split(",")[0].strip() | |
| if request.client: | |
| return request.client.host | |
| return None | |
| # --------------------------------------------------------------------------- | |
| # Public identity object (the "anonymous user" container) | |
| # --------------------------------------------------------------------------- | |
| class _Anon: | |
| """Marker returned for the public tier — no DB row, just a fingerprint.""" | |
| # --------------------------------------------------------------------------- | |
| # Core dependency — require_auth(min_tier=...) | |
| # --------------------------------------------------------------------------- | |
| async def require_auth( | |
| request: Request, | |
| response: Response, | |
| min_tier: str = AuthTier.PUBLIC.value, | |
| db: AsyncSession = Depends(get_db_session), | |
| credentials: HTTPAuthorizationCredentials | None = Depends(bearer_scheme), | |
| ) -> AuthContext: | |
| """ | |
| Unified auth dependency. | |
| Resolution order: | |
| 1. X-API-Key + X-API-Secret headers -> API tier | |
| 2. Authorization: Bearer <jwt> -> Registered tier | |
| 3. Otherwise -> Public tier (fingerprint) | |
| """ | |
| min_rank = _TIER_ORDER.get(min_tier, 0) | |
| settings = get_settings() | |
| ip = _client_ip(request) | |
| ua = request.headers.get("User-Agent") | |
| lang = request.headers.get("Accept-Language") | |
| fingerprint = client_fingerprint(ip=ip, user_agent=ua, accept_language=lang) | |
| # ---- 1. API key path ---- | |
| api_key_id_hdr = request.headers.get("X-API-Key") | |
| api_secret_hdr = request.headers.get("X-API-Secret") | |
| if api_key_id_hdr and api_secret_hdr: | |
| result = await db.execute( | |
| select(ApiKey).where( | |
| ApiKey.key_id == api_key_id_hdr, ApiKey.status == APIKeyStatus.ACTIVE.value | |
| ) | |
| ) | |
| key_rec = result.scalar_one_or_none() | |
| if key_rec is None or not verify_api_secret(api_secret_hdr, key_rec.secret_hash): | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.AUTH_FAILED, | |
| success=False, | |
| api_key_id=key_rec.id if key_rec else None, | |
| fingerprint=fingerprint, | |
| ip=ip, | |
| user_agent=ua, | |
| error="invalid_api_key", | |
| ) | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Invalid or revoked API credentials.", | |
| ) | |
| if min_rank > _TIER_ORDER[AuthTier.API.value]: | |
| raise HTTPException( | |
| status_code=status.HTTP_403_FORBIDDEN, | |
| detail="This endpoint requires a higher tier.", | |
| ) | |
| # Bump daily counter, attach rate headers | |
| info = await check_and_increment(tier=AuthTier.API.value, key_suffix=key_rec.key_id) | |
| if not info.allowed: | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.RATE_LIMIT_EXCEEDED, | |
| success=False, | |
| api_key_id=key_rec.id, | |
| fingerprint=fingerprint, | |
| ip=ip, | |
| user_agent=ua, | |
| ) | |
| _attach_rate_headers(response, info) | |
| raise HTTPException( | |
| status_code=status.HTTP_429_TOO_MANY_REQUESTS, | |
| detail=f"Daily rate limit exceeded ({info.limit} requests/day).", | |
| headers=_rate_headers(info), | |
| ) | |
| key_rec.requests_today = info.used | |
| key_rec.last_used_at = datetime.now(timezone.utc) | |
| await db.flush() | |
| _attach_rate_headers(response, info) | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.AUTH_FAILED if not info.allowed else AuthEventType.AUTH_FAILED, | |
| success=True, | |
| api_key_id=key_rec.id, | |
| fingerprint=fingerprint, | |
| ip=ip, | |
| user_agent=ua, | |
| ) | |
| return AuthContext( | |
| tier=AuthTier.API.value, | |
| fingerprint=fingerprint, | |
| api_key=_api_key_to_out(key_rec), | |
| rate=_to_rate_info(info), | |
| ) | |
| # ---- 2. JWT path ---- | |
| if credentials is not None: | |
| try: | |
| payload = decode_token(credentials.credentials) | |
| except jwt.ExpiredSignatureError: | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Token expired.", | |
| headers={"WWW-Authenticate": "Bearer"}, | |
| ) | |
| except jwt.InvalidTokenError as exc: | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail=f"Invalid token: {exc}", | |
| ) | |
| if payload.get("type") != "access": | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Wrong token type; use the access token.", | |
| ) | |
| user_id = payload.get("sub") | |
| result = await db.execute(select(User).where(User.id == user_id)) | |
| user = result.scalar_one_or_none() | |
| if user is None or user.status != UserStatus.ACTIVE.value: | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="User not found or inactive.", | |
| ) | |
| if min_rank > _TIER_ORDER[AuthTier.REGISTERED.value]: | |
| raise HTTPException( | |
| status_code=status.HTTP_403_FORBIDDEN, | |
| detail="This endpoint requires a higher tier.", | |
| ) | |
| info = await check_and_increment( | |
| tier=AuthTier.REGISTERED.value, key_suffix=str(user.id) | |
| ) | |
| if not info.allowed: | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.RATE_LIMIT_EXCEEDED, | |
| success=False, | |
| user_id=user.id, | |
| fingerprint=fingerprint, | |
| ip=ip, | |
| user_agent=ua, | |
| ) | |
| _attach_rate_headers(response, info) | |
| raise HTTPException( | |
| status_code=status.HTTP_429_TOO_MANY_REQUESTS, | |
| detail=f"Daily rate limit exceeded ({info.limit} requests/day). " | |
| f"Generate an API key for higher limits.", | |
| headers=_rate_headers(info), | |
| ) | |
| _attach_rate_headers(response, info) | |
| return AuthContext( | |
| tier=AuthTier.REGISTERED.value, | |
| fingerprint=fingerprint, | |
| user=UserOut.model_validate(user), | |
| rate=_to_rate_info(info), | |
| ) | |
| # ---- 3. Public / anonymous ---- | |
| if min_rank > _TIER_ORDER[AuthTier.PUBLIC.value]: | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Authentication required for this endpoint.", | |
| headers={"WWW-Authenticate": "Bearer"}, | |
| ) | |
| info = await check_and_increment(tier=AuthTier.PUBLIC.value, key_suffix=fingerprint) | |
| if not info.allowed: | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.RATE_LIMIT_EXCEEDED, | |
| success=False, | |
| fingerprint=fingerprint, | |
| ip=ip, | |
| user_agent=ua, | |
| ) | |
| _attach_rate_headers(response, info) | |
| raise HTTPException( | |
| status_code=status.HTTP_429_TOO_MANY_REQUESTS, | |
| detail="Public rate limit exceeded. Sign up for 25x more queries.", | |
| headers=_rate_headers(info), | |
| ) | |
| _attach_rate_headers(response, info) | |
| return AuthContext( | |
| tier=AuthTier.PUBLIC.value, | |
| fingerprint=fingerprint, | |
| rate=_to_rate_info(info), | |
| ) | |
| # --------------------------------------------------------------------------- | |
| # Helpers for headers | |
| # --------------------------------------------------------------------------- | |
| def _to_rate_info(info: RateInfo) -> RateLimitInfo: | |
| return RateLimitInfo( | |
| tier=info.tier, | |
| limit=info.limit, | |
| used=info.used, | |
| remaining=info.remaining, | |
| reset_at=info.reset_at, | |
| soft_warning=info.soft_warning, | |
| ) | |
| def _rate_headers(info: RateInfo) -> dict[str, str]: | |
| return { | |
| "X-RateLimit-Limit": str(info.limit), | |
| "X-RateLimit-Remaining": str(info.remaining), | |
| "X-RateLimit-Reset": str(int(info.reset_at.timestamp())), | |
| "X-RateLimit-Tier": info.tier, | |
| } | |
| def _attach_rate_headers(response: Response, info: RateInfo) -> None: | |
| for k, v in _rate_headers(info).items(): | |
| response.headers[k] = v | |
| def _api_key_to_out(rec: ApiKey) -> dict: | |
| return { | |
| "id": rec.id, | |
| "key_id": rec.key_id, | |
| "name": rec.name, | |
| "description": rec.description, | |
| "tier": rec.tier, | |
| "quota_per_day": rec.quota_per_day, | |
| "requests_today": rec.requests_today, | |
| "last_used_at": rec.last_used_at, | |
| "status": rec.status, | |
| "created_at": rec.created_at, | |
| "expires_at": rec.expires_at, | |
| } | |
| # --------------------------------------------------------------------------- | |
| # Endpoints | |
| # --------------------------------------------------------------------------- | |
| async def register( | |
| request: Request, | |
| payload: RegisterRequest, | |
| db: AsyncSession = Depends(get_db_session), | |
| ) -> RegisterResponse: | |
| """ | |
| Step 1 of the OTP flow. Always returns success unless the email is | |
| obviously invalid — we don't leak whether an account exists. | |
| """ | |
| settings = get_settings() | |
| email = payload.email.lower() | |
| lang = payload.language | |
| code = await issue_otp(db, email=email, purpose=OTPPurpose.SIGNUP) | |
| email_msg = render_otp_email(code, language=lang) | |
| email_msg.to = email | |
| await send_email(email_msg) | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.REGISTER_REQUEST, | |
| success=True, | |
| ip=_client_ip(request), | |
| user_agent=request.headers.get("User-Agent"), | |
| metadata={"language": lang, "email_domain": email.split("@")[-1]}, | |
| ) | |
| return RegisterResponse( | |
| status="otp_sent", | |
| message=( | |
| "Verification code sent. Check your inbox." | |
| if lang == "en" | |
| else "Msimbo umetumwa. Angalia barua pepe yako." | |
| ), | |
| expires_in_seconds=settings.otp_expire_minutes * 60, | |
| ) | |
| async def verify( | |
| request: Request, | |
| response: Response, | |
| payload: VerifyRequest, | |
| db: AsyncSession = Depends(get_db_session), | |
| ) -> TokenResponse: | |
| """ | |
| Step 2 of the OTP flow. Verifies the code, creates the user if | |
| needed, and issues access + refresh tokens. | |
| """ | |
| settings = get_settings() | |
| email = payload.email.lower() | |
| ok = await verify_otp(db, email=email, code=payload.otp, purpose=OTPPurpose.SIGNUP) | |
| if not ok: | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.VERIFY_FAILED, | |
| success=False, | |
| ip=_client_ip(request), | |
| user_agent=request.headers.get("User-Agent"), | |
| error="invalid_or_expired_otp", | |
| ) | |
| raise HTTPException( | |
| status_code=status.HTTP_400_BAD_REQUEST, | |
| detail="Invalid or expired verification code.", | |
| ) | |
| # Find or create user | |
| result = await db.execute(select(User).where(User.email == email)) | |
| user = result.scalar_one_or_none() | |
| if user is None: | |
| user = User( | |
| email=email, | |
| email_verified=True, | |
| tier=AuthTier.REGISTERED.value, | |
| status=UserStatus.ACTIVE.value, | |
| ) | |
| db.add(user) | |
| await db.flush() | |
| else: | |
| user.email_verified = True | |
| user.status = UserStatus.ACTIVE.value | |
| user.last_login_at = datetime.now(timezone.utc) | |
| await db.flush() | |
| access, _ = create_access_token(str(user.id), user.tier, user.email) | |
| refresh, _ = create_refresh_token(str(user.id)) | |
| # Persist the refresh-token hash so we can rotate / revoke it | |
| sess = Session( | |
| user_id=user.id, | |
| refresh_token_hash=hash_refresh_token(refresh), | |
| user_agent=request.headers.get("User-Agent"), | |
| ip_address=_client_ip(request), | |
| expires_at=refresh_token_expiry(), | |
| ) | |
| db.add(sess) | |
| await db.flush() | |
| # Set httpOnly refresh cookie | |
| response.set_cookie( | |
| key=REFRESH_COOKIE_NAME, | |
| value=refresh, | |
| httponly=True, | |
| secure=settings.is_production, | |
| samesite="lax", | |
| max_age=settings.jwt_refresh_token_expire_days * 24 * 60 * 60, | |
| path="/auth", | |
| ) | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.VERIFY_SUCCESS, | |
| success=True, | |
| user_id=user.id, | |
| ip=_client_ip(request), | |
| user_agent=request.headers.get("User-Agent"), | |
| ) | |
| rate = await check_and_increment( | |
| tier=AuthTier.REGISTERED.value, key_suffix=str(user.id) | |
| ) | |
| return TokenResponse( | |
| access_token=access, | |
| expires_in=settings.jwt_access_token_expire_minutes * 60, | |
| user=UserOut.model_validate(user), | |
| rate_info={ | |
| "tier": rate.tier, | |
| "limit": rate.limit, | |
| "used": rate.used, | |
| "remaining": rate.remaining, | |
| "reset_at": rate.reset_at.isoformat(), | |
| }, | |
| ) | |
| async def refresh( | |
| request: Request, | |
| response: Response, | |
| db: AsyncSession = Depends(get_db_session), | |
| ukweli_refresh_token: str | None = Cookie(default=None), | |
| ) -> TokenResponse: | |
| """ | |
| Rotate a refresh token. Old token is revoked; a new pair is issued. | |
| """ | |
| settings = get_settings() | |
| raw = ukweli_refresh_token | |
| if raw is None: | |
| body = await request.body() | |
| # Allow JSON body fallback for clients that can't send cookies | |
| import json | |
| try: | |
| raw = json.loads(body or b"{}").get("refresh_token") | |
| except Exception: | |
| raw = None | |
| if not raw: | |
| raise HTTPException( | |
| status_code=status.HTTP_401_UNAUTHORIZED, | |
| detail="Missing refresh token.", | |
| ) | |
| try: | |
| payload = decode_token(raw) | |
| except jwt.ExpiredSignatureError: | |
| raise HTTPException(status_code=401, detail="Refresh token expired.") | |
| except jwt.InvalidTokenError: | |
| raise HTTPException(status_code=401, detail="Invalid refresh token.") | |
| if payload.get("type") != "refresh": | |
| raise HTTPException(status_code=401, detail="Wrong token type.") | |
| incoming_hash = hash_refresh_token(raw) | |
| result = await db.execute( | |
| select(Session).where( | |
| Session.refresh_token_hash == incoming_hash, | |
| Session.revoked_at.is_(None), | |
| ) | |
| ) | |
| sess = result.scalar_one_or_none() | |
| if sess is None or sess.expires_at < datetime.now(timezone.utc): | |
| raise HTTPException(status_code=401, detail="Refresh token revoked or expired.") | |
| user = ( | |
| await db.execute(select(User).where(User.id == sess.user_id)) | |
| ).scalar_one_or_none() | |
| if user is None or user.status != UserStatus.ACTIVE.value: | |
| raise HTTPException(status_code=401, detail="User inactive.") | |
| # Rotate | |
| sess.revoked_at = datetime.now(timezone.utc) | |
| new_access, _ = create_access_token(str(user.id), user.tier, user.email) | |
| new_refresh, _ = create_refresh_token(str(user.id)) | |
| db.add( | |
| Session( | |
| user_id=user.id, | |
| refresh_token_hash=hash_refresh_token(new_refresh), | |
| user_agent=request.headers.get("User-Agent"), | |
| ip_address=_client_ip(request), | |
| expires_at=refresh_token_expiry(), | |
| ) | |
| ) | |
| response.set_cookie( | |
| key=REFRESH_COOKIE_NAME, | |
| value=new_refresh, | |
| httponly=True, | |
| secure=settings.is_production, | |
| samesite="lax", | |
| max_age=settings.jwt_refresh_token_expire_days * 24 * 60 * 60, | |
| path="/auth", | |
| ) | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.LOGIN_REFRESH, | |
| success=True, | |
| user_id=user.id, | |
| ip=_client_ip(request), | |
| user_agent=request.headers.get("User-Agent"), | |
| ) | |
| rate = await check_and_increment( | |
| tier=AuthTier.REGISTERED.value, key_suffix=str(user.id) | |
| ) | |
| return TokenResponse( | |
| access_token=new_access, | |
| expires_in=settings.jwt_access_token_expire_minutes * 60, | |
| user=UserOut.model_validate(user), | |
| rate_info={ | |
| "tier": rate.tier, | |
| "limit": rate.limit, | |
| "used": rate.used, | |
| "remaining": rate.remaining, | |
| "reset_at": rate.reset_at.isoformat(), | |
| }, | |
| ) | |
| async def logout( | |
| request: Request, | |
| response: Response, | |
| db: AsyncSession = Depends(get_db_session), | |
| ukweli_refresh_token: str | None = Cookie(default=None), | |
| ): | |
| """Revoke the current refresh token + clear the cookie.""" | |
| if ukweli_refresh_token: | |
| incoming_hash = hash_refresh_token(ukweli_refresh_token) | |
| result = await db.execute( | |
| select(Session).where( | |
| Session.refresh_token_hash == incoming_hash, | |
| Session.revoked_at.is_(None), | |
| ) | |
| ) | |
| sess = result.scalar_one_or_none() | |
| if sess is not None: | |
| sess.revoked_at = datetime.now(timezone.utc) | |
| await _audit( | |
| db, | |
| event_type=AuthEventType.LOGOUT, | |
| success=True, | |
| user_id=sess.user_id, | |
| ip=_client_ip(request), | |
| user_agent=request.headers.get("User-Agent"), | |
| ) | |
| response.delete_cookie(REFRESH_COOKIE_NAME, path="/auth") | |
| return Response(status_code=status.HTTP_204_NO_CONTENT) | |
| async def me( | |
| auth: AuthContext = Depends(require_auth), | |
| ) -> dict: | |
| """ | |
| Return the current identity. Always 200 — public callers get a | |
| bare-bones response with just `tier` and `rate_info`. | |
| """ | |
| return { | |
| "tier": auth.tier, | |
| "user": auth.user.model_dump(mode="json") if auth.user else None, | |
| "rate_info": auth.rate.model_dump(mode="json"), | |
| } | |