"""Bounded Android Local SMS Gateway driver."""

from __future__ import annotations

import time
from urllib.parse import urljoin, urlparse

import requests

from .base import ProviderConfigurationError, ProviderResult


DEFAULT_ALLOWED_ORIGINS = ("http://10.8.0.2:8082",)


class AndroidLocalDriver:
    def __init__(self, *, session=None, allowed_origins=DEFAULT_ALLOWED_ORIGINS) -> None:
        self.session = session or requests.Session()
        self.allowed_origins = tuple(origin.rstrip("/") for origin in allowed_origins)

    def validate_configuration(self, config: dict) -> None:
        base_url = str(config.get("base_url") or "").rstrip("/")
        endpoint = str(config.get("endpoint") or "/")
        parsed = urlparse(base_url)
        origin = f"{parsed.scheme}://{parsed.hostname}:{parsed.port or (80 if parsed.scheme == 'http' else 443)}"
        if parsed.username or parsed.password or parsed.query or parsed.fragment:
            raise ProviderConfigurationError("Gateway URL contains forbidden components")
        if parsed.scheme != "http" or origin not in self.allowed_origins:
            raise ProviderConfigurationError("Gateway origin is not allowlisted")
        if not endpoint.startswith("/") or urlparse(endpoint).netloc or ".." in endpoint:
            raise ProviderConfigurationError("Gateway endpoint must be a local absolute path")

    def _url(self, config: dict) -> str:
        self.validate_configuration(config)
        return urljoin(str(config["base_url"]).rstrip("/") + "/", str(config.get("endpoint") or "/").lstrip("/"))

    def health_check(self, config: dict, secret: str | None) -> ProviderResult:
        started = time.monotonic()
        try:
            response = self.session.get(self._url(config), timeout=self._timeout(config), allow_redirects=False)
            healthy = 200 <= response.status_code < 400
            return ProviderResult(healthy, "active" if healthy else "offline", str(response.status_code),
                                  None if healthy else "provider_unavailable", response.status_code >= 500,
                                  int((time.monotonic() - started) * 1000))
        except (requests.RequestException, ProviderConfigurationError) as exc:
            category = "invalid_configuration" if isinstance(exc, ProviderConfigurationError) else "network_error"
            return ProviderResult(False, "offline", error_category=category,
                                  retryable=category == "network_error", latency_ms=int((time.monotonic() - started) * 1000))

    def send(self, config: dict, secret: str | None, message: dict) -> ProviderResult:
        if not secret:
            return ProviderResult(False, "failed", error_category="credential_error")
        started = time.monotonic()
        try:
            response = self.session.post(
                self._url(config), json={"to": message["recipient"], "message": message["text"]},
                headers={"Authorization": secret, "Content-Type": "application/json"},
                timeout=self._timeout(config), allow_redirects=False,
            )
            code = str(response.status_code)
            latency = int((time.monotonic() - started) * 1000)
            if 200 <= response.status_code < 300:
                return ProviderResult(True, "accepted", code, latency_ms=latency)
            if response.status_code in (401, 403):
                return ProviderResult(False, "failed", code, "credential_error", False, latency)
            return ProviderResult(False, "failed", code, "provider_unavailable", response.status_code >= 500, latency)
        except ProviderConfigurationError:
            return ProviderResult(False, "failed", error_category="invalid_configuration")
        except (requests.Timeout, requests.ConnectionError):
            return ProviderResult(False, "failed", error_category="network_error", retryable=True,
                                  latency_ms=int((time.monotonic() - started) * 1000))
        except requests.RequestException:
            return ProviderResult(False, "failed", error_category="provider_unavailable", retryable=False,
                                  latency_ms=int((time.monotonic() - started) * 1000))

    @staticmethod
    def _timeout(config: dict) -> tuple[float, float]:
        return (float(config.get("connect_timeout", 3)), float(config.get("read_timeout", 10)))
