diff --git a/app/controller/add_knowledge_route.py b/app/controller/add_knowledge_route.py index 02a4da3..6745362 100644 --- a/app/controller/add_knowledge_route.py +++ b/app/controller/add_knowledge_route.py @@ -1,14 +1,15 @@ import asyncio import json -from http.client import HTTPException 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.output_parsers import StrOutputParser +from starlette import status 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.knowledge_utils import knowledge_service 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") + @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") # async def knowledge_embed_text(text_list: List[str]): # id_list = await next_id(len(text_list))