165 lines
7.4 KiB
Python
165 lines
7.4 KiB
Python
"""Meeting-language regression tests without media or model execution."""
|
|
|
|
import json
|
|
import tempfile
|
|
import unittest
|
|
from functools import partial
|
|
from pathlib import Path
|
|
from unittest.mock import Mock, patch
|
|
|
|
from src.meeting_lab.llm.ollama import OllamaGeneration
|
|
from src.meeting_lab.models.meeting_context import load_meeting_context
|
|
from src.meeting_lab.orchestration.mvp import regenerate_mvp_protocol
|
|
from src.meeting_lab.protocol.direct_protocol_prompt import build_direct_protocol_prompt
|
|
from src.meeting_lab.protocol.generate_direct_protocol import (
|
|
estimate_input_tokens,
|
|
generate_direct_protocol,
|
|
select_transcript_input,
|
|
)
|
|
|
|
|
|
class MeetingLanguageTests(unittest.TestCase):
|
|
def test_diarized_plain_fallback_retains_english_instruction(self) -> None:
|
|
segments = [
|
|
{
|
|
"start": index,
|
|
"end": index + 1,
|
|
"text": "Word.",
|
|
"speaker_id": f"SPEAKER_{index % 2:02d}",
|
|
}
|
|
for index in range(200)
|
|
]
|
|
plain = " ".join(segment["text"] for segment in segments)
|
|
budget = estimate_input_tokens(
|
|
build_direct_protocol_prompt(plain, meeting_language="en")
|
|
)
|
|
selected = select_transcript_input(
|
|
{"text": plain, "segments": segments, "speaker_labels_anonymous": True},
|
|
None,
|
|
meeting_language="en",
|
|
safe_input_token_budget=budget,
|
|
)
|
|
self.assertEqual(selected.representation, "plain_transcript_fallback")
|
|
self.assertIn("Write the meeting protocol in English.", selected.prompt)
|
|
self.assertEqual(selected.text, plain)
|
|
|
|
def test_generation_and_regeneration_preserve_language_and_inputs(self) -> None:
|
|
for language, expected in (
|
|
("de", "German"),
|
|
("en", "English"),
|
|
(None, "German"),
|
|
):
|
|
for diarized in (False, True):
|
|
with (
|
|
self.subTest(language=language, diarized=diarized),
|
|
tempfile.TemporaryDirectory() as directory,
|
|
):
|
|
run = Path(directory)
|
|
context_path = run / "context" / "meeting_context.yaml"
|
|
context_path.parent.mkdir()
|
|
context_text = (
|
|
'schema_version: "1"\nmeeting:\n'
|
|
" meeting_id: test\n title: Mixed terminology\n"
|
|
+ (f" language: {language}\n" if language else "")
|
|
+ ' notes: "Freigabe für Product X"\n'
|
|
"participants:\n - participant_id: person\n"
|
|
' display_name: "Jörg Müller"\n'
|
|
"speaker_mappings:\n SPEAKER_00: person\n"
|
|
)
|
|
context_path.write_text(context_text, encoding="utf-8")
|
|
transcript = run / ("diarization" if diarized else "transcript")
|
|
transcript.mkdir()
|
|
transcript /= (
|
|
"transcript_diarized.json" if diarized else "transcript.json"
|
|
)
|
|
transcript.write_text(
|
|
json.dumps(
|
|
{
|
|
"text": "I will check Product X.",
|
|
"speaker_labels_anonymous": diarized,
|
|
"segments": [
|
|
{
|
|
"id": 0,
|
|
"start": 0,
|
|
"end": 1,
|
|
"speaker_id": "SPEAKER_00",
|
|
"text": "I will check Product X.",
|
|
}
|
|
],
|
|
}
|
|
),
|
|
encoding="utf-8",
|
|
)
|
|
original = transcript.read_bytes()
|
|
call = Mock(
|
|
return_value=OllamaGeneration(
|
|
text="# Meeting Protocol",
|
|
raw_response={},
|
|
client_wall_time_seconds=0.1,
|
|
)
|
|
)
|
|
generate = partial(
|
|
generate_direct_protocol,
|
|
model_check=Mock(return_value={}),
|
|
generation_call=call,
|
|
)
|
|
result = generate(transcript, context_path)
|
|
self.assertIn(
|
|
f"Write the meeting protocol in {expected}.",
|
|
result.exact_prompt,
|
|
)
|
|
self.assertNotIn("in deutscher Sprache", result.exact_prompt)
|
|
self.assertNotIn("auf Deutsch", result.exact_prompt)
|
|
self.assertIn("Jörg Müller", result.exact_prompt)
|
|
self.assertIn("Mixed terminology", result.exact_prompt)
|
|
self.assertIn("I will check Product X.", result.transcript_input)
|
|
self.assertEqual(
|
|
result.runtime_metadata["output_language"], language or "de"
|
|
)
|
|
self.assertEqual(
|
|
context_path.read_text(encoding="utf-8"), context_text
|
|
)
|
|
context = load_meeting_context(context_path)
|
|
with (
|
|
patch(
|
|
"src.meeting_lab.orchestration.mvp.generate_direct_protocol",
|
|
side_effect=generate,
|
|
),
|
|
patch(
|
|
"src.meeting_lab.orchestration.mvp.transcribe_audio",
|
|
side_effect=AssertionError("Retranscription is forbidden"),
|
|
),
|
|
):
|
|
regenerate_mvp_protocol(run, meeting_context=context)
|
|
metadata = json.loads(
|
|
(run / "protocol" / "runtime_metadata.json").read_text()
|
|
)
|
|
self.assertEqual(metadata["output_language"], language or "de")
|
|
self.assertIn(
|
|
f"Write the meeting protocol in {expected}.",
|
|
call.call_args.args[2],
|
|
)
|
|
self.assertEqual(transcript.read_bytes(), original)
|
|
self.assertEqual(
|
|
load_meeting_context(context_path).data, context.data
|
|
)
|
|
self.assertEqual(context.speaker_mappings, {"SPEAKER_00": "person"})
|
|
|
|
def test_no_context_defaults_to_german(self) -> None:
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
transcript = Path(directory) / "transcript.json"
|
|
transcript.write_text('{"text": "English source text."}')
|
|
result = generate_direct_protocol(
|
|
transcript,
|
|
model_check=Mock(return_value={}),
|
|
generation_call=Mock(
|
|
return_value=OllamaGeneration(
|
|
text="Protocol",
|
|
raw_response={},
|
|
client_wall_time_seconds=0.1,
|
|
)
|
|
),
|
|
)
|
|
self.assertIn("Write the meeting protocol in German.", result.exact_prompt)
|
|
self.assertEqual(result.runtime_metadata["output_language"], "de")
|