Add post-diarization speaker review workflow
This commit is contained in:
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user