import asyncio import hashlib import json import os import re from contextlib import asynccontextmanager from pathlib import Path from dotenv import load_dotenv from fastapi import FastAPI, HTTPException, UploadFile from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import FileResponse from fastapi.staticfiles import StaticFiles from pydantic import BaseModel from sse_starlette.sse import EventSourceResponse from .annotate import (CompletionCache, ProviderError, build_prompt, cache_key, stream_completion) from .bootstrap import ensure_corpus from .equations import anchor_equations, extract_equations from .figures import extract_figures from .jobs import STAGES, Job, JobRegistry, emit from .parse import parse_pdf, parse_text from .references import parse_references from .retrieve import Corpus, embed_batch, retrieve_for_paragraphs from .schemas import Paragraph load_dotenv(Path(__file__).resolve().parent.parent / ".env") # --------------------------------------------------------------------------- # Startup / shutdown lifespan # --------------------------------------------------------------------------- _DEFAULT_DATA_DIR = Path(__file__).resolve().parent.parent / "data" # Repo root is two levels up from this file (astroparse_api/main.py → backend/ → repo root) _REPO_ROOT = Path(__file__).resolve().parent.parent.parent @asynccontextmanager async def lifespan(app: FastAPI): if os.environ.get("ASTROPARSE_SKIP_CORPUS") != "1": corpus_dir = Path( os.environ.get("ASTROPARSE_CORPUS_DIR", str(_DEFAULT_DATA_DIR / "corpus")) ) faiss_path = Path( os.environ.get("ASTROPARSE_FAISS_PATH", str(_DEFAULT_DATA_DIR / "astroparse_fp16.faiss")) ) beacon_repo = os.environ.get("BEACON_CORPUS_REPO", "kiyer/beacon_corpus") await asyncio.to_thread(ensure_corpus, corpus_dir, faiss_path, beacon_repo) await asyncio.to_thread(get_corpus) yield app = FastAPI(lifespan=lifespan) app.add_middleware( CORSMiddleware, allow_origins=["http://localhost:5173"], allow_methods=["*"], allow_headers=["*"], ) registry = JobRegistry() cache = CompletionCache( os.environ.get( "ASTROPARSE_CACHE_PATH", str(Path(__file__).resolve().parent.parent / "annotations_cache.sqlite"), ) ) _corpus: Corpus | None = None # Figures output directory — override via env var for testing _DEFAULT_FIGURES_DIR = Path(__file__).resolve().parent.parent / "data" / "figures" # Validation patterns for the figures endpoint _HASH_RE = re.compile(r"^[0-9a-f]{64}$") _FIGID_RE = re.compile(r"^(f|eq)\d+$") def get_corpus() -> Corpus: """Lazy singleton; tests monkeypatch this function.""" global _corpus if _corpus is None: _corpus = Corpus.load() return _corpus def get_figures_dir() -> Path: """Return the figures base directory, from env or default.""" return Path(os.environ.get("ASTROPARSE_FIGURES_DIR", str(_DEFAULT_FIGURES_DIR))) _TEXT_SUFFIXES = {".txt", ".md", ".markdown"} def _is_text_upload(file_name: str, content_type: str) -> bool: """Return True when the upload is plaintext/markdown (not a PDF).""" suffix = Path(file_name).suffix.lower() return suffix in _TEXT_SUFFIXES or content_type.startswith("text/") async def run_pipeline(job: Job, pdf_bytes: bytes, file_name: str, content_type: str = ""): try: async def stage(i, status, log=None): await emit(job, "stage", {"index": i, "name": STAGES[i], "status": status, "log": log}) is_text = _is_text_upload(file_name, content_type) # Stage 0: Parse await stage(0, "active") if is_text: paper, raw_refs = await asyncio.to_thread(parse_text, pdf_bytes, file_name) await stage(0, "done", "plaintext manuscript — no leaves to cut") else: paper, raw_refs = await asyncio.to_thread(parse_pdf, pdf_bytes, file_name) await stage(0, "done", f"{paper.pages} leaves read") # Defensive defaults — overwritten below if corpus loads successfully references_out: list[dict] = [] cite_index: dict[str, list[str]] = {} # Stage 1: Clean (already done in parse — emit done immediately) await stage(1, "active") await stage(1, "done", "headers and footers set aside") # Stage 2: Segment await stage(2, "active") await stage(2, "done", f"{len(paper.paragraphs)} paragraphs ruled") # Compute content hash for figures storage pdf_hash = hashlib.sha256(pdf_bytes).hexdigest() if is_text: # Text path: no figure or equation extraction figures = [] fig_embs = [] equations_out = [] else: # Extract figures (synchronous via pymupdf — cheap, no network) figures_dir = get_figures_dir() / pdf_hash figures_dir.mkdir(parents=True, exist_ok=True) figures = await asyncio.to_thread(extract_figures, pdf_bytes, figures_dir) eq_raw = await asyncio.to_thread(extract_equations, pdf_bytes, figures_dir) equations_with_anchor = anchor_equations(eq_raw, paper.paragraphs) equations_out = [ {"id": e["id"], "page": e["page"], "tag": e["tag"], "afterPara": e["afterPara"], "hasImage": e["hasImage"]} for e in equations_with_anchor ] # Stage 3: Embed paragraphs + captions in a single OpenAI call await stage(3, "active") key = os.environ.get("OPENAI_API_KEY", "") # Build combined text list: paragraphs first, then figure captions para_texts = [p.text for p in paper.paragraphs] if is_text: all_texts = para_texts else: fig_captions = [f.caption for f in figures] all_texts = para_texts + fig_captions all_embs = await asyncio.to_thread(embed_batch, all_texts, key) # Split embeddings back into paragraph and figure parts para_embs = all_embs[:len(para_texts)] if not is_text: fig_embs = all_embs[len(para_texts):] await stage(3, "done", f"{len(all_embs)} passages committed to memory") # Stage 4: Load corpus index await stage(4, "active") corpus = await asyncio.to_thread(get_corpus) await stage(4, "done", "lexicon ready") # Parse references section (uses corpus sets for bibcode/arXiv matching) parsed_refs, cite_index = parse_references( raw_refs, corpus.arxiv_ids, corpus.bibcodes ) references_out = [ { "id": r.id, "raw": r.raw, "short": r.short, "year": r.year, "bibcode": r.bibcode, "arxiv": r.arxiv, "corpusMatch": r.corpus_match, } for r in parsed_refs ] # Stage 5: Retrieve for paragraphs await stage(5, "active") lit_papers, lit_by_para = await asyncio.to_thread( retrieve_for_paragraphs, corpus, paper.paragraphs, para_embs ) if is_text: lit_by_fig = {} else: # Also retrieve for figure captions; merge into global lit_papers fig_pseudo_paras = [ Paragraph(id=fig.id, section=fig.label, text=fig.caption) for fig in figures ] if fig_pseudo_paras: fig_lit_papers, lit_by_fig = await asyncio.to_thread( retrieve_for_paragraphs, corpus, fig_pseudo_paras, fig_embs ) # Merge: first-write-wins (paragraphs take priority) for k, v in fig_lit_papers.items(): if k not in lit_papers: lit_papers[k] = v else: lit_by_fig = {} await stage(5, "done", f"{len(lit_papers)} works retrieved from the library") # Stage 6: Figures stage await stage(6, "active") if is_text: await stage(6, "done", "0 figures set in plates") else: await stage(6, "done", f"{len(figures)} figures set in plates") await emit(job, "result", { "paper": paper.model_dump(), "litPapers": {k: v.model_dump() for k, v in lit_papers.items()}, "litByPara": lit_by_para, "pdfHash": pdf_hash, "figures": [f.model_dump() for f in figures], "litByFig": lit_by_fig, "equations": equations_out, "references": references_out, "citeIndex": cite_index, }) except Exception as e: # noqa: BLE001 — job boundary await emit(job, "error", {"message": str(e)[:300]}) @app.post("/api/parse", status_code=202) async def parse_endpoint(file: UploadFile): pdf_bytes = await file.read() job = registry.create() asyncio.create_task(run_pipeline( job, pdf_bytes, file.filename or "manuscript.pdf", file.content_type or "", )) return {"jobId": job.id} @app.get("/api/jobs/{job_id}/events") async def job_events(job_id: str): job = registry.get(job_id) if job is None: raise HTTPException(404, "unknown job") async def gen(): while True: item = await job.queue.get() if item is None: break yield {"event": item["event"], "data": json.dumps(item["data"])} return EventSourceResponse(gen()) @app.get("/api/figures/{pdf_hash}/{fig_id}.png") async def serve_figure(pdf_hash: str, fig_id: str): """Serve a figure PNG; validates path params to prevent traversal.""" if not _HASH_RE.match(pdf_hash) or not _FIGID_RE.match(fig_id): raise HTTPException(404, "not found") png_path = get_figures_dir() / pdf_hash / f"{fig_id}.png" if not png_path.is_file(): raise HTTPException(404, "not found") return FileResponse(str(png_path), media_type="image/png") class PaperContext(BaseModel): title: str = "" sectionOutline: list[str] = [] opening: str = "" prevParagraph: str = "" nextParagraph: str = "" class AnnotateRequest(BaseModel): paragraphId: str paragraph: str section: str mode: str lit: list[dict] = [] provider: str model: str key: str paperContext: PaperContext | None = None @app.post("/api/annotate") async def annotate(req: AnnotateRequest): async def gen(): if not req.key: yield { "event": "error", "data": json.dumps({"message": "no API key set — open the annotator settings"}), } return ctx = req.paperContext.model_dump() if req.paperContext else None ck = cache_key(req.paragraph, req.mode, [l["id"] for l in req.lit], req.model, paper_context=ctx) hit = cache.get(ck) if hit is not None: yield {"event": "done", "data": json.dumps({"text": hit, "cached": True})} return prompt = build_prompt(req.paragraph, req.section, req.mode, req.lit, paper_context=ctx) full = [] try: async for chunk in stream_completion(req.provider, req.model, prompt, req.key): full.append(chunk) yield {"event": "token", "data": json.dumps({"text": chunk})} except ProviderError as e: yield {"event": "error", "data": json.dumps({"message": str(e)})} return except Exception as e: # noqa: BLE001 — network/timeout errors must not escape: # an escaped exception breaks the SSE stream AND (for gemini) can put the # key-bearing URL into server logs via the traceback. Emit a sanitized event. yield {"event": "error", "data": json.dumps({"message": f"could not reach {req.provider} ({type(e).__name__})"})} return text = "".join(full).strip() cache.put(ck, text) yield {"event": "done", "data": json.dumps({"text": text, "cached": False})} return EventSourceResponse(gen()) # --------------------------------------------------------------------------- # Static frontend serving — mount after all API routes so /api/* takes precedence # --------------------------------------------------------------------------- _static_dir = Path( os.environ.get("BEACON_STATIC_DIR", str(_REPO_ROOT / "frontend" / "dist")) ) if _static_dir.is_dir(): app.mount("/", StaticFiles(directory=str(_static_dir), html=True), name="spa")