袁聪的第二次提交,项目已完整
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Protocol
|
||||
|
||||
|
||||
@@ -11,7 +12,21 @@ class VectorSearchResult:
|
||||
self.degraded = degraded
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class VectorHit:
|
||||
"""标准化后的向量命中;`memory_uuid` 用于回表,`score` 为集合内相似度。"""
|
||||
|
||||
memory_uuid: str
|
||||
score: float
|
||||
|
||||
|
||||
class VectorMemoryAdapter:
|
||||
"""Milvus 读边界:查询失败一律降级,绝不让向量库故障冒泡到主链路。"""
|
||||
|
||||
# `memory_uuid` 的候选字段名,覆盖实体列与常见别名。
|
||||
_UUID_FIELDS = ("memory_uuid", "uuid", "memory_id", "id")
|
||||
_SCORE_FIELDS = ("score", "distance", "similarity")
|
||||
|
||||
def __init__(self, client: VectorClient, collection: str) -> None:
|
||||
self.client = client
|
||||
self.collection = collection
|
||||
@@ -26,3 +41,80 @@ class VectorMemoryAdapter:
|
||||
return VectorSearchResult(hits)
|
||||
except Exception:
|
||||
return VectorSearchResult([], degraded=True)
|
||||
|
||||
def parse_hits(self, hits: Any) -> list[VectorHit]:
|
||||
"""把 pymilvus 的嵌套命中结构折叠为 `VectorHit` 列表(纯函数,不抛异常)。
|
||||
|
||||
pymilvus 返回 `[[{id, distance, entity}]]`(每查询一组),行内既有实体字典
|
||||
也有命名为 `id` / `distance` 的字段,这里两种形态都吸收;没有可识别
|
||||
`memory_uuid` 的行被丢弃而不是猜造标识,避免召回无法回表的假引用。
|
||||
"""
|
||||
parsed: list[VectorHit] = []
|
||||
for row in self._flatten(hits):
|
||||
hit = self._parse_hit(row)
|
||||
if hit is not None:
|
||||
parsed.append(hit)
|
||||
return parsed
|
||||
|
||||
def _flatten(self, hits: Any) -> list[Any]:
|
||||
rows: list[Any] = []
|
||||
for group in hits if isinstance(hits, (list, tuple)) else [hits]:
|
||||
if isinstance(group, (list, tuple)):
|
||||
rows.extend(group)
|
||||
else:
|
||||
rows.append(group)
|
||||
return rows
|
||||
|
||||
def _parse_hit(self, row: Any) -> VectorHit | None:
|
||||
if isinstance(row, dict):
|
||||
fields = dict(row)
|
||||
elif isinstance(row, (list, tuple)) and len(row) >= 2:
|
||||
# 无实体字典时的裸数组形态:约定 (主键, 距离)。
|
||||
return self._hit_from_pair(str(row[0]), row[1])
|
||||
else:
|
||||
fields = self._object_fields(row)
|
||||
entity = fields.get("entity")
|
||||
if isinstance(entity, dict):
|
||||
merged = dict(entity)
|
||||
merged.update({key: value for key, value in fields.items() if key != "entity"})
|
||||
fields = merged
|
||||
uuid = self._first_string(fields, self._UUID_FIELDS)
|
||||
if uuid is None:
|
||||
return None
|
||||
return VectorHit(memory_uuid=uuid, score=self._score(fields))
|
||||
|
||||
def _hit_from_pair(self, raw_id: str, raw_score: Any) -> VectorHit | None:
|
||||
if not raw_id:
|
||||
return None
|
||||
return VectorHit(memory_uuid=raw_id, score=self._as_float(raw_score) or 0.0)
|
||||
|
||||
def _object_fields(self, row: Any) -> dict[str, Any]:
|
||||
fields: dict[str, Any] = {}
|
||||
for name in (*self._UUID_FIELDS, *self._SCORE_FIELDS, "entity"):
|
||||
if hasattr(row, name):
|
||||
fields[name] = getattr(row, name)
|
||||
return fields
|
||||
|
||||
def _first_string(self, fields: dict[str, Any], names: tuple[str, ...]) -> str | None:
|
||||
for name in names:
|
||||
value = fields.get(name)
|
||||
if value is None:
|
||||
continue
|
||||
if isinstance(value, str) and value.strip():
|
||||
return value.strip()
|
||||
if isinstance(value, int):
|
||||
return str(value)
|
||||
return None
|
||||
|
||||
def _score(self, fields: dict[str, Any]) -> float:
|
||||
for name in self._SCORE_FIELDS:
|
||||
score = self._as_float(fields.get(name))
|
||||
if score is not None:
|
||||
return score
|
||||
return 0.0
|
||||
|
||||
def _as_float(self, value: Any) -> float | None:
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
Reference in New Issue
Block a user