import json import logging from typing import Any, Literal, TypedDict import sqlalchemy as sa from pydantic import AwareDatetime, BaseModel, ValidationError from app.adapters import aws_secrets_manager from app.adapters.db import db from app.config import settings from app.core.encrypter import FernetEncrypter, decrypt, encrypt from app.dsp.enums import DSPClientName from app.dsp.models import DSPClient from app.fandata.enums import FanCredentialsStatus from app.fandata.models import FanCredentials from app.fandata.types import FanBatch from app.pipeline import services from app.pipeline.enums import RunSource from app.pipeline.exceptions import PipelineAlreadyRunningError from app.pipeline.fan_collect import collect from app.pipeline.fan_fanout import fanout logger = logging.getLogger(__name__) # --------------------------------------------------------------------------- # fan_fanout # --------------------------------------------------------------------------- class FanFanoutResult(TypedDict): run_id: str skipped: Literal[False] class FanFanoutSkippedResult(TypedDict): run_id: None skipped: Literal[True] reason: str class FanFanoutEvent(BaseModel): dsp_client: DSPClientName token_status: FanCredentialsStatus | None = FanCredentialsStatus.active not_collected_since: AwareDatetime | None = None max_consecutive_failures: int | None = None limit: int | None = None def fan_fanout_handler( event: dict[str, Any], _context: Any ) -> FanFanoutResult | FanFanoutSkippedResult: logger.info( "fan_fanout invoked", extra={"event": json.dumps(event, default=str)}, ) try: params = FanFanoutEvent.model_validate(event) except ValidationError as exc: logger.error("Invalid event payload: %s", exc) raise ValueError(f"Invalid event payload: {exc}") from exc with db.autocommit(): client = DSPClient.query.where( DSPClient.name == params.dsp_client ).one_or_none() if client is None: raise ValueError(f"DSP client '{params.dsp_client}' not found") try: with db.transaction(): run = services.create_queued_run( source=RunSource.scheduled, dsp_client_id=client.id, token_status=params.token_status, not_collected_since=params.not_collected_since, max_consecutive_failures=params.max_consecutive_failures, limit=params.limit, ) except PipelineAlreadyRunningError as exc: logger.warning( "Pipeline already running for %s — skipping", params.dsp_client, extra={"dsp_client": params.dsp_client}, ) return {"run_id": None, "skipped": True, "reason": str(exc)} logger.info( "Run created: run_id=%s client=%s", run.id, params.dsp_client, extra={"run_id": str(run.id), "dsp_client": params.dsp_client}, ) fanout(run.id) return {"run_id": str(run.id), "skipped": False} # --------------------------------------------------------------------------- # fan_collect # --------------------------------------------------------------------------- class BatchItemFailure(TypedDict): itemIdentifier: str class FanCollectResult(TypedDict): batchItemFailures: list[BatchItemFailure] def fan_collect_handler(event: dict[str, Any], _context: Any) -> FanCollectResult: records = event.get("Records", []) logger.info( "fan_collect invoked: %d records", len(records), extra={"record_count": len(records)}, ) failed: list[BatchItemFailure] = [] for record in records: message_id: str = record.get("messageId", "unknown") try: batch = FanBatch.model_validate_json(record["body"]) except (KeyError, ValidationError) as exc: logger.error( "Failed to parse SQS record %s: %s", message_id, exc, extra={"message_id": message_id}, ) failed.append({"itemIdentifier": message_id}) continue try: collect(batch=batch) except Exception: logger.exception( "Unexpected error processing record %s run_id=%s", message_id, batch.run_id, extra={"message_id": message_id, "run_id": str(batch.run_id)}, ) failed.append({"itemIdentifier": message_id}) if failed: logger.warning( "%d/%d records failed", len(failed), len(records), extra={"failed": len(failed), "total": len(records)}, ) return {"batchItemFailures": failed} # --------------------------------------------------------------------------- # prepare_rotation # --------------------------------------------------------------------------- class PrepareRotationResult(TypedDict): new_key_id: str total_keys: int def prepare_rotation_handler( _event: dict[str, Any], _context: Any ) -> PrepareRotationResult: logger.info("prepare_rotation invoked") if settings.encrypter_backend != "fernet": raise ValueError("prepare_rotation requires fernet backend") new_key = FernetEncrypter.generate_key() current_keys = aws_secrets_manager.get_fernet_keys() updated_keys = [new_key, *current_keys] aws_secrets_manager.put_fernet_keys(updated_keys) new_key_id = new_key.partition(":")[0] logger.info( "prepare_rotation complete: new_key_id=%s total_keys=%d", new_key_id, len(updated_keys), extra={"new_key_id": new_key_id, "total_keys": len(updated_keys)}, ) return { "new_key_id": new_key_id, "total_keys": len(updated_keys), } # --------------------------------------------------------------------------- # rotate_keys # --------------------------------------------------------------------------- class RotateKeysEvent(BaseModel): after_id: int = 0 batch_size: int = 5000 class RotateKeysResult(TypedDict): rotated: int errors: int next_after_id: int | None done: bool finalized: bool def _try_finalize( fernet_keys: list[str], active_prefix: str, done: bool, errors: int ) -> bool: if not done or errors: return False with db.autocommit(): has_remaining = FanCredentials.query.where( ~FanCredentials.refresh_token_encrypted.startswith(active_prefix) ).exists() if has_remaining: return False aws_secrets_manager.put_fernet_keys([fernet_keys[0]]) logger.info("rotate_keys finalized: removed %d old key(s)", len(fernet_keys) - 1) return True def rotate_keys_handler(event: dict[str, Any], _context: Any) -> RotateKeysResult: logger.info("rotate_keys invoked", extra={"event": event}) try: params = RotateKeysEvent.model_validate(event) except ValidationError as exc: raise ValueError(f"Invalid event payload: {exc}") from exc if settings.encrypter_backend != "fernet": logger.info("rotate_keys skipped: backend is not Fernet") return { "rotated": 0, "errors": 0, "next_after_id": None, "done": True, "finalized": False, } fernet_keys = aws_secrets_manager.get_fernet_keys() if len(fernet_keys) < 2: raise ValueError( "rotate_keys requires at least 2 keys in Secrets Manager — run prepare_rotation first" ) if fernet_keys[0] != settings.fernet_keys[0]: raise ValueError( "FERNET_KEYS env var is out of sync with Secrets Manager — redeploy Lambda after prepare_rotation" ) active_prefix = settings.fernet_keys[0].partition(":")[0] + ":" with db.autocommit(): page = ( FanCredentials.query.where(FanCredentials.id > params.after_id) .where(~FanCredentials.refresh_token_encrypted.startswith(active_prefix)) .order_by(FanCredentials.id) .limit(params.batch_size) .all() ) if not page: aws_secrets_manager.put_fernet_keys([fernet_keys[0]]) logger.info( "rotate_keys finalized: removed %d old key(s)", len(fernet_keys) - 1 ) return { "rotated": 0, "errors": 0, "next_after_id": None, "done": True, "finalized": True, } updates: list[tuple[int, str]] = [] errors = 0 for fan_credentials in page: try: updates.append( ( fan_credentials.id, encrypt(decrypt(fan_credentials.refresh_token_encrypted)), ) ) except Exception: logger.exception( "Failed to re-encrypt fan_credentials id=%s", fan_credentials.id, extra={"fan_credentials_id": fan_credentials.id}, ) errors += 1 if updates: with db.transaction(): db.session.execute( sa.update(FanCredentials), [ {"id": fan_id, "refresh_token_encrypted": token} for fan_id, token in updates ], ) done = len(page) < params.batch_size next_after_id = page[-1].id if not done else None finalized = _try_finalize(fernet_keys, active_prefix, done, errors) logger.info( "rotate_keys batch complete: rotated=%d errors=%d done=%s finalized=%s", len(updates), errors, done, finalized, extra={ "rotated": len(updates), "errors": errors, "done": done, "finalized": finalized, "next_after_id": next_after_id, }, ) return { "rotated": len(updates), "errors": errors, "next_after_id": next_after_id, "done": done, "finalized": finalized, }