DocDoeAI / app /services /generation_cache.py
asnannp's picture
Deploy backend cd4237ff: support routes + rate limit + exam_date nullable + upload 413 fix
7c6ffa6
Raw
History Blame Contribute Delete
5.45 kB
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