Add post-diarization speaker review workflow

This commit is contained in:
2026-09-17 19:16:57 +02:00
parent bd48aec5aa
commit 2215414282
8 changed files with 240 additions and 24 deletions
+3
View File
@@ -104,9 +104,12 @@ def test_paired_anonymous_generation_and_mapped_regeneration(
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()
+93 -1
View File
@@ -6,6 +6,7 @@ from types import SimpleNamespace
from typing import Any
import pytest
import yaml
from mka.application.config import AppSettings
from mka.application.meeting_service import (
@@ -99,6 +100,24 @@ class FakeMeetingLab:
)
)
return SimpleNamespace(exit_code=2, run_dir=self.run_dir, protocol_path=None)
if config["stop_after_diarization"]:
write_speaker_review_artifacts(self.run_dir)
(self.run_dir / "run_metadata.json").write_text(
json.dumps(
{
"status": "awaiting_speaker_review",
"diarization": {"speaker_count": 2},
}
),
encoding="utf-8",
)
progress_sink(
SimpleNamespace(
stage="diarization", status="completed", elapsed_seconds=0.3,
progress=None, message=None,
)
)
return SimpleNamespace(exit_code=0, run_dir=self.run_dir, protocol_path=None)
protocol = self.run_dir / "protocol.md"
protocol.write_text("# Generated protocol\n", encoding="utf-8")
return SimpleNamespace(exit_code=0, run_dir=self.run_dir, protocol_path=protocol)
@@ -387,7 +406,7 @@ def test_process_propagates_enabled_diarization_and_progress(tmp_path: Path) ->
audio.write_bytes(b"audio")
events = []
service.process(
outcome = service.process(
audio,
meeting(),
participants(),
@@ -397,6 +416,9 @@ def test_process_propagates_enabled_diarization_and_progress(tmp_path: Path) ->
assert gateway.config_values is not None
assert gateway.config_values["diarization"] == "gpu"
assert gateway.config_values["stop_after_diarization"] is True
assert outcome.awaiting_speaker_review is True
assert outcome.detected_speaker_count == 2
assert [(event.stage, event.status) for event in events[:2]] == [
("preparing", "started"),
("preparing", "completed"),
@@ -404,6 +426,33 @@ def test_process_propagates_enabled_diarization_and_progress(tmp_path: Path) ->
assert events[2].progress == 0.25
def test_diarization_checkpoint_defers_protocol_but_disabled_run_does_not(tmp_path: Path) -> None:
service, gateway = make_service(tmp_path)
audio = tmp_path / "meeting.wav"
audio.write_bytes(b"audio")
checkpoint = service.process(
audio, meeting(), participants(), ProcessingOptions(diarization_enabled=True)
)
assert checkpoint.succeeded
assert checkpoint.awaiting_speaker_review
assert checkpoint.original_protocol is None
assert checkpoint.protocol_path is None
normal_root = tmp_path / "without-diarization"
normal_root.mkdir()
normal_service, normal_gateway = make_service(normal_root)
normal_audio = normal_root / "meeting.wav"
normal_audio.write_bytes(b"audio")
normal = normal_service.process(
normal_audio, meeting(), participants(), ProcessingOptions(diarization_enabled=False)
)
assert normal.succeeded
assert not normal.awaiting_speaker_review
assert normal.protocol_path is not None
assert normal_gateway.config_values["stop_after_diarization"] is False
def test_process_propagates_ordered_diarization_container_args(tmp_path: Path) -> None:
service, gateway = make_service(tmp_path)
service.settings = replace(
@@ -537,6 +586,49 @@ def test_protocol_regeneration_allows_unmapped_and_rejects_duplicate_participant
)
@pytest.mark.parametrize(
"mappings",
[
{"SPEAKER_00": "martin"},
{"SPEAKER_00": "martin", "SPEAKER_01": "anna"},
],
)
def test_protocol_generation_uses_partial_and_complete_confirmed_mappings(
tmp_path: Path, mappings: dict[str, str]
) -> None:
service, gateway = make_service(tmp_path)
write_speaker_review_artifacts(gateway.run_dir)
outcome = service.regenerate_protocol(gateway.run_dir, mappings)
assert outcome.succeeded
assert gateway.regeneration is not None
assert gateway.regeneration["meeting_context"].data["speaker_mappings"] == mappings
def test_failed_protocol_generation_keeps_checkpoint_mapping_for_retry(tmp_path: Path) -> None:
service, gateway = make_service(tmp_path)
source = write_speaker_review_artifacts(gateway.run_dir)
original_source = source.read_bytes()
original_regeneration = gateway.regenerate_protocol
def fail_regeneration(*args, **kwargs):
raise RuntimeError("model unavailable")
gateway.regenerate_protocol = fail_regeneration
with pytest.raises(RuntimeError, match="model unavailable"):
service.regenerate_protocol(gateway.run_dir, {"SPEAKER_00": "martin"})
persisted = yaml.safe_load(
(gateway.run_dir / "context" / "meeting_context.yaml").read_text(encoding="utf-8")
)
assert persisted["speaker_mappings"] == {"SPEAKER_00": "martin"}
gateway.regenerate_protocol = original_regeneration
retry = service.regenerate_protocol(gateway.run_dir, {"SPEAKER_00": "martin"})
assert retry.succeeded
assert source.read_bytes() == original_source
def test_fallback_attribution_loss_is_read_from_runtime_metadata(tmp_path: Path) -> None:
run_dir = tmp_path / "run"
protocol_dir = run_dir / "protocol"
+6 -4
View File
@@ -80,7 +80,7 @@ def test_widget_reruns_filter_release_and_update_status():
assert not app.exception
def test_anonymous_regeneration_remains_allowed(monkeypatch):
def test_checkpoint_allows_anonymous_protocol_generation(monkeypatch):
from mka.application.meeting_service import SpeakerMappingReview, SpeakerReview
from mka.ui import streamlit_app as ui
@@ -89,6 +89,8 @@ def test_anonymous_regeneration_remains_allowed(monkeypatch):
run_dir=Path("/tmp/mapping-test"),
speaker_attribution_available=True,
original_protocol="Anonymous protocol",
awaiting_speaker_review=True,
detected_speaker_count=1,
)
review = SpeakerMappingReview((SpeakerReview("SPEAKER_00", ()),), (("a", "A"),), {})
calls = []
@@ -124,9 +126,9 @@ def test_anonymous_regeneration_remains_allowed(monkeypatch):
_render_result()
app = AppTest.from_function(result_app, args=(outcome,)).run()
regenerate = app.button[1]
assert not regenerate.disabled
generate = next(button for button in app.button if button.label == "Generate protocol")
assert not generate.disabled
assert app.warning
regenerate.click().run()
generate.click().run()
assert calls == [{}]
assert not app.exception