Files
kbqa-system/backend/src/kbqa/chat/service.py
T
2026-07-13 16:58:27 +08:00

221 lines
8.8 KiB
Python

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": 24,
"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)