from dataclasses import dataclass from pathlib import Path import random import re import fitz from huggingface_hub import HfApi, hf_hub_download from tqdm import tqdm @dataclass class Contract: contract_id: str remote_path: str pdf_path: Path pages: list[dict] @property def text(self) -> str: return "\n\n".join(f"[PAGE {p['page']}] {p['text']}" for p in self.pages) def normalize_text(text: str) -> str: text = text.replace("\x00", " ") text = re.sub(r"-\s*\n\s*", "", text) text = re.sub(r"\s+", " ", text) return text.strip() def extract_pages(pdf_path: Path) -> list[dict]: with fitz.open(pdf_path) as document: return [ {"page": i, "text": normalize_text(page.get_text("text"))} for i, page in enumerate(document, start=1) if normalize_text(page.get_text("text")) ] def download_contracts(dataset_id: str, pdf_dir: Path, limit: int, seed: int) -> list[Contract]: api = HfApi() files = list(api.list_repo_files(repo_id=dataset_id, repo_type="dataset")) pdfs = sorted(p for p in files if p.lower().endswith(".pdf") and "full_contract" in p.lower()) if len(pdfs) < limit: raise RuntimeError(f"Found only {len(pdfs)} CUAD PDFs, requested {limit}.") rng = random.Random(seed) selected = pdfs.copy() rng.shuffle(selected) selected = sorted(selected[:limit]) contracts = [] for remote_path in tqdm(selected, desc="Downloading contracts"): local_path = Path(hf_hub_download( repo_id=dataset_id, repo_type="dataset", filename=remote_path, local_dir=str(pdf_dir), )) contracts.append(Contract( contract_id=Path(remote_path).stem, remote_path=remote_path, pdf_path=local_path, pages=extract_pages(local_path), )) return contracts KEYWORD_GROUPS = { "termination": ["terminate", "termination", "expiration", "renewal", "notice", "breach"], "confidentiality": ["confidential", "non-disclosure", "proprietary information", "trade secret"], "liability": ["limitation of liability", "liable", "liability", "damages", "indemnif", "hold harmless"], } def retrieve_context(contract: Contract, max_pages: int = 16, max_chars: int = 40000) -> str: pages = contract.pages scored = [] for page in pages: lower = page["text"].lower() score = sum(lower.count(term) for terms in KEYWORD_GROUPS.values() for term in terms) scored.append((score, page["page"], page["text"])) chosen = {p["page"]: p["text"] for p in pages[:3]} for score, number, text in sorted(scored, reverse=True): if score > 0: chosen[number] = text if len(chosen) >= max_pages: break return "\n\n".join(f"[PAGE {number}] {chosen[number]}" for number in sorted(chosen))[:max_chars]