"""Compare 3 sources of text for the same Common Crawl pages: the WET text Common Crawl
ships, trafilatura with FineWeb's settings, and Resiliparse's main-content extraction.
Each output then goes through FineWeb's own quality filter from datatrove.
Run with: uv run --with datatrove --with spacy --with trafilatura --with resiliparse python compare.py"""
import base64
import gzip
import json
import time

import trafilatura
from datatrove.data import Document
from datatrove.pipeline.filters import FineWebQualityFilter
from resiliparse.extract.html2text import extract_plain_text
from resiliparse.parse.encoding import bytes_to_str, detect_encoding

quality = FineWebQualityFilter()


def with_trafilatura(raw):
    # The settings datatrove's Trafilatura extractor uses for FineWeb.
    return trafilatura.extract(raw, favor_precision=True, include_comments=False, deduplicate=True) or ""


def with_resiliparse(raw):
    return extract_plain_text(bytes_to_str(raw, detect_encoding(raw)), main_content=True)


def check(text, doc_id):
    """FineWeb's quality filter: True, or the reason the page was removed."""
    if not text.strip():
        return "empty"
    result = quality.filter(Document(text=text, id=doc_id))
    return True if result is True else result[1]


pages = [json.loads(line) for line in gzip.open("sample.jsonl.gz", "rt")]
pages = [p for p in pages if "wet" in p]
seconds = {"trafilatura": 0.0, "resiliparse": 0.0}
rows = []
for i, page in enumerate(pages):
    raw = base64.b64decode(page["html_b64"])
    texts = {"wet": page["wet"]}
    for name, fn in [("trafilatura", with_trafilatura), ("resiliparse", with_resiliparse)]:
        start = time.perf_counter()
        texts[name] = fn(raw)
        seconds[name] += time.perf_counter() - start
    rows.append({
        "url": page["url"],
        "languages": page.get("languages"),
        "html_bytes": len(raw),
        **{f"{k}_words": len(v.split()) for k, v in texts.items()},
        **{f"{k}_fineweb": check(v, str(i)) for k, v in texts.items()},
    })

json.dump({"seconds": seconds, "pages": rows}, open("results.json", "w"), indent=1)
print(len(rows), "pages")
print("extraction seconds:", {k: round(v, 1) for k, v in seconds.items()})
