Files
meeting-lab/tests/test_diarization.py
T

223 lines
8.2 KiB
Python

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