Add post-diarization review checkpoint
This commit is contained in:
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user