39 lines
1.6 KiB
Python
39 lines
1.6 KiB
Python
from collections.abc import AsyncIterator
|
|
|
|
from fastapi import FastAPI
|
|
from httpx import ASGITransport, AsyncClient
|
|
from sqlalchemy.ext.asyncio import AsyncEngine, AsyncSession, async_sessionmaker
|
|
|
|
from kbqa.api.errors import register_exception_handlers
|
|
from kbqa.api.middleware import RequestIDMiddleware
|
|
from kbqa.chat.routes import router
|
|
from kbqa.database import get_session
|
|
|
|
|
|
async def test_session_api_persists_and_deletes(sqlite_engine: AsyncEngine) -> None:
|
|
factory = async_sessionmaker(sqlite_engine, expire_on_commit=False)
|
|
app = FastAPI()
|
|
app.add_middleware(RequestIDMiddleware)
|
|
register_exception_handlers(app)
|
|
app.include_router(router, prefix="/api/v1")
|
|
|
|
async def session_override() -> AsyncIterator[AsyncSession]:
|
|
async with factory() as session:
|
|
yield session
|
|
|
|
app.dependency_overrides[get_session] = session_override
|
|
async with AsyncClient(transport=ASGITransport(app=app), base_url="http://test") as client:
|
|
created = await client.post("/api/v1/sessions", json={"title": "持久化测试"})
|
|
assert created.status_code == 201
|
|
session_id = created.json()["id"]
|
|
|
|
listed = await client.get("/api/v1/sessions")
|
|
assert listed.json()["items"][0]["id"] == session_id
|
|
messages = await client.get(f"/api/v1/sessions/{session_id}/messages")
|
|
assert messages.json()["items"] == []
|
|
|
|
deleted = await client.delete(f"/api/v1/sessions/{session_id}")
|
|
assert deleted.status_code == 204
|
|
missing = await client.get(f"/api/v1/sessions/{session_id}/messages")
|
|
assert missing.status_code == 404
|