#!/usr/bin/env python3
"""Sandboxed integration tests for automatic Codex usage-reset redemption."""

from __future__ import annotations

import base64
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
import json
import os
from pathlib import Path
import subprocess
import tempfile
import threading
import time
import unittest
import uuid

REPO = Path(__file__).resolve().parents[1]


def jwt(claims: dict) -> str:
    payload = base64.urlsafe_b64encode(json.dumps(claims).encode()).decode().rstrip("=")
    return f"header.{payload}.signature"


class ResetHandler(BaseHTTPRequestHandler):
    def do_GET(self) -> None:  # noqa: N802 - BaseHTTPRequestHandler contract
        state = self.server.state
        if self.path == "/wham/usage":
            state["usage_gets"] += 1
            body = state["usage"]
        elif self.path == "/wham/rate-limit-reset-credits":
            state["credit_gets"] += 1
            body = state["credits"]
        else:
            self.send_error(404)
            return
        encoded = json.dumps(body).encode()
        self.send_response(200)
        self.send_header("Content-Type", "application/json")
        self.send_header("Content-Length", str(len(encoded)))
        self.end_headers()
        self.wfile.write(encoded)

    def do_POST(self) -> None:  # noqa: N802 - BaseHTTPRequestHandler contract
        state = self.server.state
        if self.path != "/wham/rate-limit-reset-credits/consume":
            self.send_error(404)
            return
        length = int(self.headers.get("Content-Length") or 0)
        state["posts"].append(json.loads(self.rfile.read(length)))
        status, body = state["responses"].pop(0)
        if status == 200 and body.get("code") == "reset":
            state["credits"]["available_count"] -= 1
        if "credits_after_post" in state:
            state["credits"] = state["credits_after_post"]
        encoded = json.dumps(body).encode()
        self.send_response(status)
        self.send_header("Content-Type", "application/json")
        self.send_header("Content-Length", str(len(encoded)))
        self.end_headers()
        self.wfile.write(encoded)

    def log_message(self, _format: str, *_args) -> None:
        return


class CodexPoolSandbox:
    """A one-account codex pool with a fake usage/credits endpoint."""

    def setUp(self) -> None:
        self.temp = tempfile.TemporaryDirectory(prefix="multiacc-reset-")
        self.pool = Path(self.temp.name) / "pool"
        account = self.pool / "acct-01"
        account.mkdir(parents=True)
        manifest = {"version": 1, "threshold": 90, "accounts": [
            {"id": "acct-01", "email": "test@example.invalid", "home": "mac"}]}
        (self.pool / "accounts.json").write_text(json.dumps(manifest), encoding="utf-8")
        claims = {"exp": int(time.time()) + 86400,
                  "https://api.openai.com/auth": {"chatgpt_account_id": "account-test"}}
        auth = {"auth_mode": "chatgpt", "tokens": {
            "access_token": jwt(claims), "refresh_token": "fake-refresh",
            "account_id": "account-test"}}
        (account / "auth.json").write_text(json.dumps(auth), encoding="utf-8")
        self.server = ThreadingHTTPServer(("127.0.0.1", 0), ResetHandler)
        self.server.state = {"usage_gets": 0, "credit_gets": 0, "posts": [],
                             "usage": {}, "credits": {}, "responses": []}
        self.thread = threading.Thread(target=self.server.serve_forever, daemon=True)
        self.thread.start()

    def tearDown(self) -> None:
        self.server.shutdown()
        self.server.server_close()
        self.thread.join(timeout=5)
        self.temp.cleanup()


class CodexResetIntegrationTest(CodexPoolSandbox, unittest.TestCase):
    def run_limits(self) -> subprocess.CompletedProcess:
        base = f"http://127.0.0.1:{self.server.server_port}/wham"
        env = os.environ.copy()
        env.update({"CODEX_ACCOUNTS_ROOT": str(self.pool),
                    "CODEX_MULTIACC_NO_SYNC": "1", "CODEX_MULTIACC_MIN_FETCH": "0",
                    "CODEX_MULTIACC_AUTO_RESET": "1", "PYTHONDONTWRITEBYTECODE": "1",
                    "CODEX_MULTIACC_USAGE_URL": f"{base}/usage"})
        env.pop("CODEX_MULTIACC_RESET_CREDITS_URL", None)
        env.pop("CODEX_MULTIACC_RESET_CONSUME_URL", None)
        return subprocess.run([REPO / "bin/codex-accounts", "limits", "--force"],
                              capture_output=True, text=True, env=env, timeout=20, check=False)

    def test_redeems_at_five_percent_remaining_once(self) -> None:
        now = int(time.time())
        self.server.state["usage"] = {"plan_type": "pro", "rate_limit": {
            "allowed": True, "primary_window": {"used_percent": 95,
                "limit_window_seconds": 18000, "reset_at": now + 3600}},
            "rate_limit_reset_credits": {"available_count": 2}}
        self.server.state["credits"] = {"available_count": 2, "credits": [
            {"id": "later", "status": "available", "expires_at": "2026-09-02T00:00:00Z"},
            {"id": "sooner", "status": "available", "expires_at": "2026-09-01T00:00:00Z"}]}
        self.server.state["responses"] = [(200, {"code": "reset", "windows_reset": 2})]
        result = self.run_limits()
        self.assertEqual(result.returncode, 0, result.stderr)
        self.assertIn("usage reset redeemed automatically at 95% used", result.stdout)
        self.assertEqual(len(self.server.state["posts"]), 1)
        request = self.server.state["posts"][0]
        self.assertEqual(request["credit_id"], "sooner")
        uuid.UUID(request["redeem_request_id"])
        state = json.loads((self.pool / "acct-01/.usage-reset.json").read_text())
        self.assertEqual(state["state"], "complete")
        marker = self.pool / "acct-01/.limited"
        self.assertFalse(marker.exists())
        marker.write_text(f"{now + 3600}\nreason=client-rate-limit\n", encoding="utf-8")
        (self.pool / "acct-01/.usage-reset.json").unlink()
        self.server.state["responses"] = [(200, {"code": "already_redeemed"})]
        result = self.run_limits()
        self.assertEqual(result.returncode, 0, result.stderr)
        self.assertFalse(marker.exists())
        self.assertEqual(self.server.state["posts"][0]["redeem_request_id"],
                         self.server.state["posts"][1]["redeem_request_id"])
        self.run_limits()
        self.assertEqual(len(self.server.state["posts"]), 2)

    def test_does_not_redeem_below_the_pools_threshold(self) -> None:
        now = int(time.time())
        self.server.state["usage"] = {"rate_limit": {"allowed": True,
            "primary_window": {"used_percent": 89, "limit_window_seconds": 18000,
                               "reset_at": now + 3600}},
            "rate_limit_reset_credits": {"available_count": 1}}
        self.server.state["credits"] = {"available_count": 1, "credits": []}
        result = self.run_limits()
        self.assertEqual(result.returncode, 0, result.stderr)
        self.assertEqual(self.server.state["credit_gets"], 1)
        self.assertEqual(self.server.state["posts"], [])
        self.assertFalse((self.pool / "acct-01/.usage-reset.json").exists())

    def test_an_account_parked_at_the_pools_threshold_uses_its_reset(self) -> None:
        # Operator rule (2026-09-22): a PARKED account uses its reset. The pool parks at
        # 90%; waiting for 95% left accounts at 90-94% idle for days with credits unused.
        now = int(time.time())
        self.server.state["usage"] = {"rate_limit": {"allowed": True,
            "primary_window": {"used_percent": 92, "limit_window_seconds": 604800,
                               "reset_at": now + 4 * 86400}}}
        self.server.state["credits"] = {"available_count": 1, "credits": []}
        self.server.state["responses"] = [(200, {"code": "reset", "windows_reset": 1})]
        self.assertEqual(self.run_limits().returncode, 0)
        self.assertEqual(len(self.server.state["posts"]), 1)
        self.assertFalse((self.pool / "acct-01/.limited").exists())

    def test_a_client_reported_park_uses_its_reset(self) -> None:
        # The rollout scan parked the account on a real 429 its telemetry does not show.
        now = int(time.time())
        marker = self.pool / "acct-01/.limited"
        marker.write_text(f"{now + 3 * 86400}\nbucket=client:7d percent=100 "
                          f"marked_at=2020-01-01T00:00:00Z reason=client-rate-limit\n")
        self.server.state["usage"] = {"rate_limit": {"allowed": True,
            "primary_window": {"used_percent": 40, "limit_window_seconds": 604800,
                               "reset_at": now + 3 * 86400}}}
        self.server.state["credits"] = {"available_count": 1, "credits": []}
        self.server.state["responses"] = [(200, {"code": "reset", "windows_reset": 1})]
        self.assertEqual(self.run_limits().returncode, 0)
        self.assertEqual(len(self.server.state["posts"]), 1)
        self.assertFalse(marker.exists())

    def test_a_park_that_lifts_within_the_hour_keeps_its_reset(self) -> None:
        now = int(time.time())
        (self.pool / "acct-01/.limited").write_text(
            f"{now + 1200}\nbucket=client:5h percent=100 marked_at=2020-01-01T00:00:00Z "
            f"reason=client-rate-limit\n")
        self.server.state["usage"] = {"rate_limit": {"allowed": True,
            "primary_window": {"used_percent": 40, "limit_window_seconds": 604800,
                               "reset_at": now + 3 * 86400}}}
        self.server.state["credits"] = {"available_count": 1, "credits": []}
        self.assertEqual(self.run_limits().returncode, 0)
        self.assertEqual(self.server.state["posts"], [])

    def test_retries_a_lost_response_with_the_same_idempotency_key(self) -> None:
        now = int(time.time())
        self.server.state["usage"] = {"rate_limit": {"allowed": False,
            "limit_reached": True, "primary_window": {"used_percent": 100,
                "limit_window_seconds": 604800, "reset_at": now + 86400}},
            "rate_limit_reset_credits": {"available_count": 1}}
        self.server.state["credits"] = {"available_count": 1, "credits": []}
        self.server.state["responses"] = [
            (500, {"error": "response lost"}),
            (200, {"code": "already_redeemed", "windows_reset": 2})]
        first = self.run_limits()
        self.assertIn("retry is idempotent", first.stdout)
        self.assertEqual(json.loads((self.pool / "acct-01/.usage-reset.json").read_text())
                         ["state"], "pending")
        second = self.run_limits()
        self.assertEqual(second.returncode, 0, second.stderr)
        self.assertEqual(len(self.server.state["posts"]), 2)
        keys = [request["redeem_request_id"] for request in self.server.state["posts"]]
        self.assertEqual(keys[0], keys[1])
        self.assertEqual(self.server.state["credit_gets"], 4)


class CodexResetWatermarkTest(CodexPoolSandbox, unittest.TestCase):
    """A redemption answers every rejection before it: bin/codex's rollout scan must not
    re-create a park from one, whether or not a marker existed (review, 2026-09-22)."""

    run_limits = CodexResetIntegrationTest.run_limits

    def test_redemption_writes_the_watermark_even_without_a_marker(self) -> None:
        now = int(time.time())
        account = self.pool / "acct-01"
        self.assertFalse((account / ".limited").exists())
        (account / ".client-limit-cleared").write_text("5\n", encoding="utf-8")
        self.server.state["usage"] = {"plan_type": "pro", "rate_limit": {"allowed": True,
            "primary_window": {"used_percent": 96, "limit_window_seconds": 604800,
                               "reset_at": now + 86400}}}
        self.server.state["credits"] = {"available_count": 1, "credits": [
            {"id": "c1", "status": "available", "expires_at": "2099-01-01T00:00:00Z"}]}
        self.server.state["responses"] = [(200, {"code": "reset", "windows_reset": 1})]
        self.assertEqual(self.run_limits().returncode, 0)
        mark = int((account / ".client-limit-cleared").read_text().split()[0])
        self.assertGreaterEqual(mark, now)

    def test_the_watermark_never_moves_backwards(self) -> None:
        now = int(time.time())
        account = self.pool / "acct-01"
        future = now + 7200
        (account / ".client-limit-cleared").write_text(f"{future}\n", encoding="utf-8")
        self.server.state["usage"] = {"plan_type": "pro", "rate_limit": {"allowed": True,
            "primary_window": {"used_percent": 96, "limit_window_seconds": 604800,
                               "reset_at": now + 86400}}}
        self.server.state["credits"] = {"available_count": 1, "credits": []}
        self.server.state["responses"] = [(200, {"code": "reset", "windows_reset": 1})]
        self.assertEqual(self.run_limits().returncode, 0)
        self.assertEqual(int((account / ".client-limit-cleared").read_text().split()[0]), future)


class CodexMarkerNamesOneBucketTest(CodexPoolSandbox, unittest.TestCase):
    """The marker's epoch is the reset of the bucket its detail line NAMES.

    The codex writer used to pair the highest-PERCENT bucket's name with the
    LATEST reset over every offender, so a five-hour window at 100% resetting
    tonight was written under a seven-day epoch — and everything that trusts the
    marker parked the account for a week over a window that refills in hours.
    #20 fixed this for the claude pool and left codex behind.
    """

    def run_limits(self) -> subprocess.CompletedProcess:  # no auto-reset: it clears markers
        base = f"http://127.0.0.1:{self.server.server_port}/wham"
        env = os.environ.copy()
        env.update({"CODEX_ACCOUNTS_ROOT": str(self.pool),
                    "CODEX_MULTIACC_NO_SYNC": "1", "CODEX_MULTIACC_MIN_FETCH": "0",
                    "CODEX_MULTIACC_AUTO_RESET": "0", "PYTHONDONTWRITEBYTECODE": "1",
                    "CODEX_MULTIACC_USAGE_URL": f"{base}/usage"})
        return subprocess.run([REPO / "bin/codex-accounts", "limits", "--force"],
                              capture_output=True, text=True, env=env, timeout=20, check=False)

    def _marker(self):
        text = (self.pool / "acct-01/.limited").read_text(encoding="utf-8").splitlines()
        return int(text[0]), text[1]

    def test_two_offenders_write_one_bucket_with_its_own_reset(self) -> None:
        now = int(time.time())
        session_reset, weekly_reset = now + 3600, now + 6 * 86400
        self.server.state["usage"] = {"plan_type": "pro", "rate_limit": {
            "allowed": True,
            # 100% and back in an hour …
            "primary_window": {"used_percent": 100, "limit_window_seconds": 18000,
                               "reset_at": session_reset},
            # … beside 95% that is six days out. Both are over the threshold.
            "secondary_window": {"used_percent": 95, "limit_window_seconds": 604800,
                                 "reset_at": weekly_reset}}}
        self.assertEqual(self.run_limits().returncode, 0)
        epoch, detail = self._marker()
        # The longest-lived offender is named, and the epoch is ITS reset — the
        # account really is excluded until then, and a reader that takes the
        # detail's percent gets the percent of the very bucket the epoch belongs
        # to rather than a different bucket's.
        self.assertEqual(epoch, weekly_reset)
        self.assertIn("percent=95", detail)
        self.assertIn(time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime(weekly_reset)), detail)
        self.assertNotIn("percent=100", detail)

    def test_a_single_offender_still_names_itself(self) -> None:
        now = int(time.time())
        reset = now + 4 * 86400
        self.server.state["usage"] = {"plan_type": "pro", "rate_limit": {
            "allowed": True,
            "primary_window": {"used_percent": 91, "limit_window_seconds": 604800,
                               "reset_at": reset},
            "secondary_window": {"used_percent": 12, "limit_window_seconds": 18000,
                                 "reset_at": now + 900}}}
        self.assertEqual(self.run_limits().returncode, 0)
        epoch, detail = self._marker()
        self.assertEqual(epoch, reset)
        self.assertIn("percent=91", detail)


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