44 lines
1.3 KiB
Python
44 lines
1.3 KiB
Python
import pytest
|
|
import asyncio
|
|
from unittest.mock import AsyncMock
|
|
from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession
|
|
from sqlalchemy.orm import sessionmaker
|
|
|
|
# 测试数据库配置
|
|
TEST_DATABASE_URL = "postgresql+asyncpg://postgres:12345@localhost:5432/test_ai_talk_callback"
|
|
|
|
@pytest.fixture(scope="session")
|
|
def event_loop():
|
|
"""创建一个事件循环实例用于测试"""
|
|
loop = asyncio.get_event_loop_policy().new_event_loop()
|
|
yield loop
|
|
loop.close()
|
|
|
|
@pytest.fixture
|
|
async def test_engine():
|
|
"""创建测试数据库引擎"""
|
|
from sqlalchemy.ext.asyncio import create_async_engine
|
|
engine = create_async_engine(TEST_DATABASE_URL, echo=False)
|
|
yield engine
|
|
await engine.dispose()
|
|
|
|
@pytest.fixture
|
|
async def test_db_session(test_engine):
|
|
"""创建测试数据库会话"""
|
|
from app.database import Base
|
|
|
|
# 创建所有表
|
|
async with test_engine.begin() as conn:
|
|
await conn.run_sync(Base.metadata.create_all)
|
|
|
|
# 创建会话
|
|
async_session = sessionmaker(
|
|
test_engine, class_=AsyncSession, expire_on_commit=False
|
|
)
|
|
|
|
async with async_session() as session:
|
|
yield session
|
|
|
|
# 清理:删除所有表
|
|
async with test_engine.begin() as conn:
|
|
await conn.run_sync(Base.metadata.drop_all) |