feat: add speaker mapping and workflow improvements

This commit is contained in:
2026-08-25 15:38:58 +02:00
parent adf454d77d
commit dd8a618719
14 changed files with 1544 additions and 10 deletions
+8
View File
@@ -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:
+78
View File
@@ -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
+212
View File
@@ -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")
+136
View File
@@ -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"
)
+68
View File
@@ -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