#!/usr/bin/env python3
"""Measure structured-field availability in two pinned official XML snapshots.

Python 3.9+; standard library only. No requests to a screening service. No names,
document numbers, dates of birth, addresses, or per-record rows are emitted.
"""
import argparse
import collections
import csv
import datetime as dt
import hashlib
import json
import pathlib
import re
import xml.etree.ElementTree as ET

VERSION = "1.0.0"
PLACEHOLDERS = {"", "na", "n/a", "unknown", "none", "-", "not known", "not available", "nil"}
MONTHS = {m: i for i, m in enumerate("Jan Feb Mar Apr May Jun Jul Aug Sep Oct Nov Dec".split(), 1)}


def text(value):
    return " ".join((value or "").split())


def present(value):
    return text(value).casefold() not in PLACEHOLDERS


def values(element, path):
    return [text(e.text) for e in element.findall(path) if present(e.text)]


def is_full_date(value, source):
    """Only an entire, valid day/month/year value; never ranges or circa dates."""
    value = text(value)
    if source == "ofac-sdn":
        match = re.fullmatch(r"(\d{1,2}) ([A-Z][a-z]{2}) (\d{4})", value)
        if not match or match[2] not in MONTHS:
            return False
        day, month, year = int(match[1]), MONTHS[match[2]], int(match[3])
    else:
        match = re.fullmatch(r"(\d{2})/(\d{2})/(\d{4})", value)
        if not match:
            return False
        day, month, year = map(int, match.groups())
    try:
        dt.date(year, month, day)
        return True
    except ValueError:
        return False


def parse_root(path):
    root = ET.parse(path).getroot()
    for element in root.iter():
        element.tag = element.tag.split("}")[-1]
    return root


def ofac_flags(record):
    aliases = [a for a in record.findall("akaList/aka") if
               any(present(a.findtext(k)) for k in ("firstName", "lastName"))]
    flags = {
        "alias_present": bool(aliases),
        "strong_alias_present": any(text(a.findtext("category")).casefold() == "strong" for a in aliases),
        "weak_alias_present": any(text(a.findtext("category")).casefold() == "weak" for a in aliases),
        "address_country_present": bool(values(record, "addressList/address/country")),
    }
    if record.findtext("sdnType") == "Individual":
        dobs = values(record, "dateOfBirthList/dateOfBirthItem/dateOfBirth")
        # An idList also holds Gender, Website and sanctions-risk text. It is not
        # an identifier-presence field. Select explicitly passport-typed entries.
        passport = any("passport" in text(e.findtext("idType")).casefold() and
                       present(e.findtext("idNumber")) for e in record.findall("idList/id"))
        national_id = any(text(e.findtext("idType")) == "National ID No." and
                          present(e.findtext("idNumber")) for e in record.findall("idList/id"))
        nationality = bool(values(record, "nationalityList/nationality/country"))
        flags.update(individual_flags(dobs, passport, national_id, nationality, "ofac-sdn"))
    return flags


def uk_flags(record):
    name_records = [n for n in record.findall("Names/Name") if
                    any(present(n.findtext(f"Name{i}")) for i in range(1, 7))]
    aliases = [n for n in name_records if text(n.findtext("NameType")).casefold() == "alias"]
    variations = [n for n in name_records if text(n.findtext("NameType")).casefold() == "primary name variation"]
    flags = {
        "alias_present": bool(aliases),
        "primary_name_variation_present": bool(variations),
        "alias_or_primary_variation_present": bool(aliases or variations),
        "good_quality_alias_present": any(text(n.findtext("AliasStrength")).casefold() == "good quality a.k.a" for n in aliases),
        "low_quality_alias_present": any(text(n.findtext("AliasStrength")).casefold() == "low quality a.k.a" for n in aliases),
        "address_country_present": bool(values(record, "Addresses/Address/AddressCountry")),
        "non_latin_name_present": bool(values(record, "NonLatinNames/NonLatinName/NameNonLatinScript")),
    }
    if record.findtext("IndividualEntityShip") == "Individual":
        prefix = "IndividualDetails/Individual/"
        dobs = values(record, prefix + "DOBs/DOB")
        passport = bool(values(record, prefix + "PassportDetails/Passport/PassportNumber"))
        national_id = bool(values(record, prefix + "NationalIdentifierDetails/NationalIdentifier/NationalIdentifierNumber"))
        nationality = bool(values(record, prefix + "Nationalities/Nationality"))
        flags.update(individual_flags(dobs, passport, national_id, nationality, "uk-sanctions"))
    if record.findtext("IndividualEntityShip") == "Entity":
        flags["business_registration_number_present"] = bool(values(record, "EntityDetails/Entity/BusinessRegistrationNumbers/BusinessRegistrationNumber"))
    if record.findtext("IndividualEntityShip") == "Ship":
        flags["imo_number_present"] = bool(values(record, "ShipDetails/Ship/IMONumbers/IMONumber"))
    return flags


def individual_flags(dobs, passport, national_id, nationality, source):
    full = any(is_full_date(d, source) for d in dobs)
    return {
        "dob_present": bool(dobs),
        "dob_at_least_one_full_calendar_date": full,
        "dob_present_no_full_calendar_date": bool(dobs) and not full,
        "dob_absent": not dobs,
        "multiple_dob_values": len(dobs) > 1,
        "passport_number_present": passport,
        "national_id_source_field_present": national_id,
        "nationality_present": nationality,
        "neither_dob_nor_passport_number": not dobs and not passport,
    }


def analyze(source, path):
    root = parse_root(path)
    if source == "ofac-sdn":
        assert root.tag == "sdnList", "Unexpected OFAC root"
        records = root.findall("sdnEntry")
        declared = int(root.findtext("publshInformation/Record_Count"))
        assert declared == len(records), "OFAC declared/parsed record count mismatch"
        source_date = root.findtext("publshInformation/Publish_Date")
        type_path, id_path, parser = "sdnType", "uid", ofac_flags
        expected_types = {"Individual", "Entity", "Vessel", "Aircraft"}
    else:
        assert root.tag == "Designations", "Unexpected UK root"
        records = root.findall("Designation")
        declared = None
        source_date = root.findtext("DateGenerated")
        type_path, id_path, parser = "IndividualEntityShip", "UniqueID", uk_flags
        expected_types = {"Individual", "Entity", "Ship"}
        name_types = {text(e.text).casefold() for e in root.findall(".//NameType")}
        assert name_types <= {"primary name", "primary name variation", "alias"}, "Unexpected UK name type"
    assert records, "Empty source"
    ids = [text(r.findtext(id_path)) for r in records]
    assert all(ids) and len(ids) == len(set(ids)), "Missing or duplicate source record IDs"
    types = collections.Counter(text(r.findtext(type_path)) for r in records)
    assert set(types) <= expected_types, "Unexpected source entity type"
    counts, denominators = collections.Counter(), collections.Counter()
    for record in records:
        entity_type = text(record.findtext(type_path))
        flags = parser(record)
        if entity_type == "Individual":
            assert sum(flags[k] for k in ("dob_at_least_one_full_calendar_date", "dob_present_no_full_calendar_date", "dob_absent")) == 1
            assert not flags["neither_dob_nor_passport_number"] or flags["dob_absent"]
        for metric, flag in flags.items():
            assert type(flag) is bool
            key = (entity_type, metric)
            counts[key] += int(flag)
            denominators[key] += 1
    rows = []
    for (entity_type, metric), denominator in sorted(denominators.items()):
        count = counts[(entity_type, metric)]
        assert 0 <= count <= denominator == types[entity_type]
        rows.append({"source": source, "entity_type": entity_type, "metric": metric,
                     "count": count, "denominator": denominator, "percent": round(100 * count / denominator, 2)})
    return {"source": source, "source_date_literal": source_date, "parsed_record_count": len(records),
            "publisher_declared_record_count": declared, "record_types": dict(sorted(types.items())),
            "all_ids_present_and_unique": True, "metrics": rows}


def run(ofac, uk, manifest_path, out):
    manifest = json.loads(pathlib.Path(manifest_path).read_text())
    results = []
    for source, path in (("ofac-sdn", ofac), ("uk-sanctions", uk)):
        actual = hashlib.sha256(pathlib.Path(path).read_bytes()).hexdigest()
        expected = next(x["sha256"] for x in manifest["snapshots"] if x["key"] == source)
        if actual != expected:
            raise ValueError(f"{source}: SHA-256 differs from pinned manifest; do not label new data as this study")
        results.append(analyze(source, path))
    output = {"study": "Structured identifiers and names in two official sanctions snapshots",
              "analysis_version": VERSION, "snapshot_capture_date": "2026-09-30",
              "unit": "source record, not unique real-world person or entity", "sources": results}
    out = pathlib.Path(out)
    out.mkdir(parents=True, exist_ok=True)
    (out / "results.json").write_text(json.dumps(output, indent=2, ensure_ascii=False) + "\n")
    with (out / "metrics.csv").open("w", newline="") as handle:
        writer = csv.DictWriter(handle, fieldnames=["source", "entity_type", "metric", "count", "denominator", "percent"])
        writer.writeheader()
        for source in results:
            writer.writerows(source["metrics"])
    print(json.dumps({r["source"]: {"record_count": r["parsed_record_count"], "types": r["record_types"]} for r in results}, indent=2))


if __name__ == "__main__":
    cli = argparse.ArgumentParser(description=__doc__)
    cli.add_argument("--ofac", required=True, help="Path to pinned OFAC SDN XML")
    cli.add_argument("--uk", required=True, help="Path to pinned UK Sanctions List XML")
    cli.add_argument("--manifest", default=str(pathlib.Path(__file__).with_name("source-manifest.json")))
    cli.add_argument("--out", required=True, help="Directory for aggregate CSV/JSON")
    args = cli.parse_args()
    run(args.ofac, args.uk, args.manifest, args.out)
