100 lines
3.3 KiB
Python
100 lines
3.3 KiB
Python
import asyncio
|
||
from typing import List, Optional
|
||
|
||
from llama_index.core import Document, VectorStoreIndex
|
||
from llama_index.vector_stores.milvus import MilvusVectorStore
|
||
|
||
from app.config.env import env
|
||
from app.utils.llm_utils import create_embeddings, create_llama_index_llm
|
||
|
||
|
||
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("hello")
|
||
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,
|
||
)
|
||
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, query: str, top_k: int = 5) -> List[dict]:
|
||
"""异步向量搜索"""
|
||
# 创建查询引擎,设置top_k参数
|
||
vector_index = VectorStoreIndex.from_vector_store(
|
||
vector_store=await self.get_vector_store(),
|
||
embed_model=self.embeddings,
|
||
)
|
||
query_engine = vector_index.as_query_engine(
|
||
streaming=False,
|
||
similarity_top_k=top_k, # 这里正确使用top_k参数
|
||
llm=create_llama_index_llm()
|
||
)
|
||
|
||
# 异步执行查询
|
||
response = await query_engine.aquery(query)
|
||
|
||
return {
|
||
"answer": str(response),
|
||
"sources": [
|
||
{
|
||
"text": node.node.text,
|
||
"metadata": node.node.metadata,
|
||
"score": node.score
|
||
}
|
||
for node in response.source_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()
|