from collections.abc import AsyncGenerator, Generator from typing import Any import pytest from httpx import ASGITransport, AsyncClient from sqlalchemy.ext.asyncio import ( AsyncEngine, AsyncSession, async_sessionmaker, create_async_engine, ) from sqlalchemy.pool import NullPool from app.config import settings from app.db import get_db_session from app.main import app from app.models import Base TEST_DATABASE_URL = settings.database_url if "/headquarter_test" not in TEST_DATABASE_URL: TEST_DATABASE_URL = TEST_DATABASE_URL.replace("/headquarter", "/headquarter_test") if TEST_DATABASE_URL.startswith("postgresql://"): TEST_DATABASE_URL = TEST_DATABASE_URL.replace("postgresql://", "postgresql+asyncpg://", 1) @pytest.fixture(scope="session") def event_loop() -> Generator[Any, None, None]: import asyncio loop = asyncio.get_event_loop_policy().new_event_loop() yield loop loop.close() @pytest.fixture(scope="session") async def db_engine() -> AsyncGenerator[AsyncEngine, None]: engine = create_async_engine(TEST_DATABASE_URL, echo=False, poolclass=NullPool) async with engine.begin() as conn: await conn.run_sync(Base.metadata.create_all) yield engine async with engine.begin() as conn: await conn.run_sync(Base.metadata.drop_all) await engine.dispose() @pytest.fixture async def db_session( db_engine: AsyncEngine, ) -> AsyncGenerator[async_sessionmaker[AsyncSession], None]: async with db_engine.connect() as connection: trans = await connection.begin_nested() testing_session_local = async_sessionmaker( connection, class_=AsyncSession, expire_on_commit=False ) async def override_get_db() -> AsyncGenerator[AsyncSession, None]: async with testing_session_local() as session: yield session app.dependency_overrides[get_db_session] = override_get_db original_db_url = settings.database_url settings.database_url = TEST_DATABASE_URL yield testing_session_local settings.database_url = original_db_url app.dependency_overrides.pop(get_db_session, None) await trans.rollback() @pytest.fixture async def client( db_session: async_sessionmaker[AsyncSession], ) -> AsyncGenerator[AsyncClient, None]: async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as ac: yield ac @pytest.fixture async def auth_client( client: AsyncClient, ) -> AsyncGenerator[AsyncClient, None]: original_debug = settings.debug original_bypass = settings.auth_dev_bypass settings.debug = True settings.auth_dev_bypass = True yield client settings.debug = original_debug settings.auth_dev_bypass = original_bypass