#!/usr/bin/env python3 """Run one gold-corpus scenario through the existing extraction flow.""" from __future__ import annotations import argparse import json import sys import time from pathlib import Path from typing import Any import requests REPO_ROOT = Path(__file__).resolve().parents[1] if str(REPO_ROOT) not in sys.path: sys.path.insert(0, str(REPO_ROOT)) from src.meeting_lab.extraction.extract_chunks import ( DEFAULT_ENDPOINT, EXTRACTION_CATEGORIES, build_prompt, call_ollama, normalize_current_schema, parse_json_response, ) REQUIRED_KEYS = set(EXTRACTION_CATEGORIES) def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser( description="Run one gold-standard transcript through extraction." ) parser.add_argument("scenario", type=Path, help="Gold scenario directory") parser.add_argument("--model", required=True, help="Ollama model name") parser.add_argument( "--endpoint", default=DEFAULT_ENDPOINT, help=f"Ollama generate endpoint (default: {DEFAULT_ENDPOINT})", ) parser.add_argument( "--timeout", type=int, default=1800, help="HTTP timeout in seconds (default: 1800)", ) parser.add_argument( "--temperature", type=float, default=0.0, help="Sampling temperature (default: 0.0)", ) parser.add_argument( "--num-predict", type=int, default=8192, help="Maximum generated tokens (default: 8192)", ) parser.add_argument( "--num-ctx", type=int, default=32768, help="Context window tokens (default: 32768)", ) return parser.parse_args() def scenario_paths(scenario_dir: Path) -> tuple[Path, Path, Path]: if not scenario_dir.is_dir(): raise FileNotFoundError(f"Scenario directory not found: {scenario_dir}") transcript_path = scenario_dir / "transcript.txt" expected_path = scenario_dir / "expected.json" actual_path = scenario_dir / "actual.json" if not transcript_path.is_file(): raise FileNotFoundError(f"Missing transcript.txt: {transcript_path}") if not expected_path.is_file(): raise FileNotFoundError(f"Missing expected.json: {expected_path}") return transcript_path, expected_path, actual_path def read_json_object(path: Path) -> dict[str, Any]: try: data = json.loads(path.read_text(encoding="utf-8")) except json.JSONDecodeError as exc: raise ValueError(f"Invalid JSON in {path}: {exc}") from exc if not isinstance(data, dict): raise ValueError(f"JSON file must contain an object: {path}") return data def validate_required_keys(data: dict[str, Any], path: Path) -> None: missing = sorted(REQUIRED_KEYS - set(data)) if missing: raise ValueError(f"Missing required keys in {path}: {', '.join(missing)}") def format_items(items: Any) -> list[str]: if not isinstance(items, list): return [f""] if not items: return [""] return [str(item) for item in items] def print_decision_comparison( scenario_dir: Path, model: str, expected: dict[str, Any], actual: dict[str, Any], runtime: float, ) -> None: expected_decisions = expected.get("decisions", []) actual_decisions = actual.get("decisions", []) print(f"Scenario: {scenario_dir}") print(f"Model: {model}") print(f"Expected decision count: {len(expected_decisions)}") print(f"Actual decision count: {len(actual_decisions)}") print("Expected decisions:") for item in format_items(expected_decisions): print(f"- {item}") print("Actual decisions:") for item in format_items(actual_decisions): print(f"- {item}") print(f"Runtime: {runtime:.2f}s") def run_gold_test( scenario_dir: Path, model: str, endpoint: str, timeout: int, temperature: float, num_predict: int | None, num_ctx: int | None, ) -> Path: transcript_path, expected_path, actual_path = scenario_paths(scenario_dir) expected = read_json_object(expected_path) validate_required_keys(expected, expected_path) transcript = transcript_path.read_text(encoding="utf-8-sig").strip() if not transcript: raise ValueError(f"The transcript is empty: {transcript_path}") prompt = build_prompt(transcript_path.name, transcript) started = time.perf_counter() raw_text, _metadata = call_ollama( endpoint=endpoint, model=model, prompt=prompt, timeout=timeout, temperature=temperature, num_predict=num_predict, num_ctx=num_ctx, ) try: parsed = parse_json_response(raw_text) except (json.JSONDecodeError, ValueError) as exc: raw_path = actual_path.with_suffix(".raw.txt") raw_path.write_text(raw_text + "\n", encoding="utf-8") raise ValueError(f"Model output was not valid JSON. Raw output: {raw_path}") from exc actual = normalize_current_schema(parsed) actual_path.write_text( json.dumps(actual, ensure_ascii=False, indent=2) + "\n", encoding="utf-8", ) actual_from_disk = read_json_object(actual_path) validate_required_keys(actual_from_disk, actual_path) runtime = time.perf_counter() - started print_decision_comparison( scenario_dir=scenario_dir, model=model, expected=expected, actual=actual_from_disk, runtime=runtime, ) return actual_path def main() -> int: args = parse_args() try: actual_path = run_gold_test( scenario_dir=args.scenario, model=args.model, endpoint=args.endpoint, timeout=args.timeout, temperature=args.temperature, num_predict=args.num_predict, num_ctx=args.num_ctx, ) except requests.ConnectionError: print( "Error: Ollama is not reachable. Is `ollama serve` running?", file=sys.stderr, ) return 1 except requests.Timeout: print("Error: The Ollama request timed out.", file=sys.stderr) return 1 except requests.HTTPError as exc: print(f"Error: Ollama returned an HTTP error: {exc}", file=sys.stderr) return 1 except (OSError, UnicodeError, ValueError) as exc: print(f"Error: {exc}", file=sys.stderr) return 1 print(f"Actual JSON: {actual_path}") return 0 if __name__ == "__main__": raise SystemExit(main())