| 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 |
|
|