feat: add speaker mapping and workflow improvements
This commit is contained in:
@@ -39,6 +39,14 @@ def test_environment_defaults_to_no_diarization_container_args(monkeypatch) -> N
|
||||
settings = AppSettings.from_environment()
|
||||
|
||||
assert settings.diarization_container_args == ()
|
||||
assert settings.glossary_database == Path("data/database/glossary.sqlite3")
|
||||
|
||||
|
||||
def test_environment_configures_glossary_database(monkeypatch, tmp_path: Path) -> None:
|
||||
database = tmp_path / "terms.sqlite3"
|
||||
monkeypatch.setenv("MKA_GLOSSARY_DATABASE", str(database))
|
||||
|
||||
assert AppSettings.from_environment().glossary_database == database
|
||||
|
||||
|
||||
def test_environment_parses_multiple_ordered_container_args(monkeypatch) -> None:
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
import sqlite3
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
from mka.application.glossary import GlossaryConflictError, GlossaryRepository
|
||||
|
||||
|
||||
def repository(tmp_path: Path) -> GlossaryRepository:
|
||||
result = GlossaryRepository(tmp_path / "database" / "glossary.sqlite3")
|
||||
result.initialize()
|
||||
return result
|
||||
|
||||
|
||||
def test_initialization_creates_empty_versioned_database(tmp_path: Path) -> None:
|
||||
glossary = repository(tmp_path)
|
||||
|
||||
assert glossary.database_path.is_file()
|
||||
assert glossary.list() == []
|
||||
with sqlite3.connect(glossary.database_path) as connection:
|
||||
assert connection.execute("PRAGMA user_version").fetchone()[0] == 1
|
||||
tables = {
|
||||
row[0]
|
||||
for row in connection.execute("SELECT name FROM sqlite_master WHERE type = 'table'")
|
||||
}
|
||||
assert {"glossary_entries", "glossary_aliases"} <= tables
|
||||
|
||||
|
||||
def test_create_read_update_and_deactivate_with_multiple_aliases(tmp_path: Path) -> None:
|
||||
glossary = repository(tmp_path)
|
||||
created = glossary.create(
|
||||
"Secugrid HS",
|
||||
"product",
|
||||
aliases=("Sikirgut", "Secugrid H S"),
|
||||
description="Canonical core product name",
|
||||
)
|
||||
|
||||
assert glossary.get(created.id).aliases == ("Secugrid H S", "Sikirgut")
|
||||
assert glossary.list("sikir")[0].canonical_term == "Secugrid HS"
|
||||
|
||||
updated = glossary.update(
|
||||
created.id,
|
||||
"Secugrid HS",
|
||||
"technical_term",
|
||||
aliases=("Sekugrid HS",),
|
||||
description="Updated",
|
||||
is_active=False,
|
||||
)
|
||||
|
||||
assert updated.category == "technical_term"
|
||||
assert updated.aliases == ("Sekugrid HS",)
|
||||
assert not updated.is_active
|
||||
assert glossary.list(active_only=True) == []
|
||||
|
||||
|
||||
def test_terms_and_aliases_are_unique_case_insensitively_across_entries(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
glossary = repository(tmp_path)
|
||||
glossary.create("PBAT", "acronym", aliases=("Polybutylene adipate terephthalate",))
|
||||
|
||||
with pytest.raises(GlossaryConflictError):
|
||||
glossary.create("pbat", "material")
|
||||
with pytest.raises(GlossaryConflictError):
|
||||
glossary.create("Other", "other", aliases=("PBAT",))
|
||||
with pytest.raises(GlossaryConflictError):
|
||||
glossary.create("Polybutylene adipate terephthalate", "material")
|
||||
|
||||
|
||||
def test_delete_removes_entry_and_aliases(tmp_path: Path) -> None:
|
||||
glossary = repository(tmp_path)
|
||||
entry = glossary.create("Luminy", "product", aliases=("Lumini",))
|
||||
|
||||
glossary.delete(entry.id)
|
||||
|
||||
assert glossary.list() == []
|
||||
with sqlite3.connect(glossary.database_path) as connection:
|
||||
assert connection.execute("SELECT COUNT(*) FROM glossary_aliases").fetchone()[0] == 0
|
||||
@@ -5,6 +5,8 @@ from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
|
||||
from mka.application.config import AppSettings
|
||||
from mka.application.meeting_service import (
|
||||
MeetingDetails,
|
||||
@@ -29,6 +31,7 @@ class FakeMeetingLab:
|
||||
self.context_data: dict[str, Any] | None = None
|
||||
self.config_values: dict[str, Any] | None = None
|
||||
self.fail = False
|
||||
self.regeneration: dict[str, Any] | None = None
|
||||
|
||||
def create_context(self, data: dict[str, Any]) -> FakeContext:
|
||||
self.context_data = data
|
||||
@@ -99,6 +102,36 @@ class FakeMeetingLab:
|
||||
protocol.write_text("# Generated protocol\n", encoding="utf-8")
|
||||
return SimpleNamespace(exit_code=0, run_dir=self.run_dir, protocol_path=protocol)
|
||||
|
||||
def regenerate_protocol(
|
||||
self,
|
||||
run_dir: Path,
|
||||
meeting_context: Any,
|
||||
progress_sink: Any,
|
||||
**options: Any,
|
||||
) -> Any:
|
||||
self.regeneration = {
|
||||
"run_dir": run_dir,
|
||||
"meeting_context": meeting_context,
|
||||
"options": options,
|
||||
}
|
||||
progress_sink(
|
||||
SimpleNamespace(
|
||||
stage="protocol_generation",
|
||||
status="started",
|
||||
elapsed_seconds=0.0,
|
||||
progress=None,
|
||||
message=None,
|
||||
)
|
||||
)
|
||||
protocol = Path(run_dir) / "protocol.md"
|
||||
protocol.write_text("# Regenerated protocol\n", encoding="utf-8")
|
||||
protocol_dir = Path(run_dir) / "protocol"
|
||||
protocol_dir.mkdir(exist_ok=True)
|
||||
(protocol_dir / "runtime_metadata.json").write_text(
|
||||
json.dumps({"speaker_attribution_available": True}), encoding="utf-8"
|
||||
)
|
||||
return SimpleNamespace(exit_code=0, run_dir=run_dir, protocol_path=protocol)
|
||||
|
||||
|
||||
def make_service(tmp_path: Path) -> tuple[MeetingProcessingService, FakeMeetingLab]:
|
||||
model = tmp_path / "model.bin"
|
||||
@@ -107,6 +140,7 @@ def make_service(tmp_path: Path) -> tuple[MeetingProcessingService, FakeMeetingL
|
||||
settings = AppSettings(
|
||||
data_root=tmp_path / "meetings",
|
||||
whisper_model=model,
|
||||
glossary_database=tmp_path / "glossary.sqlite3",
|
||||
whisper_executable="/opt/whisper-cli",
|
||||
protocol_model="test:model",
|
||||
diarization_mode="gpu",
|
||||
@@ -139,6 +173,60 @@ def participants() -> list[ParticipantInput]:
|
||||
]
|
||||
|
||||
|
||||
def write_speaker_review_artifacts(run_dir: Path) -> Path:
|
||||
diarization_dir = run_dir / "diarization"
|
||||
context_dir = run_dir / "context"
|
||||
diarization_dir.mkdir(parents=True)
|
||||
context_dir.mkdir()
|
||||
transcript_path = diarization_dir / "transcript_diarized.json"
|
||||
transcript_path.write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"speaker_labels_anonymous": True,
|
||||
"segments": [
|
||||
{
|
||||
"speaker_id": "SPEAKER_01",
|
||||
"text": "I will prepare all raw materials before Wednesday.",
|
||||
},
|
||||
{"speaker_id": "SPEAKER_00", "text": "Yes."},
|
||||
{
|
||||
"speaker_id": "SPEAKER_00",
|
||||
"text": "We will run the production trial on Wednesday.",
|
||||
},
|
||||
{
|
||||
"speaker_id": "SPEAKER_00",
|
||||
"text": "The trial requires the complete production team.",
|
||||
},
|
||||
{"speaker_id": None, "text": "Unassigned text."},
|
||||
],
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
(context_dir / "meeting_context.yaml").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"schema_version": "1",
|
||||
"meeting": {
|
||||
"meeting_id": "speaker-review",
|
||||
"title": "Speaker review",
|
||||
"language": "en",
|
||||
},
|
||||
"participants": [
|
||||
{"participant_id": "martin", "display_name": "Martin"},
|
||||
{"participant_id": "anna", "display_name": "Anna"},
|
||||
],
|
||||
"speaker_mappings": {},
|
||||
"mentioned_people": [],
|
||||
"organization": {"departments": []},
|
||||
"known_entities": {},
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
return transcript_path
|
||||
|
||||
|
||||
def test_build_context_uses_actual_v1_shape(tmp_path: Path) -> None:
|
||||
service, gateway = make_service(tmp_path)
|
||||
|
||||
@@ -171,6 +259,40 @@ def test_build_context_preserves_explicit_speaker_mapping(tmp_path: Path) -> Non
|
||||
assert gateway.context_data["speaker_mappings"] == {"SPEAKER_00": "martin"}
|
||||
|
||||
|
||||
def test_build_context_includes_only_active_authoritative_glossary_terms(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
service, gateway = make_service(tmp_path)
|
||||
service.glossary.create("Secugrid HS", "product", aliases=("Sikirgut", "Secugrid H S"))
|
||||
inactive = service.glossary.create("Old Name", "other")
|
||||
service.glossary.set_active(inactive.id, False)
|
||||
|
||||
service.build_context(meeting(), participants())
|
||||
|
||||
assert gateway.context_data is not None
|
||||
assert gateway.context_data["known_entities"] == {
|
||||
"Authoritative terminology": ["Secugrid HS (aliases: Secugrid H S, Sikirgut)"]
|
||||
}
|
||||
rules = gateway.context_data["context_rules"]
|
||||
assert "do not invent matches" in rules["glossary_canonical_spelling"]
|
||||
assert "inside compounds" in rules["glossary_core_terms"]
|
||||
assert "Old Name" not in str(gateway.context_data)
|
||||
|
||||
|
||||
def test_glossary_is_rendered_into_meeting_lab_protocol_context(tmp_path: Path) -> None:
|
||||
meeting_context = pytest.importorskip("src.meeting_lab.models.meeting_context")
|
||||
service, gateway = make_service(tmp_path)
|
||||
service.glossary.create("PBAT", "acronym", aliases=("P B A T",))
|
||||
context = service.build_context(meeting(), participants())
|
||||
|
||||
real_context = meeting_context.create_meeting_context(context.data)
|
||||
prompt_context = meeting_context.render_meeting_context_for_prompt(real_context)
|
||||
|
||||
assert "Authoritative terminology: PBAT (aliases: P B A T)" in prompt_context
|
||||
assert "Use canonical glossary spellings" in prompt_context
|
||||
assert "use canonical core terms inside compounds" in prompt_context
|
||||
|
||||
|
||||
def test_build_context_translates_mentioned_only_person(tmp_path: Path) -> None:
|
||||
service, gateway = make_service(tmp_path)
|
||||
people = participants() + [
|
||||
@@ -331,6 +453,96 @@ def test_failure_reports_stage_and_preserves_run_dir(tmp_path: Path) -> None:
|
||||
assert (gateway.run_dir / "run_metadata.json").is_file()
|
||||
|
||||
|
||||
def test_speaker_review_lists_detected_labels_participants_and_excerpts(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
service, gateway = make_service(tmp_path)
|
||||
source = write_speaker_review_artifacts(gateway.run_dir)
|
||||
|
||||
review = service.load_speaker_mapping_review(gateway.run_dir, excerpts_per_speaker=2)
|
||||
|
||||
assert review is not None
|
||||
assert [speaker.speaker_label for speaker in review.speakers] == [
|
||||
"SPEAKER_00",
|
||||
"SPEAKER_01",
|
||||
]
|
||||
assert review.speakers[0].excerpts == (
|
||||
"We will run the production trial on Wednesday.",
|
||||
"The trial requires the complete production team.",
|
||||
)
|
||||
assert review.participants == (("martin", "Martin"), ("anna", "Anna"))
|
||||
assert review.current_mappings == {}
|
||||
assert "SPEAKER_00" in source.read_text(encoding="utf-8")
|
||||
|
||||
|
||||
def test_protocol_only_regeneration_persists_mappings_without_rewriting_transcript(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
service, gateway = make_service(tmp_path)
|
||||
service.glossary.create("ENLYZE", "organization", aliases=("Enlyse",))
|
||||
source = write_speaker_review_artifacts(gateway.run_dir)
|
||||
original_source = source.read_bytes()
|
||||
events = []
|
||||
|
||||
outcome = service.regenerate_protocol(
|
||||
gateway.run_dir,
|
||||
{"SPEAKER_00": "martin"},
|
||||
progress_sink=events.append,
|
||||
)
|
||||
|
||||
assert outcome.succeeded
|
||||
assert outcome.original_protocol == "# Regenerated protocol\n"
|
||||
assert outcome.speaker_attribution_available is True
|
||||
assert gateway.context_data is not None
|
||||
assert gateway.context_data["speaker_mappings"] == {"SPEAKER_00": "martin"}
|
||||
assert gateway.context_data["known_entities"] == {
|
||||
"Authoritative terminology": ["ENLYZE (aliases: Enlyse)"]
|
||||
}
|
||||
assert gateway.regeneration is not None
|
||||
assert gateway.regeneration["meeting_context"].data["speaker_mappings"] == {
|
||||
"SPEAKER_00": "martin"
|
||||
}
|
||||
assert gateway.regeneration["options"]["protocol_num_ctx"] == 32_768
|
||||
assert events[0].stage == "protocol_generation"
|
||||
assert source.read_bytes() == original_source
|
||||
|
||||
|
||||
def test_protocol_regeneration_allows_unmapped_and_rejects_duplicate_participant(
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
service, gateway = make_service(tmp_path)
|
||||
write_speaker_review_artifacts(gateway.run_dir)
|
||||
|
||||
outcome = service.regenerate_protocol(gateway.run_dir, {})
|
||||
|
||||
assert outcome.succeeded
|
||||
assert gateway.context_data is not None
|
||||
assert gateway.context_data["speaker_mappings"] == {}
|
||||
|
||||
with pytest.raises(ValueError, match="only one speaker"):
|
||||
service.regenerate_protocol(
|
||||
gateway.run_dir,
|
||||
{"SPEAKER_00": "martin", "SPEAKER_01": "martin"},
|
||||
)
|
||||
|
||||
|
||||
def test_fallback_attribution_loss_is_read_from_runtime_metadata(tmp_path: Path) -> None:
|
||||
run_dir = tmp_path / "run"
|
||||
protocol_dir = run_dir / "protocol"
|
||||
protocol_dir.mkdir(parents=True)
|
||||
(protocol_dir / "runtime_metadata.json").write_text(
|
||||
json.dumps(
|
||||
{
|
||||
"speaker_attribution_available": False,
|
||||
"speaker_attribution_loss_reason": "plain_transcript_fallback",
|
||||
}
|
||||
),
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
assert MeetingProcessingService._speaker_attribution_available(run_dir) is False
|
||||
|
||||
|
||||
def test_uploaded_source_is_preserved_in_meeting_directory(tmp_path: Path) -> None:
|
||||
service, _ = make_service(tmp_path)
|
||||
source = SimpleNamespace(getbuffer=lambda: b"source audio")
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
import json
|
||||
from datetime import date
|
||||
|
||||
import pytest
|
||||
|
||||
from mka.application.meeting_service import ParticipantInput
|
||||
from mka.application.run_inputs import (
|
||||
RunInputJsonError,
|
||||
RunInputState,
|
||||
export_run_inputs,
|
||||
import_run_inputs,
|
||||
run_input_filename,
|
||||
)
|
||||
|
||||
|
||||
def populated_state() -> RunInputState:
|
||||
return RunInputState(
|
||||
title="F&E Technik Abteilungs Jour Fixe",
|
||||
description="Review the production trial.",
|
||||
language="de",
|
||||
meeting_date=date(2026, 8, 25),
|
||||
participants=(
|
||||
ParticipantInput(
|
||||
participant_id="martin",
|
||||
display_name="Martin Tazl",
|
||||
role="Project lead",
|
||||
organization="Engineering",
|
||||
),
|
||||
ParticipantInput(
|
||||
participant_id="alex",
|
||||
display_name="Alexander Funk",
|
||||
attendance_status="mentioned_only",
|
||||
),
|
||||
),
|
||||
audio_normalization=False,
|
||||
diarization_enabled=True,
|
||||
source_file_name="2026-08-25_jour_fixe.wav",
|
||||
)
|
||||
|
||||
|
||||
def test_current_form_state_serializes_with_schema_and_all_supported_values() -> None:
|
||||
document = json.loads(export_run_inputs(populated_state()))
|
||||
|
||||
assert document["schema_version"] == 1
|
||||
assert document["meeting"]["title"] == "F&E Technik Abteilungs Jour Fixe"
|
||||
assert document["meeting"]["description"] == "Review the production trial."
|
||||
assert document["meeting"]["language"] == "de"
|
||||
assert document["meeting"]["date"] == "2026-08-25"
|
||||
assert document["meeting"]["participants"][1] == {
|
||||
"participant_id": "alex",
|
||||
"display_name": "Alexander Funk",
|
||||
"role": "",
|
||||
"organization": "",
|
||||
"attendance_status": "mentioned_only",
|
||||
}
|
||||
assert document["processing"] == {
|
||||
"audio_normalization": False,
|
||||
"diarization_enabled": True,
|
||||
}
|
||||
|
||||
|
||||
def test_export_import_round_trip_restores_form_and_participants() -> None:
|
||||
original = populated_state()
|
||||
|
||||
restored = import_run_inputs(export_run_inputs(original))
|
||||
|
||||
assert restored == original
|
||||
assert restored.participants[0].participant_id == "martin"
|
||||
assert restored.participants[0].display_name == "Martin Tazl"
|
||||
|
||||
|
||||
def test_missing_optional_fields_use_current_defaults() -> None:
|
||||
restored = import_run_inputs('{"schema_version": 1}')
|
||||
|
||||
defaults = RunInputState.defaults()
|
||||
assert restored == defaults
|
||||
|
||||
|
||||
def test_unknown_safe_fields_are_ignored() -> None:
|
||||
restored = import_run_inputs(
|
||||
json.dumps(
|
||||
{
|
||||
"schema_version": 1,
|
||||
"meeting": {"title": "Known", "future_field": {"value": 1}},
|
||||
"processing": {"future_toggle": True},
|
||||
"future_section": [1, 2, 3],
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
assert restored.title == "Known"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("content", ["not-json", "[]", b"\xff"])
|
||||
def test_malformed_json_is_rejected(content: str | bytes) -> None:
|
||||
with pytest.raises(RunInputJsonError, match="Malformed|top-level"):
|
||||
import_run_inputs(content)
|
||||
|
||||
|
||||
def test_unsupported_future_schema_is_rejected() -> None:
|
||||
with pytest.raises(RunInputJsonError, match="Unsupported.*schema_version"):
|
||||
import_run_inputs('{"schema_version": 2}')
|
||||
|
||||
|
||||
def test_media_contents_are_never_serialized() -> None:
|
||||
exported = export_run_inputs(populated_state())
|
||||
|
||||
assert "2026-08-25_jour_fixe.wav" in exported
|
||||
assert "audio_bytes" not in exported
|
||||
assert "base64" not in exported
|
||||
|
||||
|
||||
def test_source_filename_must_not_be_a_machine_specific_path() -> None:
|
||||
with pytest.raises(RunInputJsonError, match="filename, not a path"):
|
||||
import_run_inputs('{"schema_version": 1, "source_file_name": "/tmp/meeting.wav"}')
|
||||
|
||||
|
||||
def test_invalid_or_duplicate_participants_are_not_silently_reinterpreted() -> None:
|
||||
document = {
|
||||
"schema_version": 1,
|
||||
"meeting": {
|
||||
"participants": [
|
||||
{"participant_id": "martin", "display_name": "Martin"},
|
||||
{"participant_id": "martin", "display_name": "Someone else"},
|
||||
]
|
||||
},
|
||||
}
|
||||
|
||||
with pytest.raises(RunInputJsonError, match="Duplicate participant_id"):
|
||||
import_run_inputs(json.dumps(document))
|
||||
|
||||
|
||||
def test_export_filename_is_readable_and_machine_independent() -> None:
|
||||
assert run_input_filename("F&E Technik Jour Fixe") == (
|
||||
"meeting-inputs-f-e-technik-jour-fixe.json"
|
||||
)
|
||||
@@ -0,0 +1,68 @@
|
||||
from datetime import date
|
||||
|
||||
from mka.application.meeting_service import ParticipantInput
|
||||
from mka.application.run_inputs import RunInputState
|
||||
from mka.ui import streamlit_app
|
||||
|
||||
|
||||
def test_alias_input_accepts_lines_and_commas() -> None:
|
||||
assert streamlit_app._parse_aliases("Sikirgut\nSekugrid HS, Secugrid H S") == (
|
||||
"Sikirgut",
|
||||
"Sekugrid HS",
|
||||
"Secugrid H S",
|
||||
)
|
||||
|
||||
|
||||
def test_regenerated_protocol_widget_value_is_deferred_until_next_run(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
state = {"edited_protocol": "old protocol"}
|
||||
monkeypatch.setattr(streamlit_app.st, "session_state", state)
|
||||
|
||||
streamlit_app._queue_edited_protocol("regenerated protocol")
|
||||
|
||||
assert state["edited_protocol"] == "old protocol"
|
||||
assert state["pending_edited_protocol"] == "regenerated protocol"
|
||||
|
||||
streamlit_app._apply_pending_edited_protocol()
|
||||
|
||||
assert state["edited_protocol"] == "regenerated protocol"
|
||||
assert "pending_edited_protocol" not in state
|
||||
|
||||
|
||||
def test_imported_inputs_are_applied_via_pending_state_before_widgets(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
existing_upload = object()
|
||||
state = {"source_media_0": existing_upload, "source_media_widget_generation": 0}
|
||||
monkeypatch.setattr(streamlit_app.st, "session_state", state)
|
||||
imported = RunInputState(
|
||||
title="Imported meeting",
|
||||
description="Imported context",
|
||||
language="en",
|
||||
meeting_date=date(2026, 8, 25),
|
||||
participants=(ParticipantInput("martin", "Martin"),),
|
||||
audio_normalization=False,
|
||||
diarization_enabled=True,
|
||||
source_file_name="meeting.wav",
|
||||
)
|
||||
|
||||
streamlit_app._queue_run_inputs(imported)
|
||||
|
||||
assert "meeting_title" not in state
|
||||
assert state["pending_run_inputs"] == imported
|
||||
|
||||
streamlit_app._apply_pending_run_inputs()
|
||||
|
||||
assert state["meeting_title"] == "Imported meeting"
|
||||
assert state["meeting_description"] == "Imported context"
|
||||
assert state["meeting_language"] == "en"
|
||||
assert state["meeting_has_date"] is True
|
||||
assert state["participants"][0]["participant_id"] == "martin"
|
||||
assert state["audio_normalization"] is False
|
||||
assert state["diarization_enabled"] is True
|
||||
assert state["imported_source_file_name"] == "meeting.wav"
|
||||
assert state["source_media_widget_generation"] == 1
|
||||
assert "source_media_1" not in state
|
||||
assert "pending_run_inputs" not in state
|
||||
assert state["source_media_0"] is existing_upload
|
||||
Reference in New Issue
Block a user