feat: stream dual-agent answers over sse

This commit is contained in:
gqt
2026-07-13 11:43:23 +08:00
parent 321cf13c16
commit 84e6f1f51a
6 changed files with 269 additions and 2 deletions
+2 -1
View File
@@ -1,8 +1,9 @@
from fastapi import APIRouter from fastapi import APIRouter
from kbqa.chat.routes import router as sessions_router from kbqa.chat.routes import chat_router, router as sessions_router
from kbqa.documents.routes import router as documents_router from kbqa.documents.routes import router as documents_router
api_router = APIRouter(prefix="/api/v1") api_router = APIRouter(prefix="/api/v1")
api_router.include_router(documents_router) api_router.include_router(documents_router)
api_router.include_router(sessions_router) api_router.include_router(sessions_router)
api_router.include_router(chat_router)
+28 -1
View File
@@ -1,6 +1,7 @@
from typing import Annotated from typing import Annotated
from fastapi import APIRouter, Depends, Query, Response, status from fastapi import APIRouter, Depends, Query, Request, Response, status
from fastapi.responses import StreamingResponse
from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.ext.asyncio import AsyncSession
from kbqa.api.errors import AppError from kbqa.api.errors import AppError
@@ -9,6 +10,7 @@ from kbqa.chat.schemas import (
MessagePage, MessagePage,
MessageResponse, MessageResponse,
MessageSourceResponse, MessageSourceResponse,
ChatStreamRequest,
SessionCreate, SessionCreate,
SessionPage, SessionPage,
SessionResponse, SessionResponse,
@@ -16,6 +18,7 @@ from kbqa.chat.schemas import (
from kbqa.database import get_session from kbqa.database import get_session
router = APIRouter(prefix="/sessions", tags=["sessions"]) router = APIRouter(prefix="/sessions", tags=["sessions"])
chat_router = APIRouter(prefix="/chat", tags=["chat"])
SessionDependency = Annotated[AsyncSession, Depends(get_session)] SessionDependency = Annotated[AsyncSession, Depends(get_session)]
@@ -91,3 +94,27 @@ async def delete_session(session_id: str, session: SessionDependency) -> Respons
if not await ChatRepository(session).delete_session(session_id): if not await ChatRepository(session).delete_session(session_id):
raise AppError("SESSION_NOT_FOUND", "会话不存在", 404) raise AppError("SESSION_NOT_FOUND", "会话不存在", 404)
return Response(status_code=status.HTTP_204_NO_CONTENT) return Response(status_code=status.HTTP_204_NO_CONTENT)
@chat_router.post("/stream")
async def stream_chat(
payload: ChatStreamRequest,
request: Request,
session: SessionDependency,
) -> StreamingResponse:
if await ChatRepository(session).get_session(payload.session_id) is None:
raise AppError("SESSION_NOT_FOUND", "会话不存在", 404)
query = payload.query.strip()
if not query:
raise AppError("VALIDATION_ERROR", "问题不能为空", 422)
service = request.app.state.chat_service
service.reserve(payload.session_id)
return StreamingResponse(
service.stream(payload.session_id, query, request.state.request_id),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"X-Accel-Buffering": "no",
"X-Request-ID": request.state.request_id,
},
)
+5
View File
@@ -9,6 +9,11 @@ class SessionCreate(BaseModel):
title: str | None = Field(default=None, max_length=200) title: str | None = Field(default=None, max_length=200)
class ChatStreamRequest(BaseModel):
session_id: str
query: str = Field(min_length=1, max_length=10_000)
class SessionResponse(BaseModel): class SessionResponse(BaseModel):
model_config = ConfigDict(from_attributes=True) model_config = ConfigDict(from_attributes=True)
+220
View File
@@ -0,0 +1,220 @@
import asyncio
import logging
import re
from collections.abc import AsyncIterator, Sequence
from contextlib import suppress
from typing import Any
from langchain.messages import AIMessage, AIMessageChunk, HumanMessage
from sqlalchemy.ext.asyncio import async_sessionmaker
from kbqa.agents.context import ChatRunContext
from kbqa.agents.factory import AgentBundle, create_agent_bundle
from kbqa.api.errors import AppError
from kbqa.chat.repository import ChatRepository, MessageSourceInput
from kbqa.chat.sse import encode_event, encode_ping
from kbqa.config import Settings
from kbqa.rag.retriever import KnowledgeRetriever
INSUFFICIENT_ANSWER = "知识库中没有足够资料回答这个问题。"
logger = logging.getLogger(__name__)
class ChatService:
def __init__(
self,
settings: Settings,
session_factory: async_sessionmaker,
retriever: KnowledgeRetriever,
) -> None:
self.settings = settings
self.session_factory = session_factory
self.retriever = retriever
self._active_sessions: set[str] = set()
self._agents: AgentBundle | None = None
def reserve(self, session_id: str) -> None:
if session_id in self._active_sessions:
raise AppError("SESSION_BUSY", "当前会话正在生成回答", 409)
self._active_sessions.add(session_id)
def _get_agents(self) -> AgentBundle:
if self._agents is None:
self._agents = create_agent_bundle(self.settings)
return self._agents
async def stream(
self,
session_id: str,
query: str,
request_id: str,
) -> AsyncIterator[bytes]:
assistant_id: str | None = None
graph_stream = None
try:
async with self.session_factory() as session:
repository = ChatRepository(session)
if await repository.get_session(session_id) is None:
yield encode_event(
"error",
{
"code": "SESSION_NOT_FOUND",
"message": "会话不存在",
"retryable": False,
},
)
return
_, assistant = await repository.begin_turn(session_id, query)
assistant_id = assistant.id
history = await repository.load_recent_messages(
session_id, self.settings.history_message_limit
)
messages = [
HumanMessage(content=message.content)
if message.role == "user"
else AIMessage(content=message.content)
for message in history
if message.id != assistant_id
]
sources: list[dict[str, Any]] = []
answer_parts: list[str] = []
answering = False
no_evidence = False
yield encode_event("status", {"phase": "researching"})
graph_stream = self._get_agents().main_agent.astream(
{"messages": messages},
context=ChatRunContext(
retriever=self.retriever,
request_id=request_id,
session_id=session_id,
assistant_message_id=assistant_id,
),
config={
"recursion_limit": 12,
"tags": ["main-agent"],
"metadata": {
"request_id": request_id,
"session_id": session_id,
"assistant_message_id": assistant_id,
},
},
stream_mode=["messages", "custom", "updates"],
version="v2",
)
iterator = graph_stream.__aiter__()
pending = asyncio.create_task(anext(iterator))
while True:
done, _ = await asyncio.wait({pending}, timeout=15.0)
if not done:
yield encode_ping()
continue
try:
chunk = pending.result()
except StopAsyncIteration:
break
pending = asyncio.create_task(anext(iterator))
chunk_type = chunk.get("type")
data = chunk.get("data")
if chunk_type == "custom" and isinstance(data, dict):
if data.get("kind") == "no_evidence":
no_evidence = True
pending.cancel()
with suppress(asyncio.CancelledError):
await pending
await graph_stream.aclose()
graph_stream = None
break
if data.get("kind") == "sources":
sources = list(data.get("sources", []))
yield encode_event("sources", {"sources": sources})
yield encode_event("status", {"phase": "answering"})
answering = True
continue
if chunk_type != "messages" or not answering:
continue
token, _metadata = data
if not isinstance(token, AIMessageChunk):
continue
text = self._message_text(token.content)
if text:
answer_parts.append(text)
yield encode_event("token", {"content": text})
if no_evidence:
answer = INSUFFICIENT_ANSWER
yield encode_event("status", {"phase": "answering"})
yield encode_event("token", {"content": answer})
else:
answer = "".join(answer_parts).strip()
self._validate_citations(answer, sources)
source_inputs = [
MessageSourceInput(
document_id=str(source["document_id"]),
chunk_id=str(source["chunk_id"]),
citation_label=str(source["label"]),
score=float(source["score"]),
)
for source in sources
]
async with self.session_factory() as session:
await ChatRepository(session).complete_answer(assistant_id, answer, source_inputs)
yield encode_event("done", {"message_id": assistant_id})
except asyncio.CancelledError:
if assistant_id is not None:
await self._mark_failed(assistant_id, "CLIENT_DISCONNECTED")
raise
except AppError as exc:
if assistant_id is not None:
await self._mark_failed(assistant_id, exc.code)
yield encode_event(
"error",
{"code": exc.code, "message": exc.message, "retryable": exc.retryable},
)
except Exception:
logger.exception("Chat stream failed request_id=%s", request_id)
if assistant_id is not None:
await self._mark_failed(assistant_id, "MODEL_UPSTREAM_ERROR")
yield encode_event(
"error",
{
"code": "MODEL_UPSTREAM_ERROR",
"message": "模型调用失败,请稍后重试",
"retryable": True,
},
)
finally:
if graph_stream is not None:
await graph_stream.aclose()
self._active_sessions.discard(session_id)
async def _mark_failed(self, message_id: str, error_code: str) -> None:
async with self.session_factory() as session:
await ChatRepository(session).fail_answer(message_id, error_code)
@staticmethod
def _message_text(content: Any) -> str:
if isinstance(content, str):
return content
if isinstance(content, Sequence):
parts: list[str] = []
for block in content:
if isinstance(block, str):
parts.append(block)
elif isinstance(block, dict) and block.get("type") == "text":
parts.append(str(block.get("text", "")))
return "".join(parts)
return ""
@staticmethod
def _validate_citations(answer: str, sources: list[dict[str, Any]]) -> None:
if not answer:
raise AppError("EMPTY_MODEL_ANSWER", "模型没有返回有效回答", 502, retryable=True)
allowed = {str(source["label"]) for source in sources}
used = set(re.findall(r"\[(S\d+)]", answer))
if sources and not used:
raise AppError("MISSING_CITATION", "模型回答缺少知识库引用", 502, retryable=True)
if not used.issubset(allowed):
raise AppError("INVALID_CITATION", "模型回答包含无效引用", 502, retryable=True)
+11
View File
@@ -0,0 +1,11 @@
import json
from collections.abc import Mapping
def encode_event(event: str, data: Mapping[str, object]) -> bytes:
payload = json.dumps(data, ensure_ascii=False, separators=(",", ":"))
return f"event: {event}\ndata: {payload}\n\n".encode()
def encode_ping() -> bytes:
return b": ping\n\n"
+3
View File
@@ -7,6 +7,7 @@ from fastapi import FastAPI
from pymilvus import MilvusClient from pymilvus import MilvusClient
from kbqa.config import get_settings from kbqa.config import get_settings
from kbqa.chat.service import ChatService
from kbqa.database import Database from kbqa.database import Database
from kbqa.documents.service import DocumentService from kbqa.documents.service import DocumentService
from kbqa.indexing.worker import IndexWorker from kbqa.indexing.worker import IndexWorker
@@ -53,6 +54,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
vector_store, vector_store,
vector_mutation_lock, vector_mutation_lock,
) )
chat_service = ChatService(settings, database.session_factory, retriever)
app.state.settings = settings app.state.settings = settings
app.state.database = database app.state.database = database
app.state.milvus_client = milvus_client app.state.milvus_client = milvus_client
@@ -61,6 +63,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
app.state.retriever = retriever app.state.retriever = retriever
app.state.index_worker = worker app.state.index_worker = worker
app.state.document_service = document_service app.state.document_service = document_service
app.state.chat_service = chat_service
await worker.start() await worker.start()
await document_service.recover_deletions() await document_service.recover_deletions()
try: try: