feat: 机器人问答接口:/knowledge/qa/stream
This commit is contained in:
@@ -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))
|
||||||
|
|||||||
Reference in New Issue
Block a user