diff --git a/tests/api/helpers.py b/tests/api/helpers.py index e4c3eef..8a78df0 100644 --- a/tests/api/helpers.py +++ b/tests/api/helpers.py @@ -1,7 +1,12 @@ +import contextlib import typing import uuid +from collections.abc import Iterator +import sqlalchemy as sa from httpx import AsyncClient +from sqlalchemy import event +from sqlalchemy.ext.asyncio import AsyncConnection, AsyncSession async def register(client: AsyncClient, username: str) -> int: @@ -33,3 +38,20 @@ async def send(client: AsyncClient, chat_id: int, text: str, key: uuid.UUID | No json={"idempotency_key": str(key or uuid.uuid4()), "text": text}, ) return response.json() + + +@contextlib.contextmanager +def count_statements(session: AsyncSession) -> Iterator[list[str]]: + """Collect every SQL statement sent on the test connection behind ``session`` while the block runs.""" + connection: typing.Final = session.bind + assert isinstance(connection, AsyncConnection) + statements: typing.Final[list[str]] = [] + + def record(_conn: sa.Connection, _cursor: object, statement: str, *_: object) -> None: + statements.append(statement) + + event.listen(connection.sync_connection, "before_cursor_execute", record) + try: + yield statements + finally: + event.remove(connection.sync_connection, "before_cursor_execute", record) diff --git a/tests/api/test_chat_listing_api.py b/tests/api/test_chat_listing_api.py index 7db98c6..081b986 100644 --- a/tests/api/test_chat_listing_api.py +++ b/tests/api/test_chat_listing_api.py @@ -1,6 +1,8 @@ import pytest from httpx import AsyncClient +from sqlalchemy.ext.asyncio import AsyncSession +from tests.api.helpers import count_statements as _count_statements from tests.api.helpers import create_direct_chat as _create_direct_chat from tests.api.helpers import login as _login from tests.api.helpers import register as _register @@ -37,6 +39,31 @@ async def test_listing_populates_unread_count_and_last_message(client: AsyncClie assert item["last_message"]["text"] == "two" +async def test_listing_runs_the_same_number_of_statements_for_one_chat_as_for_many( + client: AsyncClient, db_session: AsyncSession +) -> None: + one_chat_id, _ = await _create_direct_chat(client) + await _send(client, one_chat_id, "hello") + partner_ids = [await _register(client, f"partner{index}") for index in range(5)] + await _register(client, "carol") + for partner_id in partner_ids: + chat = await client.post("/api/chats/", json={"chat_type": "direct", "member_ids": [partner_id]}) + await _send(client, chat.json()["id"], "hello") + + with _count_statements(db_session) as many_chats_statements: + many_chats_response = await client.get("/api/chats/") + await _login(client, "alice") + with _count_statements(db_session) as one_chat_statements: + one_chat_response = await client.get("/api/chats/") + + many_chats_items = many_chats_response.json()["items"] + one_chat_items = one_chat_response.json()["items"] + assert len(many_chats_items) == 5 + assert len(one_chat_items) == 1 + assert all(item["last_message"] is not None for item in many_chats_items + one_chat_items) + assert len(many_chats_statements) == len(one_chat_statements), many_chats_statements + + @pytest.mark.usefixtures("db_session") async def test_chat_with_no_messages_has_null_last_message_and_zero_unread(client: AsyncClient) -> None: await _create_direct_chat(client) diff --git a/tests/api/test_statement_counter.py b/tests/api/test_statement_counter.py new file mode 100644 index 0000000..f14a40e --- /dev/null +++ b/tests/api/test_statement_counter.py @@ -0,0 +1,19 @@ +import sqlalchemy as sa +from sqlalchemy.ext.asyncio import AsyncSession + +from app.database.resources import create_database_engine +from tests.api.helpers import count_statements as _count_statements + + +async def test_count_statements_ignores_other_connections(db_session: AsyncSession) -> None: + engine = create_database_engine() + try: + with _count_statements(db_session) as statements: + async with engine.connect() as other_connection: + await other_connection.execute(sa.text("SELECT 1")) + await db_session.execute(sa.text("SELECT 2")) + finally: + await engine.dispose() + + assert "SELECT 1" not in statements + assert "SELECT 2" in statements