110 lines
3.1 KiB
Python
110 lines
3.1 KiB
Python
"""Shared pytest fixtures for all tests."""
|
|||
|
|
|
||
|
|
import asyncio
|
||
|
|
from typing import AsyncGenerator, Generator
|
||
|
|
|
||
|
|
import pytest
|
||
|
|
from httpx import ASGITransport, AsyncClient
|
||
|
|
from sqlalchemy.ext.asyncio import AsyncSession, create_async_engine
|
||
|
|
from sqlalchemy.orm import sessionmaker
|
||
|
|
|
||
|
|
from src.database import Base
|
||
|
|
from src.main import app
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.fixture(scope="session")
|
||
|
|
def event_loop() -> Generator:
|
||
|
|
"""Create an event loop for the test session."""
|
||
|
|
loop = asyncio.get_event_loop_policy().new_event_loop()
|
||
|
|
yield loop
|
||
|
|
loop.close()
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.fixture(scope="function")
|
||
|
|
async def test_engine():
|
||
|
|
"""Create an in-memory SQLite database engine for testing."""
|
||
|
|
engine = create_async_engine(
|
||
|
|
"sqlite+aiosqlite:///:memory:",
|
||
|
|
echo=False,
|
||
|
|
)
|
||
|
|
|
||
|
|
# Create all tables
|
||
|
|
async with engine.begin() as conn:
|
||
|
|
await conn.run_sync(Base.metadata.create_all)
|
||
|
|
|
||
|
|
yield engine
|
||
|
|
|
||
|
|
# Cleanup
|
||
|
|
await engine.dispose()
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.fixture(scope="function")
|
||
|
|
async def async_session(test_engine) -> AsyncGenerator[AsyncSession, None]:
|
||
|
|
"""Provide an async database session for tests."""
|
||
|
|
async_session_maker_local = sessionmaker(
|
||
|
|
test_engine,
|
||
|
|
class_=AsyncSession,
|
||
|
|
expire_on_commit=False,
|
||
|
|
)
|
||
|
|
|
||
|
|
async with async_session_maker_local() as session:
|
||
|
|
yield session
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.fixture(scope="function")
|
||
|
|
async def patched_session_maker(test_engine):
|
||
|
|
"""Patch async_session_maker for resolver tests."""
|
||
|
|
test_session_maker = sessionmaker(
|
||
|
|
test_engine,
|
||
|
|
class_=AsyncSession,
|
||
|
|
expire_on_commit=False,
|
||
|
|
)
|
||
|
|
|
||
|
|
# Patch the session maker in both places
|
||
|
|
import src.database
|
||
|
|
import src.resolvers.member
|
||
|
|
|
||
|
|
original_db = src.database.async_session_maker
|
||
|
|
original_resolver = src.resolvers.member.async_session_maker
|
||
|
|
|
||
|
|
src.database.async_session_maker = test_session_maker
|
||
|
|
src.resolvers.member.async_session_maker = test_session_maker
|
||
|
|
|
||
|
|
yield test_session_maker
|
||
|
|
|
||
|
|
# Restore originals
|
||
|
|
src.database.async_session_maker = original_db
|
||
|
|
src.resolvers.member.async_session_maker = original_resolver
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.fixture(scope="function")
|
||
|
|
async def graphql_client(test_engine) -> AsyncGenerator[AsyncClient, None]:
|
||
|
|
"""Provide a test client for GraphQL API testing."""
|
||
|
|
# Override the database dependency to use test database
|
||
|
|
from src.database import async_session_maker as original_session_maker
|
||
|
|
|
||
|
|
test_session_maker = sessionmaker(
|
||
|
|
test_engine,
|
||
|
|
class_=AsyncSession,
|
||
|
|
expire_on_commit=False,
|
||
|
|
)
|
||
|
|
|
||
|
|
# Monkey patch the session maker for this test
|
||
|
|
import src.database
|
||
|
|
original = src.database.async_session_maker
|
||
|
|
src.database.async_session_maker = test_session_maker
|
||
|
|
|
||
|
|
# Also update the import in resolvers
|
||
|
|
import src.resolvers.member
|
||
|
|
src.resolvers.member.async_session_maker = test_session_maker
|
||
|
|
|
||
|
|
async with AsyncClient(
|
||
|
|
transport=ASGITransport(app=app),
|
||
|
|
base_url="http://test"
|
||
|
|
) as client:
|
||
|
|
yield client
|
||
|
|
|
||
|
|
# Restore original session maker
|
||
|
|
src.database.async_session_maker = original
|
||
|
|
src.resolvers.member.async_session_maker = original
|