优化代码:使用Milvus连接池

为了优化系统资源使用,我们可以使用Milvus的连接池功能。以下是优化后的代码,主要改进点包括:

  1. 使用全局连接池管理Milvus连接
  2. 添加连接池配置参数
  3. 实现连接池的初始化和关闭
  4. 优化检索流程
from langchain_community.vectorstores import Milvus
from langchain.embeddings.base import Embeddings
from langchain.chains import RetrievalQA
from langchain.llms import Ollama
import requests
from typing import List, Dict, Any
from pymilvus import connections
import threading

# 全局连接池管理
class MilvusConnectionPool:
    _instance = None
    _lock = threading.Lock()
    
    def __new__(cls):
        if cls._instance is None:
            with cls._lock:
                if cls._instance is None:
                    cls._instance = super().__new__(cls)
                    cls._instance._connections = {}
        return cls._instance
    
    def get_connection(self, alias: str, host: str, port: str, **kwargs):
        """获取或创建Milvus连接"""
        if alias not in self._connections:
            connections.connect(alias=alias, host=host, port=port, **kwargs)
            self._connections[alias] = True
        return connections.get_connection(alias)
    
    def close_all(self):
        """关闭所有连接"""
        for alias in list(self._connections.keys()):
            try:
                connections.disconnect(alias)
                del self._connections[alias]
            except Exception as e:
                print(f"Error disconnecting {alias}: {str(e)}")

class NomicEmbedText(Embeddings):
    def __init__(self, base_url="http://10.80.0.230:11434"):
        self.base_url = base_url

    def embed_query(self, text):
        response = requests.post(
            f"{self.base_url}/api/embeddings",
            json={"model": "nomic-embed-text:v1.5", "prompt": text}
        )
        if response.status_code != 200:
            raise ValueError(f"嵌入生成失败: {response.text}")
        return response.json()["embedding"]

    def embed_documents(self, texts):
        embeddings = []
        for text in texts:
            embedding = self.embed_query(text)
            embeddings.append(embedding)
        return embeddings

def load_milvus_vector_store(collection_name="oceanx_ecm"):
    """加载Milvus向量存储,使用连接池"""
    # 初始化连接池
    connection_pool = MilvusConnectionPool()
    connection_pool.get_connection(
        alias="default",
        host="10.80.0.230",
        port="19530",
        pool_size=10,  # 连接池大小
        max_retries=3,  # 最大重试次数
        timeout=30  # 超时时间(秒)
    )
    
    embeddings = NomicEmbedText()
    vector_store = Milvus(
        collection_name=collection_name,
        connection_args={"host": "10.80.0.230", "port": "19530"},
        embedding_function=embeddings,
        auto_id=True,
        consistency_level="Strong",  # 一致性级别
        search_params={"nprobe": 16},  # 搜索参数
        index_params={
            "metric_type": "L2",
            "index_type": "IVF_FLAT",
            "params": {"nlist": 1024}
        }
    )
    return vector_store

def get_readable_doc_ids(user_id: int) -> List[str]:
    """获取用户有权限访问的文档ID"""
    response = requests.get(f"http://your-permission-api/users/{user_id}/readable_docs")
    if response.status_code == 200:
        return response.json()["data"]
    return []

def setup_llm_model():
    """设置LLM模型"""
    return Ollama(
        model="deepseek-r1:1.5b", 
        base_url="http://10.80.0.230:11434",
        temperature=0.7,
        top_p=0.9,
        timeout=60
    )

def retrieve_relevant_documents(vector_store: Milvus, question: str, user_id: int) -> List[Dict[str, Any]]:
    """检索相关文档,使用连接池"""
    try:
        # 第一次检索获取所有可能相关的文档
        retriever = vector_store.as_retriever(search_kwargs={"k": 10})
        all_retrieved_docs = retriever.get_relevant_documents(question)

        # 获取用户有权限访问的文档ID
        readable_doc_ids = get_readable_doc_ids(user_id)

        # 根据权限过滤文档
        if readable_doc_ids:
            filtered_docs = [
                {
                    "content": doc.page_content,
                    "metadata": doc.metadata
                }
                for doc in all_retrieved_docs
                if doc.metadata.get("doc_id") in readable_doc_ids
            ]
        else:
            filtered_docs = [
                {
                    "content": doc.page_content,
                    "metadata": doc.metadata
                }
                for doc in all_retrieved_docs
            ]

        return filtered_docs
    except Exception as e:
        print(f"检索文档时出错: {str(e)}")
        return []

def generate_answer(llm: Ollama, question: str, documents: List[Dict[str, Any]]) -> Dict[str, Any]:
    """生成回答"""
    if not documents:
        return {
            "answer": "没有找到您有权限访问的相关文档",
            "sources": []
        }

    # 构建上下文
    context = "\n\n".join([doc["content"] for doc in documents])

    # 构建提示词
    prompt = f"""
    基于以下提供的上下文信息,请回答用户的问题。
    如果上下文信息不足以回答问题,请回答"根据现有信息无法回答该问题"。

    上下文:
    {context}

    问题: {question}
    """

    # 调用LLM生成回答
    try:
        answer = llm(prompt)
    except Exception as e:
        answer = f"生成回答时出错: {str(e)}"

    # 准备响应
    response = {
        "answer": answer,
        "sources": [
            {
                "doc_id": doc["metadata"].get("doc_id", "未知"),
                "content": doc["content"][:200] + "..."  # 只显示部分内容
            }
            for doc in documents
        ]
    }

    return response

def main():
    try:
        # 初始化组件
        vector_store = load_milvus_vector_store()
        llm = setup_llm_model()
        user_id = 123

        # 用户问题
        question = "关于公司法律条款的解释"

        # 1. 检索相关文档
        relevant_docs = retrieve_relevant_documents(vector_store, question, user_id)

        # 2. 生成回答
        response = generate_answer(llm, question, relevant_docs)

        print("回答:", response["answer"])
        print("\n来源文档:")
        for source in response["sources"]:
            print(f"- 文档ID: {source['doc_id']}")
            print(f"  内容摘要: {source['content']}\n")
            
    finally:
        # 程序结束时关闭所有连接
        MilvusConnectionPool().close_all()

if __name__ == "__main__":
    main()

主要优化点说明:

  1. 连接池管理

    • 实现了MilvusConnectionPool单例类来管理所有Milvus连接
    • 使用线程锁确保线程安全
    • 提供统一的连接获取和关闭接口
  2. 连接参数优化

    • 添加了连接池大小配置(pool_size=10)
    • 设置了连接超时(timeout=30)和重试机制(max_retries=3)
    • 确保程序退出时关闭所有连接
  3. Milvus配置优化

    • 添加了索引参数(index_params)和搜索参数(search_params)
    • 设置了强一致性级别(consistency_level="Strong")
    • 启用了自动ID生成(auto_id=True)
  4. 错误处理增强

    • 在各个关键步骤添加了异常处理
    • 确保资源释放
  5. 性能优化

    • 减少了不必要的连接创建
    • 复用现有连接提高性能

这种实现方式可以显著减少系统资源消耗,特别是在高并发场景下,避免了频繁创建和销毁连接的开销。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐