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