Files
meeting-lab/src/meeting_lab/orchestration/mvp.py
T

484 lines
19 KiB
Python

"""Reusable audio-to-direct-protocol MVP orchestration."""
from __future__ import annotations
import json
import re
import shutil
import sys
import time
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from pathlib import Path
from typing import Any
from src.meeting_lab.audio import prepare_audio
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,
DEFAULT_NUM_CTX,
DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
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"
ffmpeg_executable: str = "ffmpeg"
audio_normalization: bool = True
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
glossary_aliases: Mapping[str, str] | None = None
protocol_num_thread: int | None = None
protocol_num_ctx: int = DEFAULT_NUM_CTX
protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET
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 regenerate_mvp_protocol(
run_dir: Path,
*,
meeting_context: ContextInput,
model: str = DEFAULT_MODEL,
ollama_endpoint: str = DEFAULT_ENDPOINT,
glossary_aliases: Mapping[str, str] | None = None,
protocol_num_thread: int | None = None,
protocol_num_ctx: int = DEFAULT_NUM_CTX,
protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
progress_sink: ProgressSink | None = None,
) -> MvpRunResult:
"""Regenerate only protocol artifacts from an existing completed run."""
started = time.perf_counter()
run_dir = Path(run_dir)
context = _effective_context(meeting_context)
if context is None:
raise ValueError("Meeting Context is required for protocol regeneration.")
if protocol_num_thread is not None and (
type(protocol_num_thread) is not int or protocol_num_thread <= 0
):
raise ValueError("Protocol Ollama thread count must be a positive integer.")
if protocol_num_ctx <= 0:
raise ValueError("Protocol Ollama context size must be positive.")
if protocol_safe_input_token_budget <= 0:
raise ValueError("Protocol safe input token budget must be positive.")
diarized_transcript = run_dir / "diarization" / "transcript_diarized.json"
plain_transcript = run_dir / "transcript" / "transcript.json"
transcript_path = (
diarized_transcript if diarized_transcript.is_file() else plain_transcript
)
if not transcript_path.is_file():
raise FileNotFoundError(
f"Existing run has no protocol transcript artifact: {run_dir}"
)
context_path = run_dir / "context" / "meeting_context.yaml"
write_meeting_context(context, context_path)
_emit(progress_sink, "protocol_generation", "started", started)
try:
result = generate_direct_protocol(
transcript_path,
context_path,
model=model,
endpoint=ollama_endpoint,
num_ctx=protocol_num_ctx,
num_thread=protocol_num_thread,
glossary_aliases=glossary_aliases,
safe_input_token_budget=protocol_safe_input_token_budget,
)
protocol_path = _persist_protocol(run_dir, result)
except Exception as exc:
_emit(
progress_sink,
"failed",
"failed",
started,
message=f"protocol_generation: {type(exc).__name__}: {exc}",
)
raise
_emit(progress_sink, "protocol_generation", "completed", started)
_emit(progress_sink, "completed", "completed", started)
return MvpRunResult(0, run_dir, protocol_path)
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.")
if config.protocol_safe_input_token_budget <= 0:
raise ValueError("Protocol safe input token budget must be positive.")
if config.protocol_num_thread is not None and (
type(config.protocol_num_thread) is not int or config.protocol_num_thread <= 0
):
raise ValueError("Protocol Ollama thread count must be a positive integer.")
if config.protocol_num_ctx <= 0:
raise ValueError("Protocol Ollama context size must be positive.")
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)
transcript_input = getattr(result, "transcript_input", None)
if transcript_input is not None:
(protocol_dir / "transcript_input.txt").write_text(
transcript_input, encoding="utf-8"
)
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,
"audio_preparation": 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()),
"audio_preparation": None,
"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,
},
)
preparation_started = time.perf_counter()
current_stage = "audio_preparation"
stage_started = preparation_started
prepared_audio = prepare_audio(
config.audio_file,
audio_dir / "prepared.wav",
ffmpeg_executable=config.ffmpeg_executable,
normalization_enabled=config.audio_normalization,
)
stage_runtimes["audio_preparation"] = round(
time.perf_counter() - preparation_started, 3
)
current_stage = "preparing"
preparation_metadata = prepared_audio.metadata()
metadata["audio_preparation"] = preparation_metadata
_write_json(audio_dir / "preparation_metadata.json", preparation_metadata)
_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,
"format": config.audio_file.suffix.lower().removeprefix("."),
"prepared_audio": preparation_metadata,
},
)
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(
prepared_audio.prepared_path,
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(
prepared_audio.prepared_path,
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,
num_ctx=config.protocol_num_ctx,
num_thread=config.protocol_num_thread,
glossary_aliases=config.glossary_aliases,
safe_input_token_budget=config.protocol_safe_input_token_budget,
)
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",
"audio_preparation": "audio_preparation",
"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)