from copy import deepcopy
import csv
from decimal import localcontext
import io
import json
from pathlib import Path
import subprocess
import sys
from tempfile import TemporaryDirectory
import unittest

from compare import InvalidInput, compare, load_cases, report


HERE = Path(__file__).resolve().parent
CASES = load_cases(HERE / "cases.json")
BY_ID = {case["id"]: case for case in CASES}


class ComparisonTests(unittest.TestCase):
    def base(self):
        return deepcopy(BY_ID["equal_subtotal_other_seller"])

    def assert_hold(self, result, reason):
        self.assertEqual(result["decision"], "hold")
        self.assertIn(reason, result["reasons"])
        for field in ("own_subtotal", "featured_subtotal", "own_minus_featured", "featured_is_own_seller"):
            self.assertIsNone(result[field])

    def test_lower_item_price_can_have_higher_item_plus_shipping_subtotal(self):
        result = compare(BY_ID["lower_item_higher_subtotal"])
        # Independently specified: (19 + 4) - (21 + 0) = 2 USD.
        self.assertEqual(result["decision"], "compare")
        self.assertEqual((result["own_subtotal"], result["featured_subtotal"], result["own_minus_featured"]),
                         ("23.00", "21.00", "2.00"))
        self.assertFalse(result["featured_is_own_seller"])

    def test_equal_subtotal_does_not_establish_seller_identity(self):
        different = compare(BY_ID["equal_subtotal_other_seller"])
        same = compare(BY_ID["same_seller"])
        self.assertEqual((different["own_minus_featured"], same["own_minus_featured"]), ("0.00", "0.00"))
        self.assertFalse(different["featured_is_own_seller"])
        self.assertTrue(same["featured_is_own_seller"])
        self.assertEqual(different["featured_seller_id"], "seller-B")
        self.assertEqual(same["featured_seller_id"], "seller-A")

    def test_exact_decimal_arithmetic_ignores_low_caller_precision(self):
        case = self.base()
        case["own"].update(item_price="0.10", shipping_price="0.20")
        case["featured"].update(item_price="0.29", shipping_price="0.00")
        with localcontext() as context:
            context.prec = 2
            result = compare(case)
        self.assertEqual((result["own_subtotal"], result["own_minus_featured"]), ("0.30", "0.01"))
        case["own"].update(item_price="999999999999.99", shipping_price="999999999999.99")
        case["featured"].update(item_price="0", shipping_price="0")
        self.assertEqual(compare(case)["own_minus_featured"], "1999999999999.98")
        case["own"], case["featured"] = case["featured"], case["own"]
        case["own"].update(source_type="own_offer", seller_id="seller-A")
        case["featured"].update(source_type="featured_offer", seller_id="seller-B")
        self.assertEqual(compare(case)["own_minus_featured"], "-1999999999999.98")

    def test_known_zero_shipping_is_accepted_but_missing_shipping_holds(self):
        case = self.base()
        self.assertEqual(compare(case)["decision"], "compare")
        for side in ("own", "featured"):
            changed = deepcopy(case)
            changed[side]["shipping_price"] = None
            self.assert_hold(compare(changed), f"{side}_missing_shipping_price")

    def test_published_hold_cases_have_independently_stated_reasons_and_no_gap(self):
        expected = {
            "missing_shipping": "own_missing_shipping_price",
            "membership_mismatch": "mismatched_membership",
            "location_mismatch": "mismatched_location_context",
            "currency_mismatch": "mismatched_currency",
            "wrong_own_seller": "wrong_own_seller",
            "failed_featured": "featured_failed",
            "missing_featured": "featured_missing",
            "external_reference": "featured_wrong_source_type",
            "expected_price": "featured_wrong_source_type",
            "lowest_price": "featured_wrong_source_type",
            "time_skew": "observation_skew_exceeds_policy",
        }
        for identifier, reason in expected.items():
            with self.subTest(case=identifier):
                self.assert_hold(compare(BY_ID[identifier]), reason)

    def test_each_context_mismatch_blocks_numerical_comparison(self):
        for field, value in (
            ("asin", "B0SYNTH002"), ("marketplace", "CA"), ("currency", "CAD"),
            ("condition", "Used"), ("customer_type", "Business"), ("membership", "PRIME"),
            ("location_context", "US:94105:representative"), ("shipping_service", "EXPEDITED"),
        ):
            with self.subTest(field=field):
                case = self.base()
                case["featured"][field] = value
                self.assert_hold(compare(case), f"mismatched_{field}")

    def test_matching_records_outside_declared_scope_still_hold(self):
        for field, value in (("marketplace", "CA"), ("currency", "CAD"),
                             ("condition", "Used"), ("customer_type", "Business")):
            with self.subTest(field=field):
                case = self.base()
                for side in ("own", "featured"):
                    case[side][field] = value
                self.assert_hold(compare(case), f"unsupported_{field}")

    def test_wrong_own_source_or_identity_holds_while_other_featured_seller_is_valid(self):
        case = self.base()
        case["own"]["source_type"] = "featured_offer"
        self.assert_hold(compare(case), "own_wrong_source_type")
        case = self.base()
        case["expected_own_seller_id"] = "seller-C"
        self.assert_hold(compare(case), "wrong_own_seller")
        self.assertEqual(compare(self.base())["decision"], "compare")

    def test_unavailable_observations_and_missing_context_hold(self):
        for side in ("own", "featured"):
            for value, suffix in ((None, "missing"), ({"status": "absent"}, "absent"), ({"status": "failed"}, "failed")):
                case = self.base()
                case[side] = value
                with self.subTest(side=side, status=suffix):
                    self.assert_hold(compare(case), f"{side}_{suffix}")
            for field in ("seller_id", "asin", "marketplace", "currency", "condition", "customer_type",
                          "membership", "location_context", "shipping_service", "observed_at", "item_price"):
                case = self.base()
                case[side][field] = None
                with self.subTest(side=side, missing=field):
                    self.assert_hold(compare(case), f"{side}_missing_{field}")

    def test_skew_policy_is_inclusive_configurable_and_not_a_freshness_claim(self):
        case = self.base()
        case["featured"]["observed_at"] = "2026-09-16T10:05:00Z"
        self.assertEqual(compare(case)["decision"], "compare")
        case["featured"]["observed_at"] = "2026-09-16T10:05:01Z"
        self.assert_hold(compare(case), "observation_skew_exceeds_policy")
        self.assertEqual(compare(case, max_skew_seconds=301)["decision"], "compare")
        case["featured"]["observed_at"] = "2026-09-16T09:54:59Z"
        self.assert_hold(compare(case), "observation_skew_exceeds_policy")
        for side in ("own", "featured"):
            case[side]["observed_at"] = "2000-01-01T00:00:00Z"
        self.assertEqual(compare(case, max_skew_seconds=0)["decision"], "compare")
        for invalid in (-1, 86_401, True, "300"):
            with self.subTest(policy=invalid), self.assertRaises(InvalidInput):
                compare(case, invalid)

    def test_invalid_money_is_rejected_instead_of_rounded_or_coerced(self):
        for value in (-1, 1.0, True, "-0.00", "-1", "+1", " 1", "1 ", "01.00", ".50",
                      "1.001", "1e2", "NaN", "Infinity", "1000000000000", ""):
            with self.subTest(value=value), self.assertRaises(InvalidInput):
                case = self.base()
                case["own"]["item_price"] = value
                compare(case)

    def test_malformed_normalized_data_and_duplicate_ids_are_rejected(self):
        for mutation in ("missing_field", "unknown_field", "invalid_time", "failed_with_prices"):
            case = self.base()
            if mutation == "missing_field":
                del case["own"]["shipping_price"]
            elif mutation == "unknown_field":
                case["own"]["tax"] = "0.00"
            elif mutation == "invalid_time":
                case["own"]["observed_at"] = "2026-02-31T00:00:00Z"
            else:
                case["own"]["status"] = "failed"
            with self.subTest(mutation=mutation), self.assertRaises(InvalidInput):
                compare(case)
        with self.assertRaises(InvalidInput):
            report([self.base(), self.base()])
        with TemporaryDirectory() as directory:
            path = Path(directory) / "bad.json"
            for text in ('{"schema_version":1,"schema_version":1,"cases":[]}',
                         '{"schema_version":true,"cases":[]}', "[]"):
                path.write_text(text)
                with self.subTest(json=text), self.assertRaises(InvalidInput):
                    load_cases(path)

    def test_cli_json_and_csv_preserve_cases_gaps_and_empty_holds(self):
        json_run = subprocess.run([sys.executable, "-B", "compare.py"], cwd=HERE, capture_output=True, text=True, timeout=10)
        csv_run = subprocess.run([sys.executable, "-B", "compare.py", "--csv"], cwd=HERE, capture_output=True, text=True, timeout=10)
        self.assertEqual((json_run.returncode, csv_run.returncode), (0, 0))
        json_rows = json.loads(json_run.stdout)["results"]
        csv_rows = list(csv.DictReader(io.StringIO(csv_run.stdout)))
        self.assertEqual(len(json_rows), 14)
        self.assertEqual(len(csv_rows), 14)
        self.assertEqual(sum(row["decision"] == "compare" for row in json_rows), 3)
        for a, b in zip(json_rows, csv_rows, strict=True):
            self.assertEqual((a["id"], a["decision"]), (b["id"], b["decision"]))
            self.assertEqual(b["own_minus_featured"], a["own_minus_featured"] or "")
            self.assertEqual(b["reasons"], ";".join(a["reasons"]))

    def test_cli_policy_override_changes_only_the_eligible_time_case(self):
        run = subprocess.run([sys.executable, "-B", "compare.py", "--max-skew-seconds", "360"],
                             cwd=HERE, capture_output=True, text=True, timeout=10)
        self.assertEqual(run.returncode, 0, run.stderr)
        output = json.loads(run.stdout)
        rows = {row["id"]: row for row in output["results"]}
        self.assertEqual(output["max_skew_seconds"], 360)
        self.assertEqual(rows["time_skew"]["decision"], "compare")
        self.assertEqual(rows["time_skew"]["own_minus_featured"], "0.00")
        self.assertEqual(sum(row["decision"] == "compare" for row in rows.values()), 4)


if __name__ == "__main__":
    unittest.main()
