from threading import Lock from typing import Any, cast from pymilvus import DataType, MilvusClient # type: ignore[import-untyped] from zbt.core.config import Settings from zbt.core.errors import AppError from zbt.domains.agent.memory import ( CustomerMemory, CustomerMemoryHit, CustomerMemoryIndex, ) from zbt.infrastructure.embedding.bge_model import load_bge_m3_model class MilvusCustomerMemoryIndex(CustomerMemoryIndex): """使用 BGE-M3 对显式长期记忆进行按客户隔离的语义检索。""" def __init__(self, settings: Settings) -> None: self._settings = settings self._collection_name = ( f"{settings.milvus_collection_prefix}customer_memories" ) self._client: MilvusClient | None = None self._embedding_model: Any | None = None self._client_lock = Lock() self._embedding_model_lock = Lock() self._collection_lock = Lock() self._collection_ready = False def upsert(self, memory: CustomerMemory) -> None: try: client = self._get_client() self._ensure_collection(client) client.upsert( collection_name=self._collection_name, data=[ { "id": memory.id, "owner_id": memory.owner_id, "category": memory.category, "content": memory.content, "vector": self._embed([memory.content])[0], } ], ) client.flush(collection_name=self._collection_name) except Exception as error: raise AppError( "CUSTOMER_MEMORY_INDEX_FAILED", "长期记忆写入失败,请检查本地Milvus和BGE-M3配置", 503, retryable=True, ) from error def search( self, *, owner_id: str, query: str, limit: int, ) -> list[CustomerMemoryHit]: try: client = self._get_client() self._ensure_collection(client) results = client.search( collection_name=self._collection_name, data=[self._embed([query])[0]], filter=f'owner_id == "{self._escape_filter(owner_id)}"', limit=limit, output_fields=["owner_id"], search_params={"metric_type": "COSINE"}, ) return [ CustomerMemoryHit( memory_id=str(hit["id"]), score=float(hit.get("distance", 0.0)), ) for hit in (results[0] if results else []) ] except Exception as error: raise AppError( "CUSTOMER_MEMORY_SEARCH_FAILED", "长期记忆检索失败,请检查本地Milvus和BGE-M3配置", 503, retryable=True, ) from error def delete(self, memory_id: str) -> None: try: client = self._get_client() self._ensure_collection(client) client.delete( collection_name=self._collection_name, ids=[memory_id], ) except Exception as error: raise AppError( "CUSTOMER_MEMORY_DELETE_FAILED", "长期记忆删除失败,请检查本地Milvus配置", 503, retryable=True, ) from error def _get_client(self) -> MilvusClient: if self._client is None: with self._client_lock: if self._client is None: self._client = MilvusClient( uri=self._settings.milvus_uri, token=self._settings.milvus_token, ) return self._client def _get_embedding_model(self) -> Any: if self._embedding_model is None: with self._embedding_model_lock: if self._embedding_model is None: self._embedding_model = load_bge_m3_model( self._settings.bge_m3_model_path, self._settings.bge_m3_device, ) return self._embedding_model def _embed(self, texts: list[str]) -> list[list[float]]: result = self._get_embedding_model().encode( texts, return_dense=True, return_sparse=False, return_colbert_vecs=False, ) vectors = result.get("dense_vecs") if vectors is None: raise AppError("EMBEDDING_EMPTY", "BGE-M3未返回稠密向量", 503) raw_vectors = vectors.tolist() if hasattr(vectors, "tolist") else list(vectors) return cast(list[list[float]], raw_vectors) def _ensure_collection(self, client: MilvusClient) -> None: if self._collection_ready: return with self._collection_lock: if self._collection_ready: return if not client.has_collection(self._collection_name): self._create_collection(client) client.load_collection( collection_name=self._collection_name, timeout=30.0, ) self._collection_ready = True def _create_collection(self, client: MilvusClient) -> None: schema = MilvusClient.create_schema( auto_id=False, enable_dynamic_field=False, ) schema.add_field( field_name="id", datatype=DataType.VARCHAR, is_primary=True, max_length=64, ) schema.add_field( field_name="owner_id", datatype=DataType.VARCHAR, max_length=64, ) schema.add_field( field_name="category", datatype=DataType.VARCHAR, max_length=32, ) schema.add_field( field_name="content", datatype=DataType.VARCHAR, max_length=512, ) schema.add_field( field_name="vector", datatype=DataType.FLOAT_VECTOR, dim=self._settings.embedding_dimension, ) index_params = MilvusClient.prepare_index_params() index_params.add_index( field_name="vector", index_type="AUTOINDEX", metric_type="COSINE", ) client.create_collection( collection_name=self._collection_name, schema=schema, index_params=index_params, ) @staticmethod def _escape_filter(value: str) -> str: return value.replace("\\", "\\\\").replace('"', '\\"')