447 lines
16 KiB
Python
447 lines
16 KiB
Python
"""FastAPI dependencies (db session, storage, publisher, service-token auth)."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from collections.abc import AsyncIterator
|
|
from typing import Annotated, Any
|
|
from uuid import UUID
|
|
|
|
from fastapi import Depends, Header, HTTPException, Query, Request, status
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from src.contract_check.core.api_keys import hash_api_key
|
|
from src.contract_check.core.auth import (
|
|
AuthError,
|
|
TokenExpiredError,
|
|
TokenInvalidError,
|
|
verify_access_token,
|
|
)
|
|
from src.contract_check.core.auth_refresh import RefreshTokenStore
|
|
from src.contract_check.core.config import get_settings
|
|
from src.contract_check.core.db.models import PasskeyCredential, User
|
|
from src.contract_check.core.db.repositories import (
|
|
ApiKeyRepository,
|
|
DocumentRepository,
|
|
PasskeyRepository,
|
|
UserRepository,
|
|
)
|
|
from src.contract_check.core.db.repositories.service_tokens import ServiceTokenRepository
|
|
from src.contract_check.core.db.session import get_session
|
|
from src.contract_check.core.logging import get_logger
|
|
from src.contract_check.core.mq.publisher import Publisher
|
|
from src.contract_check.core.notifications.publisher import NotificationPublisher
|
|
from src.contract_check.core.passkeys import PasskeyChallengeStore
|
|
from src.contract_check.core.rate_limit import RateLimiter, RateLimitResult
|
|
from src.contract_check.core.s3.port import Storage
|
|
from src.contract_check.core.security.passwords import hash_password_async
|
|
from src.contract_check.core.tokens import hash_token
|
|
|
|
log = get_logger(__name__)
|
|
|
|
|
|
async def get_db_session() -> AsyncIterator[AsyncSession]:
|
|
async for session in get_session():
|
|
yield session
|
|
|
|
|
|
AsyncSessionDep = Annotated[AsyncSession, Depends(get_db_session)]
|
|
|
|
|
|
def get_storage(request: Request) -> Storage:
|
|
storage: Storage = request.app.state.storage
|
|
return storage
|
|
|
|
|
|
StorageDep = Annotated[Storage, Depends(get_storage)]
|
|
|
|
|
|
def get_publisher(request: Request) -> Publisher:
|
|
publisher: Publisher = request.app.state.publisher
|
|
return publisher
|
|
|
|
|
|
PublisherDep = Annotated[Publisher, Depends(get_publisher)]
|
|
|
|
|
|
def get_notification_publisher(request: Request) -> NotificationPublisher:
|
|
publisher: NotificationPublisher = request.app.state.notification_publisher
|
|
return publisher
|
|
|
|
|
|
NotificationPublisherDep = Annotated[NotificationPublisher, Depends(get_notification_publisher)]
|
|
|
|
|
|
def get_redis(request: Request) -> Any:
|
|
"""Return the app-state async Redis client.
|
|
|
|
Set on app.state.redis in the api lifespan. Used by the refresh-token store.
|
|
"""
|
|
redis: Any = request.app.state.redis
|
|
return redis
|
|
|
|
|
|
def get_refresh_store(request: Request) -> RefreshTokenStore:
|
|
"""Build a RefreshTokenStore from the app-state Redis client.
|
|
|
|
Raises 503 if Redis is unavailable — webUI auth cannot function without it.
|
|
"""
|
|
redis: Any = getattr(request.app.state, "redis", None)
|
|
if redis is None:
|
|
raise HTTPException(status_code=503, detail="refresh-token store unavailable")
|
|
settings = get_settings()
|
|
return RefreshTokenStore(redis, ttl_seconds=settings.jwt_refresh_ttl_days * 24 * 3600)
|
|
|
|
|
|
RefreshStoreDep = Annotated[RefreshTokenStore, Depends(get_refresh_store)]
|
|
|
|
|
|
def get_passkey_challenge_store(request: Request) -> PasskeyChallengeStore:
|
|
"""Build a PasskeyChallengeStore from the app-state Redis client.
|
|
|
|
Raises 503 if Redis is unavailable — WebAuthn ceremonies cannot function
|
|
without server-side challenge storage.
|
|
"""
|
|
redis: Any = getattr(request.app.state, "redis", None)
|
|
if redis is None:
|
|
raise HTTPException(status_code=503, detail="passkey challenge store unavailable")
|
|
settings = get_settings()
|
|
return PasskeyChallengeStore(redis, ttl_seconds=settings.passkey_challenge_ttl_seconds)
|
|
|
|
|
|
PasskeyChallengeStoreDep = Annotated[PasskeyChallengeStore, Depends(get_passkey_challenge_store)]
|
|
|
|
|
|
async def require_service_token(
|
|
session: AsyncSessionDep,
|
|
authorization: Annotated[str | None, Header()] = None,
|
|
) -> None:
|
|
"""Validate `Authorization: Bearer <token>` against service_tokens table.
|
|
|
|
Raises 401 on missing/invalid/revoked token. Logs last_used_at on success.
|
|
"""
|
|
if not authorization or not authorization.lower().startswith("bearer "):
|
|
raise HTTPException(status_code=401, detail="Missing or invalid Authorization header")
|
|
|
|
raw = authorization[7:].strip()
|
|
token_hash = hash_token(raw)
|
|
repo = ServiceTokenRepository(session)
|
|
token_id = await repo.get_id_by_hash(token_hash)
|
|
if token_id is None:
|
|
raise HTTPException(status_code=401, detail="Invalid or revoked token")
|
|
|
|
await repo.bump_last_used(token_id)
|
|
await session.commit()
|
|
|
|
|
|
AuthDep = Annotated[None, Depends(require_service_token)]
|
|
|
|
|
|
async def get_or_create_user_for_telegram(
|
|
session: AsyncSession,
|
|
telegram_id: int,
|
|
*,
|
|
profile: dict[str, object] | None = None,
|
|
verified: bool = True,
|
|
) -> User:
|
|
"""Fetch or create a user identified by telegram_id.
|
|
|
|
On creation stores the profile snapshot, verification flag, and binding time.
|
|
On existing row optionally refreshes the profile snapshot.
|
|
"""
|
|
repo = UserRepository(session)
|
|
user = await repo.get_by_telegram_id(telegram_id)
|
|
if user is not None:
|
|
if profile:
|
|
await repo.update_telegram_profile(user.id, profile=profile, verified=verified)
|
|
await session.commit()
|
|
user.telegram_verified = verified
|
|
return user
|
|
|
|
user = await repo.create_telegram_user(telegram_id, profile=profile, verified=verified)
|
|
await session.commit()
|
|
return user
|
|
|
|
|
|
async def bind_telegram_to_user(
|
|
session: AsyncSession,
|
|
user_id: UUID,
|
|
telegram_id: int,
|
|
) -> None:
|
|
"""Link a Telegram id to an existing user (web -> Telegram).
|
|
|
|
Raises HTTPException 409 if the telegram_id is already bound to another user.
|
|
"""
|
|
repo = UserRepository(session)
|
|
if await repo.telegram_id_exists_for_other_user(telegram_id, user_id):
|
|
raise HTTPException(status_code=409, detail="Telegram id already bound to another account")
|
|
|
|
await repo.bind_telegram(user_id, telegram_id)
|
|
await session.commit()
|
|
|
|
|
|
async def set_user_password(session: AsyncSession, user_id: UUID, password: str) -> None:
|
|
"""Set/rotate a web password for a user (Telegram -> web UI access)."""
|
|
hashed = await hash_password_async(password)
|
|
await UserRepository(session).set_password(user_id, hashed)
|
|
await session.commit()
|
|
|
|
|
|
def get_rate_limiter(request: Request) -> RateLimiter:
|
|
"""Return the app-state rate limiter (Redis or in-memory fallback)."""
|
|
limiter: RateLimiter = request.app.state.rate_limiter
|
|
return limiter
|
|
|
|
|
|
RateLimiterDep = Annotated[RateLimiter, Depends(get_rate_limiter)]
|
|
|
|
|
|
def _auth_rate_limit_headers(result: RateLimitResult) -> dict[str, str]:
|
|
retry = int(result.retry_after_sec or 1)
|
|
return {"Retry-After": str(retry)}
|
|
|
|
|
|
async def require_auth_rate_limit(
|
|
request: Request,
|
|
rate_limiter: RateLimiter,
|
|
*,
|
|
email: str | None = None,
|
|
) -> None:
|
|
"""Throttle public auth endpoints by IP and, optionally, by recipient email.
|
|
|
|
Mail-sending endpoints should pass ``email`` so a single IP cannot bomb an
|
|
inbox; other endpoints pass only the IP bucket.
|
|
"""
|
|
settings = get_settings()
|
|
ip = request.client.host if request.client else "unknown"
|
|
ip_key = f"auth:ip:{ip}"
|
|
ip_result = await rate_limiter.allow(ip_key, settings.auth_rate_limit_ip_rps)
|
|
if not ip_result.allowed:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
|
detail="rate limit exceeded",
|
|
headers=_auth_rate_limit_headers(ip_result),
|
|
)
|
|
if email:
|
|
email_key = f"auth:email:{email.lower().strip()}"
|
|
email_result = await rate_limiter.allow(email_key, settings.auth_rate_limit_email_rps)
|
|
if not email_result.allowed:
|
|
raise HTTPException(
|
|
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
|
detail="rate limit exceeded",
|
|
headers=_auth_rate_limit_headers(email_result),
|
|
)
|
|
|
|
|
|
AuthRateLimiterDep = Annotated[RateLimiter, Depends(get_rate_limiter)]
|
|
|
|
|
|
async def fetch_user_by_email(session: AsyncSession, email: str) -> User | None:
|
|
"""Fetch a user by email (case-sensitive — normalize upstream). Returns None if not found."""
|
|
return await UserRepository(session).get_by_email(email)
|
|
|
|
|
|
async def fetch_passkey_credentials_for_user(
|
|
session: AsyncSession, user_id: UUID
|
|
) -> list[PasskeyCredential]:
|
|
"""All passkey credentials of a user, newest first."""
|
|
return await PasskeyRepository(session).list_for_user(user_id)
|
|
|
|
|
|
async def fetch_passkey_by_credential_id(
|
|
session: AsyncSession, credential_id: str
|
|
) -> PasskeyCredential | None:
|
|
"""Look up a passkey credential by its base64url credential id."""
|
|
result = await session.execute(
|
|
select(PasskeyCredential).where(PasskeyCredential.credential_id == credential_id)
|
|
)
|
|
return result.scalars().first()
|
|
|
|
|
|
async def fetch_user_by_id_full(session: AsyncSession, user_id: UUID) -> User | None:
|
|
"""Fetch a user by UUID including web-auth columns."""
|
|
return await UserRepository(session).get_by_id(user_id)
|
|
|
|
|
|
async def create_email_user(
|
|
session: AsyncSession, *, email: str, name: str | None, password_hash: str
|
|
) -> User:
|
|
"""Insert a new email/password user with 0 credits and return it."""
|
|
return await UserRepository(session).create_email_user(
|
|
email=email, name=name, password_hash=password_hash
|
|
)
|
|
|
|
|
|
def get_rate_limiter(request: Request) -> RateLimiter:
|
|
"""Return the app-state rate limiter (Redis in prod, Memory in tests)."""
|
|
limiter: RateLimiter = request.app.state.rate_limiter
|
|
return limiter
|
|
|
|
|
|
RateLimiterDep = Annotated[RateLimiter, Depends(get_rate_limiter)]
|
|
|
|
|
|
class ApiKeyAuth:
|
|
"""Validated B2B API key + its owning user id."""
|
|
|
|
def __init__(self, api_key_id: UUID, user_id: UUID, rate_limit_rps: int) -> None:
|
|
self.api_key_id = api_key_id
|
|
self.user_id = user_id
|
|
self.rate_limit_rps = rate_limit_rps
|
|
|
|
|
|
async def require_api_key(
|
|
session: AsyncSessionDep,
|
|
rate_limiter: RateLimiterDep,
|
|
x_api_key: Annotated[str | None, Header(alias="X-API-Key")] = None,
|
|
) -> ApiKeyAuth:
|
|
"""Validate `X-API-Key` header, rate-limit, and return key metadata.
|
|
|
|
Raises 401 on missing/invalid/revoked key. Raises 429 when rate-limited
|
|
or monthly quota exhausted.
|
|
"""
|
|
if not x_api_key:
|
|
raise HTTPException(status_code=401, detail="Missing X-API-Key header")
|
|
|
|
key_hash = hash_api_key(x_api_key)
|
|
repo = ApiKeyRepository(session)
|
|
key = await repo.get_by_hash(key_hash)
|
|
if key is None:
|
|
raise HTTPException(status_code=401, detail="Invalid API key")
|
|
|
|
if key.revoked:
|
|
raise HTTPException(status_code=401, detail="Revoked API key")
|
|
|
|
await repo.maybe_reset_monthly_quota(key.id)
|
|
|
|
# Check monthly quota (if configured) after potential reset.
|
|
monthly_quota, monthly_used = await repo.get_monthly_quota_state(key.id)
|
|
if monthly_quota is not None and int(monthly_used) >= int(monthly_quota):
|
|
raise HTTPException(status_code=429, detail="Monthly quota exceeded")
|
|
|
|
# Apply token-bucket rate limit per key id.
|
|
limit = int(key.rate_limit_rps or get_settings().b2b_default_rate_limit_rps)
|
|
rl_result = await rate_limiter.allow(f"rate_limit:{key.id}", limit)
|
|
if not rl_result.allowed:
|
|
retry_after = max(1, int(rl_result.retry_after_sec or 1))
|
|
raise HTTPException(
|
|
status_code=429,
|
|
detail="Rate limit exceeded",
|
|
headers={"Retry-After": str(retry_after)},
|
|
)
|
|
|
|
await repo.bump_last_used(key.id)
|
|
await session.commit()
|
|
|
|
return ApiKeyAuth(api_key_id=key.id, user_id=key.user_id, rate_limit_rps=limit)
|
|
|
|
|
|
ApiKeyAuthDep = Annotated[ApiKeyAuth, Depends(require_api_key)]
|
|
|
|
|
|
async def fetch_document_status_for_user(
|
|
session: AsyncSession, document_id: UUID, user_id: UUID
|
|
) -> dict[str, Any] | None:
|
|
"""Fetch document + report scoped to a specific user (B2B API)."""
|
|
status = await DocumentRepository(session).get_with_report_by_id_for_user(document_id, user_id)
|
|
if status is None:
|
|
return None
|
|
return {
|
|
"id": status.id,
|
|
"status": status.status,
|
|
"stage": status.stage,
|
|
"filename": status.filename,
|
|
"created_at": status.created_at,
|
|
"markdown": status.markdown,
|
|
"content_json": status.content_json,
|
|
"model_used": status.model_used,
|
|
"prompt_tokens": status.prompt_tokens,
|
|
"eval_tokens": status.eval_tokens,
|
|
"latency_ms": status.latency_ms,
|
|
"mime": status.mime,
|
|
"bytes": status.bytes_,
|
|
"report_created_at": status.report_created_at,
|
|
}
|
|
|
|
|
|
class CurrentUser:
|
|
"""Authenticated user extracted from a Bearer JWT."""
|
|
|
|
def __init__(self, user_id: UUID, telegram_id: int) -> None:
|
|
self.user_id = user_id
|
|
self.telegram_id = telegram_id
|
|
|
|
|
|
async def resolve_user_from_token(session: AsyncSession, token: str) -> CurrentUser:
|
|
"""Verify a raw access JWT and return the user (shared auth core)."""
|
|
try:
|
|
claims = verify_access_token(token)
|
|
except TokenExpiredError as exc:
|
|
raise HTTPException(status_code=401, detail="Token expired") from exc
|
|
except TokenInvalidError as exc:
|
|
raise HTTPException(status_code=401, detail="Invalid token") from exc
|
|
except AuthError as exc:
|
|
raise HTTPException(status_code=401, detail=str(exc)) from exc
|
|
|
|
# Ensure the user still exists (defense in depth: tokens are stateless,
|
|
# but a deleted user should not be able to use them).
|
|
if not await UserRepository(session).exists(claims.sub):
|
|
raise HTTPException(status_code=401, detail="User not found")
|
|
|
|
return CurrentUser(user_id=claims.sub, telegram_id=claims.telegram_id)
|
|
|
|
|
|
async def require_current_user(
|
|
session: AsyncSessionDep,
|
|
authorization: Annotated[str | None, Header()] = None,
|
|
) -> CurrentUser:
|
|
"""Validate `Authorization: Bearer <user-jwt>` and return the user.
|
|
|
|
This is the common auth gate for bot, web, and Mini App users.
|
|
"""
|
|
if not authorization or not authorization.lower().startswith("bearer "):
|
|
raise HTTPException(status_code=401, detail="Missing or invalid Authorization header")
|
|
|
|
token = authorization[7:].strip()
|
|
return await resolve_user_from_token(session, token)
|
|
|
|
|
|
CurrentUserDep = Annotated[CurrentUser, Depends(require_current_user)]
|
|
|
|
|
|
async def require_current_user_header_or_query(
|
|
session: AsyncSessionDep,
|
|
authorization: Annotated[str | None, Header()] = None,
|
|
access_token: Annotated[str | None, Query()] = None,
|
|
) -> CurrentUser:
|
|
"""Auth for EventSource-compatible endpoints (SSE).
|
|
|
|
Native `EventSource` cannot set request headers, so besides the usual
|
|
`Authorization: Bearer` header this also accepts the access JWT via the
|
|
`access_token` query parameter:
|
|
|
|
new EventSource(`/api/v1/reports/{id}/events?access_token=<jwt>`)
|
|
|
|
Caveat: query strings may end up in proxy access logs — pass short-lived
|
|
access tokens only, never refresh tokens.
|
|
"""
|
|
token: str | None = None
|
|
if authorization and authorization.lower().startswith("bearer "):
|
|
token = authorization[7:].strip()
|
|
elif access_token:
|
|
token = access_token.strip()
|
|
|
|
if not token:
|
|
raise HTTPException(
|
|
status_code=401,
|
|
detail="Missing auth: use Authorization header or access_token query param",
|
|
)
|
|
return await resolve_user_from_token(session, token)
|
|
|
|
|
|
EventSourceUserDep = Annotated[CurrentUser, Depends(require_current_user_header_or_query)]
|
|
|
|
|
|
async def get_credits(session: AsyncSession, user_id: UUID) -> int:
|
|
credits = await UserRepository(session).get_credits(user_id)
|
|
return credits if credits is not None else 0
|