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()