"""#1970 PR3 叙事前置服务 narrative_service 单元测试(不依赖真实 PG/OSS/CosyVoice)。""" from __future__ import annotations from dataclasses import dataclass, field from types import SimpleNamespace from typing import Any import pytest from apps.api.app.services import narrative_service as ns from apps.api.app.services.narrative_service import ( NarrativeError, _resolve_voice, _save_tts_job_as_voice_asset, prepare_narrative_voice, ) from packages.adapters.sqlalchemy_impl.models import ScriptModel from packages.domain.tts_job import TTSJob, TTSJobStatus # ── fakes ────────────────────────────────────────────────────────────────── @dataclass class FakeProfile: id: str = "prof-1" user_id: str = "u1" voice_id: str = "cv-voice-1" class FakeCloneRepo: def __init__(self, profile: FakeProfile | None = None): self._profile = profile def get(self, pid: str) -> FakeProfile | None: if self._profile and self._profile.id == pid: return self._profile return None class FakeQuery: def __init__(self, script: ScriptModel | None): self._script = script def filter(self, *conditions): # 服务端写 filter(...).filter(...) 链式调用;归属/ID 已在 FakeDb 构造时过滤 return self def first(self): return self._script class FakeDb: def __init__(self, script: ScriptModel | None, *, current_user: str = "u1", query_script_id: str = "script-1"): self._script = script self._current_user = current_user self._query_script_id = query_script_id def query(self, model): visible = self._script if visible is not None and (visible.user_id != self._current_user or visible.id != self._query_script_id): visible = None return FakeQuery(visible) def _make_script(*, user_id: str = "u1", content: str = "这是一段口播文案", title: str = "测试文案", tags=None): return ScriptModel( id="script-1", user_id=user_id, title=title, content=content, segments=[], tags=tags if tags is not None else ["带货"], ) @dataclass class FakeLibrary: id: str = "lib-voice" project_id: str = "p1" kind: Any = field(default_factory=lambda: SimpleNamespace(value="voice")) @dataclass class FakeProject: id: str = "p1" class FakeProjectRepo: def __init__(self, projects=None): self._projects = projects if projects is not None else [FakeProject()] def find_accessible_projects(self, user_id): return self._projects class FakeLibraryRepo: def __init__(self, libs=None): self._libs = libs if libs is not None else [FakeLibrary()] self.created: list = [] def find_by_project(self, project_id): return list(self._libs) def create(self, library): self.created.append(library) return library @dataclass class FakeAsset: id: str = "asset-new" duration: float | None = 12.0 class FakeAssetRepo: def __init__(self): self.created: list = [] def create(self, asset): wrapped = FakeAsset(id="asset-new", duration=getattr(asset, "duration", None)) self.created.append(asset) return wrapped class FakeStorage: def __init__(self, *, fail_download: bool = False): self.fail_download = fail_download self.uploaded: list = [] def download_asset(self, source, dest_path) -> bool: if self.fail_download: return False dest_path.write_bytes(b"FAKEAUDIO") return True def upload_file(self, path, key, content_type="", **kwargs): self.uploaded.append((key, content_type)) def delete_file(self, key): pass class FakeTTSRepo: def __init__(self, job: TTSJob): self.job = job self.saved: list[TTSJob] = [] def create(self, job: TTSJob) -> TTSJob: self.saved.append(job) self.job = job return job def update(self, job: TTSJob) -> TTSJob: self.job = job return job def get(self, job_id: str) -> TTSJob | None: return self.job if self.job.id == job_id else None class FakeCosyVoice: pass def _make_completed_job() -> TTSJob: job = TTSJob.create( user_id="u1", input_text="这是一段口播文案", voice_id="cv-voice-1", voice_clone_profile_id="", format="mp3", sample_rate=22050, ) job.mark_processing() job.mark_completed( output_audio_url="https://oss/tts/output/job-1.mp3", output_audio_key="tts/output/job-1.mp3", duration=12.5, ) return job # ── _resolve_voice ───────────────────────────────────────────────────────── class TestResolveVoice: def test_preset_returns_id_directly_when_no_profile(self): voice_id, clone_id = _resolve_voice( user_id="u1", tts_voice_id="longxiaochun", tts_voice_source="preset", voice_clone_repository=FakeCloneRepo(None), ) assert voice_id == "longxiaochun" assert clone_id == "" def test_preset_id_that_is_clone_profile_uuid_resolves(self): repo = FakeCloneRepo(FakeProfile()) voice_id, clone_id = _resolve_voice( user_id="u1", tts_voice_id="prof-1", tts_voice_source="preset", voice_clone_repository=repo, ) assert voice_id == "cv-voice-1" assert clone_id == "prof-1" def test_clone_source(self): voice_id, clone_id = _resolve_voice( user_id="u1", tts_voice_id="prof-1", tts_voice_source="clone", voice_clone_repository=FakeCloneRepo(FakeProfile()), ) assert voice_id == "cv-voice-1" assert clone_id == "prof-1" def test_clone_missing_404(self): with pytest.raises(NarrativeError) as ei: _resolve_voice( user_id="u1", tts_voice_id="nope", tts_voice_source="clone", voice_clone_repository=FakeCloneRepo(None), ) assert ei.value.status_code == 404 def test_clone_other_user_403(self): repo = FakeCloneRepo(FakeProfile(user_id="someone-else")) with pytest.raises(NarrativeError) as ei: _resolve_voice( user_id="u1", tts_voice_id="prof-1", tts_voice_source="clone", voice_clone_repository=repo, ) assert ei.value.status_code == 403 def test_clone_not_ready_400(self): repo = FakeCloneRepo(FakeProfile(voice_id="")) with pytest.raises(NarrativeError) as ei: _resolve_voice( user_id="u1", tts_voice_id="prof-1", tts_voice_source="clone", voice_clone_repository=repo, ) assert ei.value.status_code == 400 # ── save asset ───────────────────────────────────────────────────────────── class TestSaveVoiceAsset: def _deps(self, **storage_kw): return dict( user_id="u1", name="测试配音", project_repository=FakeProjectRepo(), asset_library_repository=FakeLibraryRepo(), asset_repository=FakeAssetRepo(), storage_service=FakeStorage(**storage_kw), ) def test_save_creates_asset(self): job = _make_completed_job() deps = self._deps() asset = _save_tts_job_as_voice_asset(job=job, **deps) assert asset.id == "asset-new" assert deps["asset_repository"].created[0].mime_type == "audio/mpeg" assert deps["storage_service"].uploaded[0][0] == "uploads/voice/tts/" + job.id + ".mp3" def test_no_project_raises(self): job = _make_completed_job() deps = self._deps() deps["project_repository"] = FakeProjectRepo(projects=[]) with pytest.raises(NarrativeError): _save_tts_job_as_voice_asset(job=job, **deps) def test_download_fail_raises_502(self): job = _make_completed_job() deps = self._deps(fail_download=True) with pytest.raises(NarrativeError) as ei: _save_tts_job_as_voice_asset(job=job, **deps) assert ei.value.status_code == 502 def test_job_without_output_raises(self): job = TTSJob.create(user_id="u1", input_text="x", voice_id="v", voice_clone_profile_id="") with pytest.raises(NarrativeError) as ei: _save_tts_job_as_voice_asset(job=job, **self._deps()) assert ei.value.status_code == 502 # ── prepare_narrative_voice 主流程(monkeypatch workflow) ───────────────── class TestPrepareNarrativeVoice: def _deps(self, db_script=None, *, has_script=True, clone_profile=None, storage_fail=False, points_enabled=False): job = _make_completed_job() return dict( db=FakeDb(db_script if db_script is not None else (_make_script() if has_script else None)), user_id="u1", script_id="script-1", tts_voice_id="longxiaochun", tts_voice_source="preset", tts_repository=FakeTTSRepo(job), cosyvoice_service=FakeCosyVoice(), voice_clone_repository=FakeCloneRepo(clone_profile), asset_repository=FakeAssetRepo(), asset_library_repository=FakeLibraryRepo(), project_repository=FakeProjectRepo(), storage_service=FakeStorage(fail_download=storage_fail), points_enabled=points_enabled, ) def test_success_returns_context(self, monkeypatch): captured = {} class FakeWorkflow: def __init__(self, *, repository, cosyvoice_service): captured["repo"] = repository self._repo = repository def start_synthesis(self, job_id): job = self._repo.get(job_id) job.mark_processing() job.mark_completed( output_audio_url="https://oss/tts/output/x.mp3", output_audio_key="tts/output/x.mp3", duration=12.5, ) return job def poll_and_process_synthesis(self, job_id, timeout=120.0): return self._repo.get(job_id) monkeypatch.setattr(ns, "TTSWorkflowService", FakeWorkflow) ctx = prepare_narrative_voice(**self._deps()) assert ctx.voice_asset_id == "asset-new" assert ctx.tts_job_id assert ctx.audio_duration == pytest.approx(12.5) assert ctx.script.tags == ["带货"] def test_script_missing_404(self): deps = self._deps(has_script=False) with pytest.raises(NarrativeError) as ei: prepare_narrative_voice(**deps) assert ei.value.status_code == 404 def test_script_other_user_404(self): deps = self._deps(db_script=_make_script(user_id="other")) with pytest.raises(NarrativeError) as ei: prepare_narrative_voice(**deps) assert ei.value.status_code == 404 def test_empty_content_400(self): deps = self._deps(db_script=_make_script(content=" ")) with pytest.raises(NarrativeError) as ei: prepare_narrative_voice(**deps) assert ei.value.status_code == 400 def test_synth_failure_raises_502(self, monkeypatch): class FailingWorkflow: def __init__(self, *, repository, cosyvoice_service): self._repo = repository def start_synthesis(self, job_id): raise RuntimeError("cosyvoice down") def process_synthesis_failure(self, job_id, error): return None monkeypatch.setattr(ns, "TTSWorkflowService", FailingWorkflow) with pytest.raises(NarrativeError) as ei: prepare_narrative_voice(**self._deps()) assert ei.value.status_code == 502 assert "配音合成失败" in ei.value.message def test_points_insufficient_402(self, monkeypatch): class FakePoints: def deduct_points(self, *a, **k): return {"success": False, "balance": 0} monkeypatch.setattr(ns, "PointsService", lambda: FakePoints()) deps = self._deps(points_enabled=True) with pytest.raises(NarrativeError) as ei: prepare_narrative_voice(**deps) assert ei.value.status_code == 402 def test_points_refund_on_failure(self, monkeypatch): class FakePoints: def __init__(self): self.refunded = 0 def deduct_points(self, *a, **k): return {"success": True, "balance": 100} def refund_points(self, user_id, amount, source, db, ref_id="", **k): self.refunded += amount points = FakePoints() monkeypatch.setattr(ns, "PointsService", lambda: points) class FailingWorkflow: def __init__(self, *, repository, cosyvoice_service): pass def start_synthesis(self, job_id): raise RuntimeError("boom") def process_synthesis_failure(self, job_id, error): return None monkeypatch.setattr(ns, "TTSWorkflowService", FailingWorkflow) deps = self._deps(points_enabled=True) with pytest.raises(NarrativeError): prepare_narrative_voice(**deps) assert points.refunded > 0 def test_clone_source_resolves_profile(self, monkeypatch): captured = {} class FakeWorkflow: def __init__(self, *, repository, cosyvoice_service): self._repo = repository captured["cosy"] = cosyvoice_service def start_synthesis(self, job_id): job = self._repo.get(job_id) captured["voice_id"] = job.voice_id job.mark_processing() job.mark_completed( output_audio_url="https://oss/tts/output/x.mp3", output_audio_key="tts/output/x.mp3", duration=12.5, ) return job def poll_and_process_synthesis(self, job_id, timeout=120.0): return self._repo.get(job_id) monkeypatch.setattr(ns, "TTSWorkflowService", FakeWorkflow) deps = self._deps(clone_profile=FakeProfile()) deps["tts_voice_id"] = "prof-1" deps["tts_voice_source"] = "clone" prepare_narrative_voice(**deps) assert captured["voice_id"] == "cv-voice-1" if __name__ == "__main__": import pytest as _pytest _pytest.main([__file__, "-q"])