Files

139 lines
5.9 KiB
Python

"""Paired Assistant/Lab smoke with only expensive media/model boundaries mocked."""
import json
import shutil
from dataclasses import replace
from unittest.mock import Mock
import pytest
from mka.application.meeting_service import MeetingProcessingService, ProcessingOptions
from mka.integrations.meeting_lab import MeetingLabGateway
from test_meeting_service import make_service, meeting, participants
mvp = pytest.importorskip("src.meeting_lab.orchestration.mvp")
from src.meeting_lab.audio import PreparedAudio # noqa: E402
from src.meeting_lab.diarization.backend import DiarizationResult # noqa: E402
from src.meeting_lab.llm.ollama import OllamaGeneration # noqa: E402
from src.meeting_lab.protocol.generate_direct_protocol import generate_direct_protocol # noqa: E402
from src.meeting_lab.transcription.whisper import TranscriptionResult # noqa: E402
def fake_prepare(source, destination, **kwargs):
destination.parent.mkdir(parents=True, exist_ok=True)
shutil.copyfile(source, destination)
return PreparedAudio(source, source.suffix[1:], destination, "ffmpeg", "ffmpeg")
def fake_transcribe(audio, model, output, language, **kwargs):
output.mkdir(parents=True, exist_ok=True)
raw, transcript, text, metadata = [
output / n
for n in ("whisper_raw.json", "transcript.json", "transcript.txt", "runtime_metadata.json")
]
raw.write_text("{}")
transcript.write_text(
json.dumps(
{
"text": "Lumini project discussion.",
"segments": [{"id": 0, "start": 0, "end": 2, "text": "Lumini project discussion."}],
}
)
)
text.write_text("Lumini project discussion.")
metadata.write_text("{}")
return TranscriptionResult(output, raw, transcript, text, metadata, 0.1)
def fake_diarize(audio, output, mode, **kwargs):
output.mkdir(parents=True, exist_ok=True)
paths = [
output / name
for name in (
"metadata.json",
"diarization.rttm",
"exclusive_diarization.rttm",
"turns.json",
"exclusive_turns.json",
)
]
metadata = {"speaker_count": 1, "runtime_seconds": 0.1}
paths[0].write_text(json.dumps(metadata))
for p in paths[1:]:
p.write_text("[]")
paths[-1].write_text(json.dumps([{"start": 0, "end": 2, "speaker_id": "SPEAKER_00"}]))
return DiarizationResult(output, *paths, metadata)
@pytest.mark.parametrize("language,expected", [("de", "German"), ("en", "English")])
@pytest.mark.parametrize(
"profile,threads", [("auto", None), ("fast", 16), ("efficient", 10), ("powersave", 4)]
)
def test_paired_anonymous_generation_and_mapped_regeneration(
tmp_path, monkeypatch, language, expected, profile, threads
):
template, _ = make_service(tmp_path)
service = MeetingProcessingService(template.settings, MeetingLabGateway())
service.glossary.create("Luminy", "product", aliases=("Lumini",))
calls = []
def model_call(endpoint, model, prompt, **kwargs):
calls.append((prompt, kwargs))
return OllamaGeneration(
{"response": "# Mock protocol", "done": True}, "# Mock protocol", 0.1
)
def generate(transcript, context, **kwargs):
return generate_direct_protocol(
transcript, context, **kwargs, model_check=lambda *_: {}, generation_call=model_call
)
prepare = Mock(side_effect=fake_prepare)
transcribe = Mock(side_effect=fake_transcribe)
diarize = Mock(side_effect=fake_diarize)
monkeypatch.setattr(mvp, "prepare_audio", prepare)
monkeypatch.setattr(mvp, "transcribe_audio", transcribe)
monkeypatch.setattr(mvp, "diarize_audio", diarize)
monkeypatch.setattr(mvp, "generate_direct_protocol", generate)
audio = tmp_path / "sample.wav"
audio.write_bytes(b"mock audio")
outcome = service.process(
audio,
replace(meeting(), language=language),
participants(),
ProcessingOptions(diarization_enabled=True, performance_profile=profile),
)
assert outcome.succeeded
assert outcome.awaiting_speaker_review
assert prepare.call_args.kwargs["normalization_enabled"] is True
assert transcribe.call_args.args[3] == language
root = outcome.run_dir
assert calls == []
service.regenerate_protocol(root, {}, performance_profile=profile)
first = root / "protocol/generations/001"
before = {p.name: p.read_bytes() for p in first.iterdir()}
original = (root / "diarization/transcript_diarized.json").read_bytes()
first_meta = json.loads((first / "runtime_metadata.json").read_text())
assert first_meta["speaker_mapping"] == {}
assert first_meta["output_language"] == language
assert "SPEAKER_00" in (first / "transcript_input.txt").read_text()
assert "timing" in json.loads((root / "run_metadata.json").read_text())
service.regenerate_protocol(root, {"SPEAKER_00": "martin"}, performance_profile=profile)
assert prepare.call_count == transcribe.call_count == diarize.call_count == 1
assert (root / "diarization/transcript_diarized.json").read_bytes() == original
assert {p.name: p.read_bytes() for p in first.iterdir()} == before
second = root / "protocol/generations/002"
metadata = json.loads((second / "runtime_metadata.json").read_text())
assert metadata["speaker_mapping"] == {"SPEAKER_00": "martin"}
assert metadata["speaker_mapping_names"] == {"SPEAKER_00": "Martin"}
assert metadata["output_language"] == language
assert metadata["num_thread"] == threads
assert metadata["glossary_aliases_configured"]["Lumini"] == "Luminy"
assert metadata["glossary_replacements"] == []
assert "Lumini" in (second / "transcript_input.txt").read_text()
assert (root / "protocol.md").resolve() == second / "protocol.md"
for prompt, options in calls:
assert f"Write the meeting protocol in {expected}." in prompt
assert "Luminy" in prompt
assert options["num_thread"] == threads