223 lines
8.2 KiB
Python
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()
|