128 lines
5.3 KiB
Python
128 lines
5.3 KiB
Python
"""知识检索**读路径**的 Milvus 适配器(Task 6)。
|
||
|
||
与写路径的物理隔离(`app/infrastructure/milvus_knowledge_writer.py`):
|
||
|
||
- 本模块**只读**(`search` + `describe_collection`),绝不 import 写路径模块 ——
|
||
检索进程不持有写客户端,向量库故障不会从写侧传染到问答主链路,反之亦然;
|
||
- 检索侧只接受 `ALLOWED_COLLECTIONS` 内的集合名,校验由调用方
|
||
(`KnowledgeRetrievalService._assert_collections_allowed`)在**任何网络调用之前**完成,
|
||
本适配器只做"第二道闸"式的兜底拒绝。
|
||
|
||
连接是**惰性**的:`__init__` 不连 Milvus,首次检索才 `import pymilvus` 并建
|
||
`AsyncMilvusClient`。任何连接/查询失败统一转成 `RecoverableAgentError`,
|
||
由 `KnowledgeRetrievalService` 捕获后降级 MySQL LIKE —— 本层不吞异常、不返回空结果,
|
||
避免把"向量库挂了"伪装成"知识库没有内容"。
|
||
"""
|
||
|
||
from __future__ import annotations
|
||
|
||
from typing import Any
|
||
|
||
from app.core.errors import ForbiddenAgentError, RecoverableAgentError
|
||
from app.core.knowledge_contracts import ALLOWED_COLLECTIONS
|
||
|
||
#: 返回给上层的标量字段(不含 `embedding`:向量不回传,省带宽也避免误用)。
|
||
DEFAULT_OUTPUT_FIELDS: tuple[str, ...] = (
|
||
"knowledge_id",
|
||
"title",
|
||
"snippet",
|
||
"tags",
|
||
"version",
|
||
"intent",
|
||
)
|
||
|
||
#: 一次检索最多返回的候选数下限/上限,防调用方传入异常 top_k。
|
||
MIN_TOP_K = 1
|
||
MAX_TOP_K = 50
|
||
|
||
|
||
class MilvusKnowledgeClient:
|
||
"""知识向量读边界:`search` 失败一律 `RecoverableAgentError`,由上层降级。"""
|
||
|
||
def __init__(self, uri: str, token: str = "") -> None:
|
||
self._uri = uri
|
||
self._token = token
|
||
self._client: Any = None
|
||
|
||
async def _ensure(self) -> Any:
|
||
if self._client is None:
|
||
try:
|
||
from pymilvus import AsyncMilvusClient # type: ignore[import-untyped]
|
||
except ImportError as exc: # pragma: no cover - 依赖已声明,缺装是环境问题
|
||
raise RecoverableAgentError("pymilvus 未安装,无法检索知识向量") from exc
|
||
try:
|
||
self._client = AsyncMilvusClient(uri=self._uri, token=self._token or None)
|
||
except Exception as exc:
|
||
raise RecoverableAgentError("Milvus 读客户端初始化失败") from exc
|
||
return self._client
|
||
|
||
async def search(
|
||
self,
|
||
*,
|
||
collection: str,
|
||
vector: list[float],
|
||
top_k: int,
|
||
output_fields: tuple[str, ...] | list[str] = DEFAULT_OUTPUT_FIELDS,
|
||
) -> list[dict[str, Any]]:
|
||
"""向量检索,返回**扁平化**的命中行(每行含标量字段与 `score`)。
|
||
|
||
pymilvus 返回的是 `[[{id, distance, entity}]]`(每查询一组),这里折叠为单层列表;
|
||
距离字段名可能是 `distance` 或 `score`,两种都吸收(COSINE 越大越相似)。
|
||
"""
|
||
if collection not in ALLOWED_COLLECTIONS:
|
||
raise ForbiddenAgentError(f"知识集合不在白名单内:{collection}")
|
||
if not vector:
|
||
raise RecoverableAgentError("检索向量不能为空")
|
||
client = await self._ensure()
|
||
try:
|
||
raw = await client.search(
|
||
collection_name=collection,
|
||
data=[list(vector)],
|
||
limit=max(MIN_TOP_K, min(int(top_k), MAX_TOP_K)),
|
||
output_fields=list(output_fields),
|
||
search_params={"metric_type": "COSINE"},
|
||
)
|
||
except Exception as exc:
|
||
raise RecoverableAgentError(f"知识向量检索失败:{collection}") from exc
|
||
return self.parse_hits(raw)
|
||
|
||
@staticmethod
|
||
def parse_hits(raw: Any) -> list[dict[str, Any]]:
|
||
"""把 pymilvus 的嵌套命中结构折叠为 `list[dict]`(纯函数,不抛异常)。"""
|
||
rows: list[dict[str, Any]] = []
|
||
for group in raw if isinstance(raw, list | tuple) else [raw]:
|
||
items = group if isinstance(group, list | tuple) else [group]
|
||
for item in items:
|
||
row = MilvusKnowledgeClient._as_row(item)
|
||
if row is not None:
|
||
rows.append(row)
|
||
return rows
|
||
|
||
@staticmethod
|
||
def _as_row(item: Any) -> dict[str, Any] | None:
|
||
if isinstance(item, dict):
|
||
fields = dict(item)
|
||
elif hasattr(item, "entity") or hasattr(item, "id"):
|
||
fields = {}
|
||
for name in (*DEFAULT_OUTPUT_FIELDS, "id", "distance", "score", "entity"):
|
||
if hasattr(item, name):
|
||
fields[name] = getattr(item, name)
|
||
else:
|
||
return None
|
||
entity = fields.pop("entity", None)
|
||
if isinstance(entity, dict):
|
||
merged = dict(entity)
|
||
merged.update({key: value for key, value in fields.items() if value is not None})
|
||
fields = merged
|
||
knowledge_id = fields.get("knowledge_id") or fields.get("id")
|
||
if knowledge_id is None or not str(knowledge_id).strip():
|
||
# 没有主键的行无法回表校验,丢弃而不是猜造标识(避免假引用)。
|
||
return None
|
||
fields["knowledge_id"] = str(knowledge_id).strip()
|
||
return fields
|
||
|
||
async def close(self) -> None:
|
||
if self._client is not None:
|
||
client, self._client = self._client, None
|
||
await client.close()
|