"""Formula extraction from PDF pages. Two backends are supported: 1. **MathPix** (cloud): POST the PDF to ``api.mathpix.com/v3/pdf`` and poll for the result JSON. Requires ``MATHPIX_APP_ID`` and ``MATHPIX_APP_KEY`` environment variables (or ``.env`` file). 2. **pix2tex** (local, MIT licence): crop each detected math bounding-box at 3× resolution and run the ``LatexOCR`` model locally. Used when MathPix credentials are absent. Model (~116 MB) is downloaded on first call. The public entry point :func:`extract_formulas` routes to whichever backend is available. """ from __future__ import annotations import logging import unicodedata from pathlib import Path from typing import Any from codex.config import get_settings from codex.models import FormulaChunk logger = logging.getLogger(__name__) # Unicode categories that indicate mathematical symbols. _MATH_CATEGORIES: frozenset[str] = frozenset({"Sm", "So"}) # Bbox size thresholds — avoids cropping tiny index fragments. _MIN_HEIGHT_PT = 15.0 _MIN_WIDTH_PT = 40.0 # Minimum number of math-category characters in a text block before we # consider it a formula candidate. _MIN_MATH_SYMBOLS = 5 # Minimum ratio of math/total chars (as float 0–1). _MIN_MATH_RATIO = 0.15 # --------------------------------------------------------------------------- # Internal helpers # --------------------------------------------------------------------------- def _is_math_char(ch: str) -> bool: """Return True if *ch* belongs to a mathematical Unicode category.""" return unicodedata.category(ch) in _MATH_CATEGORIES def _math_ratio(text: str) -> float: """Return the fraction of characters in *text* that are math chars.""" if not text: return 0.0 math_count = sum(1 for ch in text if _is_math_char(ch)) return math_count / len(text) def _math_symbol_count(text: str) -> int: """Return the absolute count of math-category characters in *text*.""" return sum(1 for ch in text if _is_math_char(ch)) # --------------------------------------------------------------------------- # pix2tex singleton # --------------------------------------------------------------------------- _pix2tex_model: Any = None def _get_pix2tex_model() -> Any: """Lazily load and return the pix2tex LatexOCR singleton.""" global _pix2tex_model # noqa: PLW0603 if _pix2tex_model is None: from pix2tex.cli import LatexOCR _pix2tex_model = LatexOCR() return _pix2tex_model # --------------------------------------------------------------------------- # pix2tex backend # --------------------------------------------------------------------------- def _extract_formulas_pix2tex(pdf_path: str) -> list[FormulaChunk]: """Extract LaTeX formulas from *pdf_path* using pix2tex (local OCR). Algorithm --------- 1. Open the PDF with PyMuPDF (fitz). 2. For each page, gather text blocks; keep only blocks whose math-symbol ratio exceeds the threshold and whose bounding box is larger than the minimum size. 3. Crop the region at 3× resolution (better OCR quality). 4. Run pix2tex ``LatexOCR`` on the cropped image. 5. Wrap the result in a :class:`~codex.models.FormulaChunk`. """ import fitz from PIL import Image results: list[FormulaChunk] = [] paper_id = Path(pdf_path).stem model = _get_pix2tex_model() try: doc: Any = fitz.open(pdf_path) except Exception: logger.exception("fitz.open failed for %s", pdf_path) return results scale = fitz.Matrix(3, 3) try: for page_num in range(len(doc)): page = doc[page_num] blocks: list[Any] = page.get_text("blocks") for block in blocks: # block = (x0, y0, x1, y1, text, block_no, block_type) if len(block) < 5: continue x0, y0, x1, y1 = block[0], block[1], block[2], block[3] text: str = block[4] # Size filter width_pt = x1 - x0 height_pt = y1 - y0 if width_pt < _MIN_WIDTH_PT or height_pt < _MIN_HEIGHT_PT: continue # Math content filter if ( _math_symbol_count(text) < _MIN_MATH_SYMBOLS or _math_ratio(text) < _MIN_MATH_RATIO ): continue # Crop at 3× for OCR quality clip_rect = fitz.Rect(x0, y0, x1, y1) crop_pix = page.get_pixmap(matrix=scale, clip=clip_rect) img = Image.frombytes( "RGB", (crop_pix.width, crop_pix.height), crop_pix.samples, ) try: latex: str = model(img) except Exception: logger.debug("pix2tex failed on page %d block, skipping", page_num + 1) continue if not latex or not latex.strip(): continue context = text[:200].replace("\n", " ") results.append( FormulaChunk( paper_id=paper_id, page=page_num + 1, raw_latex=latex.strip(), context=context, ) ) finally: doc.close() return results # --------------------------------------------------------------------------- # MathPix backend # --------------------------------------------------------------------------- def _extract_formulas_mathpix( pdf_path: str, app_id: str, app_key: str, ) -> list[FormulaChunk]: """Extract LaTeX formulas using the MathPix v3 PDF API. The function uploads the PDF, polls for results and parses the returned ``latex_simplified`` list into :class:`~codex.models.FormulaChunk` objects. Retries with exponential back-off (via ``tenacity``) on transient HTTP errors. """ import httpx from tenacity import retry, stop_after_attempt, wait_exponential paper_id = Path(pdf_path).stem headers = { "app_id": app_id, "app_key": app_key, } @retry( stop=stop_after_attempt(4), wait=wait_exponential(multiplier=1, min=2, max=30), reraise=True, ) def _post_pdf() -> dict[str, Any]: with open(pdf_path, "rb") as fh: resp = httpx.post( "https://api.mathpix.com/v3/pdf", headers=headers, files={"file": fh}, timeout=60.0, ) resp.raise_for_status() return resp.json() # type: ignore[no-any-return] submit_response = _post_pdf() pdf_id: str = submit_response.get("pdf_id", "") if not pdf_id: logger.warning("MathPix did not return a pdf_id — aborting") return [] # Poll until complete import time for _ in range(60): poll_resp = httpx.get( f"https://api.mathpix.com/v3/pdf/{pdf_id}.mmd", headers=headers, timeout=30.0, ) if poll_resp.status_code == 200: break time.sleep(2) else: logger.warning("MathPix polling timed out for pdf_id=%s", pdf_id) return [] # Parse response — MathPix MMD contains \[ ... \] or $...$ blocks mmd_text: str = poll_resp.text results: list[FormulaChunk] = [] lines = mmd_text.splitlines() in_block = False block_lines: list[str] = [] page = 1 for line in lines: if line.strip().startswith("\\["): in_block = True block_lines = [line] elif in_block: block_lines.append(line) if line.strip().endswith("\\]"): latex = "\n".join(block_lines).strip() results.append( FormulaChunk( paper_id=paper_id, page=page, raw_latex=latex, context="", ) ) in_block = False block_lines = [] # crude page tracking via MMD page markers if line.strip().startswith("