feat: add post-diarization speaker mapping workflow
This commit is contained in:
@@ -142,6 +142,65 @@ class MvpApiTests(unittest.TestCase):
|
||||
],
|
||||
)
|
||||
self.assertTrue(all(event.progress is None for event in events))
|
||||
self.assertEqual(protocol_generator.call_args.kwargs["num_ctx"], 32_768)
|
||||
self.assertEqual(
|
||||
protocol_generator.call_args.kwargs["safe_input_token_budget"],
|
||||
29_000,
|
||||
)
|
||||
|
||||
def test_protocol_only_regeneration_reuses_diarized_artifacts(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
run_dir = root / "existing-run"
|
||||
diarization_dir = run_dir / "diarization"
|
||||
diarization_dir.mkdir(parents=True)
|
||||
source = diarization_dir / "transcript_diarized.json"
|
||||
source.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"text": "SPEAKER_00: Existing statement.\n",
|
||||
"segments": [
|
||||
{
|
||||
"start": 0.0,
|
||||
"end": 1.0,
|
||||
"speaker_id": "SPEAKER_00",
|
||||
"text": "Existing statement.",
|
||||
}
|
||||
],
|
||||
"speaker_labels_anonymous": True,
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
source_before = source.read_bytes()
|
||||
mapped_context = context_data()
|
||||
mapped_context["speaker_mappings"] = {"SPEAKER_00": "person-1"}
|
||||
with (
|
||||
patch.object(mvp_api, "prepare_audio") as preparation,
|
||||
patch.object(mvp_api, "transcribe_audio") as transcription,
|
||||
patch.object(mvp_api, "diarize_audio") as diarization,
|
||||
patch.object(
|
||||
mvp_api, "generate_direct_protocol", side_effect=fake_protocol
|
||||
) as protocol,
|
||||
):
|
||||
result = mvp_api.regenerate_mvp_protocol(
|
||||
run_dir,
|
||||
meeting_context=mapped_context,
|
||||
model="qwen3.8:27b",
|
||||
protocol_num_ctx=32_768,
|
||||
protocol_safe_input_token_budget=29_000,
|
||||
)
|
||||
|
||||
self.assertEqual(result.exit_code, 0)
|
||||
self.assertEqual(result.protocol_path, run_dir / "protocol.md")
|
||||
preparation.assert_not_called()
|
||||
transcription.assert_not_called()
|
||||
diarization.assert_not_called()
|
||||
self.assertEqual(protocol.call_args.args[0], source)
|
||||
self.assertEqual(protocol.call_args.kwargs["num_ctx"], 32_768)
|
||||
self.assertEqual(source.read_bytes(), source_before)
|
||||
persisted = load_meeting_context(run_dir / "context/meeting_context.yaml")
|
||||
self.assertEqual(persisted.speaker_mappings, {"SPEAKER_00": "person-1"})
|
||||
|
||||
def test_failure_emits_terminal_failure_event(self):
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
|
||||
@@ -2,9 +2,10 @@ import json
|
||||
import tempfile
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import Mock
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from src.meeting_lab.diarization.alignment import diarized_transcript_text
|
||||
from src.meeting_lab.llm import ollama
|
||||
from src.meeting_lab.llm.ollama import OllamaGeneration
|
||||
from src.meeting_lab.protocol.direct_protocol_prompt import (
|
||||
COMPACT_DIARIZED_PROTOCOL_INSTRUCTION,
|
||||
@@ -136,6 +137,81 @@ class ProtocolInputBudgetTests(unittest.TestCase):
|
||||
self.assertEqual(result.transcript_input, compact)
|
||||
self.assertEqual(call.call_count, 1)
|
||||
|
||||
def test_mapped_speakers_and_statements_reach_final_ollama_payload(self) -> None:
|
||||
diarized = {
|
||||
"text": "",
|
||||
"segments": [
|
||||
{
|
||||
"start": 0.0,
|
||||
"end": 1.0,
|
||||
"speaker_id": "SPEAKER_00",
|
||||
"text": "We will run the trial on Wednesday.",
|
||||
},
|
||||
{
|
||||
"start": 1.0,
|
||||
"end": 2.0,
|
||||
"speaker_id": "SPEAKER_01",
|
||||
"text": "I will prepare the raw materials before then.",
|
||||
},
|
||||
{
|
||||
"start": 2.0,
|
||||
"end": 3.0,
|
||||
"speaker_id": "SPEAKER_00",
|
||||
"text": "Good. Anna owns the material preparation.",
|
||||
},
|
||||
],
|
||||
"speaker_labels_anonymous": True,
|
||||
"alignment_source": "exclusive_diarization",
|
||||
}
|
||||
context = {
|
||||
"schema_version": "1",
|
||||
"meeting": {
|
||||
"meeting_id": "speaker-test",
|
||||
"title": "Speaker test",
|
||||
"language": "en",
|
||||
},
|
||||
"participants": [
|
||||
{"participant_id": "martin", "display_name": "Martin"},
|
||||
{"participant_id": "anna", "display_name": "Anna"},
|
||||
],
|
||||
"speaker_mappings": {"SPEAKER_00": "martin", "SPEAKER_01": "anna"},
|
||||
"mentioned_people": [],
|
||||
"organization": {"departments": []},
|
||||
"known_entities": {},
|
||||
}
|
||||
response = Mock()
|
||||
response.raise_for_status.return_value = None
|
||||
response.json.return_value = {"response": "# Meeting Protocol\n", "done": True}
|
||||
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
transcript = self._write(root, diarized)
|
||||
context_path = root / "context.yaml"
|
||||
context_path.write_text(json.dumps(context), encoding="utf-8")
|
||||
with patch.object(ollama.requests, "post", return_value=response) as post:
|
||||
result = generate_direct_protocol(
|
||||
transcript,
|
||||
context_path,
|
||||
model="qwen3.8:27b",
|
||||
model_check=Mock(return_value={}),
|
||||
)
|
||||
|
||||
prompt = post.call_args.kwargs["json"]["prompt"]
|
||||
self.assertEqual(result.exact_prompt, prompt)
|
||||
self.assertIn("- SPEAKER_00: Martin (participant_id: martin)", prompt)
|
||||
self.assertIn("- SPEAKER_01: Anna (participant_id: anna)", prompt)
|
||||
self.assertIn("SPEAKER_00: We will run the trial on Wednesday.", prompt)
|
||||
self.assertIn(
|
||||
"SPEAKER_01: I will prepare the raw materials before then.", prompt
|
||||
)
|
||||
self.assertIn("SPEAKER_00: Good. Anna owns the material preparation.", prompt)
|
||||
self.assertNotIn("Martin: We will run the trial on Wednesday.", prompt)
|
||||
self.assertNotIn("Anna: I will prepare the raw materials before then.", prompt)
|
||||
self.assertIn("autoritativen SPEAKER_XX-zu-Teilnehmer-Zuordnungen", prompt)
|
||||
self.assertIn("Ich-Zusage", prompt)
|
||||
self.assertIn("nur erwähnten Personen", prompt)
|
||||
self.assertIn("keine persönliche Verantwortung", prompt)
|
||||
|
||||
def test_plain_fallback_selected_when_diarized_compact_exceeds_budget(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
@@ -180,6 +256,13 @@ class ProtocolInputBudgetTests(unittest.TestCase):
|
||||
)
|
||||
self.assertTrue(result.runtime_metadata["fallback_used"])
|
||||
self.assertEqual(result.transcript_input, plain)
|
||||
self.assertNotIn("SPEAKER_00", result.exact_prompt)
|
||||
self.assertNotIn("SPEAKER_01", result.exact_prompt)
|
||||
self.assertFalse(result.runtime_metadata["speaker_attribution_available"])
|
||||
self.assertEqual(
|
||||
result.runtime_metadata["speaker_attribution_loss_reason"],
|
||||
"plain_transcript_fallback",
|
||||
)
|
||||
|
||||
def test_oversized_plain_transcript_fails_before_any_network_call(self) -> None:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
@@ -225,6 +308,7 @@ class ProtocolInputBudgetTests(unittest.TestCase):
|
||||
)
|
||||
self.assertFalse(result.runtime_metadata["fallback_used"])
|
||||
self.assertFalse(result.runtime_metadata["diarization_enabled"])
|
||||
self.assertIsNone(result.runtime_metadata["speaker_attribution_available"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user