Add post-diarization review checkpoint

This commit is contained in:
2026-09-17 19:16:49 +02:00
parent 874ab5b68d
commit 405fa6cdf4
2 changed files with 85 additions and 0 deletions
+42
View File
@@ -9,6 +9,7 @@ from unittest.mock import patch
from scripts import run_mvp_meeting as cli
from src.meeting_lab.audio import PreparedAudio
from src.meeting_lab.diarization.backend import DiarizationResult
from src.meeting_lab.llm.ollama import OllamaGeneration
from src.meeting_lab.protocol.generate_direct_protocol import generate_direct_protocol
from src.meeting_lab.models.meeting_context import load_meeting_context
@@ -82,6 +83,27 @@ def fake_prepare(source, destination, **kwargs):
return PreparedAudio(source, source.suffix.removeprefix("."), destination, "ffmpeg", "ffmpeg")
def fake_diarize(audio, output, mode, **kwargs):
output.mkdir(parents=True, exist_ok=True)
metadata = output / "metadata.json"
ordinary = output / "diarization.rttm"
exclusive = output / "exclusive_diarization.rttm"
turns = output / "turns.json"
exclusive_turns = output / "exclusive_turns.json"
details = {"speaker_count": 2, "runtime_seconds": 0.1}
metadata.write_text(json.dumps(details), encoding="utf-8")
ordinary.write_text("", encoding="utf-8")
exclusive.write_text("", encoding="utf-8")
turns.write_text("[]", encoding="utf-8")
exclusive_turns.write_text(
json.dumps([
{"start": 0, "end": 0.5, "speaker_id": "SPEAKER_00"},
{"start": 0.5, "end": 1, "speaker_id": "SPEAKER_01"},
]), encoding="utf-8"
)
return DiarizationResult(output, metadata, ordinary, exclusive, turns, exclusive_turns, details)
def fake_protocol(transcript, context, **kwargs):
rendered_context = load_meeting_context(context)
assert rendered_context.meeting_id == "programmatic-test"
@@ -107,6 +129,26 @@ class MvpApiTests(unittest.TestCase):
model="test:model",
)
def test_post_diarization_checkpoint_skips_protocol_generation(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
config = replace(self.config(root), diarization="cpu", stop_after_diarization=True)
with (
patch.object(mvp_api, "prepare_audio", side_effect=fake_prepare),
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(mvp_api, "diarize_audio", side_effect=fake_diarize),
patch.object(mvp_api, "generate_direct_protocol") as generate,
):
result = mvp_api.run_mvp_meeting(config, meeting_context=context_data())
self.assertEqual(result.exit_code, 0)
self.assertIsNone(result.protocol_path)
generate.assert_not_called()
metadata = json.loads((result.run_dir / "run_metadata.json").read_text())
self.assertEqual(metadata["status"], "awaiting_speaker_review")
self.assertEqual(metadata["speaker_review"]["speaker_count"], 2)
self.assertTrue((result.run_dir / "diarization/transcript_diarized.json").is_file())
def test_programmatic_context_is_persisted_without_source_yaml(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)