memory_index.py 6.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197
  1. from threading import Lock
  2. from typing import Any, cast
  3. from pymilvus import DataType, MilvusClient # type: ignore[import-untyped]
  4. from zbt.core.config import Settings
  5. from zbt.core.errors import AppError
  6. from zbt.domains.agent.memory import (
  7. CustomerMemory,
  8. CustomerMemoryHit,
  9. CustomerMemoryIndex,
  10. )
  11. from zbt.infrastructure.embedding.bge_model import load_bge_m3_model
  12. class MilvusCustomerMemoryIndex(CustomerMemoryIndex):
  13. """使用 BGE-M3 对显式长期记忆进行按客户隔离的语义检索。"""
  14. def __init__(self, settings: Settings) -> None:
  15. self._settings = settings
  16. self._collection_name = (
  17. f"{settings.milvus_collection_prefix}customer_memories"
  18. )
  19. self._client: MilvusClient | None = None
  20. self._embedding_model: Any | None = None
  21. self._client_lock = Lock()
  22. self._embedding_model_lock = Lock()
  23. self._collection_lock = Lock()
  24. self._collection_ready = False
  25. def upsert(self, memory: CustomerMemory) -> None:
  26. try:
  27. client = self._get_client()
  28. self._ensure_collection(client)
  29. client.upsert(
  30. collection_name=self._collection_name,
  31. data=[
  32. {
  33. "id": memory.id,
  34. "owner_id": memory.owner_id,
  35. "category": memory.category,
  36. "content": memory.content,
  37. "vector": self._embed([memory.content])[0],
  38. }
  39. ],
  40. )
  41. client.flush(collection_name=self._collection_name)
  42. except Exception as error:
  43. raise AppError(
  44. "CUSTOMER_MEMORY_INDEX_FAILED",
  45. "长期记忆写入失败,请检查本地Milvus和BGE-M3配置",
  46. 503,
  47. retryable=True,
  48. ) from error
  49. def search(
  50. self,
  51. *,
  52. owner_id: str,
  53. query: str,
  54. limit: int,
  55. ) -> list[CustomerMemoryHit]:
  56. try:
  57. client = self._get_client()
  58. self._ensure_collection(client)
  59. results = client.search(
  60. collection_name=self._collection_name,
  61. data=[self._embed([query])[0]],
  62. filter=f'owner_id == "{self._escape_filter(owner_id)}"',
  63. limit=limit,
  64. output_fields=["owner_id"],
  65. search_params={"metric_type": "COSINE"},
  66. )
  67. return [
  68. CustomerMemoryHit(
  69. memory_id=str(hit["id"]),
  70. score=float(hit.get("distance", 0.0)),
  71. )
  72. for hit in (results[0] if results else [])
  73. ]
  74. except Exception as error:
  75. raise AppError(
  76. "CUSTOMER_MEMORY_SEARCH_FAILED",
  77. "长期记忆检索失败,请检查本地Milvus和BGE-M3配置",
  78. 503,
  79. retryable=True,
  80. ) from error
  81. def delete(self, memory_id: str) -> None:
  82. try:
  83. client = self._get_client()
  84. self._ensure_collection(client)
  85. client.delete(
  86. collection_name=self._collection_name,
  87. ids=[memory_id],
  88. )
  89. except Exception as error:
  90. raise AppError(
  91. "CUSTOMER_MEMORY_DELETE_FAILED",
  92. "长期记忆删除失败,请检查本地Milvus配置",
  93. 503,
  94. retryable=True,
  95. ) from error
  96. def _get_client(self) -> MilvusClient:
  97. if self._client is None:
  98. with self._client_lock:
  99. if self._client is None:
  100. self._client = MilvusClient(
  101. uri=self._settings.milvus_uri,
  102. token=self._settings.milvus_token,
  103. )
  104. return self._client
  105. def _get_embedding_model(self) -> Any:
  106. if self._embedding_model is None:
  107. with self._embedding_model_lock:
  108. if self._embedding_model is None:
  109. self._embedding_model = load_bge_m3_model(
  110. self._settings.bge_m3_model_path,
  111. self._settings.bge_m3_device,
  112. )
  113. return self._embedding_model
  114. def _embed(self, texts: list[str]) -> list[list[float]]:
  115. result = self._get_embedding_model().encode(
  116. texts,
  117. return_dense=True,
  118. return_sparse=False,
  119. return_colbert_vecs=False,
  120. )
  121. vectors = result.get("dense_vecs")
  122. if vectors is None:
  123. raise AppError("EMBEDDING_EMPTY", "BGE-M3未返回稠密向量", 503)
  124. raw_vectors = vectors.tolist() if hasattr(vectors, "tolist") else list(vectors)
  125. return cast(list[list[float]], raw_vectors)
  126. def _ensure_collection(self, client: MilvusClient) -> None:
  127. if self._collection_ready:
  128. return
  129. with self._collection_lock:
  130. if self._collection_ready:
  131. return
  132. if not client.has_collection(self._collection_name):
  133. self._create_collection(client)
  134. client.load_collection(
  135. collection_name=self._collection_name,
  136. timeout=30.0,
  137. )
  138. self._collection_ready = True
  139. def _create_collection(self, client: MilvusClient) -> None:
  140. schema = MilvusClient.create_schema(
  141. auto_id=False,
  142. enable_dynamic_field=False,
  143. )
  144. schema.add_field(
  145. field_name="id",
  146. datatype=DataType.VARCHAR,
  147. is_primary=True,
  148. max_length=64,
  149. )
  150. schema.add_field(
  151. field_name="owner_id",
  152. datatype=DataType.VARCHAR,
  153. max_length=64,
  154. )
  155. schema.add_field(
  156. field_name="category",
  157. datatype=DataType.VARCHAR,
  158. max_length=32,
  159. )
  160. schema.add_field(
  161. field_name="content",
  162. datatype=DataType.VARCHAR,
  163. max_length=512,
  164. )
  165. schema.add_field(
  166. field_name="vector",
  167. datatype=DataType.FLOAT_VECTOR,
  168. dim=self._settings.embedding_dimension,
  169. )
  170. index_params = MilvusClient.prepare_index_params()
  171. index_params.add_index(
  172. field_name="vector",
  173. index_type="AUTOINDEX",
  174. metric_type="COSINE",
  175. )
  176. client.create_collection(
  177. collection_name=self._collection_name,
  178. schema=schema,
  179. index_params=index_params,
  180. )
  181. @staticmethod
  182. def _escape_filter(value: str) -> str:
  183. return value.replace("\\", "\\\\").replace('"', '\\"')