109 lines
4.2 KiB
Python
109 lines
4.2 KiB
Python
"""Tests for the diarize_media MCP tool (server.handle_diarize_media).
|
|
|
|
Diarization itself (pyannote.audio) is monkeypatched so these tests run
|
|
without the optional [diarization] extra or a HuggingFace token — matching
|
|
the existing TestDetectBeatsHandler pattern in test_media_intel.py.
|
|
"""
|
|
|
|
import json
|
|
|
|
import pytest
|
|
|
|
|
|
def _write_tiny_wav(path: str, seconds: float = 1.0) -> None:
|
|
import struct
|
|
import wave
|
|
|
|
n_frames = int(44100 * seconds)
|
|
with wave.open(path, "w") as f:
|
|
f.setnchannels(1)
|
|
f.setsampwidth(2)
|
|
f.setframerate(44100)
|
|
f.writeframes(struct.pack("<%dh" % n_frames, *([0] * n_frames)))
|
|
|
|
|
|
class TestDiarizeMediaHandler:
|
|
async def test_reports_when_pyannote_unavailable(self, tmp_path, monkeypatch):
|
|
import server_tools.voice as server_mod
|
|
from server import handle_diarize_media
|
|
|
|
wav = tmp_path / "clip.wav"
|
|
_write_tiny_wav(str(wav))
|
|
monkeypatch.setattr(
|
|
server_mod, "diarization_capability", lambda token: (False, "Diarização indisponível: componente pyannote.audio ausente.")
|
|
)
|
|
result = await handle_diarize_media({"media_path": str(wav)})
|
|
text = result[0].text
|
|
assert "indisponível" in text.lower() or "unavailable" in text.lower()
|
|
assert "diarization" in text.lower()
|
|
|
|
async def test_rejects_disallowed_extension(self, tmp_path):
|
|
from server import handle_diarize_media
|
|
|
|
bad = tmp_path / "clip.txt"
|
|
bad.write_text("not audio")
|
|
with pytest.raises(ValueError):
|
|
await handle_diarize_media({"media_path": str(bad)})
|
|
|
|
async def test_writes_diarization_json_and_reports(self, tmp_path, monkeypatch):
|
|
import server_tools._shared as _shared_mod
|
|
import server_tools.voice as server_mod
|
|
from server import handle_diarize_media
|
|
|
|
wav = tmp_path / "clip.wav"
|
|
_write_tiny_wav(str(wav), seconds=2.0)
|
|
|
|
fake_transcript = {
|
|
"language": "en",
|
|
"duration": 2.0,
|
|
"text": "hello world",
|
|
"segments": [
|
|
{"text": "hello", "start": 0.0, "end": 1.0},
|
|
{"text": "world", "start": 1.0, "end": 2.0},
|
|
],
|
|
"words": [
|
|
{"word": "hello", "start": 0.0, "end": 0.5, "confidence": 0.9},
|
|
{"word": "world", "start": 1.0, "end": 1.5, "confidence": 0.9},
|
|
],
|
|
}
|
|
monkeypatch.setattr(_shared_mod, "transcribe", lambda *a, **k: fake_transcript)
|
|
monkeypatch.setattr(server_mod, "diarization_capability", lambda token: (True, "ok"))
|
|
monkeypatch.setattr(
|
|
server_mod,
|
|
"diarize",
|
|
lambda path, token, num_speakers="": [(0.0, 1.0, "A"), (1.0, 2.0, "B")],
|
|
)
|
|
|
|
result = await handle_diarize_media({"media_path": str(wav), "hf_token": "fake-token"})
|
|
text = result[0].text
|
|
assert "Speaker" in text or "speaker" in text.lower()
|
|
|
|
json_path = tmp_path / "clip_diarization.json"
|
|
assert str(json_path) in text
|
|
data = json.loads(json_path.read_text())
|
|
assert len(data["speakers"]) == 2
|
|
assert data["words"][0]["speaker_id"] == "SPEAKER_00"
|
|
assert data["words"][1]["speaker_id"] == "SPEAKER_01"
|
|
|
|
async def test_reports_when_diarization_fails(self, tmp_path, monkeypatch):
|
|
import server_tools._shared as _shared_mod
|
|
import server_tools.voice as server_mod
|
|
from server import handle_diarize_media
|
|
|
|
wav = tmp_path / "clip.wav"
|
|
_write_tiny_wav(str(wav))
|
|
|
|
fake_transcript = {
|
|
"language": "en",
|
|
"duration": 1.0,
|
|
"text": "hi",
|
|
"segments": [{"text": "hi", "start": 0.0, "end": 1.0}],
|
|
"words": [{"word": "hi", "start": 0.0, "end": 0.5, "confidence": 0.9}],
|
|
}
|
|
monkeypatch.setattr(_shared_mod, "transcribe", lambda *a, **k: fake_transcript)
|
|
monkeypatch.setattr(server_mod, "diarization_capability", lambda token: (True, "ok"))
|
|
monkeypatch.setattr(server_mod, "diarize", lambda *a, **k: None)
|
|
|
|
result = await handle_diarize_media({"media_path": str(wav), "hf_token": "fake-token"})
|
|
assert "failed" in result[0].text.lower()
|