feat: add route /knowledge/recall/stream

This commit is contained in:
martsforever
2025-09-07 22:49:27 +08:00
parent 502b2d77c6
commit a36d2c2df4
+29
View File
@@ -1,12 +1,17 @@
import asyncio import asyncio
import json
from http.client import HTTPException 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
from langchain_core.messages import HumanMessage
from langgraph.prebuilt import create_react_agent
from starlette.requests import Request from starlette.requests import Request
from starlette.responses import StreamingResponse
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.milvus_utils import milvus_service, KnowledgeQueryParam from app.utils.milvus_utils import milvus_service, KnowledgeQueryParam
@@ -15,6 +20,30 @@ def add_knowledge_route(app):
async def knowledge_search(param: KnowledgeQueryParam): async def knowledge_search(param: KnowledgeQueryParam):
return await milvus_service.async_search(param) return await milvus_service.async_search(param)
@app.post("/knowledge/recall/stream")
async def knowledge_search(param: KnowledgeQueryParam):
search_response = await milvus_service.async_search(param)
agent = create_react_agent(create_react_agent(
model=create_llm(),
prompt=f"""你需要根据如下内容来回答用户的问题:{search_response.answer}"""
))
async def generator_function():
# 先把检索结果返回前端
yield f'data: {json.dumps({"type": "retrieve", "data": search_response}, ensure_ascii=False)}\n\n'
async for chunk in agent.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'
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))