"""Zero-dependency SpicyAPI reference client for Python 3.11+.

This is copyable repository example code, not a published PyPI package. The
ten-minute polling deadline is a local safety bound, not a production SLA.
"""

from __future__ import annotations

import json
import os
import random
import socket
import time
import urllib.error
import urllib.parse
import urllib.request
from typing import Any, Callable, Mapping

API_BASE_URL = "https://api.spicyapi.ai/api/v1"
REQUEST_TIMEOUT_SECONDS = 30.0
WAIT_TIMEOUT_SECONDS = 10.0 * 60.0
TERMINAL_STATES = frozenset({"succeeded", "failed", "canceled", "expired"})
ACTIVE_STATES = frozenset({"queued", "running"})
RETRYABLE_HTTP = frozenset({408, 429, 500, 502, 503, 504})
RETRYABLE_CODES = frozenset({429, 500, 50301})

TransportResult = tuple[int, Mapping[str, str], bytes]
Transport = Callable[[str, str, Mapping[str, str], bytes | None, float], TransportResult]


class SpicyApiError(RuntimeError):
    """HTTP, envelope, or network failure with support correlation fields."""

    def __init__(
        self,
        message: str,
        *,
        status: int = 0,
        code: int | None = None,
        request_id: str = "",
    ) -> None:
        super().__init__(message)
        self.status = status
        self.code = code
        self.request_id = request_id


class SpicyTimeoutError(TimeoutError):
    """A local request or polling deadline elapsed; remote state is unknown."""

    def __init__(self, message: str, *, task_id: str | None = None) -> None:
        super().__init__(message)
        self.task_id = task_id


def _urllib_transport(
    method: str,
    url: str,
    headers: Mapping[str, str],
    body: bytes | None,
    timeout: float,
) -> TransportResult:
    request = urllib.request.Request(url, data=body, headers=dict(headers), method=method)
    try:
        with urllib.request.urlopen(request, timeout=timeout) as response:
            return response.status, dict(response.headers.items()), response.read()
    except urllib.error.HTTPError as error:
        return error.code, dict(error.headers.items()), error.read()


class SpicyClient:
    def __init__(
        self,
        api_key: str | None = None,
        *,
        base_url: str = API_BASE_URL,
        transport: Transport = _urllib_transport,
        sleep: Callable[[float], None] = time.sleep,
        monotonic: Callable[[], float] = time.monotonic,
        random_value: Callable[[], float] = random.random,
        request_timeout_seconds: float = REQUEST_TIMEOUT_SECONDS,
        wait_timeout_seconds: float = WAIT_TIMEOUT_SECONDS,
        max_retries: int = 3,
    ) -> None:
        self.api_key = api_key or os.environ.get("SPICY_API_KEY", "")
        if not self.api_key:
            raise ValueError("SPICY_API_KEY is required")
        self.base_url = base_url.rstrip("/")
        self.transport = transport
        self.sleep = sleep
        self.monotonic = monotonic
        self.random_value = random_value
        self.request_timeout_seconds = request_timeout_seconds
        self.wait_timeout_seconds = wait_timeout_seconds
        self.max_retries = max_retries

    def list_models(
        self,
        *,
        modality: str | None = None,
        provider: str | None = None,
        task: str | None = None,
        search: str | None = None,
        include_schema: bool | None = None,
        include_examples: bool | None = None,
    ) -> dict[str, Any]:
        values: dict[str, str] = {}
        for key, value in {
            "modality": modality,
            "provider": provider,
            "task": task,
            "search": search,
            "includeSchema": None if include_schema is None else str(int(include_schema)),
            "includeExamples": None if include_examples is None else str(int(include_examples)),
        }.items():
            if value is not None:
                values[key] = value
        suffix = f"?{urllib.parse.urlencode(values)}" if values else ""
        return self._request("GET", f"/models{suffix}")

    def get_model(self, model: str) -> dict[str, Any]:
        if not model:
            raise ValueError("model is required")
        encoded = urllib.parse.quote(model, safe="")
        return self._request("GET", f"/models/{encoded}")

    def create_task(
        self,
        *,
        model: str,
        input_data: Mapping[str, Any],
        idempotency_key: str,
        callback_url: str | None = None,
        mature: bool | None = None,
    ) -> dict[str, Any]:
        if not idempotency_key.strip():
            raise ValueError("Idempotency-Key is required")
        body: dict[str, Any] = {"model": model, "input": dict(input_data)}
        if callback_url is not None:
            body["callBackUrl"] = callback_url
        if mature is not None:
            body["mature"] = mature
        return self._request(
            "POST",
            "/jobs/createTask",
            body=body,
            headers={"Idempotency-Key": idempotency_key},
        )

    def get_task(
        self,
        task_id: str,
        *,
        timeout_seconds: float | None = None,
    ) -> dict[str, Any]:
        if not task_id:
            raise ValueError("task_id is required")
        query = urllib.parse.urlencode({"taskId": task_id})
        return self._request(
            "GET",
            f"/jobs/recordInfo?{query}",
            timeout_seconds=timeout_seconds,
        )

    def retry_task(self, task_id: str, idempotency_key: str) -> dict[str, Any]:
        if not task_id:
            raise ValueError("task_id is required")
        if not idempotency_key.strip():
            raise ValueError("Idempotency-Key is required")
        return self._request(
            "POST",
            "/jobs/retry",
            body={"taskId": task_id},
            headers={"Idempotency-Key": idempotency_key},
        )

    def wait_for_terminal(
        self,
        task_id: str,
        *,
        timeout_seconds: float | None = None,
    ) -> dict[str, Any]:
        total_timeout = self.wait_timeout_seconds if timeout_seconds is None else timeout_seconds
        deadline = self.monotonic() + total_timeout
        interval = 2.0

        while self.monotonic() < deadline:
            remaining = deadline - self.monotonic()
            task = self.get_task(
                task_id,
                timeout_seconds=min(self.request_timeout_seconds, remaining),
            )
            state = task.get("state")
            if state in TERMINAL_STATES:
                return task
            if state not in ACTIVE_STATES:
                raise SpicyApiError(f"unknown task state: {state!r}", status=200, code=200)

            delay = min(self._jitter(interval), max(0.0, deadline - self.monotonic()))
            if delay > 0:
                self.sleep(delay)
            interval = min(interval * 1.5, 15.0)

        raise SpicyTimeoutError(
            f"task {task_id} exceeded the local {total_timeout}s polling deadline; "
            "its remote state is unknown",
            task_id=task_id,
        )

    def _request(
        self,
        method: str,
        path: str,
        *,
        body: Mapping[str, Any] | None = None,
        headers: Mapping[str, str] | None = None,
        timeout_seconds: float | None = None,
    ) -> Any:
        request_headers = {
            "Accept": "application/json",
            "Authorization": f"Bearer {self.api_key}",
            **({"Content-Type": "application/json"} if body is not None else {}),
            **dict(headers or {}),
        }
        encoded_body = None if body is None else json.dumps(body, separators=(",", ":")).encode()
        timeout = self.request_timeout_seconds if timeout_seconds is None else timeout_seconds
        request_deadline = self.monotonic() + timeout

        for attempt in range(self.max_retries + 1):
            remaining = request_deadline - self.monotonic()
            if remaining <= 0:
                raise SpicyTimeoutError(
                    f"request exceeded the local {timeout}s timeout"
                )
            try:
                status, response_headers, raw = self.transport(
                    method,
                    f"{self.base_url}{path}",
                    request_headers,
                    encoded_body,
                    min(self.request_timeout_seconds, remaining),
                )
            except (TimeoutError, socket.timeout) as error:
                if attempt < self.max_retries:
                    self.sleep(min(self._retry_delay(attempt), max(0.0, request_deadline - self.monotonic())))
                    continue
                raise SpicyTimeoutError(
                    f"request exceeded the local {timeout}s timeout"
                ) from error
            except (urllib.error.URLError, OSError) as error:
                if attempt < self.max_retries:
                    self.sleep(min(self._retry_delay(attempt), max(0.0, request_deadline - self.monotonic())))
                    continue
                raise SpicyApiError(f"network request failed: {error}") from error

            try:
                envelope = json.loads(raw)
            except (json.JSONDecodeError, UnicodeDecodeError) as error:
                if attempt < self.max_retries and status in RETRYABLE_HTTP:
                    self.sleep(min(
                        self._retry_delay(attempt, response_headers.get("Retry-After")),
                        max(0.0, request_deadline - self.monotonic()),
                    ))
                    continue
                raise SpicyApiError(
                    "response was not valid JSON", status=status
                ) from error

            code = envelope.get("code") if isinstance(envelope, dict) else None
            message = envelope.get("msg") if isinstance(envelope, dict) else None
            request_id = envelope.get("request_id", "") if isinstance(envelope, dict) else ""
            retryable = status in RETRYABLE_HTTP or code in RETRYABLE_CODES
            if (not 200 <= status < 300 or code != 200) and attempt < self.max_retries and retryable:
                self.sleep(min(
                    self._retry_delay(attempt, response_headers.get("Retry-After")),
                    max(0.0, request_deadline - self.monotonic()),
                ))
                continue
            if not 200 <= status < 300 or code != 200:
                raise SpicyApiError(
                    message if isinstance(message, str) else f"request failed with HTTP {status}",
                    status=status,
                    code=code if isinstance(code, int) else None,
                    request_id=request_id if isinstance(request_id, str) else "",
                )
            if "data" not in envelope:
                raise SpicyApiError(
                    "successful envelope omitted data",
                    status=status,
                    code=code,
                    request_id=request_id,
                )
            return envelope["data"]

        raise AssertionError("retry loop exited unexpectedly")

    def _jitter(self, seconds: float) -> float:
        return seconds * (0.8 + self.random_value() * 0.4)

    def _retry_delay(self, attempt: int, retry_after: str | None = None) -> float:
        exponential = min(0.5 * (2**attempt), 8.0)
        try:
            server_delay = max(0.0, float(retry_after)) if retry_after is not None else 0.0
        except ValueError:
            server_delay = 0.0
        return max(self._jitter(exponential), server_delay)
