#!/usr/bin/env python3
"""
Tests for round-12 explicit per-request timeouts in ``asc_common.request``.

Without an explicit ``timeout`` kwarg, ``requests.request(...)`` blocks
forever on a hung TCP socket. CI then ticks until the runner's 6h hard
limit and a human pages on. The tests here pin:

  * default ``(connect, read)`` = (10.0, 30.0)
  * env overrides ``ASC_REQUEST_TIMEOUT_CONNECT_SEC`` /
    ``ASC_REQUEST_TIMEOUT_READ_SEC`` parse + take effect
  * invalid / negative env values fall back to the defaults so a typo
    can't disable the timeout
  * the timeout is forwarded to ``requests.request`` on every attempt
    (including retries)
"""

from __future__ import annotations

import os
import sys
import unittest
from pathlib import Path
from unittest import mock

sys.path.insert(0, str(Path(__file__).resolve().parent))

import asc_common  # noqa: E402


class _OkResponse:
    """Minimal stand-in for a 200 OK response."""

    status_code = 200
    text = ""

    def json(self) -> dict:
        return {"data": []}


class RequestTimeoutTests(unittest.TestCase):
    """``request()`` must pass an explicit ``(connect, read)`` timeout."""

    def _capture_timeout_via_request(self, env: dict) -> tuple[float, float]:
        captured: dict = {}

        def fake_request(method, url, **kwargs):
            captured["timeout"] = kwargs.get("timeout")
            return _OkResponse()

        with mock.patch.dict(os.environ, env, clear=False), \
                mock.patch.object(asc_common.requests, "request",
                                  side_effect=fake_request):
            asc_common.request("GET", "/apps/111", "tok")

        return captured.get("timeout")

    def test_default_timeouts_are_10s_connect_30s_read(self):
        timeout = self._capture_timeout_via_request({
            "ASC_REQUEST_TIMEOUT_CONNECT_SEC": "",
            "ASC_REQUEST_TIMEOUT_READ_SEC": "",
        })
        self.assertEqual(timeout, (10.0, 30.0))

    def test_env_override_takes_effect(self):
        timeout = self._capture_timeout_via_request({
            "ASC_REQUEST_TIMEOUT_CONNECT_SEC": "5",
            "ASC_REQUEST_TIMEOUT_READ_SEC": "60",
        })
        self.assertEqual(timeout, (5.0, 60.0))

    def test_invalid_env_falls_back_to_default(self):
        """A typo (e.g. ``CONNECT=ten``) MUST NOT disable the timeout."""
        timeout = self._capture_timeout_via_request({
            "ASC_REQUEST_TIMEOUT_CONNECT_SEC": "ten",
            "ASC_REQUEST_TIMEOUT_READ_SEC": "thirty",
        })
        self.assertEqual(timeout, (10.0, 30.0))

    def test_non_positive_env_falls_back_to_default(self):
        """Zero/negative values would disable the timeout; default instead."""
        timeout = self._capture_timeout_via_request({
            "ASC_REQUEST_TIMEOUT_CONNECT_SEC": "0",
            "ASC_REQUEST_TIMEOUT_READ_SEC": "-5",
        })
        self.assertEqual(timeout, (10.0, 30.0))

    def test_timeout_forwarded_on_every_retry_attempt(self):
        """A 503 retry must pass the same timeout on the retry call."""
        captured_timeouts: list = []
        responses = iter([_Status(503), _OkResponse()])

        def fake_request(method, url, **kwargs):
            captured_timeouts.append(kwargs.get("timeout"))
            return next(responses)

        with mock.patch.dict(
                os.environ,
                {"ASC_REQUEST_TIMEOUT_CONNECT_SEC": "",
                 "ASC_REQUEST_TIMEOUT_READ_SEC": ""},
                clear=False), \
                mock.patch.object(asc_common.requests, "request",
                                  side_effect=fake_request), \
                mock.patch.object(asc_common.time, "sleep"):
            asc_common.request("GET", "/apps/111", "tok")

        self.assertEqual(len(captured_timeouts), 2)
        self.assertEqual(captured_timeouts[0], (10.0, 30.0))
        self.assertEqual(captured_timeouts[1], (10.0, 30.0))


class _Status:
    """Stand-in for a non-2xx response that should trigger a retry."""

    def __init__(self, code: int) -> None:
        self.status_code = code
        self.text = f"<error {code}>"


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