"""Loopback regressions for real POST forwarding, response loss and effects."""

import asyncio
import csv
import io
import unittest

import httpx
from retry_jobs_demo import PAYLOAD, RetryFixture, make_client, matrix_csv, request_headers, run_all


class RetryJobTests(unittest.IsolatedAsyncioTestCase):
    async def asyncTearDown(self):
        pending = [task for task in asyncio.all_tasks() if task is not asyncio.current_task()]
        self.assertEqual(pending, [], "test left asynchronous work behind")

    async def test_exact_case_matrix_and_response_loss_order(self):
        report = await run_all()
        expected = {
            "blind_retry": (2, 0, 2, 201), "ignored_stable_key": (2, 0, 2, 201),
            "supported_stable_key": (2, 0, 1, 201), "new_key_on_retry": (2, 0, 2, 201),
            "changed_payload": (2, 0, 1, 409), "reconcile_confirmed_effect": (1, 1, 1, 200),
        }
        self.assertEqual({case["case"] for case in report["cases"]}, set(expected))
        for case in report["cases"]:
            with self.subTest(case=case["case"]):
                self.assertEqual((case["origin_post_receipts"], case["origin_get_receipts"],
                                  case["applied_jobs"], case["second_action_status"]), expected[case["case"]])
                self.assertEqual(case["proxy_post_receipts"], case["origin_post_receipts"])
                self.assertEqual(case["proxy_get_receipts"], case["origin_get_receipts"])
                self.assertEqual(case["first_client_error"], "RemoteProtocolError")
                events = case["events"]
                chain = ("explicit_post", "request_received", "job_applied", "response_sent",
                         "complete_origin_response_received", "response_dropped_before_client_headers", "request_failed")
                positions = [next(e["sequence"] for e in events if e["event"] == name) for name in chain]
                self.assertEqual(positions, sorted(positions))
                self.assertEqual(sum(e["event"] == "response_dropped_before_client_headers" for e in events), 1)
                self.assertEqual(sum(e["event"] == "explicit_post" for e in events), case["origin_post_receipts"])
                self.assertEqual(sum(e["event"] == "job_applied" for e in events), case["applied_jobs"])
                self.assertEqual(case["applied_jobs"], len(case["jobs"]))
                self.assertEqual(case["dropped_origin_response"]["body"]["job"], case["jobs"][0])
                origin_keys = [e["key"] for e in events if e["side"] == "origin" and e["event"] == "request_received" and e["method"] == "POST"]
                if case["case"] == "ignored_stable_key":
                    self.assertEqual(origin_keys, ["key-1", "key-1"])
                if case["case"] == "new_key_on_retry":
                    self.assertEqual(origin_keys, ["key-1", "key-2"])
                self.assertEqual(case["cleanup"], {"handlers": 0, "writers": 0, "forced_cancellations": 0,
                                                    "client_closed": True, "errors": []})
                if case["case"] == "supported_stable_key":
                    self.assertEqual(case["client_received_result"], case["dropped_origin_response"])
                if case["case"] == "reconcile_confirmed_effect":
                    self.assertEqual(case["client_received_result"]["body"]["jobs"], case["jobs"])
        rows = list(csv.DictReader(io.StringIO(matrix_csv(report))))
        self.assertEqual(len(rows), 6)
        for row, case in zip(rows, report["cases"]):
            self.assertEqual(row["case"], case["case"])
            for field in ("proxy_post_receipts", "proxy_get_receipts", "origin_post_receipts", "origin_get_receipts", "applied_jobs"):
                self.assertEqual(int(row[field]), case[field])

    async def test_same_key_concurrency_replays_one_atomic_effect(self):
        fixture = RetryFixture(True, drop_first_post=False, gate_posts=2)
        async with fixture:
            async with make_client(fixture) as client:
                responses = await asyncio.wait_for(asyncio.gather(
                    client.post(fixture.origin_url + "/jobs", headers=request_headers("same-key"), json=PAYLOAD),
                    client.post(fixture.origin_url + "/jobs", headers=request_headers("same-key"),
                                json={"quantity": 1, "item": "synthetic-widget"}),
                ), 5.0)
                self.assertEqual([r.status_code for r in responses], [201, 201])
                self.assertEqual(responses[0].content, responses[1].content)
                self.assertEqual(len(fixture.jobs), 1)
                self.assertEqual(fixture.receipts["origin"]["POST"], 2)
                origin_posts = [e["sequence"] for e in fixture.events if e["side"] == "origin" and e["event"] == "request_received"]
                applied = next(e["sequence"] for e in fixture.events if e["event"] == "job_applied")
                self.assertLess(max(origin_posts), applied)  # both arrived before the first effect
                self.assertEqual(sum(e["event"] == "saved_result_replayed" for e in fixture.events), 1)
        fixture.assert_clean()
        self.assertTrue(client.is_closed)

    async def test_distinct_operation_keys_allow_identical_payloads(self):
        fixture = RetryFixture(True, drop_first_post=False, gate_posts=2)
        async with fixture:
            async with make_client(fixture) as client:
                responses = await asyncio.wait_for(asyncio.gather(*(
                    client.post(fixture.origin_url + "/jobs", headers=request_headers(f"key-{i}", f"operation-{i}"), json=PAYLOAD)
                    for i in (1, 2)
                )), 5.0)
                self.assertEqual([r.status_code for r in responses], [201, 201])
                self.assertNotEqual(responses[0].json()["job"]["id"], responses[1].json()["job"]["id"])
                self.assertEqual(len(fixture.jobs), 2)
                self.assertEqual([job["payload"] for job in fixture.jobs], [PAYLOAD, PAYLOAD])
                self.assertEqual({job["operation_ref"] for job in fixture.jobs}, {"operation-1", "operation-2"})
        fixture.assert_clean()

    async def test_same_key_changed_payload_conflicts_under_concurrency(self):
        fixture = RetryFixture(True, drop_first_post=False, gate_posts=2)
        payloads = [PAYLOAD, {**PAYLOAD, "quantity": 2}]
        async with fixture:
            async with make_client(fixture) as client:
                responses = await asyncio.wait_for(asyncio.gather(*(
                    client.post(fixture.origin_url + "/jobs", headers=request_headers("same-key"), json=payload)
                    for payload in payloads
                )), 5.0)
                self.assertEqual(sorted(r.status_code for r in responses), [201, 409])
                self.assertEqual(len(fixture.jobs), 1)
                winner = next(i for i, response in enumerate(responses) if response.status_code == 201)
                self.assertEqual(fixture.jobs[0]["payload"], payloads[winner])
                replay = await client.post(fixture.origin_url + "/jobs", headers=request_headers("same-key"), json=payloads[winner])
                self.assertEqual(replay.content, responses[winner].content)
                self.assertEqual(len(fixture.jobs), 1)
                wrong_ref = await client.post(fixture.origin_url + "/jobs", headers=request_headers("same-key", "different-operation"), json=payloads[winner])
                self.assertEqual(wrong_ref.status_code, 409)
                self.assertEqual(len(fixture.jobs), 1)
        fixture.assert_clean()

    async def test_fixture_refuses_an_external_destination(self):
        fixture = RetryFixture(True, drop_first_post=False)
        async with fixture:
            async with make_client(fixture) as client:
                with self.assertRaises(httpx.RemoteProtocolError):
                    await client.post("http://example.invalid/never-requested", headers=request_headers(), json=PAYLOAD)
        self.assertEqual(fixture.receipts, {"proxy": {"POST": 0, "GET": 0}, "origin": {"POST": 0, "GET": 0}})
        self.assertEqual(fixture.jobs, [])
        self.assertEqual(fixture.errors, [{"side": "proxy", "error": "AssertionError"}])
        self.assertFalse(fixture.writers or fixture.tasks or fixture.forced_cancellations)
        self.assertFalse(any(server.is_serving() for server in fixture.servers))


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