#!/usr/bin/env python3 """Run the demo example papers end-to-end and save replay traces.""" from __future__ import annotations import json import os import shutil import sys import time from datetime import datetime, timezone from pathlib import Path SRC = Path(__file__).resolve().parent REPO_ROOT = SRC.parent for extra in (SRC, REPO_ROOT / "src"): extra_str = str(extra) if extra_str not in sys.path: sys.path.insert(0, extra_str) from dotenv import load_dotenv load_dotenv(REPO_ROOT / ".env") load_dotenv(REPO_ROOT.parent / "dryrun" / ".env", override=False) # Align Gemini env aliases used across the repo. if not os.getenv("GEMINI_API_KEY"): for alt in ("GOOGLE_GENAI_API_KEY", "GOOGLE_API_KEY"): if os.getenv(alt): os.environ["GEMINI_API_KEY"] = os.environ[alt] break import runner as runner_module from runner import PipelineConfig from common.paper_package import load_paper_package from step_08_annotation.pipeline import TwoPassAnnotationPipeline from streamlit_config import EXAMPLES REPLAY_ROOT = REPO_ROOT / "replay_traces" WORK_ROOT = REPO_ROOT / "hf_space" / "runs" / "replay_build" def _env(name: str, default: str) -> str: return (os.getenv(name) or default).strip() def _load_json(path: Path): if not path.exists(): return None try: return json.loads(path.read_text(encoding="utf-8")) except Exception: return None def _copy_if_exists(src: Path, dst: Path) -> bool: if not src.exists(): return False dst.parent.mkdir(parents=True, exist_ok=True) if src.is_dir(): if dst.exists(): shutil.rmtree(dst) shutil.copytree(src, dst) else: shutil.copy2(src, dst) return True def _run_annotation(paper_dir: Path, annotation_root: Path) -> tuple[dict | None, Path | None, str | None]: discovery = _load_json(paper_dir / "usage_discovery_from_contributions.json") or {} clusters = discovery.get("clusters") or [] if not clusters: return None, None, "No refined downstream usage clusters; annotation skipped." llm_provider = _env("LLM_PROVIDER", "gemini") llm_model = _env("LLM_MODEL", "gemini-3.1-pro-preview") formatter_model = _env("ANNOTATION_FORMATTER_MODEL", "gemini/gemini-3.1-pro-preview") judge_model = _env("ANNOTATION_JUDGE_MODEL", "gemini/gemini-3.1-pro-preview") candidate_count = int(_env("ANNOTATION_CANDIDATE_COUNT", "3")) paper = load_paper_package(paper_dir) pipeline = TwoPassAnnotationPipeline( provider=llm_provider, model=llm_model, formatter_model=formatter_model or None, judge_model=judge_model or None, output_root=annotation_root, annotator_id="replay_trace_builder", candidate_count=max(1, candidate_count), formatter_max_attempts=3, include_reference_examples=True, prompt_profile="full", ) result = pipeline.run(paper) return result.result, result.run_dir, None def _package_trace( *, label: str, arxiv_id: str, paper_input: str, job_dir: Path, paper_dir: Path, events: list[str], status: str, annotation_run_dir: Path | None, annotation_skipped_reason: str | None, pipeline_failed_reason: str | None, pipeline_stopped_reason: str | None, ) -> Path: out_dir = REPLAY_ROOT / arxiv_id if out_dir.exists(): shutil.rmtree(out_dir) out_dir.mkdir(parents=True, exist_ok=True) paper_out = out_dir / "processed_papers" / arxiv_id paper_out.mkdir(parents=True, exist_ok=True) # Core replay payloads (keep disk footprint manageable). keep_files = [ "paper_metadata.json", "usage_contexts.json", "usage_context_labels.json", "usage_uses_extends_verified.json", "usage_citing_paragraphs.json", "usage_contributions.json", "usage_discovery_from_contributions.json", ] for name in keep_files: _copy_if_exists(paper_dir / name, paper_out / name) _copy_if_exists(job_dir / "logs", out_dir / "logs") _copy_if_exists(job_dir / "summary.txt", out_dir / "summary.txt") _copy_if_exists(job_dir / "run_config.json", out_dir / "run_config.json") _copy_if_exists(job_dir / "input_ids.json", out_dir / "input_ids.json") annotation_payload_path = None if annotation_run_dir and annotation_run_dir.exists(): ann_dst = out_dir / "two_pass_outputs" / annotation_run_dir.name _copy_if_exists(annotation_run_dir, ann_dst) payload = ann_dst / "pass_2_ui_payload.json" if payload.exists(): annotation_payload_path = str(payload.relative_to(out_dir)) discovery = _load_json(paper_out / "usage_discovery_from_contributions.json") or {} contributions = _load_json(paper_out / "usage_contributions.json") or {} payload = None if annotation_payload_path: payload = _load_json(out_dir / annotation_payload_path) public_export = { "citation_clusters": (discovery or {}).get("clusters") or [], "target_contribution_decompositions": (payload or {}).get("claims") or [], } (out_dir / "scipaths_run_results.json").write_text( json.dumps(public_export, indent=2, ensure_ascii=False), encoding="utf-8", ) meta = { "label": label, "arxiv_id": arxiv_id, "paper_input": paper_input, "status": status, "built_at": datetime.now(timezone.utc).isoformat(), "source_job_dir": str(job_dir), "paper_dir": str((out_dir / "processed_papers" / arxiv_id).relative_to(out_dir)), "annotation_payload_path": annotation_payload_path, "annotation_skipped_reason": annotation_skipped_reason, "pipeline_failed_reason": pipeline_failed_reason, "pipeline_stopped_reason": pipeline_stopped_reason, "events": events, "cluster_count": len((discovery or {}).get("clusters") or []), "contribution_count": len((contributions or {}).get("contributions") or []), "claim_count": len((payload or {}).get("claims") or []) if isinstance(payload, dict) else 0, } (out_dir / "replay_meta.json").write_text(json.dumps(meta, indent=2, ensure_ascii=False), encoding="utf-8") return out_dir def run_one(label: str, paper_input: str) -> dict: if not os.getenv("GEMINI_API_KEY"): raise SystemExit("GEMINI_API_KEY is required to build replay traces.") arxiv_id = runner_module.parse_arxiv_id(paper_input) print(f"\n=== Building replay trace for {label} ({arxiv_id}) ===", flush=True) cfg = PipelineConfig( repo_root=REPO_ROOT, source_root=REPO_ROOT / "src" / "processed_papers", paper_input=paper_input, llm_provider=_env("LLM_PROVIDER", "gemini"), llm_model=_env("LLM_MODEL", "gemini-3.1-pro-preview"), llm_model_step4=_env("LLM_MODEL_STEP4", "gemini-3-flash-preview"), model_path="Deep-Citation/Workspace/acl_scicite_wksp_trl/best_model.pt", model_data_dir="Deep-Citation/Data", model_class_def="Deep-Citation/Data/class_def.json", model_lm="scibert", device="cpu", embedding_model="sentence-transformers/all-mpnet-base-v2", ) events: list[str] = [] artifact_path = None pipeline_failed_reason = None pipeline_stopped_reason = None t0 = time.time() for line, maybe_artifact in runner_module.run_pipeline(cfg, WORK_ROOT): if line: print(line, flush=True) events.append(line) if line.startswith("Pipeline stopped:"): pipeline_stopped_reason = line if "failed" in line.lower(): pipeline_failed_reason = line if maybe_artifact: artifact_path = maybe_artifact if not artifact_path: raise RuntimeError(f"Pipeline produced no artifact for {arxiv_id}") job_dir = Path(str(artifact_path)).with_suffix("") paper_dir = job_dir / "processed_papers" / arxiv_id annotation_run_dir = None annotation_skipped_reason = None if pipeline_failed_reason: status = "Failed" annotation_skipped_reason = f"{pipeline_failed_reason} Annotation was not run." elif pipeline_stopped_reason: status = "Stopped" annotation_skipped_reason = f"{pipeline_stopped_reason} Annotation was not run." else: print("[annotation] starting", flush=True) try: _run_output, annotation_run_dir, skip = _run_annotation( paper_dir=paper_dir, annotation_root=job_dir / "two_pass_outputs", ) if skip: annotation_skipped_reason = skip status = "Completed" print(f"[annotation] skipped: {skip}", flush=True) else: status = "Completed" events.append(f"[annotation] complete: {annotation_run_dir}") print(f"[annotation] complete: {annotation_run_dir}", flush=True) except Exception as exc: status = "Failed" pipeline_failed_reason = f"Annotation failed: {exc}" annotation_skipped_reason = pipeline_failed_reason events.append(pipeline_failed_reason) print(pipeline_failed_reason, flush=True) out_dir = _package_trace( label=label, arxiv_id=arxiv_id, paper_input=paper_input, job_dir=job_dir, paper_dir=paper_dir, events=events, status=status, annotation_run_dir=annotation_run_dir, annotation_skipped_reason=annotation_skipped_reason, pipeline_failed_reason=pipeline_failed_reason, pipeline_stopped_reason=pipeline_stopped_reason, ) elapsed = time.time() - t0 print(f"Saved replay trace -> {out_dir} ({status}, {elapsed/60:.1f} min)", flush=True) return { "label": label, "arxiv_id": arxiv_id, "status": status, "trace_dir": str(out_dir), "elapsed_sec": elapsed, "annotation_skipped_reason": annotation_skipped_reason, "pipeline_failed_reason": pipeline_failed_reason, "pipeline_stopped_reason": pipeline_stopped_reason, } def main() -> int: REPLAY_ROOT.mkdir(parents=True, exist_ok=True) WORK_ROOT.mkdir(parents=True, exist_ok=True) results = [] for label, paper_input in EXAMPLES.items(): results.append(run_one(label, paper_input)) index = { "built_at": datetime.now(timezone.utc).isoformat(), "examples": results, } index_path = REPLAY_ROOT / "index.json" index_path.write_text(json.dumps(index, indent=2, ensure_ascii=False), encoding="utf-8") print(f"\nWrote index -> {index_path}", flush=True) for item in results: print(f"- {item['arxiv_id']}: {item['status']} -> {item['trace_dir']}", flush=True) return 0 if all(item["status"] == "Completed" for item in results) else 1 if __name__ == "__main__": raise SystemExit(main())