from copy import deepcopy
import json
from pathlib import Path
import sqlite3
import subprocess
import sys
from tempfile import TemporaryDirectory
import unittest

from consumer import (
    IdentityConflict, InvalidMessage, ListingKey, ReadResult, Store, WrongKey, refresh_once,
)


HERE = Path(__file__).resolve().parent
FIXTURE = json.loads((HERE / "fixtures/scenario.json").read_text())


class ConsumerTests(unittest.TestCase):
    def setUp(self):
        self.temporary = TemporaryDirectory()
        self.addCleanup(self.temporary.cleanup)
        self.path = Path(self.temporary.name) / "inbox.sqlite3"
        self.store = Store(self.path)
        self.addCleanup(self.store.close)
        self.event = deepcopy(FIXTURE["initial_event"])
        self.key = ListingKey.from_event(self.event)

    def ingest(self, event=None):
        return self.store.ingest(json.dumps(event or self.event))

    def test_duplicate_reordered_json_does_not_create_work_or_reopen_clean_work(self):
        self.assertEqual(self.ingest().disposition, "stored")
        self.assertEqual(refresh_once(self.store, self.key, lambda k: ReadResult(k, {"read": 1})), "applied")
        reordered = dict(reversed(list(self.event.items())))
        duplicate = self.store.ingest(json.dumps(reordered, indent=4))
        self.assertEqual((duplicate.disposition, duplicate.may_ack), ("duplicate", True))
        self.assertEqual(self.store.inbox_count(), 1)
        self.assertEqual(self.store.state(self.key), {
            "generation": 1, "applied_generation": 1, "dirty": False, "snapshot": {"read": 1},
        })
        self.assertIsNone(self.store.begin_refresh(self.key))

    def test_identity_binds_source_subscription_and_notification_id(self):
        self.ingest()
        for field, value in (("source", "other-source"), ("subscription", "other-subscription"),
                             ("notification_id", "other-notice")):
            event = {**self.event, field: value}
            self.assertEqual(self.ingest(event).disposition, "stored")
        self.assertEqual(self.store.inbox_count(), 4)
        self.assertEqual(self.store.state(self.key)["generation"], 4)

    def test_conflicting_duplicate_rejects_payload_key_type_or_time_change(self):
        self.ingest()
        for changed in (
            {"payload": {"hint": "different"}}, {"marketplace_id": "market-B"},
            {"notification_type": "listing_issues_changed"}, {"event_time": "2026-09-16T08:00:00Z"},
        ):
            with self.subTest(changed=changed), self.assertRaises(IdentityConflict):
                self.ingest({**self.event, **changed})
        self.assertEqual(self.store.inbox_count(), 1)
        self.assertEqual(self.store.state(self.key)["generation"], 1)
        self.assertIsNone(self.store.state(ListingKey("seller-A", "SKU-RED", "market-B")))

    def test_failed_work_insert_and_update_roll_back_inbox_in_same_transaction(self):
        for operation in ("INSERT", "UPDATE"):
            with self.subTest(operation=operation):
                with sqlite3.connect(self.path) as connection:
                    connection.execute(f"""CREATE TRIGGER fail_work BEFORE {operation} ON work
                        BEGIN SELECT RAISE(ABORT, 'injected work failure'); END""")
                event = {**self.event, "notification_id": f"failed-{operation}"}
                before = self.store.inbox_count()
                with self.assertRaisesRegex(sqlite3.IntegrityError, "injected"):
                    self.ingest(event)
                self.assertEqual(self.store.inbox_count(), before)
                with sqlite3.connect(self.path) as connection:
                    connection.execute("DROP TRIGGER fail_work")
                if operation == "INSERT":
                    self.assertIsNone(self.store.state(self.key))
                    self.ingest()
        self.assertEqual(self.store.state(self.key)["generation"], 1)

    def test_process_exit_after_persistence_before_ack_then_redelivery(self):
        child = subprocess.run([
            sys.executable, "-B", "-c", "import os,sys; from consumer import Store; "
            "Store(sys.argv[1]).ingest(sys.stdin.read()); os._exit(23)", str(self.path),
        ], input=json.dumps(self.event), text=True, capture_output=True, cwd=HERE, timeout=10)
        self.assertEqual(child.returncode, 23, child.stderr)
        with Store(self.path) as restarted:
            self.assertEqual(restarted.inbox_count(), 1)
            self.assertEqual(restarted.dirty_keys(), [self.key])
            duplicate = restarted.ingest(json.dumps(self.event))
            self.assertEqual((duplicate.disposition, duplicate.may_ack), ("duplicate", True))
            self.assertEqual(restarted.state(self.key)["generation"], 1)
            self.assertEqual(refresh_once(restarted, self.key, lambda k: ReadResult(k, {"read": 1})), "applied")
        with Store(self.path) as another_restart:
            self.assertEqual(another_restart.dirty_keys(), [])
            self.assertEqual(another_restart.state(self.key)["snapshot"], {"read": 1})

    def test_out_of_order_new_event_triggers_read_instead_of_applying_payload(self):
        self.ingest()
        refresh_once(self.store, self.key, lambda k: ReadResult(k, FIXTURE["initial_read"]))
        self.ingest(FIXTURE["late_event"])
        pending = self.store.state(self.key)
        self.assertEqual((pending["generation"], pending["applied_generation"], pending["dirty"]), (2, 1, True))
        self.assertEqual(pending["snapshot"], FIXTURE["initial_read"])
        reads = []

        def read(key):
            reads.append(key)
            return ReadResult(key, FIXTURE["retry_read"])

        self.assertEqual(refresh_once(self.store, self.key, read), "applied")
        self.assertEqual(reads, [self.key])
        self.assertEqual(self.store.state(self.key)["snapshot"], FIXTURE["retry_read"])

    def test_new_notification_during_read_rejects_old_completion_and_keeps_dirty(self):
        self.ingest()
        reads = []

        def racing_read(key):
            reads.append(key)
            with Store(self.path) as ingester:
                ingester.ingest(json.dumps(FIXTURE["during_read_event"]))
            return ReadResult(key, FIXTURE["discarded_read"])

        self.assertEqual(refresh_once(self.store, self.key, racing_read), "superseded")
        self.assertEqual(self.store.state(self.key), {
            "generation": 2, "applied_generation": 0, "dirty": True, "snapshot": None,
        })
        self.assertEqual(refresh_once(self.store, self.key, lambda k: ReadResult(k, FIXTURE["retry_read"])), "applied")
        self.assertEqual(reads, [self.key])
        self.assertEqual(self.store.state(self.key)["snapshot"], FIXTURE["retry_read"])

    def test_read_failure_preserves_dirty_work_and_previous_local_snapshot(self):
        self.ingest()
        refresh_once(self.store, self.key, lambda k: ReadResult(k, FIXTURE["initial_read"]))
        self.ingest(FIXTURE["during_read_event"])

        def failed_read(_):
            raise OSError("synthetic provider failure")

        with self.assertRaises(OSError):
            refresh_once(self.store, self.key, failed_read)
        with Store(self.path) as restarted:
            self.assertEqual(restarted.state(self.key), {
                "generation": 2, "applied_generation": 1, "dirty": True,
                "snapshot": FIXTURE["initial_read"],
            })

    def test_wrong_key_completion_rejected_and_seller_marketplace_sku_independent(self):
        self.ingest()
        other_keys = []
        for index, (field, value) in enumerate((
            ("seller_id", "seller-B"), ("marketplace_id", "market-B"), ("sku", "SKU-BLUE"),
        )):
            event = {**self.event, field: value, "notification_id": f"independent-{index}"}
            self.ingest(event)
            other_keys.append(ListingKey.from_event(event))
        ticket = self.store.begin_refresh(self.key)
        with self.assertRaises(WrongKey):
            self.store.complete(ticket, ReadResult(other_keys[0], {"read": 99}))
        self.assertEqual(len(self.store.dirty_keys()), 4)
        self.assertTrue(self.store.complete(ticket, ReadResult(self.key, {"read": 1})))
        for key in other_keys:
            self.assertEqual(self.store.state(key), {
                "generation": 1, "applied_generation": 0, "dirty": True, "snapshot": None,
            })

    def test_repeated_completion_cannot_replace_snapshot(self):
        self.ingest()
        ticket = self.store.begin_refresh(self.key)
        self.assertTrue(self.store.complete(ticket, ReadResult(self.key, {"read": 1})))
        self.assertFalse(self.store.complete(ticket, ReadResult(self.key, {"read": 2})))
        self.assertEqual(self.store.state(self.key)["snapshot"], {"read": 1})

    def test_strict_input_rejects_missing_scope_unknown_fields_and_malformed_json(self):
        missing_market = {k: v for k, v in self.event.items() if k != "marketplace_id"}
        nested = {}
        for _ in range(25):
            nested = {"nested": nested}
        invalid = [
            json.dumps(missing_market),
            *(json.dumps({**self.event, field: value}) for field, value in (
                ("marketplace_id", ""), ("seller_id", None), ("sku", " SKU-RED"),
                ("schema_version", True), ("notification_type", "raw_provider_type"),
                ("event_time", "2026-02-31T10:00:00Z"), ("event_time", "2026-09-16T10:00:00"),
                ("payload", []), ("extra", 1), ("payload", {"bad": float("nan")}),
                ("payload", {"bad": float("inf")}), ("payload", {"bad": "\ud800"}),
                ("payload", nested), ("payload", {"large": "x" * 65_536}),
            )),
            json.dumps(self.event).replace('"schema_version": 1', '"schema_version": 1, "schema_version": 1'),
            json.dumps(self.event).replace('"synthetic status change"', '1e999'),
            "[]", "{",
        ]
        for index, raw in enumerate(invalid):
            with self.subTest(case=index), self.assertRaises(InvalidMessage):
                self.store.ingest(raw)
        self.assertEqual(self.store.inbox_count(), 0)
        self.assertEqual(self.store.dirty_keys(), [])

    def test_invalid_read_document_cannot_clear_pending_work(self):
        self.ingest()
        for document in ([], {"bad": float("nan")}, {"large": "x" * 65_536}):
            with self.subTest(document_type=type(document)), self.assertRaises(InvalidMessage):
                self.store.complete(self.store.begin_refresh(self.key), ReadResult(self.key, document))
        self.assertTrue(self.store.state(self.key)["dirty"])
        self.assertIsNone(self.store.state(self.key)["snapshot"])


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