feat: optimize MilvusSearchResponse
This commit is contained in:
+18
-11
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user