From 502b2d77c6b2e2b0c771f4fa1ce77d8cc5f8bc44 Mon Sep 17 00:00:00 2001 From: martsforever Date: Sun, 7 Sep 2025 22:49:14 +0800 Subject: [PATCH] feat: optimize MilvusSearchResponse --- app/utils/milvus_utils.py | 29 ++++++++++++++++++----------- 1 file changed, 18 insertions(+), 11 deletions(-) diff --git a/app/utils/milvus_utils.py b/app/utils/milvus_utils.py index dfc2184..5a93783 100644 --- a/app/utils/milvus_utils.py +++ b/app/utils/milvus_utils.py @@ -19,6 +19,17 @@ class KnowledgeQueryParam(BaseModel): top_k: int = Field(default=5, description="返回结果数量") +class MilvusSearchNode(BaseModel): + text: str = Field(..., description="文档内容") + metadata: dict = Field(..., description="文档元信息") + score: float = Field(..., description="相似度分数") + + +class MilvusSearchResponse(BaseModel): + answer: str = Field(..., description="根据用户问题从检索结果中提取的内容") + nodes: List[MilvusSearchNode] = Field(..., description="搜索结果") + + class MilvusService: def __init__(self): milvus_vector_store: Optional[MilvusVectorStore] = None @@ -74,7 +85,7 @@ class MilvusService: return vector_index # 检索milvus文档 - async def async_search(self, param: KnowledgeQueryParam) -> List[dict]: + async def async_search(self, param: KnowledgeQueryParam) -> MilvusSearchResponse: """异步向量搜索""" # 创建查询引擎,设置top_k参数 vector_index = VectorStoreIndex.from_vector_store( @@ -113,20 +124,16 @@ class MilvusService: """) | create_llm() | StrOutputParser() chain_input = {"context": "\n".join([node.node.text for node in result_nodes]), "question": param.question} - print("chain_input", chain_input) + relative_content = await chain.ainvoke(chain_input) - return { - "answer": relative_content, - "sources": [ - { - "text": node.node.text, - "metadata": node.node.metadata, - "score": node.score - } + return MilvusSearchResponse( + answer=relative_content, + nodes=[ + MilvusSearchNode(text=node.node.text, metadata=node.node.metadata, score=node.score) for node in result_nodes ] - } + ) async def async_delete(self, doc_id: str): vector_index = VectorStoreIndex.from_vector_store(