"""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("")) 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)