from __future__ import annotations import hashlib import json from datetime import datetime, timezone from typing import Any from sqlalchemy.orm import Session from app.core.database import SessionLocal from app.models.generation_cache import GenerationCache from app.services.provider_quota import record_usage from app.services.provider_registry import normalize_task_type def stable_hash(value: Any) -> str: if isinstance(value, str): payload = value else: payload = json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":")) return hashlib.sha256(payload.encode("utf-8")).hexdigest() def scene_plan_cache_key( *, document_id: str, material_hash: str, video_mode: str, target_duration_seconds: int | float, language: str, evidence_level: str, prompt_version: str, ) -> str: return _key( "scene_plan", { "document_id": document_id, "material_hash": material_hash, "video_mode": video_mode, "target_duration_seconds": target_duration_seconds, "language": language, "evidence_level": evidence_level, "prompt_version": prompt_version, }, ) def tts_cache_key(*, provider: str, voice: str, language: str, narration_text: str) -> str: return _key( "tts", { "provider": provider, "voice": voice, "language": language, "narration_text": narration_text, }, ) def study_material_cache_key( *, document_id: str, content_hash: str, task_type: str, evidence_level: str, prompt_version: str, ) -> str: return _key( normalize_task_type(task_type), { "document_id": document_id, "content_hash": content_hash, "task_type": normalize_task_type(task_type), "evidence_level": evidence_level, "prompt_version": prompt_version, }, ) def image_asset_cache_key(*, prompt_hash: str, style: str, aspect_ratio: str) -> str: return _key( "image_asset", { "prompt_hash": prompt_hash, "style": style, "aspect_ratio": aspect_ratio, }, ) def get_cached_generation( *, cache_key: str, task_type: str, provider: str, user_id: str | None = None, db: Session | None = None, ) -> GenerationCache | None: session, close = _db_or_new(db) try: cached = session.get(GenerationCache, cache_key) if cached is None: return None if cached.task_type != normalize_task_type(task_type): return None if cached.expires_at is not None and _as_aware(cached.expires_at) <= datetime.now(timezone.utc): return None record_usage( provider=provider, task_type=task_type, request_units=0, response_units=0, estimated_cost_usd=0, status="success", user_id=user_id, cache_hit=True, db=session, ) if close: session.commit() session.refresh(cached) return cached finally: if close: session.close() def store_generation_cache( *, cache_key: str, task_type: str, provider: str, input_hash: str, output_json: dict[str, Any] | None = None, output_text: str | None = None, output_file_path: str | None = None, metadata_json: dict[str, Any] | None = None, expires_at: datetime | None = None, db: Session | None = None, ) -> GenerationCache: session, close = _db_or_new(db) try: cached = session.get(GenerationCache, cache_key) if cached is None: cached = GenerationCache(cache_key=cache_key) cached.task_type = normalize_task_type(task_type) cached.provider = provider cached.input_hash = input_hash cached.output_json = output_json cached.output_text = output_text cached.output_file_path = output_file_path cached.metadata_json = metadata_json or {} cached.expires_at = expires_at session.add(cached) session.commit() session.refresh(cached) return cached finally: if close: session.close() def invalidate_document_generation_cache(*, document_id: str, db: Session | None = None) -> int: session, close = _db_or_new(db) try: rows = session.query(GenerationCache).all() deleted = 0 for row in rows: metadata = row.metadata_json or {} key_mentions_document = document_id in row.cache_key metadata_mentions_document = metadata.get("document_id") == document_id if key_mentions_document or metadata_mentions_document: session.delete(row) deleted += 1 session.commit() return deleted finally: if close: session.close() def _key(prefix: str, payload: dict[str, Any]) -> str: return f"{prefix}:{stable_hash(payload)}" def _db_or_new(db: Session | None) -> tuple[Session, bool]: if db is not None: return db, False return SessionLocal(), True def _as_aware(value: datetime) -> datetime: if value.tzinfo is None: return value.replace(tzinfo=timezone.utc) return value