first commit
This commit is contained in:
Binary file not shown.
@@ -0,0 +1,235 @@
|
||||
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()
|
||||
Reference in New Issue
Block a user