DocDoeAI / tests /test_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.85 kB
from __future__ import annotations
from pathlib import Path
def test_tts_cache_hit_avoids_provider_call(client, monkeypatch, tmp_path: Path) -> None:
from app.core import config
from app.schemas.video import VideoSceneAudioInput, VideoScenePlanAudioInput
from app.services import tts_provider
monkeypatch.setenv("TTS_PROVIDER", "mock")
monkeypatch.setenv("TTS_OUTPUT_DIR", str(tmp_path / "audio"))
config.get_settings.cache_clear()
class CountingProvider(tts_provider.MockTTSProvider):
calls = 0
def generate_scene_audio(self, **kwargs):
type(self).calls += 1
return super().generate_scene_audio(**kwargs)
monkeypatch.setattr(tts_provider, "get_tts_provider", lambda provider_name=None: CountingProvider())
scene = VideoSceneAudioInput(
scene_id=1,
id=1,
type="explanation",
duration_seconds=2,
screen_text="Photosynthesis",
voice_text="Plants make food using sunlight.",
subtitle_text="Plants make food using sunlight.",
keywords=["plants"],
)
plan = VideoScenePlanAudioInput(
title="Biology",
duration_minutes=0.1,
language="English",
style="clean_explainer",
visual_style="clean_explainer",
video_format="16:9",
voice_mode="english_soft",
scenes=[scene],
)
first = tts_provider.generate_audio_for_scene_plan(
scene_plan=plan,
voice_mode="english_soft",
voice="nila",
language="English",
provider_name="mock",
video_id="video-one",
user_id="usr_demo_student",
use_cache=True,
)
second = tts_provider.generate_audio_for_scene_plan(
scene_plan=plan,
voice_mode="english_soft",
voice="nila",
language="English",
provider_name="mock",
video_id="video-two",
user_id="usr_demo_student",
use_cache=True,
)
assert CountingProvider.calls == 1
assert first["audio_files"][0]["cache_hit"] is False
assert second["audio_files"][0]["cache_hit"] is True
def test_scene_plan_cache_hit_avoids_regeneration(client) -> None:
from app.core.database import SessionLocal
from app.services.generation_cache import (
get_cached_generation,
scene_plan_cache_key,
store_generation_cache,
)
key = scene_plan_cache_key(
document_id="doc_1",
material_hash="hash-a",
video_mode="exam_study_video",
target_duration_seconds=600,
language="English",
evidence_level="document_only",
prompt_version="v1",
)
with SessionLocal() as db:
store_generation_cache(
db=db,
cache_key=key,
task_type="scene_plan",
provider="openrouter",
input_hash="hash-a",
output_json={"title": "Cached plan", "scenes": []},
)
cached = get_cached_generation(
db=db,
cache_key=key,
task_type="scene_plan",
provider="openrouter",
)
assert cached is not None
assert cached.output_json == {"title": "Cached plan", "scenes": []}
def test_cache_key_changes_when_prompt_version_changes() -> None:
from app.services.generation_cache import study_material_cache_key
v1 = study_material_cache_key(
document_id="doc_1",
content_hash="content-a",
task_type="quiz_generation",
evidence_level="document_only",
prompt_version="v1",
)
v2 = study_material_cache_key(
document_id="doc_1",
content_hash="content-a",
task_type="quiz_generation",
evidence_level="document_only",
prompt_version="v2",
)
assert v1 != v2
def test_cache_hit_records_usage_with_cache_hit_true(client) -> None:
from app.core.database import SessionLocal
from app.models.provider_usage_log import ProviderUsageLog
from app.services.generation_cache import get_cached_generation, store_generation_cache
with SessionLocal() as db:
store_generation_cache(
db=db,
cache_key="notes:abc",
task_type="notes_generation",
provider="openrouter",
input_hash="abc",
output_text="cached notes",
)
cached = get_cached_generation(
db=db,
cache_key="notes:abc",
task_type="notes_generation",
provider="openrouter",
user_id="usr_demo_student",
)
logs = db.query(ProviderUsageLog).filter(ProviderUsageLog.cache_hit.is_(True)).all()
assert cached is not None
assert len(logs) == 1
assert logs[0].request_units == 0
assert logs[0].estimated_cost_usd == 0
def test_cache_hit_records_usage_when_helper_owns_session(client) -> None:
from app.core.database import SessionLocal
from app.models.provider_usage_log import ProviderUsageLog
from app.services.generation_cache import get_cached_generation, store_generation_cache
store_generation_cache(
cache_key="notes:owned-session",
task_type="notes_generation",
provider="openrouter",
input_hash="owned-session",
output_text="cached notes",
)
cached = get_cached_generation(
cache_key="notes:owned-session",
task_type="notes_generation",
provider="openrouter",
user_id="usr_demo_student",
)
with SessionLocal() as db:
logs = (
db.query(ProviderUsageLog)
.filter(
ProviderUsageLog.task_type == "notes_generation",
ProviderUsageLog.cache_hit.is_(True),
)
.all()
)
assert cached is not None
assert any(log.provider == "openrouter" for log in logs)