feat: 机器人问答接口:/knowledge/qa/stream

This commit is contained in:
martsforever
2025-11-24 21:47:54 +08:00
parent 2b2eb7e6ab
commit a3e8d739ed
+75 -3
View File
@@ -1,14 +1,15 @@
import asyncio import asyncio
import json import json
from http.client import HTTPException
from typing import List from typing import List
from fastapi import UploadFile, File, Form from fastapi import UploadFile, File, Form, HTTPException
from langchain_core.messages import HumanMessage, SystemMessage from langchain_core.messages import HumanMessage, SystemMessage
from langchain_core.output_parsers import StrOutputParser from langchain_core.output_parsers import StrOutputParser
from starlette import status
from starlette.requests import Request from starlette.requests import Request
from starlette.responses import StreamingResponse from starlette.responses import StreamingResponse, JSONResponse
from app.general.perform_general_operation import perform_general_operation
from app.utils.db_utils import AsyncSessionDep from app.utils.db_utils import AsyncSessionDep
from app.utils.knowledge_utils import knowledge_service from app.utils.knowledge_utils import knowledge_service
from app.utils.llm_utils import create_llm from app.utils.llm_utils import create_llm
@@ -45,6 +46,77 @@ def add_knowledge_route(app):
return StreamingResponse(generator_function(), media_type="text/event-stream") return StreamingResponse(generator_function(), media_type="text/event-stream")
@app.post("/knowledge/qa/stream")
async def knowledge_search(body: dict, session: AsyncSessionDep, request: Request):
# 机器人的id,用来一会查询知识库编码以及判断机器人是否已经禁用
qaId = body.get('input').get('qaId')
# 聊天历史
messages = body.get('input').get('messages')
# 用户问题
question = body.get('input').get('question')
# 查询qa_bot信息
perform_result = await perform_general_operation(
session=session,
module='knowledge_qa_bot',
data={"id": qaId},
debug_data=[],
user=request.state.user,
type='item'
)
if 'error' in perform_result:
return JSONResponse(content={"message": perform_result.get('error')}, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR)
qa_record_bot = perform_result.get('result', None)
# 找不到机器人
if not qa_record_bot:
return JSONResponse(content={"message": f"""无法找到对应问答机器人的编号:{qaId}"""}, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR)
# 机器人已经被禁用
if qa_record_bot.get('disable') == 'Y':
return JSONResponse(content={"message": "该问答机器人已经禁用"}, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR)
# 查询rel_qa_kb信息
perform_result = await perform_general_operation(
session=session,
module='rel_qa_base',
data={"all": True, "filters": [{"id": "01", "field": "qaId", "operator": "=", "value": qaId}]},
debug_data=[],
user=request.state.user,
type='list'
)
if 'error' in perform_result:
return JSONResponse(content={"message": perform_result.get('error')}, status_code=status.HTTP_500_INTERNAL_SERVER_ERROR)
rel_qa_kb_list = perform_result.get('list', [])
kb_codes = [item.get('kbCode') for item in rel_qa_kb_list]
async def generator_function():
param = KnowledgeQueryParam(kb_code=kb_codes, question=question)
search_response = await milvus_service.async_search(param)
chain = create_llm() | StrOutputParser()
# 先把检索结果返回前端
yield f'data: {json.dumps({"type": "retrieve", "data": json.dumps(search_response.model_dump())}, ensure_ascii=False)}\n\n'
external_prompt = qa_record_bot.get('prompt', '')
print("external_prompt", external_prompt, qa_record_bot)
async for chunk in chain.astream([
SystemMessage(content=f"""你需要根据如下内容来回答用户的问题:{search_response.answer},如果内容为'在提供的资料中未找到相关信息',你需要根据你自己的知识来回答用户问题。""" + external_prompt),
*messages,
]):
yield f'data: {json.dumps({"type": "messages", "data": chunk}, ensure_ascii=False)}\n\n'
return StreamingResponse(generator_function(), media_type="text/event-stream")
# @app.post("/knowledge/embed_text") # @app.post("/knowledge/embed_text")
# async def knowledge_embed_text(text_list: List[str]): # async def knowledge_embed_text(text_list: List[str]):
# id_list = await next_id(len(text_list)) # id_list = await next_id(len(text_list))