From 405fa6cdf401825caa2536f9afe2b3d299c56592 Mon Sep 17 00:00:00 2001 From: Martin Date: Thu, 17 Sep 2026 19:16:49 +0200 Subject: [PATCH] Add post-diarization review checkpoint --- src/meeting_lab/orchestration/mvp.py | 43 ++++++++++++++++++++++++++++ tests/test_mvp_api.py | 42 +++++++++++++++++++++++++++ 2 files changed, 85 insertions(+) diff --git a/src/meeting_lab/orchestration/mvp.py b/src/meeting_lab/orchestration/mvp.py index c520dd9..8cb05b7 100644 --- a/src/meeting_lab/orchestration/mvp.py +++ b/src/meeting_lab/orchestration/mvp.py @@ -65,6 +65,7 @@ class MvpMeetingConfig: diarization_runtime: str = "native" diarization_container_image: str | None = None diarization_container_args: Sequence[str] = () + stop_after_diarization: bool = False @dataclass(frozen=True) @@ -130,6 +131,10 @@ def regenerate_mvp_protocol( ) protocol_path = persist_protocol_generation(run_dir, result, context=context) except Exception as exc: + _record_protocol_generation_attempt( + run_dir, status="failed", runtime_seconds=time.perf_counter() - started, + failure={"type": type(exc).__name__, "message": str(exc)}, + ) _emit( progress_sink, "failed", @@ -138,11 +143,40 @@ def regenerate_mvp_protocol( message=f"protocol_generation: {type(exc).__name__}: {exc}", ) raise + _record_protocol_generation_attempt( + run_dir, status="completed", runtime_seconds=time.perf_counter() - started + ) _emit(progress_sink, "protocol_generation", "completed", started) _emit(progress_sink, "completed", "completed", started) return MvpRunResult(0, run_dir, protocol_path) +def _record_protocol_generation_attempt( + run_dir: Path, + *, + status: str, + runtime_seconds: float, + failure: dict[str, str] | None = None, +) -> None: + """Record downstream attempts without changing completed upstream artifacts.""" + metadata_path = run_dir / "run_metadata.json" + if not metadata_path.is_file(): + return + metadata = json.loads(metadata_path.read_text(encoding="utf-8")) + attempts = metadata.setdefault("protocol_generation_attempts", []) + if not isinstance(attempts, list): + attempts = metadata["protocol_generation_attempts"] = [] + attempt: dict[str, Any] = { + "status": status, + "runtime_seconds": round(runtime_seconds, 3), + } + if failure is not None: + attempt["failure"] = failure + attempts.append(attempt) + metadata["status"] = "completed" if status == "completed" else "awaiting_speaker_review" + _write_json(metadata_path, metadata) + + def create_unique_run_dir( output_root: Path, meeting_name: str, @@ -419,6 +453,15 @@ def run_mvp_meeting( ) _emit(progress_sink, "diarization", "completed", overall_started) + if config.stop_after_diarization: + metadata["status"] = "awaiting_speaker_review" + metadata["speaker_review"] = { + "status": "awaiting_optional_mapping", + "speaker_count": metadata["diarization"].get("speaker_count"), + } + _emit(progress_sink, "completed", "completed", overall_started) + return MvpRunResult(0, run_dir, None) + current_stage = "protocol_generation" stage_started = time.perf_counter() _emit(progress_sink, "protocol_generation", "started", overall_started) diff --git a/tests/test_mvp_api.py b/tests/test_mvp_api.py index bc7d561..9520139 100644 --- a/tests/test_mvp_api.py +++ b/tests/test_mvp_api.py @@ -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)