diff --git a/app/controller/add_knowledge_route.py b/app/controller/add_knowledge_route.py index 0eb8ec5..e2c6c32 100644 --- a/app/controller/add_knowledge_route.py +++ b/app/controller/add_knowledge_route.py @@ -5,6 +5,8 @@ from typing import List from fastapi import UploadFile, File, Form from langchain_core.messages import HumanMessage +from langchain_core.output_parsers import StrOutputParser +from langchain_core.prompts import ChatPromptTemplate from langgraph.prebuilt import create_react_agent from starlette.requests import Request from starlette.responses import StreamingResponse @@ -21,26 +23,24 @@ def add_knowledge_route(app): return await milvus_service.async_search(param) @app.post("/knowledge/recall/stream") - async def knowledge_search(param: KnowledgeQueryParam): + async def knowledge_search(body: dict): + param = KnowledgeQueryParam(**body.get('input')) search_response = await milvus_service.async_search(param) - agent = create_react_agent(create_react_agent( - model=create_llm(), - prompt=f"""你需要根据如下内容来回答用户的问题:{search_response.answer}""" - )) + chain = ChatPromptTemplate.from_template( + f"""你需要根据如下内容来回答用户的问题:{search_response.answer}""" + ) | create_llm() | StrOutputParser() async def generator_function(): # 先把检索结果返回前端 - yield f'data: {json.dumps({"type": "retrieve", "data": search_response}, ensure_ascii=False)}\n\n' + yield f'data: {json.dumps({"type": "retrieve", "data": search_response.model_dump()}, ensure_ascii=False)}\n\n' - async for chunk in agent.astream( + async for chunk in chain.astream( {"messages": [HumanMessage(content=param.question)]}, stream_mode="messages" ): - # 流式回复 - msg_chunk = chunk[0] - yield f'data: {json.dumps({"type": "messages", "data": msg_chunk.content}, ensure_ascii=False)}\n\n' + yield f'data: {json.dumps({"type": "messages", "data": chunk}, ensure_ascii=False)}\n\n' return StreamingResponse(generator_function(), media_type="text/event-stream")