484 lines
19 KiB
Python
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)
|