"""Tests for tools/fal_common.py — shared FAL.ai SDK plumbing.

Covers: import_fal_client, _normalize_fal_queue_url_format,
_extract_http_status, _ManagedFalSyncClient (init + submit).
"""

import sys
import types
from unittest.mock import MagicMock, patch

import pytest

from tools.fal_common import (
    _ManagedFalSyncClient,
    _extract_http_status,
    _normalize_fal_queue_url_format,
    import_fal_client,
)

# ---------------------------------------------------------------------------
# import_fal_client
# ---------------------------------------------------------------------------

@pytest.fixture
def fake_fal_client(monkeypatch):
    module = types.ModuleType("fal_client")
    monkeypatch.setitem(sys.modules, "fal_client", module)
    return module

class TestImportFalClient:

    def test_pm_ensure_import_error_is_swallowed(self, monkeypatch, fake_fal_client):
        """If pm.ensure_import raises ImportError, the plain import decides."""
        monkeypatch.setattr(
            "pm.ensure_import",
            MagicMock(side_effect=ImportError("not installed")),
        )

        assert import_fal_client() is fake_fal_client

    def test_pm_ensure_other_exception_raises_import_error(self, monkeypatch):
        """A non-ImportError from pm (an install hint) surfaces as ImportError."""
        monkeypatch.setattr(
            "pm.ensure_import",
            MagicMock(side_effect=RuntimeError("install hint")),
        )
        with pytest.raises(ImportError, match="install hint"):
            import_fal_client()

    def test_pm_missing_is_swallowed(self, fake_fal_client):
        """If pm itself can't be imported, the plain import decides."""
        import builtins

        original_import = builtins.__import__

        def failing_import(name, *args, **kwargs):
            if name == "pm":
                raise ImportError("no module")
            return original_import(name, *args, **kwargs)

        with patch("builtins.__import__", side_effect=failing_import):
            result = import_fal_client()
            assert result is fake_fal_client

# ---------------------------------------------------------------------------
# _normalize_fal_queue_url_format
# ---------------------------------------------------------------------------

class TestNormalizeFalQueueUrlFormat:
    def test_adds_trailing_slash(self):
        assert (
            _normalize_fal_queue_url_format("https://queue.example.com")
            == "https://queue.example.com/"
        )

    def test_strips_trailing_slashes_then_adds_one(self):
        assert (
            _normalize_fal_queue_url_format("https://queue.example.com///")
            == "https://queue.example.com/"
        )

    def test_strips_whitespace(self):
        assert (
            _normalize_fal_queue_url_format("  https://queue.example.com  ")
            == "https://queue.example.com/"
        )

    def test_empty_string_raises(self):
        with pytest.raises(ValueError, match="Managed FAL queue origin is required"):
            _normalize_fal_queue_url_format("")

    def test_whitespace_only_raises(self):
        with pytest.raises(ValueError, match="Managed FAL queue origin is required"):
            _normalize_fal_queue_url_format("   ")

# ---------------------------------------------------------------------------
# _extract_http_status
# ---------------------------------------------------------------------------

class TestExtractHttpStatus:
    def test_returns_status_from_response_attribute(self):
        """httpx.HTTPStatusError exposes .response.status_code."""
        exc = MagicMock()
        exc.response = MagicMock()
        exc.response.status_code = 404
        assert _extract_http_status(exc) == 404

    def test_returns_status_from_exc_status_code(self):
        """fal_client wrappers expose .status_code directly."""
        exc = MagicMock()
        exc.response = None
        exc.status_code = 500
        assert _extract_http_status(exc) == 500

    def test_returns_none_when_no_response_and_no_status_code(self):
        exc = Exception("plain error")
        assert _extract_http_status(exc) is None

    def test_returns_none_when_response_status_code_not_int(self):
        exc = MagicMock()
        exc.response = MagicMock()
        exc.response.status_code = "not-int"
        del exc.status_code
        assert _extract_http_status(exc) is None

    def test_response_status_takes_precedence_over_exc_status(self):
        exc = MagicMock()
        exc.response = MagicMock()
        exc.response.status_code = 200
        exc.status_code = 500
        assert _extract_http_status(exc) == 200

    def test_falls_back_to_exc_status_when_response_status_not_int(self):
        exc = MagicMock()
        exc.response = MagicMock()
        exc.response.status_code = "bad"
        exc.status_code = 503
        assert _extract_http_status(exc) == 503

# ---------------------------------------------------------------------------
# _ManagedFalSyncClient — __init__
# ---------------------------------------------------------------------------

def _make_fal_client_mock(
    *,
    sync_client_class=None,
    client_module=None,
    http_client=None,
    default_timeout=120.0,
):
    """Build a mock fal_client module with all required attributes."""
    if sync_client_class is None:
        sync_client_class = MagicMock()
    if client_module is None:
        client_module = MagicMock()

    sync_client_instance = MagicMock()
    sync_client_instance._client = (
        http_client if http_client is not None else MagicMock()
    )
    sync_client_instance.default_timeout = default_timeout
    sync_client_class.return_value = sync_client_instance

    fal_client = MagicMock()
    fal_client.SyncClient = sync_client_class
    fal_client.client = client_module
    return fal_client, sync_client_instance

class TestManagedFalSyncClientInit:

    def test_init_raises_when_sync_client_missing(self):
        fal_client = MagicMock()
        fal_client.SyncClient = None
        fal_client.client = MagicMock()
        with pytest.raises(RuntimeError, match="fal_client.SyncClient is required"):
            _ManagedFalSyncClient(
                fal_client, key="k", queue_run_origin="https://q.example.com"
            )

    def test_init_raises_when_http_client_missing(self):
        """SyncClient._client is None → RuntimeError."""
        fal_client = MagicMock()
        sync_instance = MagicMock()
        sync_instance._client = None
        fal_client.SyncClient.return_value = sync_instance
        fal_client.client = MagicMock()
        with pytest.raises(RuntimeError, match="SyncClient._client is required"):
            _ManagedFalSyncClient(
                fal_client, key="k", queue_run_origin="https://q.example.com"
            )

# ---------------------------------------------------------------------------
# _ManagedFalSyncClient — submit
# ---------------------------------------------------------------------------

class TestManagedFalSyncClientSubmit:
    def _make_client(self, **kwargs):
        fal_client, sync_instance = _make_fal_client_mock(**kwargs)
        client = _ManagedFalSyncClient(
            fal_client, key="k", queue_run_origin="https://queue.example.com"
        )
        return client, sync_instance, fal_client

    def test_submit_basic(self):
        client, sync_instance, fal_client = self._make_client()
        response = MagicMock()
        response.json.return_value = {
            "request_id": "req-1",
            "response_url": "https://q.example.com/resp",
            "status_url": "https://q.example.com/status",
            "cancel_url": "https://q.example.com/cancel",
        }
        client._maybe_retry_request = MagicMock(return_value=response)
        client._raise_for_status = MagicMock()
        client._request_handle_class = MagicMock()

        result = client.submit("my-app", {"prompt": "hello"})

        client._maybe_retry_request.assert_called_once()
        call_args = client._maybe_retry_request.call_args
        assert call_args[0][0] is client._http_client
        assert call_args[0][1] == "POST"
        assert call_args[0][2] == "https://queue.example.com/my-app"
        assert call_args[1]["json"] == {"prompt": "hello"}
        assert call_args[1]["timeout"] == 120.0
        client._raise_for_status.assert_called_once_with(response)
        client._request_handle_class.assert_called_once()
        assert result is client._request_handle_class.return_value

    def test_submit_with_idempotency_key_is_not_retried(self):
        """The Nous gateway cannot replay an accepted submit reliably.

        A retry can replace the original billing/entitlement response with a
        misleading ``idempotency key ... without a reusable request handle``
        conflict, so keyed submits must make exactly one POST attempt.
        """
        http_client = MagicMock()
        response = MagicMock()
        response.json.return_value = {
            "request_id": "req-1",
            "response_url": "https://q.example.com/resp",
            "status_url": "https://q.example.com/status",
            "cancel_url": "https://q.example.com/cancel",
        }
        http_client.request.return_value = response
        client, _, _ = self._make_client(http_client=http_client)
        client._maybe_retry_request = MagicMock()
        client._raise_for_status = MagicMock()
        client._request_handle_class = MagicMock()

        client.submit(
            "my-app", {"prompt": "hello"},
            headers={"X-Idempotency-Key": "submission-1"},
        )

        http_client.request.assert_called_once_with(
            "POST", "https://queue.example.com/my-app",
            json={"prompt": "hello"}, timeout=120.0,
            headers={"X-Idempotency-Key": "submission-1"},
        )
        client._maybe_retry_request.assert_not_called()

        # Negative arm: an unkeyed submit keeps the SDK's retry ladder (transport errors, 429, ingress 5xx).
        client._maybe_retry_request = MagicMock(return_value=response)
        client.submit("my-app", {"prompt": "hello"})
        client._maybe_retry_request.assert_called_once()
        assert http_client.request.call_count == 1

    def test_submit_with_path(self):
        client, _, _ = self._make_client()
        response = MagicMock()
        response.json.return_value = {
            "request_id": "r",
            "response_url": "",
            "status_url": "",
            "cancel_url": "",
        }
        client._maybe_retry_request = MagicMock(return_value=response)
        client._raise_for_status = MagicMock()
        client._request_handle_class = MagicMock()

        client.submit("my-app", {}, path="sub/path")

        url = client._maybe_retry_request.call_args[0][2]
        assert url == "https://queue.example.com/my-app/sub/path"

    def test_submit_with_path_strips_leading_slash(self):
        client, _, _ = self._make_client()
        response = MagicMock()
        response.json.return_value = {
            "request_id": "r",
            "response_url": "",
            "status_url": "",
            "cancel_url": "",
        }
        client._maybe_retry_request = MagicMock(return_value=response)
        client._raise_for_status = MagicMock()
        client._request_handle_class = MagicMock()

        client.submit("my-app", {}, path="/leading/slash")

        url = client._maybe_retry_request.call_args[0][2]
        assert url == "https://queue.example.com/my-app/leading/slash"

    def test_submit_with_webhook_url(self):
        client, _, _ = self._make_client()
        response = MagicMock()
        response.json.return_value = {
            "request_id": "r",
            "response_url": "",
            "status_url": "",
            "cancel_url": "",
        }
        client._maybe_retry_request = MagicMock(return_value=response)
        client._raise_for_status = MagicMock()
        client._request_handle_class = MagicMock()

        client.submit("my-app", {}, webhook_url="https://hook.example.com/cb")

        url = client._maybe_retry_request.call_args[0][2]
        assert "fal_webhook=https%3A%2F%2Fhook.example.com%2Fcb" in url

    def test_submit_with_custom_headers(self):
        client, _, _ = self._make_client()
        response = MagicMock()
        response.json.return_value = {
            "request_id": "r",
            "response_url": "",
            "status_url": "",
            "cancel_url": "",
        }
        client._maybe_retry_request = MagicMock(return_value=response)
        client._raise_for_status = MagicMock()
        client._request_handle_class = MagicMock()

        client.submit("my-app", {}, headers={"X-Custom": "value"})

        headers = client._maybe_retry_request.call_args[1]["headers"]
        assert headers["X-Custom"] == "value"

    def test_submit_uses_custom_default_timeout(self):
        """SyncClient.default_timeout is used if present."""
        fal_client, sync_instance = _make_fal_client_mock(default_timeout=300.0)
        client = _ManagedFalSyncClient(
            fal_client, key="k", queue_run_origin="https://q.example.com"
        )
        response = MagicMock()
        response.json.return_value = {
            "request_id": "r",
            "response_url": "",
            "status_url": "",
            "cancel_url": "",
        }
        client._maybe_retry_request = MagicMock(return_value=response)
        client._raise_for_status = MagicMock()
        client._request_handle_class = MagicMock()

        client.submit("app", {})

        assert client._maybe_retry_request.call_args[1]["timeout"] == 300.0

    def test_submit_passes_request_handle_kwargs(self):
        client, _, _ = self._make_client()
        response = MagicMock()
        response.json.return_value = {
            "request_id": "req-123",
            "response_url": "https://q.example.com/resp",
            "status_url": "https://q.example.com/status",
            "cancel_url": "https://q.example.com/cancel",
        }
        client._maybe_retry_request = MagicMock(return_value=response)
        client._raise_for_status = MagicMock()
        client._request_handle_class = MagicMock()

        client.submit("app", {})

        kwargs = client._request_handle_class.call_args[1]
        assert kwargs["request_id"] == "req-123"
        assert kwargs["response_url"] == "https://q.example.com/resp"
        assert kwargs["status_url"] == "https://q.example.com/status"
        assert kwargs["cancel_url"] == "https://q.example.com/cancel"
        assert kwargs["client"] is client._http_client
