test: verify real kbqa workflow
This commit is contained in:
@@ -0,0 +1,38 @@
|
||||
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
|
||||
@@ -0,0 +1,56 @@
|
||||
from sqlalchemy import event, select
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||
|
||||
from kbqa.chat.models import ChatSession, Message
|
||||
from kbqa.database import Base
|
||||
from kbqa.documents.models import Document, IndexJob
|
||||
import kbqa.models # noqa: F401
|
||||
|
||||
|
||||
async def test_sqlite_cascades_document_and_session(tmp_path) -> None:
|
||||
engine = create_async_engine(f"sqlite+aiosqlite:///{tmp_path / 'test.db'}")
|
||||
|
||||
@event.listens_for(engine.sync_engine, "connect")
|
||||
def enable_foreign_keys(connection, _record) -> None:
|
||||
connection.execute("PRAGMA foreign_keys=ON")
|
||||
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
async with factory() as session:
|
||||
document = Document(
|
||||
id="doc",
|
||||
filename="a.txt",
|
||||
stored_path="a",
|
||||
file_type="txt",
|
||||
mime_type="text/plain",
|
||||
file_size=1,
|
||||
sha256="0" * 64,
|
||||
status="pending",
|
||||
chunk_count=0,
|
||||
)
|
||||
session.add_all(
|
||||
[
|
||||
document,
|
||||
IndexJob(id="job", document_id="doc", status="queued", attempt_count=0),
|
||||
ChatSession(id="session"),
|
||||
]
|
||||
)
|
||||
await session.commit()
|
||||
session.add(
|
||||
Message(
|
||||
id="message",
|
||||
session_id="session",
|
||||
role="user",
|
||||
content="q",
|
||||
status="completed",
|
||||
order_index=0,
|
||||
)
|
||||
)
|
||||
await session.commit()
|
||||
await session.delete(document)
|
||||
await session.delete(await session.get(ChatSession, "session"))
|
||||
await session.commit()
|
||||
assert await session.scalar(select(IndexJob).where(IndexJob.id == "job")) is None
|
||||
assert await session.scalar(select(Message).where(Message.id == "message")) is None
|
||||
await engine.dispose()
|
||||
Reference in New Issue
Block a user