Add diarization and reusable MVP meeting pipeline

This commit is contained in:
2026-08-23 19:29:47 +02:00
parent f2d21c1faf
commit 8dab928763
18 changed files with 1697 additions and 180 deletions
+360
View File
@@ -0,0 +1,360 @@
"""Reusable audio-to-direct-protocol MVP orchestration."""
from __future__ import annotations
import json
import re
import shutil
import sys
import time
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Any, Callable, Mapping, Sequence
from src.meeting_lab.diarization import (
DEFAULT_MODEL as DEFAULT_DIARIZATION_MODEL,
diarize_audio,
write_diarized_transcript,
)
from src.meeting_lab.llm.ollama import DEFAULT_ENDPOINT
from src.meeting_lab.models.meeting_context import (
MeetingContext,
create_meeting_context,
load_meeting_context,
validate_meeting_context,
write_meeting_context,
)
from src.meeting_lab.progress import ProgressEvent, ProgressSink, ProgressStatus
from src.meeting_lab.protocol.generate_direct_protocol import (
DEFAULT_MODEL,
DirectProtocolResult,
generate_direct_protocol,
load_compact_transcript,
)
from src.meeting_lab.transcription.whisper import transcribe_audio
DEFAULT_OUTPUT_ROOT = Path("meeting_data/runs")
ContextInput = MeetingContext | Mapping[str, Any]
@dataclass(frozen=True)
class MvpMeetingConfig:
audio_file: Path
whisper_model: Path
whisper_executable: str = "whisper-cli"
context_file: Path | None = None
output_root: Path = DEFAULT_OUTPUT_ROOT
language: str = "de"
threads: str | int = "auto"
model: str = DEFAULT_MODEL
ollama_endpoint: str = DEFAULT_ENDPOINT
diarization: str = "off"
diarization_runtime: str = "native"
diarization_container_image: str | None = None
diarization_container_args: Sequence[str] = ()
@dataclass(frozen=True)
class MvpRunResult:
exit_code: int
run_dir: Path | None
protocol_path: Path | None
def create_unique_run_dir(
output_root: Path,
meeting_name: str,
now: Callable[[], datetime] = datetime.now,
) -> Path:
safe_name = re.sub(r"[^A-Za-z0-9_.-]+", "_", meeting_name).strip("._-") or "meeting"
base = output_root / f"{safe_name}_{now().strftime('%Y%m%d_%H%M%S')}"
candidate = base
suffix = 1
while candidate.exists():
candidate = output_root / f"{base.name}_{suffix:02d}"
suffix += 1
candidate.mkdir(parents=True)
return candidate
def _write_json(path: Path, value: Any) -> None:
path.write_text(json.dumps(value, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
def _effective_context(value: ContextInput | None) -> MeetingContext | None:
if value is None:
return None
if isinstance(value, MeetingContext):
validate_meeting_context(value.data)
return value
if isinstance(value, Mapping):
return create_meeting_context(dict(value), source_file=Path("<programmatic>"))
raise TypeError("meeting_context must be MeetingContext, mapping, or None.")
def _validate_inputs(
config: MvpMeetingConfig, meeting_context: MeetingContext | None
) -> None:
if not config.audio_file.is_file():
raise FileNotFoundError(f"Audio file does not exist: {config.audio_file}")
if not config.whisper_model.is_file():
raise FileNotFoundError(f"Whisper model does not exist: {config.whisper_model}")
if meeting_context is not None and config.context_file is not None:
raise ValueError("Use either a context file or a programmatic Meeting Context, not both.")
if meeting_context is not None:
validate_meeting_context(meeting_context.data)
elif config.context_file is not None:
if not config.context_file.is_file():
raise FileNotFoundError(
f"Meeting Context file does not exist: {config.context_file}"
)
load_meeting_context(config.context_file)
if config.diarization not in ("off", "auto", "gpu", "cpu"):
raise ValueError(f"Unsupported diarization mode: {config.diarization}")
if config.diarization_runtime not in ("native", "container"):
raise ValueError(
f"Unsupported diarization runtime: {config.diarization_runtime}"
)
if (
config.diarization != "off"
and config.diarization_runtime == "container"
and not config.diarization_container_image
):
raise ValueError("A diarization container image is required.")
def _emit(
sink: ProgressSink | None,
stage: str,
status: ProgressStatus,
overall_started: float,
*,
message: str | None = None,
) -> None:
if sink is not None:
sink(
ProgressEvent(
stage=stage,
status=status,
elapsed_seconds=time.perf_counter() - overall_started,
message=message,
)
)
def _persist_protocol(run_dir: Path, result: DirectProtocolResult) -> Path:
protocol_dir = run_dir / "protocol"
protocol_dir.mkdir(exist_ok=True)
(protocol_dir / "exact_prompt.txt").write_text(result.exact_prompt, encoding="utf-8")
_write_json(protocol_dir / "raw_response.json", result.raw_response)
_write_json(protocol_dir / "runtime_metadata.json", result.runtime_metadata)
protocol_path = run_dir / "protocol.md"
protocol_path.write_text(result.protocol_text, encoding="utf-8")
return protocol_path
def run_mvp_meeting(
config: MvpMeetingConfig,
*,
meeting_context: ContextInput | None = None,
progress_sink: ProgressSink | None = None,
) -> MvpRunResult:
"""Run the existing MVP directly, without subprocess or GUI dependencies."""
overall_started = time.perf_counter()
validation_started = time.perf_counter()
_emit(progress_sink, "preparing", "started", overall_started)
try:
effective_context = _effective_context(meeting_context)
_validate_inputs(config, effective_context)
except Exception as exc:
_emit(
progress_sink,
"failed",
"failed",
overall_started,
message=f"preparing: {type(exc).__name__}: {exc}",
)
print(f"Error: {type(exc).__name__}: {exc}", file=sys.stderr)
return MvpRunResult(2, None, None)
validation_runtime = time.perf_counter() - validation_started
run_dir = create_unique_run_dir(config.output_root, config.audio_file.stem)
timestamp = datetime.now().astimezone().isoformat(timespec="seconds")
transcript_path = run_dir / "transcript" / "transcript.json"
protocol_path = run_dir / "protocol.md"
stage_runtimes: dict[str, float | None] = {
"validation": round(validation_runtime, 3),
"setup": None,
"whisper": None,
"transcript_validation": None,
"protocol": None,
}
if config.diarization != "off":
stage_runtimes["diarization"] = None
stage_runtimes["diarization_alignment"] = None
metadata: dict[str, Any] = {
"run_id": run_dir.name,
"timestamp": timestamp,
"input_audio": str(config.audio_file.resolve()),
"transcript_output": str(transcript_path.resolve()),
"protocol_output": str(protocol_path.resolve()),
"whisper_model": str(config.whisper_model.resolve()),
"model": config.model,
"ollama_endpoint": config.ollama_endpoint,
"status": "running",
"stage_runtimes_seconds": stage_runtimes,
"total_runtime_seconds": None,
"failure": None,
"diarization": {
"enabled": config.diarization != "off",
"backend": "pyannote.audio" if config.diarization != "off" else None,
"model": DEFAULT_DIARIZATION_MODEL if config.diarization != "off" else None,
"requested_device_mode": config.diarization,
"runtime": config.diarization_runtime if config.diarization != "off" else None,
"metadata_path": None,
"transcript_diarized": None,
},
}
current_stage = "preparing"
stage_started = time.perf_counter()
try:
audio_dir = run_dir / "audio"
transcript_dir = run_dir / "transcript"
context_dir = run_dir / "context"
protocol_dir = run_dir / "protocol"
audio_dir.mkdir()
transcript_dir.mkdir()
context_dir.mkdir()
protocol_dir.mkdir()
_write_json(
audio_dir / "input_manifest.json",
{
"source_file": str(config.audio_file.resolve()),
"filename": config.audio_file.name,
"size_bytes": config.audio_file.stat().st_size,
},
)
preserved_context: Path | None = None
if effective_context is not None:
preserved_context = context_dir / "meeting_context.yaml"
write_meeting_context(effective_context, preserved_context)
elif config.context_file is not None:
preserved_context = context_dir / "meeting_context.yaml"
shutil.copy2(config.context_file, preserved_context)
stage_runtimes["setup"] = round(time.perf_counter() - stage_started, 3)
_emit(progress_sink, "preparing", "completed", overall_started)
current_stage = "transcription"
stage_started = time.perf_counter()
_emit(progress_sink, "transcription", "started", overall_started)
transcription = transcribe_audio(
config.audio_file,
config.whisper_model,
transcript_dir,
config.language,
executable=config.whisper_executable,
threads=config.threads,
)
stage_runtimes["whisper"] = round(time.perf_counter() - stage_started, 3)
_emit(progress_sink, "transcription", "completed", overall_started)
stage_started = time.perf_counter()
load_compact_transcript(transcription.transcript_json)
stage_runtimes["transcript_validation"] = round(
time.perf_counter() - stage_started, 3
)
protocol_transcript = transcription.transcript_json
if config.diarization != "off":
current_stage = "diarization"
stage_started = time.perf_counter()
_emit(progress_sink, "diarization", "started", overall_started)
diarization_dir = run_dir / "diarization"
diarization = diarize_audio(
config.audio_file,
diarization_dir,
config.diarization,
runtime=config.diarization_runtime,
container_image=config.diarization_container_image,
container_args=config.diarization_container_args,
)
stage_runtimes["diarization"] = round(
time.perf_counter() - stage_started, 3
)
metadata["diarization"].update(
{
"actual_device": diarization.metadata.get("actual_device"),
"device_name": diarization.metadata.get("device_name"),
"runtime_seconds": diarization.metadata.get("runtime_seconds"),
"speaker_count": diarization.metadata.get("speaker_count"),
"metadata_path": str(diarization.metadata_path.resolve()),
}
)
stage_started = time.perf_counter()
protocol_transcript, diarized_text = write_diarized_transcript(
transcription.transcript_json,
diarization.exclusive_turns_json,
diarization_dir,
)
load_compact_transcript(protocol_transcript)
stage_runtimes["diarization_alignment"] = round(
time.perf_counter() - stage_started, 3
)
metadata["diarization"].update(
{
"transcript_diarized": str(protocol_transcript.resolve()),
"transcript_diarized_text": str(diarized_text.resolve()),
}
)
_emit(progress_sink, "diarization", "completed", overall_started)
current_stage = "protocol_generation"
stage_started = time.perf_counter()
_emit(progress_sink, "protocol_generation", "started", overall_started)
result = generate_direct_protocol(
protocol_transcript,
preserved_context,
model=config.model,
endpoint=config.ollama_endpoint,
)
stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3)
protocol_path = _persist_protocol(run_dir, result)
_emit(progress_sink, "protocol_generation", "completed", overall_started)
metadata["status"] = "completed"
_emit(progress_sink, "completed", "completed", overall_started)
except Exception as exc:
metadata_stage = {
"preparing": "setup",
"transcription": "whisper",
"diarization": "diarization",
"protocol_generation": "protocol",
}.get(current_stage, current_stage)
runtime_key = metadata_stage
if runtime_key in stage_runtimes and stage_runtimes[runtime_key] is None:
stage_runtimes[runtime_key] = round(time.perf_counter() - stage_started, 3)
metadata["status"] = "failed"
metadata["failure"] = {
"stage": metadata_stage,
"type": type(exc).__name__,
"message": str(exc),
}
protocol_path = None
_emit(
progress_sink,
"failed",
"failed",
overall_started,
message=f"{current_stage}: {type(exc).__name__}: {exc}",
)
print(f"Error: {type(exc).__name__}: {exc}", file=sys.stderr)
finally:
metadata["total_runtime_seconds"] = round(time.perf_counter() - overall_started, 3)
_write_json(run_dir / "run_metadata.json", metadata)
exit_code = 0 if metadata["status"] == "completed" else 2
return MvpRunResult(exit_code, run_dir, protocol_path)