Add diarization and reusable MVP meeting pipeline
This commit is contained in:
@@ -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()
|
||||
@@ -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"
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
@@ -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"),
|
||||
):
|
||||
|
||||
Reference in New Issue
Block a user