Код: Выделить всё
models.pyКод: Выделить всё
from sqlalchemy.engine import Connection
from sqlalchemy.ext.asyncio import AsyncConnection
class Paginator:
def __init__(
self,
conn: Union[Connection, AsyncConnection],
query: str,
params: dict = None,
batch_size: int = 10
):
self.conn = conn
self.query = query
self.params = params
self.batch_size = batch_size
self.current_offset = 0
self.total_count = None
async def _get_total_count_async(self) -> int:
"""Fetch the total count of records asynchronously."""
count_query = f"SELECT COUNT(*) FROM ({self.query}) as total"
query=text(count_query).bindparams(**(self.params or {}))
result = await self.conn.execute(query)
return result.scalar()
Код: Выделить всё
test_models.pyКод: Выделить всё
@pytest.fixture(scope='function')
async def async_session():
async_engine=create_async_engine('postgresql+asyncpg://localhost:5432/db')
async_session = sessionmaker(
expire_on_commit=False,
autocommit=False,
autoflush=False,
bind=async_engine,
class_=AsyncSession,
)
async with async_session() as session:
await session.begin()
yield session
await session.rollback()
@pytest.mark.asyncio
async def test_get_total_count_async(async_session):
# Prepare the paginator
paginator = Paginator(
conn=session,
query="SELECT * FROM test_table",
batch_size=2
)
# Perform the total count query asynchronously
total_count = await paginator._get_total_count_async()
# Assertion to verify the result
assert total_count == 0
Подробнее здесь: https://stackoverflow.com/questions/790 ... -on-pytest