Add diarization and reusable MVP meeting pipeline

This commit is contained in:
2026-08-23 19:29:47 +02:00
parent f2d21c1faf
commit 8dab928763
18 changed files with 1697 additions and 180 deletions
+222
View File
@@ -0,0 +1,222 @@
import json
import os
import tempfile
import unittest
from pathlib import Path
from types import SimpleNamespace
from unittest.mock import patch
from src.meeting_lab.diarization.alignment import align_transcript, write_diarized_transcript
from src.meeting_lab.diarization.backend import (
DiarizationError,
diarize_audio,
run_container_pyannote,
select_device,
)
class FakeCuda:
def __init__(self, available: bool, *, name: str = "Test GPU", failure=None):
self.available = available
self.name = name
self.failure = failure
def is_available(self):
return self.available
def get_device_name(self, index):
if self.failure:
raise self.failure
return self.name
class FakeTorch:
def __init__(self, available: bool, *, failure=None):
self.cuda = FakeCuda(available, failure=failure)
self.probes = []
def device(self, name):
return name
def zeros(self, size, *, device):
self.probes.append(device)
if self.cuda.failure:
raise self.cuda.failure
return [0]
class DeviceSelectionTests(unittest.TestCase):
def test_auto_selects_usable_gpu(self):
torch = FakeTorch(True)
self.assertEqual(select_device("auto", torch), ("cuda", "Test GPU"))
self.assertEqual(torch.probes, ["cuda"])
def test_auto_falls_back_to_cpu(self):
self.assertEqual(select_device("auto", FakeTorch(False)), ("cpu", None))
self.assertEqual(
select_device("auto", FakeTorch(True, failure=RuntimeError("probe"))),
("cpu", None),
)
def test_explicit_cpu_does_not_probe_gpu(self):
torch = FakeTorch(True)
self.assertEqual(select_device("cpu", torch), ("cpu", None))
self.assertEqual(torch.probes, [])
def test_explicit_gpu_fails_when_unavailable(self):
with self.assertRaisesRegex(DiarizationError, "GPU is unavailable"):
select_device("gpu", FakeTorch(False))
class AlignmentTests(unittest.TestCase):
def test_exclusive_overlap_assigns_anonymous_speakers(self):
transcript = {
"text": "Original unchanged text.",
"segments": [
{"id": 0, "start": 0.0, "end": 4.0, "text": "Hallo"},
{"id": 1, "start": 4.0, "end": 6.0, "text": "Antwort"},
],
}
original = json.loads(json.dumps(transcript))
turns = [
{"start": 0.0, "end": 3.0, "speaker_id": "SPEAKER_00"},
{"start": 3.0, "end": 6.0, "speaker_id": "SPEAKER_01"},
]
derived = align_transcript(transcript, turns)
self.assertEqual(transcript, original)
self.assertEqual(derived["segments"][0]["speaker_id"], "SPEAKER_00")
self.assertEqual(derived["segments"][0]["speaker_overlap_seconds"], 3.0)
self.assertEqual(derived["segments"][1]["speaker_id"], "SPEAKER_01")
self.assertIn("SPEAKER_00: Hallo", derived["text"])
self.assertTrue(derived["speaker_labels_anonymous"])
self.assertEqual(derived["alignment_source"], "exclusive_diarization")
def test_speaker_aware_transcript_is_a_separate_artifact(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
source = root / "transcript.json"
turns = root / "exclusive_turns.json"
source_text = json.dumps(
{"text": "Original", "segments": [{"start": 0, "end": 1, "text": "Hi"}]}
)
source.write_text(source_text, encoding="utf-8")
turns.write_text(
json.dumps([{"start": 0, "end": 1, "speaker_id": "SPEAKER_07"}]),
encoding="utf-8",
)
json_path, text_path = write_diarized_transcript(source, turns, root / "derived")
self.assertEqual(source.read_text(encoding="utf-8"), source_text)
self.assertNotEqual(json_path, source)
self.assertIn("SPEAKER_07", json_path.read_text(encoding="utf-8"))
self.assertIn("SPEAKER_07", text_path.read_text(encoding="utf-8"))
class ContainerAdapterTests(unittest.TestCase):
def test_container_configuration_and_metadata_do_not_persist_token(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
audio = root / "audio.wav"
output = root / "output"
audio.write_bytes(b"audio")
observed = {}
def fake_runner(command, **kwargs):
observed["command"] = command
output.mkdir(exist_ok=True)
metadata = {
"backend": "pyannote.audio",
"model": "pyannote/speaker-diarization-community-1",
"requested_device_mode": "gpu",
"actual_device": "cuda",
"credentials_persisted": False,
}
(output / "metadata.json").write_text(json.dumps(metadata))
return SimpleNamespace(returncode=0, stdout="ok", stderr="")
with patch.dict("os.environ", {"HF_TOKEN": "secret-token"}):
result = run_container_pyannote(
audio,
output,
"gpu",
image="test/image",
container_args=("--device=/dev/test",),
runner=fake_runner,
uid_getter=lambda: 2345,
gid_getter=lambda: 3456,
)
command = observed["command"]
shell_command = command[-1]
persisted = "".join(
path.read_text(encoding="utf-8")
for path in output.iterdir()
if path.is_file()
)
self.assertNotIn("secret-token", persisted)
self.assertNotIn("secret-token", command)
self.assertIn("HF_TOKEN", command)
self.assertIn("chown -R 2345:3456 /output", shell_command)
self.assertIn("chmod -R u+rwX /output", shell_command)
self.assertNotIn("1000:1000", shell_command)
device_index = command.index("--device=/dev/test")
self.assertLess(device_index, command.index("test/image"))
self.assertFalse(result.metadata["credentials_persisted"])
self.assertEqual(result.metadata["runtime_adapter"], "container")
self.assertTrue(
all(os.access(path, os.W_OK) for path in (output, *output.rglob("*")))
)
def test_unwritable_container_artifact_is_rejected_before_metadata_update(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
audio = root / "audio.wav"
output = root / "output"
audio.write_bytes(b"audio")
def fake_runner(command, **kwargs):
output.mkdir(exist_ok=True)
metadata = output / "metadata.json"
metadata.write_text("{}", encoding="utf-8")
metadata.chmod(0o444)
return SimpleNamespace(returncode=0, stdout="", stderr="")
with patch(
"src.meeting_lab.diarization.backend.os.access",
side_effect=lambda path, mode: Path(path).name != "metadata.json",
):
with self.assertRaisesRegex(DiarizationError, "not writable"):
run_container_pyannote(
audio,
output,
"cpu",
image="test/image",
runner=fake_runner,
)
def test_orchestrator_dispatches_runtime(self):
with patch(
"src.meeting_lab.diarization.backend.run_container_pyannote"
) as container:
diarize_audio(
Path("audio.wav"),
Path("out"),
"cpu",
runtime="container",
container_image="image",
container_args=("--arg",),
)
container.assert_called_once_with(
Path("audio.wav"),
Path("out"),
"cpu",
image="image",
container_args=("--arg",),
)
if __name__ == "__main__":
unittest.main()
+18
View File
@@ -81,6 +81,24 @@ class TranscriptLoadingTests(unittest.TestCase):
class GeneratorTests(unittest.TestCase):
def test_prompt_requires_contextual_discussion_density_without_transcript_replay(self) -> None:
with tempfile.TemporaryDirectory() as directory:
transcript = Path(directory) / "transcript.json"
write_transcript(transcript)
result = generate_direct_protocol(
transcript,
model_check=Mock(return_value={}),
generation_call=Mock(return_value=generation()),
)
self.assertIn("vollständiges, strukturiertes", result.exact_prompt)
self.assertIn("relevante Diskussionsverläufe", result.exact_prompt)
self.assertIn("unterschiedliche Positionen", result.exact_prompt)
self.assertIn("Entscheidungsgrundlagen", result.exact_prompt)
self.assertIn("nicht am Meeting teilgenommen haben", result.exact_prompt)
self.assertIn("keine reine Wiedergabe des Transkripts", result.exact_prompt)
self.assertIn("nicht unnötig durch Wiederholungen", result.exact_prompt)
def test_optional_context_absent_and_generation_called_once(self) -> None:
with tempfile.TemporaryDirectory() as directory:
transcript = Path(directory) / "transcript.json"
+47
View File
@@ -11,8 +11,10 @@ from src.meeting_lab.extraction.extract_chunks import (
)
from src.meeting_lab.models.meeting_context import (
MeetingContextValidationError,
create_meeting_context,
load_meeting_context,
render_meeting_context_for_prompt,
serialize_meeting_context_yaml,
validate_meeting_context,
)
@@ -192,6 +194,51 @@ class MeetingContextTests(unittest.TestCase):
self.assertNotIn("responsible: Björn", prompt_context)
self.assertNotIn("responsible: Jovana", prompt_context)
def test_existing_context_without_speaker_mappings_remains_valid(self) -> None:
self.assertEqual(self.context.speaker_mappings, {})
self.assertIsNone(self.context.participant_for_speaker("SPEAKER_00"))
def test_explicit_speaker_mapping_is_authoritative(self) -> None:
data = copy.deepcopy(self.context.data)
participant = data["participants"][0]
data["speaker_mappings"] = {"SPEAKER_03": participant["participant_id"]}
context = create_meeting_context(data)
rendered = render_meeting_context_for_prompt(context)
self.assertEqual(
context.participant_for_speaker("SPEAKER_03")["participant_id"],
participant["participant_id"],
)
self.assertIsNone(context.participant_for_speaker("SPEAKER_04"))
self.assertIn("Confirmed diarization speaker mappings (authoritative)", rendered)
self.assertIn("Unmapped SPEAKER_XX labels must remain anonymous", rendered)
def test_speaker_mapping_must_reference_existing_participant(self) -> None:
data = copy.deepcopy(self.context.data)
data["speaker_mappings"] = {"SPEAKER_00": "unknown-person"}
with self.assertRaisesRegex(MeetingContextValidationError, "unknown participant"):
validate_meeting_context(data)
def test_speaker_mapping_label_must_use_pyannote_shape(self) -> None:
data = copy.deepcopy(self.context.data)
data["speaker_mappings"] = {
"Martin": data["participants"][0]["participant_id"]
}
with self.assertRaisesRegex(MeetingContextValidationError, "speaker label"):
validate_meeting_context(data)
def test_generated_context_yaml_is_deterministic_and_round_trips(self) -> None:
first = serialize_meeting_context_yaml(self.context)
second = serialize_meeting_context_yaml(self.context)
self.assertEqual(first, second)
SCRATCH_DIR.mkdir(exist_ok=True)
path = SCRATCH_DIR / "generated_context.yaml"
path.write_text(first, encoding="utf-8")
loaded = load_meeting_context(path)
self.assertEqual(loaded.data, self.context.data)
if __name__ == "__main__":
unittest.main()
+175
View File
@@ -0,0 +1,175 @@
import json
import subprocess
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
from scripts import run_mvp_meeting as cli
from src.meeting_lab.models.meeting_context import load_meeting_context
from src.meeting_lab.orchestration import mvp as mvp_api
from src.meeting_lab.orchestration.mvp import MvpMeetingConfig, MvpRunResult
from src.meeting_lab.protocol.generate_direct_protocol import DirectProtocolResult
from src.meeting_lab.transcription.whisper import TranscriptionError, TranscriptionResult
def context_data():
return {
"schema_version": "1",
"meeting": {
"meeting_id": "programmatic-test",
"title": "Programmatic Test",
"language": "de",
"date": None,
"objective": "API prüfen",
"notes": "",
},
"participants": [
{
"participant_id": "person-1",
"display_name": "Test Person",
"aliases": [],
"role": "Projektleitung",
"department": None,
"attendance_status": "present",
"notes": None,
}
],
"speaker_mappings": {"SPEAKER_00": "person-1"},
"mentioned_people": [],
"organization": {"name": "Example", "departments": []},
"known_entities": {},
"context_rules": {"do_not_infer_responsibilities": True},
}
def fake_transcribe(audio, model, output, language, **kwargs):
output.mkdir(parents=True, exist_ok=True)
raw = output / "whisper_raw.json"
transcript = output / "transcript.json"
text = output / "transcript.txt"
metadata = output / "runtime_metadata.json"
raw.write_text('{"transcription": []}\n', encoding="utf-8")
transcript.write_text(
json.dumps(
{
"text": "Ein kurzer Besprechungstext.",
"segments": [
{
"id": 0,
"start": 0.0,
"end": 1.0,
"text": "Ein kurzer Besprechungstext.",
}
],
}
),
encoding="utf-8",
)
text.write_text("Ein kurzer Besprechungstext.\n", encoding="utf-8")
metadata.write_text('{"runtime_seconds": 0.1}\n', encoding="utf-8")
return TranscriptionResult(output, raw, transcript, text, metadata, 0.1)
def fake_protocol(transcript, context, **kwargs):
rendered_context = load_meeting_context(context)
assert rendered_context.meeting_id == "programmatic-test"
return DirectProtocolResult(
protocol_text="# Meeting Protocol\n",
exact_prompt="prompt",
model_metadata={"model": kwargs["model"]},
runtime_metadata={"request_count": 1},
raw_response={"response": "# Meeting Protocol\n"},
)
class MvpApiTests(unittest.TestCase):
def config(self, root: Path):
audio = root / "meeting.wav"
model = root / "model.bin"
audio.write_bytes(b"audio")
model.write_bytes(b"model")
return MvpMeetingConfig(
audio_file=audio,
whisper_model=model,
output_root=root / "runs",
model="test:model",
)
def test_programmatic_context_is_persisted_without_source_yaml(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
config = self.config(root)
events = []
with (
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(
mvp_api, "generate_direct_protocol", side_effect=fake_protocol
),
patch.object(subprocess, "run") as subprocess_run,
):
result = mvp_api.run_mvp_meeting(
config, meeting_context=context_data(), progress_sink=events.append
)
self.assertEqual(result.exit_code, 0)
context_path = result.run_dir / "context/meeting_context.yaml"
self.assertTrue(context_path.is_file())
persisted = load_meeting_context(context_path)
self.assertEqual(persisted.meeting_id, "programmatic-test")
self.assertEqual(persisted.speaker_mappings, {"SPEAKER_00": "person-1"})
subprocess_run.assert_not_called()
self.assertEqual(
[(event.stage, event.status) for event in events],
[
("preparing", "started"),
("preparing", "completed"),
("transcription", "started"),
("transcription", "completed"),
("protocol_generation", "started"),
("protocol_generation", "completed"),
("completed", "completed"),
],
)
self.assertTrue(all(event.progress is None for event in events))
def test_failure_emits_terminal_failure_event(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
config = self.config(root)
events = []
with patch.object(
mvp_api,
"transcribe_audio",
side_effect=TranscriptionError("stopped"),
):
result = mvp_api.run_mvp_meeting(config, progress_sink=events.append)
self.assertEqual(result.exit_code, 2)
self.assertEqual(events[-1].stage, "failed")
self.assertEqual(events[-1].status, "failed")
self.assertIn("transcription", events[-1].message)
metadata = json.loads((result.run_dir / "run_metadata.json").read_text())
self.assertEqual(metadata["status"], "failed")
def test_cli_defaults_and_wrapper_delegate_without_subprocess(self):
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
config = self.config(root)
args = cli.parse_args(
[str(config.audio_file), "--whisper-model", str(config.whisper_model)]
)
expected = MvpRunResult(0, root / "run", root / "run/protocol.md")
with patch.object(cli, "run_mvp_meeting", return_value=expected) as api:
actual = cli.run(args, context_override=context_data())
self.assertEqual(actual, (0, expected.run_dir, expected.protocol_path))
delegated = api.call_args.args[0]
self.assertEqual(delegated.diarization, "off")
self.assertEqual(delegated.language, "de")
self.assertEqual(delegated.whisper_executable, "whisper-cli")
self.assertEqual(api.call_args.kwargs["meeting_context"], context_data())
if __name__ == "__main__":
unittest.main()
+124 -12
View File
@@ -6,7 +6,9 @@ from pathlib import Path
from unittest.mock import Mock, patch
from scripts import run_mvp_meeting
from src.meeting_lab.orchestration import mvp as mvp_api
from src.meeting_lab.protocol.generate_direct_protocol import DirectProtocolResult
from src.meeting_lab.diarization.backend import DiarizationResult
from src.meeting_lab.transcription.whisper import TranscriptionError, TranscriptionResult
@@ -50,7 +52,15 @@ def fake_transcribe(
metadata = output_dir / "runtime_metadata.json"
raw.write_text('{"transcription": []}\n', encoding="utf-8")
transcript.write_text(
json.dumps({"text": "Besprechungstext.", "segments": []}) + "\n",
json.dumps(
{
"text": "Besprechungstext.",
"segments": [
{"id": 0, "start": 0.0, "end": 1.0, "text": "Besprechungstext."}
],
}
)
+ "\n",
encoding="utf-8",
)
text.write_text("Besprechungstext.\n", encoding="utf-8")
@@ -58,6 +68,44 @@ def fake_transcribe(
return TranscriptionResult(output_dir, raw, transcript, text, metadata, 1.25)
def fake_diarize(audio_path, output_dir, device_mode, **kwargs):
output_dir.mkdir(parents=True, exist_ok=True)
metadata = {
"backend": "pyannote.audio",
"model": "pyannote/speaker-diarization-community-1",
"requested_device_mode": device_mode,
"actual_device": "cuda",
"device_name": "Fake GPU",
"runtime_seconds": 2.5,
"speaker_count": 1,
"credentials_persisted": False,
}
paths = {
"metadata": output_dir / "metadata.json",
"ordinary": output_dir / "diarization.rttm",
"exclusive": output_dir / "exclusive_diarization.rttm",
"turns": output_dir / "turns.json",
"exclusive_turns": output_dir / "exclusive_turns.json",
}
paths["metadata"].write_text(json.dumps(metadata), encoding="utf-8")
paths["ordinary"].write_text("", encoding="utf-8")
paths["exclusive"].write_text("", encoding="utf-8")
paths["turns"].write_text("[]", encoding="utf-8")
paths["exclusive_turns"].write_text(
json.dumps([{"start": 0, "end": 10, "speaker_id": "SPEAKER_00"}]),
encoding="utf-8",
)
return DiarizationResult(
output_dir,
paths["metadata"],
paths["ordinary"],
paths["exclusive"],
paths["turns"],
paths["exclusive_turns"],
metadata,
)
class MvpOrchestratorTests(unittest.TestCase):
def create_inputs(self, root: Path) -> tuple[Path, Path, Path]:
audio = root / "team meeting.wav"
@@ -88,9 +136,9 @@ class MvpOrchestratorTests(unittest.TestCase):
root = Path(directory)
args = self.args(root)
with (
patch.object(run_mvp_meeting, "transcribe_audio", side_effect=fake_transcribe) as whisper,
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe) as whisper,
patch.object(
run_mvp_meeting,
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
) as protocol,
@@ -129,9 +177,9 @@ class MvpOrchestratorTests(unittest.TestCase):
["--whisper-executable", "/tools/whisper-cli", "--threads", "4"],
)
with (
patch.object(run_mvp_meeting, "transcribe_audio", side_effect=fake_transcribe) as whisper,
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe) as whisper,
patch.object(
run_mvp_meeting,
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
) as protocol,
@@ -148,17 +196,81 @@ class MvpOrchestratorTests(unittest.TestCase):
protocol.call_args.kwargs["endpoint"], "http://ollama.test:11434"
)
def test_diarization_is_off_by_default_and_preserves_protocol_input(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
args = self.args(root)
with (
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(mvp_api, "diarize_audio") as diarization,
patch.object(
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
) as protocol,
):
code, run_dir, _ = run_mvp_meeting.run(args)
self.assertEqual(code, 0)
diarization.assert_not_called()
self.assertEqual(protocol.call_args.args[0], run_dir / "transcript/transcript.json")
self.assertFalse((run_dir / "diarization").exists())
def test_diarization_cli_propagates_and_uses_derived_protocol_input(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
args = self.args(
root,
[
"--diarization", "gpu",
"--diarization-runtime", "container",
"--diarization-container-image", "rocm/test",
"--diarization-container-arg=--device=/dev/kfd",
],
)
with (
patch.object(
mvp_api, "transcribe_audio", side_effect=fake_transcribe
),
patch.object(
mvp_api, "diarize_audio", side_effect=fake_diarize
) as diarization,
patch.object(
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
) as protocol,
):
code, run_dir, _ = run_mvp_meeting.run(args)
self.assertEqual(code, 0)
self.assertEqual(diarization.call_args.args[2], "gpu")
self.assertEqual(diarization.call_args.kwargs["runtime"], "container")
self.assertEqual(
diarization.call_args.kwargs["container_args"], ("--device=/dev/kfd",)
)
derived = run_dir / "diarization/transcript_diarized.json"
self.assertEqual(protocol.call_args.args[0], derived)
self.assertIn("SPEAKER_00", derived.read_text(encoding="utf-8"))
self.assertEqual(
json.loads((run_dir / "transcript/transcript.json").read_text())["text"],
"Besprechungstext.",
)
run_metadata = json.loads((run_dir / "run_metadata.json").read_text())
self.assertTrue(run_metadata["diarization"]["enabled"])
self.assertNotIn("HF_TOKEN", json.dumps(run_metadata))
def test_whisper_failure_is_recorded_and_protocol_is_not_called(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
args = self.args(root)
with (
patch.object(
run_mvp_meeting,
mvp_api,
"transcribe_audio",
side_effect=TranscriptionError("whisper stopped"),
),
patch.object(run_mvp_meeting, "generate_direct_protocol") as protocol,
patch.object(mvp_api, "generate_direct_protocol") as protocol,
):
code, run_dir, protocol_path = run_mvp_meeting.run(args)
@@ -176,9 +288,9 @@ class MvpOrchestratorTests(unittest.TestCase):
root = Path(directory)
args = self.args(root)
with (
patch.object(run_mvp_meeting, "transcribe_audio", side_effect=fake_transcribe),
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(
run_mvp_meeting,
mvp_api,
"generate_direct_protocol",
side_effect=ValueError("generation stopped"),
),
@@ -209,9 +321,9 @@ class MvpOrchestratorTests(unittest.TestCase):
root = Path(directory)
args = self.args(root)
with (
patch.object(run_mvp_meeting, "transcribe_audio", side_effect=fake_transcribe),
patch.object(mvp_api, "transcribe_audio", side_effect=fake_transcribe),
patch.object(
run_mvp_meeting,
mvp_api,
"generate_direct_protocol",
return_value=protocol_result(),
),
@@ -237,7 +349,7 @@ class MvpOrchestratorTests(unittest.TestCase):
"--output-root", str(args.output_root),
]
with patch.object(
run_mvp_meeting,
mvp_api,
"transcribe_audio",
side_effect=TranscriptionError("failed"),
):