#!/usr/bin/env python3
"""Fetch a frozen EN/JA Wikipedia article-creation cohort from official APIs."""

from __future__ import annotations

import argparse
import contextlib
import datetime as dt
import gzip
import hashlib
import io
import json
import pathlib
import time
import urllib.error
import urllib.parse
import urllib.request

SPARQL_ENDPOINT = "https://query.wikidata.org/sparql"
EN_API = "https://en.wikipedia.org/w/api.php"
JA_API = "https://ja.wikipedia.org/w/api.php"
ROOT_CLASSES = {
    "Q7397": "software",
    "Q9143": "programming language",
    "Q188860": "software library",
    "Q271680": "software framework",
    "Q9135": "operating system",
}
USER_AGENT = "TrendiumoResearch/1.0 (https://trendiumo.com/about; towelsoftgroup@gmail.com)"


def sha256_file(file: pathlib.Path) -> str:
    digest = hashlib.sha256()
    with file.open("rb") as stream:
        for chunk in iter(lambda: stream.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


@contextlib.contextmanager
def deterministic_gzip_text(file: pathlib.Path):
    file.parent.mkdir(parents=True, exist_ok=True)
    with file.open("wb") as raw:
        with gzip.GzipFile(filename="", mode="wb", fileobj=raw, mtime=0) as compressed:
            with io.TextIOWrapper(compressed, encoding="utf-8", newline="\n") as text:
                yield text


def iso_utc(value: str) -> dt.datetime:
    parsed = dt.datetime.fromisoformat(value.replace("Z", "+00:00"))
    return parsed.astimezone(dt.timezone.utc)


class ApiClient:
    def __init__(self, raw_writer, delay: float = 0.08, retries: int = 6):
        self.raw_writer = raw_writer
        self.delay = delay
        self.retries = retries
        self.request_count = 0

    def json(self, endpoint: str, params: dict, *, method: str = "POST") -> dict:
        encoded = urllib.parse.urlencode(params).encode("utf-8")
        url = endpoint
        data = encoded
        if method == "GET":
            url = f"{endpoint}?{encoded.decode('utf-8')}"
            data = None
        for attempt in range(self.retries):
            request = urllib.request.Request(
                url,
                data=data,
                headers={
                    "User-Agent": USER_AGENT,
                    "Accept": "application/json, application/sparql-results+json",
                    "Content-Type": "application/x-www-form-urlencoded",
                },
                method=method,
            )
            try:
                with urllib.request.urlopen(request, timeout=90) as response:
                    payload = json.load(response)
                if "error" in payload:
                    raise RuntimeError(f"API error from {endpoint}: {payload['error']}")
                self.request_count += 1
                self.raw_writer.write(json.dumps({
                    "sequence": self.request_count,
                    "endpoint": endpoint,
                    "params": params,
                    "response": payload,
                }, ensure_ascii=False, sort_keys=True, separators=(",", ":")) + "\n")
                self.raw_writer.flush()
                if self.delay:
                    time.sleep(self.delay)
                return payload
            except urllib.error.HTTPError as error:
                retryable = error.code in {429, 500, 502, 503, 504}
                if not retryable or attempt == self.retries - 1:
                    raise
                retry_after = error.headers.get("Retry-After")
                wait = float(retry_after) if retry_after and retry_after.isdigit() else min(30, 2 ** attempt)
                time.sleep(wait)
            except (urllib.error.URLError, TimeoutError):
                if attempt == self.retries - 1:
                    raise
                time.sleep(min(30, 2 ** attempt))
        raise RuntimeError("API retry loop ended unexpectedly")


def class_values() -> str:
    return " ".join(f"wd:{item}" for item in ROOT_CLASSES)


def fetch_wikidata_items(client: ApiClient, start_iso: str, end_iso: str) -> tuple[dict, int]:
    count_query = f"""
SELECT (COUNT(DISTINCT ?item) AS ?count) WHERE {{
  VALUES ?class {{ {class_values()} }}
  ?item wdt:P31 ?class.
  ?item (wdt:P571|wdt:P577) ?cohortDate.
  FILTER(?cohortDate >= "{start_iso}"^^xsd:dateTime && ?cohortDate < "{end_iso}"^^xsd:dateTime)
  ?enArticle schema:about ?item; schema:isPartOf <https://en.wikipedia.org/>.
}}
""".strip()
    rows_query = f"""
SELECT DISTINCT ?item ?class ?cohortProperty ?cohortDate ?enArticle ?jaArticle WHERE {{
  VALUES ?class {{ {class_values()} }}
  ?item wdt:P31 ?class.
  {{ ?item wdt:P571 ?cohortDate. BIND("P571" AS ?cohortProperty) }}
  UNION
  {{ ?item wdt:P577 ?cohortDate. BIND("P577" AS ?cohortProperty) }}
  FILTER(?cohortDate >= "{start_iso}"^^xsd:dateTime && ?cohortDate < "{end_iso}"^^xsd:dateTime)
  ?enArticle schema:about ?item; schema:isPartOf <https://en.wikipedia.org/>.
  OPTIONAL {{ ?jaArticle schema:about ?item; schema:isPartOf <https://ja.wikipedia.org/>. }}
}}
ORDER BY ?item ?class ?cohortProperty ?cohortDate
""".strip()
    count_payload = client.json(SPARQL_ENDPOINT, {"query": count_query, "format": "json"})
    expected = int(count_payload["results"]["bindings"][0]["count"]["value"])
    rows_payload = client.json(SPARQL_ENDPOINT, {"query": rows_query, "format": "json"})
    items: dict[str, dict] = {}
    for binding in rows_payload["results"]["bindings"]:
        item_id = binding["item"]["value"].rsplit("/", 1)[-1]
        class_id = binding["class"]["value"].rsplit("/", 1)[-1]
        en_title = urllib.parse.unquote(binding["enArticle"]["value"].split("/wiki/", 1)[-1]).replace("_", " ")
        ja_value = binding.get("jaArticle", {}).get("value")
        ja_title = urllib.parse.unquote(ja_value.split("/wiki/", 1)[-1]).replace("_", " ") if ja_value else None
        item = items.setdefault(item_id, {
            "itemId": item_id,
            "classIds": set(),
            "cohortDates": set(),
            "enTitle": en_title,
            "jaTitle": ja_title,
        })
        if item["enTitle"] != en_title or item["jaTitle"] != ja_title:
            raise RuntimeError(f"Conflicting sitelinks for {item_id}")
        item["classIds"].add(class_id)
        item["cohortDates"].add((binding["cohortProperty"]["value"], binding["cohortDate"]["value"]))
    if len(items) != expected:
        raise RuntimeError(f"SPARQL completeness check failed: expected {expected}, received {len(items)}")
    return items, expected


def batches(values: list[str], size: int = 1):
    for index in range(0, len(values), size):
        yield values[index:index + size]


def resolve_title(title: str, aliases: dict[str, str]) -> str:
    current = title
    visited = set()
    while current in aliases and current not in visited:
        visited.add(current)
        current = aliases[current]
    return current


def fetch_first_revisions(client: ApiClient, endpoint: str, titles: list[str]) -> dict[str, dict]:
    output: dict[str, dict] = {}
    for title_batch in batches(titles):
        payload = client.json(endpoint, {
            "action": "query",
            "format": "json",
            "formatversion": "2",
            "prop": "revisions",
            "titles": "|".join(title_batch),
            "redirects": "1",
            "rvdir": "newer",
            "rvlimit": "1",
            "rvprop": "timestamp|tags",
        })
        query = payload.get("query", {})
        aliases: dict[str, str] = {}
        redirected_from: set[str] = set()
        for row in query.get("normalized", []):
            aliases[row["from"]] = row["to"]
        for row in query.get("redirects", []):
            aliases[row["from"]] = row["to"]
            redirected_from.add(row["from"])
        pages = {page.get("title"): page for page in query.get("pages", [])}
        for requested in title_batch:
            resolved = resolve_title(requested, aliases)
            page = pages.get(resolved, {})
            revision = (page.get("revisions") or [{}])[0]
            output[requested] = {
                "requestedTitle": requested,
                "resolvedTitle": page.get("title") or resolved,
                "redirected": any(name in redirected_from for name in {requested, aliases.get(requested)} if name),
                "pageId": page.get("pageid"),
                "missing": bool(page.get("missing", False) or not revision.get("timestamp")),
                "firstRevision": revision.get("timestamp"),
                "firstRevisionTags": revision.get("tags") or [],
            }
    return output


def write_dataset(file: pathlib.Path, rows: list[dict]) -> None:
    with deterministic_gzip_text(file) as stream:
        for row in rows:
            stream.write(json.dumps(row, ensure_ascii=False, sort_keys=True, separators=(",", ":")) + "\n")


def main() -> None:
    parser = argparse.ArgumentParser()
    parser.add_argument("--output", default="evidence/data/wikipedia-en-ja-technical-2016-2023.jsonl.gz")
    parser.add_argument("--raw-cache", default="evidence/data/raw-api-responses.jsonl.gz")
    parser.add_argument("--metadata", default="evidence/fetch-metadata.json")
    parser.add_argument("--start", default="2016-01-01T00:00:00Z")
    parser.add_argument("--end", default="2024-01-01T00:00:00Z")
    parser.add_argument("--cutoff", default="2026-07-15T00:00:00Z")
    parser.add_argument("--delay", type=float, default=0.08)
    parser.add_argument("--minimum-rows", type=int, default=150)
    args = parser.parse_args()

    output = pathlib.Path(args.output)
    raw_cache = pathlib.Path(args.raw_cache)
    metadata_file = pathlib.Path(args.metadata)
    start = iso_utc(args.start)
    end = iso_utc(args.end)
    cutoff = iso_utc(args.cutoff)
    if not start < end <= cutoff:
        raise SystemExit("Expected start < end <= cutoff")

    executed_at = dt.datetime.now(dt.timezone.utc).isoformat()
    with deterministic_gzip_text(raw_cache) as raw_writer:
        client = ApiClient(raw_writer, delay=args.delay)
        items, candidate_count = fetch_wikidata_items(client, args.start, args.end)
        en_titles = sorted({item["enTitle"] for item in items.values()})
        en_revisions = fetch_first_revisions(client, EN_API, en_titles)

        missing_en_revision = 0
        for item in items.values():
            en_revision = en_revisions.get(item["enTitle"], {})
            timestamp = en_revision.get("firstRevision")
            if not timestamp:
                missing_en_revision += 1
        if missing_en_revision:
            raise RuntimeError(f"Completeness gate failed: {missing_en_revision} English pages lack oldest revisions")
        eligible = list(items.values())

        if len(eligible) < args.minimum_rows:
            raise RuntimeError(f"Pre-execution gate failed: only {len(eligible)} eligible items")

        ja_titles = sorted({item["jaTitle"] for item in eligible if item.get("jaTitle")})
        ja_revisions = fetch_first_revisions(client, JA_API, ja_titles)
        rows = []
        for item in sorted(eligible, key=lambda row: row["itemId"]):
            en_revision = en_revisions[item["enTitle"]]
            ja_revision = ja_revisions.get(item.get("jaTitle"), {}) if item.get("jaTitle") else {}
            rows.append({
                "itemId": item["itemId"],
                "classIds": sorted(item["classIds"]),
                "cohortDates": [
                    {"property": prop, "timestamp": timestamp}
                    for prop, timestamp in sorted(item["cohortDates"])
                ],
                "enTitle": item["enTitle"],
                "enResolvedTitle": en_revision.get("resolvedTitle"),
                "enRedirected": bool(en_revision.get("redirected")),
                "enPageId": en_revision.get("pageId"),
                "enFirstRevision": en_revision.get("firstRevision"),
                "enFirstRevisionTags": en_revision.get("firstRevisionTags") or [],
                "jaSitelinkPresent": bool(item.get("jaTitle")),
                "jaTitle": item.get("jaTitle"),
                "jaResolvedTitle": ja_revision.get("resolvedTitle"),
                "jaRedirected": bool(ja_revision.get("redirected")),
                "jaPageId": ja_revision.get("pageId"),
                "jaFirstRevision": ja_revision.get("firstRevision"),
                "jaFirstRevisionTags": ja_revision.get("firstRevisionTags") or [],
            })

        request_count = client.request_count

    write_dataset(output, rows)
    created_ja = sum(bool(row["jaFirstRevision"]) for row in rows)
    metadata = {
        "schemaVersion": 2,
        "executedAt": executed_at,
        "source": "https://query.wikidata.org/sparql",
        "revisionApis": [EN_API, JA_API],
        "userAgent": USER_AGENT,
        "rootClasses": ROOT_CLASSES,
        "cohortRule": "Direct P31 membership in one of five frozen software/computing classes, an English Wikipedia sitelink, and at least one Wikidata inception (P571) or publication date (P577) in [2016-01-01, 2024-01-01).",
        "periodStart": args.start,
        "periodEndExclusive": args.end,
        "snapshotCutoff": args.cutoff,
        "candidateItemCount": candidate_count,
        "eligibleRowCount": len(rows),
        "rowCount": len(rows),
        "jaFirstRevisionPresentCount": created_ja,
        "jaFirstRevisionAbsentCount": len(rows) - created_ja,
        "missingEnglishRevisionCount": missing_en_revision,
        "requestCount": request_count,
        "completenessGuards": [
            "Count query must equal distinct item count from the ordered detail query.",
            f"Eligible cohort must contain at least {args.minimum_rows} rows.",
            "Every cohort item must have an oldest English revision timestamp.",
            "Items without a Japanese sitelink or revision remain in the dataset for right-censoring.",
        ],
        "datasetPath": str(output),
        "datasetSha256": sha256_file(output),
        "rawCachePath": str(raw_cache),
        "rawCacheSha256": sha256_file(raw_cache),
    }
    metadata["sha256"] = metadata["datasetSha256"]
    metadata_file.parent.mkdir(parents=True, exist_ok=True)
    metadata_file.write_text(json.dumps(metadata, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
    print(json.dumps({
        "eligibleRows": len(rows),
        "jaPresent": created_ja,
        "jaAbsent": len(rows) - created_ja,
        "requests": request_count,
        "datasetSha256": metadata["datasetSha256"],
    }, ensure_ascii=False))


if __name__ == "__main__":
    main()
