milvus_store.py 5.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163
  1. from __future__ import annotations
  2. from dataclasses import dataclass
  3. from typing import Any
  4. from pymilvus import DataType, MilvusClient
  5. from app.embeddings import EmbeddingProvider
  6. from app.schemas import Evidence
  7. @dataclass(frozen=True)
  8. class IndexedDocument:
  9. content: str
  10. source: str
  11. doc_type: str
  12. chunk_index: int
  13. policy_type: str
  14. category: str
  15. metadata: dict[str, Any]
  16. class MilvusVectorStore:
  17. def __init__(
  18. self,
  19. uri: str,
  20. token: str,
  21. collection_name: str,
  22. embeddings: EmbeddingProvider,
  23. ) -> None:
  24. kwargs: dict[str, Any] = {"uri": uri}
  25. if token:
  26. kwargs["token"] = token
  27. self.client = MilvusClient(**kwargs)
  28. self.collection_name = collection_name
  29. self.embeddings = embeddings
  30. def is_ready(self) -> bool:
  31. self.client.list_collections()
  32. return True
  33. def create_collection(self, recreate: bool = False) -> None:
  34. exists = self.client.has_collection(collection_name=self.collection_name)
  35. # 课程数据允许显式重建;生产环境应创建新 Collection 后切换 Alias。
  36. if exists and recreate:
  37. self.client.drop_collection(collection_name=self.collection_name)
  38. exists = False
  39. if exists:
  40. self.client.load_collection(collection_name=self.collection_name)
  41. return
  42. schema = MilvusClient.create_schema(
  43. auto_id=True,
  44. enable_dynamic_field=False,
  45. )
  46. schema.add_field("id", DataType.INT64, is_primary=True)
  47. schema.add_field("content", DataType.VARCHAR, max_length=8192)
  48. schema.add_field("source", DataType.VARCHAR, max_length=1024)
  49. schema.add_field("doc_type", DataType.VARCHAR, max_length=64)
  50. schema.add_field("chunk_index", DataType.INT64)
  51. schema.add_field("policy_type", DataType.VARCHAR, max_length=64)
  52. schema.add_field("category", DataType.VARCHAR, max_length=64)
  53. schema.add_field("metadata", DataType.JSON)
  54. schema.add_field(
  55. "embedding",
  56. DataType.FLOAT_VECTOR,
  57. # Schema 维度必须与当前 Embedding Provider 完全一致。
  58. dim=self.embeddings.dimension,
  59. )
  60. index_params = self.client.prepare_index_params()
  61. index_params.add_index(
  62. field_name="embedding",
  63. index_name="embedding_hnsw",
  64. index_type="HNSW",
  65. metric_type="COSINE",
  66. params={"M": 16, "efConstruction": 200},
  67. )
  68. index_params.add_index(
  69. field_name="policy_type",
  70. index_name="policy_type_inverted",
  71. index_type="INVERTED",
  72. )
  73. self.client.create_collection(
  74. collection_name=self.collection_name,
  75. schema=schema,
  76. index_params=index_params,
  77. )
  78. self.client.load_collection(collection_name=self.collection_name)
  79. def insert_documents(self, documents: list[IndexedDocument]) -> int:
  80. if not documents:
  81. return 0
  82. # 先批量生成向量,再用 strict zip 防止文档与向量静默错位。
  83. vectors = self.embeddings.embed_documents([item.content for item in documents])
  84. rows = []
  85. for item, vector in zip(documents, vectors, strict=True):
  86. rows.append(
  87. {
  88. "content": item.content,
  89. "source": item.source,
  90. "doc_type": item.doc_type,
  91. "chunk_index": item.chunk_index,
  92. "policy_type": item.policy_type,
  93. "category": item.category,
  94. "metadata": item.metadata,
  95. "embedding": vector,
  96. }
  97. )
  98. self.client.insert(collection_name=self.collection_name, data=rows)
  99. self.client.flush(collection_name=self.collection_name)
  100. return len(rows)
  101. def search(
  102. self,
  103. query: str,
  104. top_k: int = 5,
  105. policy_type: str = "",
  106. ) -> list[Evidence]:
  107. query_vector = self.embeddings.embed_query(query)
  108. # Filter 只由已校验的结构化字段生成,不接收任意 Milvus 表达式。
  109. if policy_type and not policy_type.replace("_", "").isalnum():
  110. raise ValueError("policy_type 非法")
  111. filter_expression = f'policy_type == "{policy_type}"' if policy_type else ""
  112. results = self.client.search(
  113. collection_name=self.collection_name,
  114. data=[query_vector],
  115. anns_field="embedding",
  116. filter=filter_expression,
  117. limit=top_k,
  118. output_fields=[
  119. "content",
  120. "source",
  121. "doc_type",
  122. "chunk_index",
  123. "policy_type",
  124. "category",
  125. "metadata",
  126. ],
  127. search_params={"metric_type": "COSINE", "params": {"ef": 64}},
  128. )
  129. evidence: list[Evidence] = []
  130. # Store 层统一转换 Evidence,Graph 不依赖 PyMilvus 的原始返回结构。
  131. for hit in results[0]:
  132. entity = hit.get("entity", {})
  133. evidence.append(
  134. Evidence(
  135. source_type="milvus",
  136. source=entity.get("source", "unknown"),
  137. content=entity.get("content", ""),
  138. score=float(hit.get("distance", 0.0)),
  139. metadata={
  140. "id": hit.get("id"),
  141. "doc_type": entity.get("doc_type"),
  142. "chunk_index": entity.get("chunk_index"),
  143. "policy_type": entity.get("policy_type"),
  144. "category": entity.get("category"),
  145. **(entity.get("metadata") or {}),
  146. },
  147. )
  148. )
  149. return evidence