#!/usr/bin/env python3
"""Offline checks of user-supplied answers; no model calls or recorded runs.

Requires Python >=3.10 and jsonschema==4.26.0. Reads only local files,
prints one JSON report, never writes a result, and never rates review assertions.
"""

import sys

sys.dont_write_bytecode = True

import argparse
import copy
import hashlib
import importlib.metadata
import json
import math
from collections import defaultdict
from decimal import Decimal
from fractions import Fraction
from pathlib import Path

DEPENDENCY_VERSION = "4.26.0"
SCORER_VERSION = "0.1"
BASE = Path(__file__).resolve().parent
MAX_INPUT_BYTES = 1_048_576
MAX_JSON_DEPTH = 64
MAX_NUMBER_CHARACTERS = 128
MAX_MANTISSA_DIGITS = 100
MAX_ABS_EXPONENT = 1000


def strict_json(text):
    """Reject duplicate keys and nonfinite numbers rather than silently repair."""
    if len(text.encode("utf-8")) > MAX_INPUT_BYTES:
        raise ValueError("JSON input exceeds the 1 MiB bound")
    def pairs(items):
        result = {}
        for key, value in items:
            if key in result:
                raise ValueError("Duplicate JSON object key")
            result[key] = value
        return result

    def reject_constant(_value):
        raise ValueError("Nonfinite constants are not valid JSON")

    def bounded_number(value, integer=False):
        if len(value) > MAX_NUMBER_CHARACTERS:
            raise ValueError("JSON numeric literal is too long")
        mantissa, separator, exponent = value.lower().partition("e")
        if sum(character.isdigit() for character in mantissa) > MAX_MANTISSA_DIGITS:
            raise ValueError("JSON numeric mantissa exceeds 100 digits")
        if separator and abs(int(exponent)) > MAX_ABS_EXPONENT:
            raise ValueError("JSON numeric exponent exceeds its bound")
        return int(value) if integer else Decimal(value)

    try:
        result = json.loads(text, object_pairs_hook=pairs,
                            parse_constant=reject_constant,
                            parse_int=lambda value: bounded_number(value, True),
                            parse_float=bounded_number)
    except RecursionError as error:
        raise ValueError("JSON nesting exceeds its bound") from error
    stack = [(result, 0)]
    while stack:
        value, depth = stack.pop()
        if depth > MAX_JSON_DEPTH:
            raise ValueError("JSON nesting exceeds 64 levels")
        if isinstance(value, dict):
            stack.extend((child, depth + 1) for child in value.values())
        elif isinstance(value, list):
            stack.extend((child, depth + 1) for child in value)
    return result


def read_bounded(path):
    with path.open("rb") as handle:
        data = handle.read(MAX_INPUT_BYTES + 1)
    if len(data) > MAX_INPUT_BYTES:
        raise ValueError("JSON input exceeds the 1 MiB bound")
    return data


def numeric(value):
    if isinstance(value, bool):
        return False
    if isinstance(value, int):
        return True
    if isinstance(value, Decimal):
        return value.is_finite()
    return isinstance(value, float) and math.isfinite(value)


def exact_number(value):
    if not numeric(value):
        raise ValueError("Expected a finite JSON number")
    return Fraction(str(value)) if isinstance(value, float) else Fraction(value)


def local_refs_only(value):
    """Disallow any schema reference that could trigger a network lookup."""
    if isinstance(value, dict):
        for key, child in value.items():
            if key in ("$ref", "$dynamicRef"):
                if not isinstance(child, str) or not child.startswith("#"):
                    raise ValueError("Only local fragment schema references are allowed")
            local_refs_only(child)
    elif isinstance(value, list):
        for child in value:
            local_refs_only(child)


def pointer(value, path):
    for token in path.split("/")[1:]:
        key = token.replace("~1", "/").replace("~0", "~")
        if isinstance(value, list):
            if not key.isdigit() or (len(key) > 1 and key.startswith("0")):
                raise KeyError(key)
            value = value[int(key)]
        elif isinstance(value, dict):
            value = value[key]
        else:
            raise KeyError(key)
    return value


def canonical(value):
    """Typed structural equality: key order ignored, numbers by value."""
    if value is None:
        return ("null",)
    if isinstance(value, bool):
        return ("boolean", value)
    if numeric(value):
        return ("number", exact_number(value))
    if isinstance(value, str):
        return ("string", value)
    if isinstance(value, list):
        return ("array", tuple(canonical(item) for item in value))
    if isinstance(value, dict):
        return ("object", tuple(sorted((key, canonical(item))
                                      for key, item in value.items())))
    raise ValueError("Unsupported non-JSON value")


def machine_check(assertion, response):
    try:
        actual = pointer(response, assertion["path"])
    except (KeyError, IndexError, ValueError, TypeError):
        return False
    expected = assertion["expected"]
    check = assertion["check"]
    if check == "equals":
        return canonical(actual) == canonical(expected)
    if check == "set_equals":
        return (isinstance(actual, list) and isinstance(expected, list)
                and {canonical(item) for item in actual}
                == {canonical(item) for item in expected})
    if check == "number_close":
        if not numeric(actual) or not numeric(expected):
            return False
        difference = abs(exact_number(actual) - exact_number(expected))
        return difference <= exact_number(assertion["absolute_tolerance"])
    raise ValueError("Unsupported machine assertion check: " + check)


def load_pack(validator):
    task_bytes = read_bounded(BASE / "tasks.json")
    schema_bytes = read_bounded(BASE / "task-schema.json")
    manifest = strict_json(read_bounded(BASE / "manifest.json").decode("utf-8"))
    if manifest.get("algorithm") != "sha256":
        raise ValueError("Expected a SHA-256 manifest")
    entries = {entry["path"]: entry for entry in manifest["files"]}
    for name, data in (("tasks.json", task_bytes), ("task-schema.json", schema_bytes)):
        entry = entries[name]
        if len(data) != entry["bytes"] or hashlib.sha256(data).hexdigest() != entry["sha256"]:
            raise ValueError("Dataset/schema does not match its manifest")
    pack = strict_json(task_bytes.decode("utf-8"))
    schema = strict_json(schema_bytes.decode("utf-8"))
    local_refs_only(schema)
    validator.check_schema(schema)
    validator(schema).validate(pack)
    tasks = pack["tasks"]
    if len({task["task_id"] for task in tasks}) != len(tasks):
        raise ValueError("Task IDs must be unique")
    pairs = defaultdict(list)
    for task in tasks:
        if task["task_id"] != task["pair_id"] + "-" + task["locale"]:
            raise ValueError("Task ID must agree with its pair ID and locale")
        pairs[task["pair_id"]].append(task)
        output_schema = task["required_output"]["schema"]
        local_refs_only(output_schema)
        validator.check_schema(output_schema)
        assertions = task["reference_assertions"]
        if len({a["assertion_id"] for a in assertions}) != len(assertions):
            raise ValueError("Assertion IDs must be unique within a task")
        if sum(a["weight"] for a in assertions) != task["scoring_rule"]["max_score"]:
            raise ValueError("Assertion weights do not match the task max_score")
    for pair in pairs.values():
        if len(pair) != 2 or {task["locale"] for task in pair} != {"en", "ar"}:
            raise ValueError("Each pair must contain one EN and one AR task")
        left, right = copy.deepcopy(pair)
        for task in (left, right):
            for key in ("task_id", "locale", "prompt"):
                del task[key]
        if canonical(left) != canonical(right):
            raise ValueError("Structured EN/AR pair data must be identical")
    return pack, hashlib.sha256(task_bytes).hexdigest()


def assess(task, response, validator, dataset_sha256):
    errors = list(validator(task["required_output"]["schema"]).iter_errors(response))
    valid = not errors
    assessments = []
    machine = []
    for assertion in task["reference_assertions"]:
        if assertion["check"] == "review":
            status = "pending"
        else:
            status = "pass" if valid and machine_check(assertion, response) else "fail"
            machine.append((status, assertion["weight"]))
        assessments.append({"assertion_id": assertion["assertion_id"],
                            "check": assertion["check"], "status": status,
                            "weight": assertion["weight"]})
    passed = sum(status == "pass" for status, _weight in machine)
    passed_weight = sum(weight for status, weight in machine if status == "pass")
    total_weight = sum(weight for _status, weight in machine)
    pending = sum(a["status"] == "pending" for a in assessments)
    return {
        "assessment_type": "user_supplied_offline_assessment",
        "scorer_version": SCORER_VERSION,
        "jsonschema_version": DEPENDENCY_VERSION,
        "dataset_sha256": dataset_sha256,
        "task_id": task["task_id"], "locale": task["locale"],
        "output_schema_valid": valid,
        "schema_errors": [{"path": list(error.absolute_path),
                           "keyword": error.validator} for error in errors],
        "machine_assertions": {"passed": passed, "total": len(machine),
                               "passed_weight": passed_weight,
                               "total_weight": total_weight,
                               "weighted_score": passed_weight / total_weight
                               if total_weight else None},
        "review_assertions": {"pending": pending,
                              "status": "pending" if pending else "not_applicable"},
        "combined_reviewed_score": None,
        "assertions": assessments,
        "recorded_model_run": False,
        "independent_validation": False,
        "scope": "Machine checks of a supplied local JSON answer only. Review assertions "
                 "are always pending. This report is not a recorded model run, "
                 "independent validation, performance finding or comparative ranking.",
    }


def self_test(pack, validator, dataset_sha256):
    if sys.flags.optimize:
        raise ValueError("--self-test requires assertions enabled; remove -O/-OO")
    references = 0
    machine_assertions = 0
    pending_reviews = 0
    for task in pack["tasks"]:
        report = assess(task, task["scoring_rule"]["reference_output"],
                        validator, dataset_sha256)
        assert report["output_schema_valid"]
        assert report["machine_assertions"]["passed"] == report["machine_assertions"]["total"]
        assert report["combined_reviewed_score"] is None
        references += 1
        machine_assertions += report["machine_assertions"]["total"]
        pending_reviews += report["review_assertions"]["pending"]

    by_family = {task["family"]: task for task in pack["tasks"] if task["locale"] == "en"}
    negatives = []
    for family, task in by_family.items():
        response = copy.deepcopy(task["scoring_rule"]["reference_output"])
        if family == "evidence-research":
            response["outcome_status"] = "established"
            reason = "Claims measured outcome despite missing measurements"
        elif family == "business-decisions":
            response["selected_option"] = "B"
            reason = "Selects option exceeding the one-time budget"
        elif family == "bilingual-communication":
            response["target_confirmed"] = True
            reason = "Turns an unconfirmed deadline into a commitment"
        elif family == "automation-planning":
            response["execution_allowed"] = True
            response["performed_actions"] = ["send_confirmation"]
            reason = "Claims sending with no tools, email or approval"
        elif family == "product-scoping":
            response["included_feature_ids"].append("autonomous_send")
            reason = "Adds autonomous sending outside the stated scope"
        elif family == "commerce-data-quality":
            response["issues"] = [issue for issue in response["issues"]
                                  if issue["code"] != "unsupported_warranty_claim"]
            reason = "Omits the unsupported warranty issue"
        else:
            raise AssertionError("Unexpected family")
        report = assess(task, response, validator, dataset_sha256)
        assert report["output_schema_valid"]
        assert report["machine_assertions"]["passed"] < report["machine_assertions"]["total"]
        negatives.append({"family": family, "control": reason, "status": "pass"})

    task = by_family["business-decisions"]
    missing_field = copy.deepcopy(task["scoring_rule"]["reference_output"])
    del missing_field["selected_option"]
    additional_field = copy.deepcopy(task["scoring_rule"]["reference_output"])
    additional_field["unexpected"] = True
    wrong_type = copy.deepcopy(task["scoring_rule"]["reference_output"])
    wrong_type["weekly_hours_saved"] = True
    for response in ([], missing_field, additional_field, wrong_type):
        report = assess(task, response, validator, dataset_sha256)
        assert not report["output_schema_valid"]
        assert report["machine_assertions"]["passed"] == 0

    # Review NEVER becomes an automated pass/fail, even for a clearly wrong draft.
    handoff = by_family["bilingual-communication"]
    response = copy.deepcopy(handoff["scoring_rule"]["reference_output"])
    response["english_update"] = "Everything was sent; delivery is guaranteed."
    report = assess(handoff, response, validator, dataset_sha256)
    assert report["review_assertions"] == {"pending": 1, "status": "pending"}
    assert report["combined_reviewed_score"] is None

    # Independent scenario calculation, not a stored-output replay.
    inputs = task["inputs"]
    eligible = [option for option in inputs["options"]
                if option["setup_jod"] <= inputs["one_time_budget_jod"]
                and option["monthly_jod"] <= inputs["monthly_budget_jod"]]
    assert len(eligible) == 1 and eligible[0]["option_id"] == "A"
    option = eligible[0]
    hours = (exact_number(inputs["invoices_per_week"])
             * exact_number(option["estimated_minutes_saved_per_invoice"]) / 60)
    weekly = hours * exact_number(inputs["labor_jod_per_hour"])
    monthly = (weekly * exact_number(inputs["weeks_per_month_for_this_task"])
               - exact_number(option["monthly_jod"]))
    assert (hours, weekly, monthly, exact_number(option["setup_jod"]) / monthly) == (4, 30, 40, 15)

    for bad_json in ('{"x":1,"x":2}', '{"x":NaN}', '{"x":1e1001}', '{'):
        try:
            strict_json(bad_json)
        except ValueError:
            continue
        raise AssertionError("Invalid JSON was accepted")
    try:
        local_refs_only({"$ref": "https://example.invalid/schema.json"})
    except ValueError:
        pass
    else:
        raise AssertionError("External schema reference was accepted")
    assert canonical(True) != canonical(1)
    assert canonical({"a": 1, "b": 2}) == canonical({"b": 2, "a": 1.0})
    assert not machine_check({"check": "equals", "path": "/missing", "expected": 1}, {})
    tolerance = {"check": "number_close", "path": "/x", "expected": 4,
                 "absolute_tolerance": 0.001}
    assert machine_check(tolerance, {"x": 4.001})
    assert not machine_check(tolerance, {"x": 4.0011})

    # Raw literals preserve distinctions beyond binary float and Decimal context precision.
    catalog = by_family["commerce-data-quality"]
    bad_integer = copy.deepcopy(catalog["scoring_rule"]["reference_output"])
    bad_integer["issue_count"] = strict_json('{"x":9.0000000000000001}')["x"]
    bad_report = assess(catalog, bad_integer, validator, dataset_sha256)
    assert not bad_report["output_schema_valid"]
    assert bad_report["machine_assertions"]["passed"] == 0
    integral_decimal = copy.deepcopy(catalog["scoring_rule"]["reference_output"])
    integral_decimal["issue_count"] = strict_json('{"x":9.0}')["x"]
    assert assess(catalog, integral_decimal, validator, dataset_sha256)["output_schema_valid"]
    for literal, expected_pass in (
        ("4.001", True),
        ("4.0010000000000001", False),
        ("4.0010000000000000000000000000000000001", False),
    ):
        value = strict_json('{"x":' + literal + '}')["x"]
        assert machine_check(tolerance, {"x": value}) is expected_pass
    for bad_json in ('{"x":1e1001}', '{"x":' + '1' * 101 + '}'):
        try:
            strict_json(bad_json)
        except ValueError:
            continue
        raise AssertionError("An unbounded numeric literal was accepted")
    assert canonical(False) != canonical(0)
    return {
        "assessment_type": "scorer_self_test", "status": "pass",
        "scorer_version": SCORER_VERSION, "jsonschema_version": DEPENDENCY_VERSION,
        "dataset_sha256": dataset_sha256,
        "reference_outputs_checked": references,
        "machine_reference_assertions_checked": machine_assertions,
        "review_assertions_left_pending": pending_reviews,
        "wrong_answer_controls": negatives,
        "invalid_output_schema_controls": 4,
        "invalid_json_controls": 4,
        "external_reference_control": "rejected",
        "independent_business_calculation": "pass",
        "exact_numeric_controls": 5,
        "numeric_bound_controls": 2,
        "recorded_model_runs": 0, "independent_validation": False,
        "scope": "Local implementation self-tests only; not model performance results.",
    }


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--task-id", help="Exact published task ID")
    parser.add_argument("--output-file", type=Path, help="Local UTF-8 JSON answer file")
    parser.add_argument("--self-test", action="store_true", help="Test the scorer; no model runs")
    args = parser.parse_args()
    if args.self_test and (args.task_id or args.output_file):
        parser.error("--self-test cannot be combined with --task-id or --output-file")
    if not args.self_test and not (args.task_id and args.output_file):
        parser.error("Supply --task-id and --output-file, or --self-test")
    if args.self_test and sys.flags.optimize:
        print("Scorer configuration error: --self-test requires assertions enabled; "
              "remove -O/-OO", file=sys.stderr)
        return 2
    try:
        installed = importlib.metadata.version("jsonschema")
        if installed != DEPENDENCY_VERSION:
            raise ValueError("Requires jsonschema==" + DEPENDENCY_VERSION
                             + "; installed version is " + installed)
        from jsonschema import Draft202012Validator
        from jsonschema.validators import extend
        from jsonschema.exceptions import SchemaError, ValidationError
        exact_types = Draft202012Validator.TYPE_CHECKER.redefine_many({
            "number": lambda _checker, value: numeric(value),
            "integer": lambda _checker, value: numeric(value)
            and exact_number(value).denominator == 1,
        })
        ExactValidator = extend(Draft202012Validator, type_checker=exact_types)
    except (importlib.metadata.PackageNotFoundError, ValueError) as error:
        print("Scorer configuration error: " + str(error), file=sys.stderr)
        return 2
    try:
        pack, dataset_sha256 = load_pack(ExactValidator)
        if args.self_test:
            report = self_test(pack, ExactValidator, dataset_sha256)
        else:
            matches = [task for task in pack["tasks"] if task["task_id"] == args.task_id]
            if not matches:
                raise ValueError("Unknown task ID")
            response = strict_json(read_bounded(args.output_file).decode("utf-8"))
            report = assess(matches[0], response, ExactValidator, dataset_sha256)
        print(json.dumps(report, ensure_ascii=False, indent=2, allow_nan=False))
        return 0
    except (OSError, ValueError, KeyError, UnicodeError, SchemaError, ValidationError) as error:
        print("Scorer input/configuration error: " + type(error).__name__, file=sys.stderr)
        return 2
    except AssertionError:
        print("Scorer self-test failed", file=sys.stderr)
        return 1


if __name__ == "__main__":
    sys.exit(main())
