from __future__ import annotations import argparse import json import os import re import sys from pathlib import Path from typing import Any import httpx REPO_ROOT = Path(__file__).resolve().parents[1] SRC_ROOT = REPO_ROOT / "src" if str(SRC_ROOT) not in sys.path: sys.path.insert(0, str(SRC_ROOT)) from summary_mcp.validators.llm_result import validate_llm_result JSON_BLOCK_RE = re.compile(r"```(?:json)?\s*(\{.*\})\s*```", re.DOTALL) def load_text(path: Path) -> str: return path.read_text(encoding="utf-8") def load_json(path: Path) -> dict[str, Any]: return json.loads(load_text(path)) def save_json(path: Path, payload: dict[str, Any]) -> None: path.write_text(json.dumps(payload, ensure_ascii=False, indent=2), encoding="utf-8") def build_summary_input(extracted: dict[str, Any]) -> dict[str, Any]: article = extracted.get('article') or {} return { 'article': { 'title': article.get('title'), 'url': article.get('url'), 'plain_text': article.get('plain_text'), 'quality_flags': article.get('quality_flags'), }, 'warnings': extracted.get('warnings', []), } def build_initial_prompt(prompt_template: str, extracted: dict[str, Any]) -> str: summary_input = build_summary_input(extracted) return ( f"{prompt_template}\n\n" "Below is the structured extracted article input. Generate the final summary JSON from it.\n\n" f"{json.dumps(summary_input, ensure_ascii=False, indent=2)}" ) def build_repair_prompt( errors: list[str], extracted: dict[str, Any], result_json: dict[str, Any], ) -> str: summary_input = build_summary_input(extracted) return ( "Please repair the following invalid summary JSON.\n\n" "Requirements:\n" "- Output valid JSON only\n" "- Keep fields that are already correct\n" "- Fix only the validator-reported errors\n" "- Do not add explanations\n\n" f"validator errors:\n{json.dumps(errors, ensure_ascii=False, indent=2)}\n\n" f"Extracted article input:\n{json.dumps(summary_input, ensure_ascii=False, indent=2)}\n\n" f"Current summary JSON:\n{json.dumps(result_json, ensure_ascii=False, indent=2)}\n" ) def extract_json_text(raw_text: str) -> str: fenced = JSON_BLOCK_RE.search(raw_text) if fenced: return fenced.group(1) stripped = raw_text.strip() start = stripped.find("{") end = stripped.rfind("}") if start == -1 or end == -1 or end <= start: raise ValueError("Model output does not contain a JSON object.") return stripped[start : end + 1] def call_llm( prompt: str, timeout_seconds: float, api_key: str | None, model: str | None, api_url: str | None, ) -> str: api_key = api_key or os.environ.get("LLM_API_KEY") or os.environ.get("OPENAI_API_KEY") model = model or os.environ.get("LLM_MODEL") or os.environ.get("OPENAI_MODEL") api_url = api_url or os.environ.get("LLM_API_URL", "https://api.openai.com/v1/chat/completions") if not api_key: raise RuntimeError("Missing LLM_API_KEY or OPENAI_API_KEY, or pass --api-key.") if not model: raise RuntimeError("Missing LLM_MODEL or OPENAI_MODEL, or pass --model.") headers = { "Authorization": f"Bearer {api_key}", "Content-Type": "application/json", } payload = { "model": model, "messages": [ { "role": "system", "content": "You are a precise JSON generator. Always output a single valid JSON object.", }, {"role": "user", "content": prompt}, ], "temperature": 0.2, } with httpx.Client(timeout=timeout_seconds) as client: response = client.post(api_url, headers=headers, json=payload) response.raise_for_status() data = response.json() try: return data["choices"][0]["message"]["content"] except (KeyError, IndexError, TypeError) as exc: raise RuntimeError(f"Unexpected LLM response shape: {json.dumps(data, ensure_ascii=False)[:1000]}") from exc def run_loop( extracted_path: Path, prompt_path: Path, output_path: Path, max_retries: int, timeout_seconds: float, api_key: str | None, model: str | None, api_url: str | None, ) -> int: extracted = load_json(extracted_path) prompt_template = load_text(prompt_path) output_path.parent.mkdir(parents=True, exist_ok=True) last_errors: list[str] = [] last_result: dict[str, Any] | None = None for attempt in range(1, max_retries + 2): if attempt == 1: prompt = build_initial_prompt(prompt_template, extracted) else: assert last_result is not None prompt = build_repair_prompt(last_errors, extracted, last_result) raw_output = call_llm(prompt, timeout_seconds, api_key, model, api_url) raw_path = output_path.with_name(f"{output_path.stem}.attempt-{attempt}.raw.txt") raw_path.write_text(raw_output, encoding="utf-8") try: result_payload = json.loads(extract_json_text(raw_output)) except (json.JSONDecodeError, ValueError) as exc: last_errors = [f"Model output is not valid JSON: {exc}"] last_result = {"raw_output": raw_output} validation_path = output_path.with_name(f"{output_path.stem}.attempt-{attempt}.validation.json") save_json( validation_path, { "valid": False, "errors": last_errors, "warnings": [], "normalized_result": None, }, ) if attempt > max_retries: output_path.write_text(raw_output, encoding="utf-8") return 1 continue attempt_path = output_path.with_name(f"{output_path.stem}.attempt-{attempt}.json") save_json(attempt_path, result_payload) save_json(output_path, result_payload) report = validate_llm_result(output_path, extracted_path) validation_path = output_path.with_name(f"{output_path.stem}.attempt-{attempt}.validation.json") save_json( validation_path, { "valid": report.valid, "errors": report.errors, "warnings": report.warnings, "normalized_result": report.normalized_result, }, ) if report.valid: return 0 last_errors = report.errors last_result = result_payload return 1 def main() -> None: parser = argparse.ArgumentParser(description="Run the minimal extraction -> LLM summary -> validation loop.") parser.add_argument("--extracted", type=Path, required=True, help="Extracted article JSON file") parser.add_argument("--prompt", type=Path, required=True, help="LLM prompt template file") parser.add_argument("--output", type=Path, required=True, help="Target path for the final summary JSON") parser.add_argument("--max-retries", type=int, default=2, help="Number of repair retries after the initial attempt") parser.add_argument("--timeout", type=float, default=60.0, help="LLM request timeout in seconds") parser.add_argument("--api-key", type=str, default=None, help="LLM API key") parser.add_argument("--model", type=str, default=None, help="LLM model name") parser.add_argument("--api-url", type=str, default=None, help="LLM chat completions API URL") args = parser.parse_args() raise SystemExit( run_loop( extracted_path=args.extracted, prompt_path=args.prompt, output_path=args.output, max_retries=args.max_retries, timeout_seconds=args.timeout, api_key=args.api_key, model=args.model, api_url=args.api_url, ) ) if __name__ == "__main__": main()