Files
ai-admin-server/app/utils/milvus_utils.py
T

149 lines
5.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import asyncio
from typing import List, Optional
from langchain_core.output_parsers import StrOutputParser
from langchain_core.prompts import ChatPromptTemplate
from llama_index.core import Document, VectorStoreIndex
from llama_index.core.vector_stores import MetadataFilters, MetadataFilter, FilterOperator
from llama_index.vector_stores.milvus import MilvusVectorStore
from pydantic import BaseModel, Field
from app.config.env import env
from app.utils.llm_utils import create_embeddings, create_llm
from app.utils.nltk_utils import load_nltk
class KnowledgeQueryParam(BaseModel):
question: str = Field(..., description="问题文本")
kb_code: str = Field(..., description="知识库编码")
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
self._milvus_vector_store = milvus_vector_store
self.embeddings = create_embeddings()
# 获取一个MilvusVectorStore实例,如果不存在就异步创建
async def get_vector_store(self):
if not self._milvus_vector_store:
print("create new_vector_store...")
self._milvus_vector_store = MilvusVectorStore(
uri=env.milvus_uri,
user=env.milvus_username,
password=env.milvus_password,
db_name=env.llama_index_database,
collection_name=env.llama_index_collection,
dim=env.llama_index_dimension,
embedding_field="embedding",
id_field="id",
similarity_metric="COSINE",
consistency_level="Strong",
overwrite=False,
)
return self._milvus_vector_store
# 检查Milvus连接是否正常
async def check_milvus_connection(self):
await self.async_search(KnowledgeQueryParam(question="hello", kb_code=""))
print("✅ Milvus connection successful:", f"Milvus://{env.milvus_username}:{env.milvus_password}@{env.milvus_uri}/{env.llama_index_database}/{env.llama_index_collection}")
async def async_create_index_from_documents(self, documents: List[Document]) -> VectorStoreIndex:
vector_store = await self.get_vector_store()
"""异步创建向量索引并插入文档"""
vector_index: VectorStoreIndex = VectorStoreIndex.from_vector_store(
vector_store=vector_store,
embed_model=self.embeddings,
)
# 预先加载NLTK语料库以避免多线程环境中的竞争条件
# 否则在多线程环境下,parser.get_nodes_from_documents 可能会出现报错信息:'WordListCorpusReader' object has no attribute '_LazyCorpusLoader__args'
load_nltk()
from llama_index.core.node_parser import SimpleNodeParser
# parser = HierarchicalNodeParser.from_defaults(chunk_sizes=[2048, 512, 128])
parser = SimpleNodeParser()
# nodes = await parser.aget_nodes_from_documents(documents)
# 使用同步方法并包装在 asyncio.to_thread 中
nodes = await asyncio.to_thread(parser.get_nodes_from_documents, documents)
await vector_index.ainsert_nodes(nodes)
return vector_index
# 检索milvus文档
async def async_search(self, param: KnowledgeQueryParam) -> MilvusSearchResponse:
"""异步向量搜索"""
# 创建查询引擎,设置top_k参数
vector_index = VectorStoreIndex.from_vector_store(
vector_store=await self.get_vector_store(),
embed_model=self.embeddings,
)
filters = MetadataFilters(filters=[MetadataFilter(key="parent_code", value=param.kb_code, operator=FilterOperator.EQ)])
retriever = vector_index.as_retriever(filters=filters, similarity_top_k=param.top_k)
result_nodes = await retriever.aretrieve(param.question)
chain = ChatPromptTemplate.from_template("""
你将收到一个用户问题和一组来自LlamaIndex的检索结果。你的任务是:
1. 分析检索结果内容
2. 提取与用户问题直接相关的信息片段
3. 基于这些相关信息生成准确、简洁且有帮助的回答
用户问题如下<question/>标签中的内容所示:
<question>
{question}
</question>
检索结果如下<context/>标签中的内容所示:
<context>
{context}
</context>
处理要求:
- 仅关注检索结果中与用户问题直接相关的内容
- 忽略任何无关或关联性较弱的信息
- 如果检索结果中没有相关信息,请明确回答:"在提供的资料中未找到相关信息"
- 不要添加检索结果之外的额外知识
- 回答时优先使用检索结果中的原文表述
""") | create_llm() | StrOutputParser()
chain_input = {"context": "\n".join([node.node.text for node in result_nodes]), "question": param.question}
relative_content = await chain.ainvoke(chain_input)
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(
vector_store=await self.get_vector_store(),
embed_model=self.embeddings,
)
await vector_index.adelete_ref_doc(doc_id)
return None
# 全局实例
milvus_service = MilvusService()