from typing import Optional, List, Dict, Any from pymilvus import MilvusClient, DataType from app.client.milvus_client_manager import milvus_client from app.conf.app_config import app_config class SummaryRepository: collection_name: str = 'message_summary' def __init__(self): self.dim = app_config.milvus.embedding_size self._ensure_connection() def _ensure_connection(self): if milvus_client.client is None: milvus_client.init() def create_collection(self, dim: int = None) -> bool: """ 创建消息摘要集合 Args: dim: 向量维度,默认从配置读取 Returns: 是否成功创建(已存在返回False) """ if dim is not None: self.dim = dim if milvus_client.has_collection(self.collection_name): return False schema = milvus_client.client.create_schema( auto_id=True, enable_dynamic_field=True, description="消息摘要向量集合" ) schema.add_field(field_name="id", datatype=DataType.INT64, is_primary=True, auto_id=True) schema.add_field(field_name="msg_time", datatype=DataType.VARCHAR, max_length=100) schema.add_field(field_name="msg_type", datatype=DataType.INT8) schema.add_field(field_name="from_user", datatype=DataType.VARCHAR, max_length=150) schema.add_field(field_name="to_user", datatype=DataType.VARCHAR, max_length=150) schema.add_field(field_name="room_id", datatype=DataType.VARCHAR, max_length=150) schema.add_field(field_name="message_context", datatype=DataType.VARCHAR, max_length=65535) schema.add_field(field_name="summary", datatype=DataType.VARCHAR, max_length=65535) schema.add_field(field_name="message_dense_vector", datatype=DataType.FLOAT_VECTOR, dim=self.dim) schema.add_field(field_name="summary_dense_vector", datatype=DataType.FLOAT_VECTOR, dim=self.dim) schema.add_field(field_name="created_at", datatype=DataType.INT64) index_params = milvus_client.client.prepare_index_params() index_params.add_index( field_name="message_dense_vector", index_type="HNSW", metric_type="COSINE", params={"M": 32, "efConstruction": 300} ) index_params.add_index( field_name="summary_dense_vector", index_type="HNSW", metric_type="COSINE", params={"M": 32, "efConstruction": 300} ) milvus_client.client.create_collection( collection_name=self.collection_name, schema=schema, index_params=index_params ) return True def insert(self, msg_time: List[str], msg_type: List[int], from_user: List[str], to_user: List[str], room_id: List[str],message_context: List[str], summary: List[str], message_dense_vector: List[List[float]], summary_dense_vector: List[List[float]]) -> List[int]: """ 插入消息摘要数据 Args: msg_time: 消息时间列表 msg_type: 消息类型列表 from_user: 发送者列表 to_user: 接收者列表 room_id: 群聊ID列表 summary: 消息摘要列表 dense_vector: 向量列表 Returns: 插入的数据ID列表 """ self._ensure_connection() self.create_collection() import time created_at = [int(time.time() * 1000)] * len(msg_time) data = [] for i in range(len(msg_time)): data.append({ "msg_time": msg_time[i], "msg_type": msg_type[i], "from_user": from_user[i], "to_user": to_user[i], "room_id": room_id[i], "message_context": message_context[i], "summary": summary[i], "message_dense_vector": message_dense_vector[i], "summary_dense_vector": summary_dense_vector[i], "created_at": created_at[i] }) result = milvus_client.client.insert( collection_name=self.collection_name, data=data ) milvus_client.client.flush(self.collection_name) return result.get("ids", []) def delete(self, expr: str): """ 删除数据 Args: expr: 删除条件表达式 """ self._ensure_connection() milvus_client.client.delete( collection_name=self.collection_name, filter=expr ) milvus_client.client.flush(self.collection_name) def drop_collection(self): """删除集合""" milvus_client.drop_collection(self.collection_name) def get_collection_stats(self) -> Dict[str, Any]: """获取集合统计信息""" self._ensure_connection() try: result = milvus_client.client.get_collection_stats(collection_name=self.collection_name) return { "num_entities": result.get("row_count", 0), "collection_name": self.collection_name } except Exception as e: print(f"获取统计信息失败: {e}") return {"num_entities": 0, "collection_name": self.collection_name} def close(self): """关闭连接""" milvus_client.close()