"""幂等创建知识检索用的三个 Milvus 集合(Task 6)。 用法:: .\\.venv\\Scripts\\python.exe tools\\setup_milvus_knowledge_collections.py 安全口径(共享 Milvus 实例,实例里还有别的项目在用的集合): 1. **幂等**:集合已存在则**直接跳过**,不重建、不覆盖、不清数据; 2. **不覆盖不同结构**:同名集合已存在但字段/维度与本脚本定义不同时,**停下来报告** 并以非零退出码结束 —— 宁可人工确认,也不动别人的数据; 3. 只处理 `ALLOWED_COLLECTIONS` 里的三个集合,集合名**不接受外部参数**。 schema 三个集合逐字相同(见实施计划 Task 6):`knowledge_id`(VARCHAR64, 主键) / `title`(256) / `snippet`(4000) / `tags`(512) / `version`(16) / `intent`(32) / `embedding`(FLOAT_VECTOR dim=1024),索引 `AUTOINDEX` + `metric_type="COSINE"`。 """ from __future__ import annotations import asyncio import sys from pathlib import Path from typing import Any sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from app.core.config import get_settings # noqa: E402 from app.core.knowledge_contracts import ALLOWED_COLLECTIONS, VECTOR_DIM # noqa: E402 PRIMARY_FIELD = "knowledge_id" VECTOR_FIELD = "embedding" INDEX_NAME = "knowledge_autoindex" #: (字段名, VARCHAR 最大长度);顺序与写入侧口径一致。 VARCHAR_FIELDS: tuple[tuple[str, int], ...] = ( (PRIMARY_FIELD, 64), ("title", 256), ("snippet", 4000), ("tags", 512), ("version", 16), ("intent", 32), ) #: 允许缺省的字段。`intent` 是**稀疏标签**:知识契约里 `None` 表示"无显式标签、由集合名推断" #: (见 `app/core/knowledge_contracts.py` 的 `intent_for_qa_id` docstring), #: 因此写入侧对普通知识**省略该字段**。若集合把它定义为 non-nullable 且无默认值, #: `upsert` 会抛 `Insert missed an field 'intent'`(已实测:106 条里 86 条普通知识全部写不进)。 NULLABLE_FIELDS = frozenset({"intent"}) def expected_fields() -> dict[str, dict[str, Any]]: """本脚本期望的字段结构:`{字段名: {type, is_primary?, max_length?, dim?}}`。""" fields: dict[str, dict[str, Any]] = {} for name, max_length in VARCHAR_FIELDS: fields[name] = { "type": "VARCHAR", "max_length": max_length, "is_primary": name == PRIMARY_FIELD, } fields[VECTOR_FIELD] = {"type": "FLOAT_VECTOR", "dim": VECTOR_DIM, "is_primary": False} return fields def describe_mismatch(described: Any) -> list[str]: """比较 Milvus `describe_collection` 结果与期望 schema,返回差异说明(一致时为空)。""" if not isinstance(described, dict): # pragma: no cover - 防御:返回结构异常 return [f"无法解析集合描述:{type(described).__name__}"] raw_fields = described.get("fields") if not isinstance(raw_fields, list): # pragma: no cover - 防御:返回结构异常 return ["集合描述缺少 fields"] actual: dict[str, dict[str, Any]] = {} for raw in raw_fields: if not isinstance(raw, dict): continue name = str(raw.get("name", "")) params = raw.get("params") if isinstance(raw.get("params"), dict) else {} actual[name] = { "type": raw.get("type"), "params": params, "is_primary": bool(raw.get("is_primary")), } problems: list[str] = [] for name, spec in expected_fields().items(): found = actual.get(name) if found is None: problems.append(f"缺少字段 {name}") continue if found["is_primary"] != spec["is_primary"]: problems.append(f"字段 {name} 主键标记不一致(实际 {found['is_primary']})") if name == VECTOR_FIELD: dim = found["params"].get("dim") if not isinstance(dim, int) or dim != VECTOR_DIM: problems.append(f"字段 {name} 维度不一致(实际 {dim},期望 {VECTOR_DIM})") continue length = found["params"].get("max_length") if length != spec["max_length"]: problems.append(f"字段 {name} 长度不一致(实际 {length},期望 {spec['max_length']})") return problems def _build_schema() -> Any: from pymilvus import DataType, MilvusClient # type: ignore[import-untyped] schema = MilvusClient.create_schema(auto_id=False, enable_dynamic_field=False) for name, max_length in VARCHAR_FIELDS: schema.add_field( field_name=name, datatype=DataType.VARCHAR, max_length=max_length, is_primary=(name == PRIMARY_FIELD), # 稀疏标签字段必须可空:写入侧对普通知识**省略**该字段(不是写空串)。 nullable=name in NULLABLE_FIELDS, ) schema.add_field(field_name=VECTOR_FIELD, datatype=DataType.FLOAT_VECTOR, dim=VECTOR_DIM) return schema def _build_index_params() -> Any: from pymilvus import MilvusClient # type: ignore[import-untyped] index_params = MilvusClient.prepare_index_params() index_params.add_index( field_name=VECTOR_FIELD, index_name=INDEX_NAME, index_type="AUTOINDEX", metric_type="COSINE", ) return index_params async def ensure_collections(uri: str, token: str = "") -> tuple[list[str], list[str], list[str]]: """确保三集合存在。返回 `(创建, 已存在, 结构冲突)`;结构冲突时不覆盖。""" from pymilvus import AsyncMilvusClient # type: ignore[import-untyped] names = sorted(ALLOWED_COLLECTIONS) created: list[str] = [] existed: list[str] = [] conflicting: list[str] = [] client = AsyncMilvusClient(uri=uri, token=token or None) try: for name in names: if await client.has_collection(name): described = await client.describe_collection(name) problems = describe_mismatch(described) if problems: conflicting.append(f"{name}: {';'.join(problems)}") print(f"[冲突] {name} 已存在但结构不同,未覆盖 -> {';'.join(problems)}") else: existed.append(name) print(f"[跳过] {name} 已存在且结构一致") continue await client.create_collection( collection_name=name, schema=_build_schema(), index_params=_build_index_params(), ) await client.load_collection(name) created.append(name) print(f"[创建] {name} 已创建并加载") finally: await client.close() return created, existed, conflicting def main() -> int: settings = get_settings() created, existed, conflicting = asyncio.run( ensure_collections(settings.milvus_uri, settings.milvus_token) ) print(f"创建 {len(created)} 个:{created}") print(f"跳过 {len(existed)} 个:{existed}") if conflicting: print("存在同名不同结构的集合,已停止且未做任何覆盖,请人工确认:") for line in conflicting: print(f" - {line}") return 2 return 0 if __name__ == "__main__": # pragma: no cover - CLI 入口 raise SystemExit(main())