| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197 |
- 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('"', '\\"')
|