| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163 |
- from __future__ import annotations
- from dataclasses import dataclass
- from typing import Any
- from pymilvus import DataType, MilvusClient
- from app.embeddings import EmbeddingProvider
- from app.schemas import Evidence
- @dataclass(frozen=True)
- class IndexedDocument:
- content: str
- source: str
- doc_type: str
- chunk_index: int
- policy_type: str
- category: str
- metadata: dict[str, Any]
- class MilvusVectorStore:
- def __init__(
- self,
- uri: str,
- token: str,
- collection_name: str,
- embeddings: EmbeddingProvider,
- ) -> None:
- kwargs: dict[str, Any] = {"uri": uri}
- if token:
- kwargs["token"] = token
- self.client = MilvusClient(**kwargs)
- self.collection_name = collection_name
- self.embeddings = embeddings
- def is_ready(self) -> bool:
- self.client.list_collections()
- return True
- def create_collection(self, recreate: bool = False) -> None:
- exists = self.client.has_collection(collection_name=self.collection_name)
- # 课程数据允许显式重建;生产环境应创建新 Collection 后切换 Alias。
- if exists and recreate:
- self.client.drop_collection(collection_name=self.collection_name)
- exists = False
- if exists:
- self.client.load_collection(collection_name=self.collection_name)
- return
- schema = MilvusClient.create_schema(
- auto_id=True,
- enable_dynamic_field=False,
- )
- schema.add_field("id", DataType.INT64, is_primary=True)
- schema.add_field("content", DataType.VARCHAR, max_length=8192)
- schema.add_field("source", DataType.VARCHAR, max_length=1024)
- schema.add_field("doc_type", DataType.VARCHAR, max_length=64)
- schema.add_field("chunk_index", DataType.INT64)
- schema.add_field("policy_type", DataType.VARCHAR, max_length=64)
- schema.add_field("category", DataType.VARCHAR, max_length=64)
- schema.add_field("metadata", DataType.JSON)
- schema.add_field(
- "embedding",
- DataType.FLOAT_VECTOR,
- # Schema 维度必须与当前 Embedding Provider 完全一致。
- dim=self.embeddings.dimension,
- )
- index_params = self.client.prepare_index_params()
- index_params.add_index(
- field_name="embedding",
- index_name="embedding_hnsw",
- index_type="HNSW",
- metric_type="COSINE",
- params={"M": 16, "efConstruction": 200},
- )
- index_params.add_index(
- field_name="policy_type",
- index_name="policy_type_inverted",
- index_type="INVERTED",
- )
- self.client.create_collection(
- collection_name=self.collection_name,
- schema=schema,
- index_params=index_params,
- )
- self.client.load_collection(collection_name=self.collection_name)
- def insert_documents(self, documents: list[IndexedDocument]) -> int:
- if not documents:
- return 0
- # 先批量生成向量,再用 strict zip 防止文档与向量静默错位。
- vectors = self.embeddings.embed_documents([item.content for item in documents])
- rows = []
- for item, vector in zip(documents, vectors, strict=True):
- rows.append(
- {
- "content": item.content,
- "source": item.source,
- "doc_type": item.doc_type,
- "chunk_index": item.chunk_index,
- "policy_type": item.policy_type,
- "category": item.category,
- "metadata": item.metadata,
- "embedding": vector,
- }
- )
- self.client.insert(collection_name=self.collection_name, data=rows)
- self.client.flush(collection_name=self.collection_name)
- return len(rows)
- def search(
- self,
- query: str,
- top_k: int = 5,
- policy_type: str = "",
- ) -> list[Evidence]:
- query_vector = self.embeddings.embed_query(query)
- # Filter 只由已校验的结构化字段生成,不接收任意 Milvus 表达式。
- if policy_type and not policy_type.replace("_", "").isalnum():
- raise ValueError("policy_type 非法")
- filter_expression = f'policy_type == "{policy_type}"' if policy_type else ""
- results = self.client.search(
- collection_name=self.collection_name,
- data=[query_vector],
- anns_field="embedding",
- filter=filter_expression,
- limit=top_k,
- output_fields=[
- "content",
- "source",
- "doc_type",
- "chunk_index",
- "policy_type",
- "category",
- "metadata",
- ],
- search_params={"metric_type": "COSINE", "params": {"ef": 64}},
- )
- evidence: list[Evidence] = []
- # Store 层统一转换 Evidence,Graph 不依赖 PyMilvus 的原始返回结构。
- for hit in results[0]:
- entity = hit.get("entity", {})
- evidence.append(
- Evidence(
- source_type="milvus",
- source=entity.get("source", "unknown"),
- content=entity.get("content", ""),
- score=float(hit.get("distance", 0.0)),
- metadata={
- "id": hit.get("id"),
- "doc_type": entity.get("doc_type"),
- "chunk_index": entity.get("chunk_index"),
- "policy_type": entity.get("policy_type"),
- "category": entity.get("category"),
- **(entity.get("metadata") or {}),
- },
- )
- )
- return evidence
|