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
+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)