File size: 2,931 Bytes
4d25872
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
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]