"""Unit tests for the push subscription store (fake asyncpg pool)."""

import pytest

from navi.push.store import PushSubscriptionStore


class FakePool:
    """Captures executed SQL + args and returns canned rows."""

    def __init__(self):
        self.queries: list[tuple[str, tuple]] = []
        self.fetchval_result = "sub-1"
        self.fetch_rows = []

    async def execute(self, query, *args):
        self.queries.append((query, args))

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

    async def fetch(self, query, *args):
        self.queries.append((query, args))
        return self.fetch_rows


@pytest.fixture
def pool():
    return FakePool()


@pytest.fixture
def store(pool):
    return PushSubscriptionStore(pool)


async def test_upsert_uses_on_conflict_and_returns_subscription(store, pool):
    sub = await store.upsert("u1", "https://ep/1", "P256DH", "AUTH", "UA")

    assert sub.id == "sub-1"
    assert sub.user_id == "u1"
    assert sub.endpoint == "https://ep/1"
    assert sub.p256dh == "P256DH"
    assert sub.auth == "AUTH"
    query, args = pool.queries[0]
    assert "ON CONFLICT (endpoint) DO UPDATE" in query
    assert args == ("sub-id-gen", "u1", "https://ep/1", "P256DH", "AUTH", "UA") or args[2] == "https://ep/1"


async def test_delete_filters_by_endpoint(store, pool):
    await store.delete("https://ep/1")
    query, args = pool.queries[0]
    assert "DELETE FROM push_subscriptions WHERE endpoint" in query
    assert args == ("https://ep/1",)


async def test_list_for_user_maps_rows(store, pool):
    class Row(dict):
        def __getitem__(self, k):
            return dict.__getitem__(self, k)

    pool.fetch_rows = [
        Row(id="a", user_id="u1", endpoint="ep1", p256dh="p1", auth="a1", user_agent=None),
        Row(id="b", user_id="u1", endpoint="ep2", p256dh="p2", auth="a2", user_agent="UA"),
    ]
    subs = await store.list_for_user("u1")
    assert [s.endpoint for s in subs] == ["ep1", "ep2"]
    assert subs[1].user_agent == "UA"
    query, args = pool.queries[0]
    assert "WHERE user_id = $1" in query


async def test_mark_pushed_updates_timestamp(store, pool):
    await store.mark_pushed("sub-9")
    query, args = pool.queries[0]
    assert "SET last_push_at = now()" in query
    assert args == ("sub-9",)