"""McpUserKey config validation + McpKeyStore / KeyResolver unit tests."""

from datetime import datetime, timezone
from unittest.mock import AsyncMock

import pytest
from cryptography.fernet import Fernet

from navi.auth.encrypt import TokenEncryptor
from navi.mcp.config import McpServerConfig, McpUserKey
from navi.mcp.keystore import KeyResolver, McpKeyStore

import navi.api.deps as deps
from tests.conftest_factory import FakeRecord


class FakePool:
    """Pool-level asyncpg fake: captures queries, returns canned results."""

    def __init__(self):
        self.queries: list[tuple[str, tuple]] = []
        self.executes: list = []
        self.fetch_returns: list = []
        self.fetchval_returns: list = []

    @staticmethod
    def _take(queue: list):
        if not queue:
            return None
        result = queue.pop(0)
        if isinstance(result, Exception):
            raise result
        return result

    async def execute(self, query, *args):
        self.queries.append((query, args))
        return self._take(self.executes) or "OK"

    async def fetchval(self, query, *args):
        self.queries.append((query, args))
        return self._take(self.fetchval_returns)

    async def fetch(self, query, *args):
        self.queries.append((query, args))
        return self._take(self.fetch_returns) or []


class TestUserKeyValidation:
    def test_header_slot(self):
        cfg = McpServerConfig(
            transport="streamable_http",
            url="https://x",
            headers={"Authorization": "Bearer default"},
            user_key=McpUserKey(header="Authorization", prefix="Bearer "),
        )
        assert cfg.accepts_user_key

    def test_env_slot(self):
        cfg = McpServerConfig(
            transport="stdio",
            command="uvx",
            env={"API_KEY": "default"},
            user_key=McpUserKey(env="API_KEY"),
        )
        assert cfg.accepts_user_key

    def test_both_destinations_rejected(self):
        with pytest.raises(Exception):
            McpServerConfig(user_key=McpUserKey(header="H", env="E"))

    def test_no_destination_rejected(self):
        with pytest.raises(Exception):
            McpServerConfig(user_key=McpUserKey())

    def test_header_requires_http_transport(self):
        with pytest.raises(Exception):
            McpServerConfig(transport="stdio", user_key=McpUserKey(header="Authorization"))

    def test_env_requires_stdio_transport(self):
        with pytest.raises(Exception):
            McpServerConfig(transport="sse", user_key=McpUserKey(env="API_KEY"))

    def test_no_user_key_is_default(self):
        assert McpServerConfig().accepts_user_key is False


class TestWithUserKey:
    def _http_cfg(self):
        return McpServerConfig(
            transport="streamable_http",
            url="https://x",
            headers={"Authorization": "Bearer default", "X-Other": "1"},
            user_key=McpUserKey(header="Authorization", prefix="Bearer "),
        )

    def test_injects_header_and_prefix(self):
        patched = self._http_cfg().with_user_key("sk-123")
        assert patched.headers["Authorization"] == "Bearer sk-123"
        assert patched.headers["X-Other"] == "1"

    def test_no_prefix_appends_raw(self):
        cfg = self._http_cfg()
        cfg.user_key.prefix = None
        assert cfg.with_user_key("sk").headers["Authorization"] == "sk"

    def test_env_injection(self):
        cfg = McpServerConfig(
            transport="stdio",
            command="uvx",
            env={"API_KEY": "default"},
            user_key=McpUserKey(env="API_KEY"),
        )
        assert cfg.with_user_key("abc").env["API_KEY"] == "abc"

    def test_original_config_untouched(self):
        cfg = self._http_cfg()
        cfg.with_user_key("sk-123")
        assert cfg.headers["Authorization"] == "Bearer default"

    def test_raises_when_no_slot(self):
        with pytest.raises(ValueError):
            McpServerConfig().with_user_key("sk")


@pytest.fixture
def encryptor(monkeypatch):
    enc = TokenEncryptor(Fernet.generate_key().decode())
    import navi.mcp.keystore as ks

    monkeypatch.setattr(ks, "get_encryptor", lambda: enc)
    return enc


class TestMcpKeyStore:
    async def test_set_persists_encrypted(self, encryptor):
        pool = FakePool()
        store = McpKeyStore(pool)
        ts = await store.set("u1", "srv", "sk-plain")
        assert isinstance(ts, datetime)
        query, args = pool.queries[0]
        assert "ON CONFLICT (user_id, server_name)" in query
        assert args[0] == "u1" and args[1] == "srv"
        assert encryptor.decrypt(args[2]) == "sk-plain"

    async def test_get_missing_returns_none(self, encryptor):
        store = McpKeyStore(FakePool())
        assert await store.get("u1", "srv") is None

    async def test_get_decrypts(self, encryptor):
        pool = FakePool()
        store = McpKeyStore(pool)
        pool.fetchval_returns.append(encryptor.encrypt("sk-plain"))
        assert await store.get("u1", "srv") == "sk-plain"

    async def test_delete_true_when_row_returned(self, encryptor):
        pool = FakePool()
        store = McpKeyStore(pool)
        pool.fetchval_returns.append("srv")
        assert await store.delete("u1", "srv") is True

    async def test_delete_false_when_no_row(self, encryptor):
        store = McpKeyStore(FakePool())
        assert await store.delete("u1", "srv") is False

    async def test_list_servers(self, encryptor):
        pool = FakePool()
        store = McpKeyStore(pool)
        ts = datetime.now(timezone.utc)
        pool.fetch_returns.append([FakeRecord(server_name="a", updated_at=ts)])
        assert await store.list_servers("u1") == {"a": ts}


class TestKeyResolver:
    @staticmethod
    def _store(key="sk-1"):
        store = AsyncMock(spec=McpKeyStore)
        store.get = AsyncMock(return_value=key)
        return store

    async def test_resolve_delegates_to_store(self, encryptor):
        store = self._store()
        resolver = KeyResolver(lambda: store)
        assert await resolver.resolve("u1", "srv") == "sk-1"
        store.get.assert_awaited_once_with("u1", "srv")

    async def test_resolve_caches_within_ttl(self, encryptor):
        store = self._store()
        clock = [1_000.0]
        resolver = KeyResolver(lambda: store, ttl=30.0, now=lambda: clock[0])
        await resolver.resolve("u1", "srv")
        await resolver.resolve("u1", "srv")
        store.get.assert_awaited_once()

        clock[0] += 31.0  # past TTL -> re-resolve
        await resolver.resolve("u1", "srv")
        assert store.get.await_count == 2

    async def test_missing_key_resolves_none(self, encryptor):
        store = self._store(key=None)
        resolver = KeyResolver(lambda: store)
        assert await resolver.resolve("u1", "srv") is None

    async def test_db_error_resolves_none(self, encryptor):
        store = AsyncMock(spec=McpKeyStore)
        store.get = AsyncMock(side_effect=RuntimeError("db down"))
        resolver = KeyResolver(lambda: store)
        assert await resolver.resolve("u1", "srv") is None

    async def test_invalidate_forces_re_resolve(self, encryptor):
        store = self._store()
        resolver = KeyResolver(lambda: store)
        await resolver.resolve("u1", "srv")
        resolver.invalidate("u1", "srv")
        await resolver.resolve("u1", "srv")
        assert store.get.await_count == 2

    async def test_invalidate_all(self, encryptor):
        store = self._store()
        resolver = KeyResolver(lambda: store)
        await resolver.resolve("u1", "srv")
        await resolver.resolve("u2", "srv")
        resolver.invalidate_all()
        await resolver.resolve("u1", "srv")
        assert store.get.await_count == 3

    async def test_on_change_fired_with_scope(self, encryptor):
        events: list = []
        resolver = KeyResolver(lambda: self._store())
        resolver.set_on_change(lambda user, srv: events.append((user, srv)))
        resolver.invalidate("u1", "srv")
        resolver.invalidate("u2", None)
        resolver.invalidate_all()
        assert events == [("u1", "srv"), ("u2", None), (None, None)]

    async def test_error_not_cached_then_recovery(self, encryptor):
        store = AsyncMock(spec=McpKeyStore)
        store.get = AsyncMock(side_effect=[RuntimeError("down"), "sk-recovered"])
        resolver = KeyResolver(lambda: store)
        assert await resolver.resolve("u1", "srv") is None
        assert await resolver.resolve("u1", "srv") == "sk-recovered"


class TestResolverWiring:
    """The lazily built process-wide resolver, on a real (fake) asyncpg pool."""

    async def test_default_resolver_awaits_the_session_pool(self, encryptor, monkeypatch):
        """Regression: the factory used to wrap `_get_pool()`'s coroutine in
        McpKeyStore, so `resolve` always raised AttributeError, swallowed it and
        silently fell back to the default plaintext credential.
        """
        import navi.mcp.keystore as ks
        from tests.conftest_factory import FakeConnection, FakePool

        conn = FakeConnection()
        conn.enqueue(encryptor.encrypt("sk-user"))
        pool = FakePool(conn)

        class _SessionStore:
            async def _get_pool(self):
                return pool

        monkeypatch.setattr(deps, "get_session_store", lambda: _SessionStore())
        monkeypatch.setattr(ks, "_resolver", None)

        resolver = ks.get_key_resolver()
        assert await resolver.resolve("u1", "srv") == "sk-user"
        assert [c[0] for c in conn.calls] == ["fetchval"]

    async def test_missing_key_resolves_none(self, encryptor, monkeypatch):
        """The store really answers `None` (no row) — not an exception that the
        resolver turned into a fallback: the query must have been issued.
        """
        import navi.mcp.keystore as ks
        from tests.conftest_factory import FakeConnection, FakePool

        conn = FakeConnection()  # nothing enqueued → no row
        pool = FakePool(conn)

        class _SessionStore:
            async def _get_pool(self):
                return pool

        monkeypatch.setattr(deps, "get_session_store", lambda: _SessionStore())
        monkeypatch.setattr(ks, "_resolver", None)

        assert await ks.get_key_resolver().resolve("u1", "srv") is None
        assert [c[0] for c in conn.calls] == ["fetchval"]