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}
检索结果如下标签中的内容所示:
{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()