#!/usr/bin/env python3
"""Offline normalized offer comparison. Synthetic inputs; no provider parser."""

import argparse
import csv
from datetime import datetime
from decimal import Decimal, localcontext
import json
from pathlib import Path
import re
import sys


CONTEXT = ("asin", "marketplace", "currency", "condition", "customer_type",
           "membership", "location_context", "shipping_service")
SCOPE = {"marketplace": "US", "currency": "USD", "condition": "New", "customer_type": "Consumer"}
SOURCES = {"own_offer", "featured_offer", "lowest_price", "foep", "external_reference"}
OBSERVED_FIELDS = {"status", "source_type", "seller_id", *CONTEXT,
                   "observed_at", "item_price", "shipping_price"}
MONEY = re.compile(r"(?:0|[1-9][0-9]{0,11})(?:\.[0-9]{1,2})?")
COLUMNS = ("id", "description", "decision", "reasons", "own_seller_id", "featured_seller_id",
           "featured_is_own_seller", "observation_skew_seconds", "own_subtotal",
           "featured_subtotal", "own_minus_featured")


class InvalidInput(ValueError):
    pass


def _text(value, name):
    if (type(value) is not str or not 1 <= len(value) <= 256
            or value != value.strip() or any(ord(c) < 32 for c in value)):
        raise InvalidInput(f"{name}: expected a nonempty bounded string")
    try:
        value.encode("utf-8")
    except UnicodeError:
        raise InvalidInput(f"{name}: invalid Unicode") from None


def _time(value):
    if type(value) is not str or not re.fullmatch(r"[0-9]{4}-[0-9]{2}-[0-9]{2}T[0-9]{2}:[0-9]{2}:[0-9]{2}Z", value):
        raise InvalidInput("observed_at: expected YYYY-MM-DDTHH:MM:SSZ")
    try:
        return datetime.strptime(value, "%Y-%m-%dT%H:%M:%SZ")
    except ValueError:
        raise InvalidInput("observed_at: invalid calendar time") from None


def _money(value):
    if type(value) is not str or not MONEY.fullmatch(value):
        raise InvalidInput("money: expected nonnegative decimal string, <=12 integer and <=2 fractional digits")
    return Decimal(value)


def _observation(value):
    if value is None:
        return
    if type(value) is not dict:
        raise InvalidInput("observation must be an object or null")
    status = value.get("status")
    if status in ("absent", "failed"):
        if set(value) != {"status"}:
            raise InvalidInput("absent/failed observations carry only status")
        return
    if status != "observed" or set(value) != OBSERVED_FIELDS:
        raise InvalidInput("observed record has missing/unknown fields or invalid status")
    if type(value["source_type"]) is not str or value["source_type"] not in SOURCES:
        raise InvalidInput("unknown source_type")
    for field in ("seller_id", *CONTEXT):
        if value[field] is not None:
            _text(value[field], field)
    if value["observed_at"] is not None:
        _time(value["observed_at"])
    for field in ("item_price", "shipping_price"):
        if value[field] is not None:
            _money(value[field])


def compare(case, max_skew_seconds=300):
    """Return a gap only after the normalized comparison contract passes."""
    if type(max_skew_seconds) is not int or not 0 <= max_skew_seconds <= 86_400:
        raise InvalidInput("max_skew_seconds must be an integer from 0 to 86400")
    if type(case) is not dict or set(case) != {"id", "description", "expected_own_seller_id", "own", "featured"}:
        raise InvalidInput("case has missing or unknown fields")
    for field in ("id", "description", "expected_own_seller_id"):
        _text(case[field], field)
    own, featured = case["own"], case["featured"]
    _observation(own)
    _observation(featured)
    result = dict.fromkeys(COLUMNS)
    result.update(id=case["id"], description=case["description"], decision="hold", reasons=[])
    reasons = result["reasons"]
    for side, observation in (("own", own), ("featured", featured)):
        if observation is None:
            reasons.append(f"{side}_missing")
        elif observation["status"] != "observed":
            reasons.append(f"{side}_{observation['status']}")
    if reasons:
        return result
    result.update(own_seller_id=own["seller_id"], featured_seller_id=featured["seller_id"])
    for side, observation, expected_source in (
        ("own", own, "own_offer"), ("featured", featured, "featured_offer"),
    ):
        if observation["source_type"] != expected_source:
            reasons.append(f"{side}_wrong_source_type")
        for field in ("seller_id", *CONTEXT, "observed_at", "item_price", "shipping_price"):
            if observation[field] is None:
                reasons.append(f"{side}_missing_{field}")
    if own["seller_id"] is not None and own["seller_id"] != case["expected_own_seller_id"]:
        reasons.append("wrong_own_seller")
    for field in CONTEXT:
        if own[field] is not None and featured[field] is not None:
            if own[field] != featured[field]:
                reasons.append(f"mismatched_{field}")
            elif field in SCOPE and own[field] != SCOPE[field]:
                reasons.append(f"unsupported_{field}")
    if own["observed_at"] is not None and featured["observed_at"] is not None:
        skew = abs(int((_time(own["observed_at"]) - _time(featured["observed_at"])).total_seconds()))
        result["observation_skew_seconds"] = skew
        if skew > max_skew_seconds:
            reasons.append("observation_skew_exceeds_policy")
    if reasons:
        return result
    # Fixed precision makes the bounded USD arithmetic independent of caller context.
    with localcontext() as context:
        context.prec = 32
        own_subtotal = _money(own["item_price"]) + _money(own["shipping_price"])
        featured_subtotal = _money(featured["item_price"]) + _money(featured["shipping_price"])
        result.update(
            decision="compare",
            featured_is_own_seller=featured["seller_id"] == case["expected_own_seller_id"],
            own_subtotal=format(own_subtotal, ".2f"),
            featured_subtotal=format(featured_subtotal, ".2f"),
            own_minus_featured=format(own_subtotal - featured_subtotal, ".2f"),
        )
    return result


def _pairs(pairs):
    result = {}
    for key, value in pairs:
        if key in result:
            raise InvalidInput("duplicate JSON object member")
        result[key] = value
    return result


def load_cases(path):
    try:
        text = Path(path).read_text(encoding="utf-8")
        if len(text.encode("utf-8")) > 262_144:
            raise InvalidInput("cases file exceeds lab size bound")
        document = json.loads(text, object_pairs_hook=_pairs)
    except (UnicodeError, ValueError, RecursionError):
        raise InvalidInput("invalid cases JSON") from None
    if (type(document) is not dict or set(document) != {"schema_version", "cases"}
            or type(document["schema_version"]) is not int or document["schema_version"] != 1
            or type(document["cases"]) is not list or not 1 <= len(document["cases"]) <= 100):
        raise InvalidInput("expected schema_version 1 and 1-100 cases")
    return document["cases"]


def report(cases, max_skew_seconds=300):
    results = [compare(case, max_skew_seconds) for case in cases]
    if len({row["id"] for row in results}) != len(results):
        raise InvalidInput("case IDs must be unique")
    return {
        "mode": "synthetic_offline_model",
        "scope": "US/USD, Consumer, New, single unit",
        "subtotal_basis": "item_plus_known_shipping_before_promotions",
        "max_skew_seconds": max_skew_seconds,
        "results": results,
    }


def write_csv(results, stream):
    writer = csv.DictWriter(stream, fieldnames=COLUMNS, lineterminator="\n")
    writer.writeheader()
    for result in results:
        row = {k: "" if v is None else v for k, v in result.items()}
        row["reasons"] = ";".join(result["reasons"])
        identity = result["featured_is_own_seller"]
        row["featured_is_own_seller"] = "" if identity is None else str(identity).lower()
        writer.writerow(row)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--cases", type=Path, default=Path(__file__).with_name("cases.json"))
    parser.add_argument("--csv", action="store_true", help="emit the same model results as CSV")
    parser.add_argument("--max-skew-seconds", type=int, default=300,
                        help="local observation-time policy, 0-86400 seconds; default 300")
    args = parser.parse_args()
    try:
        output = report(load_cases(args.cases), args.max_skew_seconds)
    except (InvalidInput, OSError) as error:
        parser.error(str(error))
    if args.csv:
        write_csv(output["results"], sys.stdout)
    else:
        print(json.dumps(output, indent=2, sort_keys=True))


if __name__ == "__main__":
    main()
