128 lines
5.4 KiB
Python
128 lines
5.4 KiB
Python
"""Milvus 长期记忆投影适配器。
|
|||
|
|
|
||
|
|
只写入已经审核的 `memory_sources`,不接受画像快照整体冒充单条记忆。
|
||
|
|
"""
|
||
|
|
|
||
|
|
from collections.abc import Awaitable, Callable
|
||
|
|
from datetime import UTC, datetime
|
||
|
|
from typing import Any, Protocol
|
||
|
|
from uuid import UUID
|
||
|
|
|
||
|
|
from app.core.conversation_privacy import sanitize_customer_service_message
|
||
|
|
from app.core.errors import RecoverableAgentError
|
||
|
|
|
||
|
|
PROFILE_COLLECTION = "user_long_term_memory_v1"
|
||
|
|
VECTOR_DIM = 1024
|
||
|
|
|
||
|
|
|
||
|
|
class MilvusProfileClient(Protocol):
|
||
|
|
async def query(self, **kwargs: Any) -> list[dict[str, Any]]: ...
|
||
|
|
|
||
|
|
async def upsert(self, **kwargs: Any) -> Any: ...
|
||
|
|
|
||
|
|
|
||
|
|
EmbeddingProvider = Callable[[str], Awaitable[list[float]]]
|
||
|
|
|
||
|
|
|
||
|
|
class MilvusProfileProjection:
|
||
|
|
"""按记忆 UUID 幂等写入长期记忆向量。"""
|
||
|
|
|
||
|
|
def __init__(
|
||
|
|
self,
|
||
|
|
client: MilvusProfileClient,
|
||
|
|
embed: EmbeddingProvider,
|
||
|
|
*,
|
||
|
|
collection: str = PROFILE_COLLECTION,
|
||
|
|
) -> None:
|
||
|
|
self._client = client
|
||
|
|
self._embed = embed
|
||
|
|
self._collection = collection
|
||
|
|
|
||
|
|
async def upsert(self, payload: dict[str, Any]) -> None:
|
||
|
|
customer_id, profile_version, sources = self._normalize(payload)
|
||
|
|
load_collection = getattr(self._client, "load_collection", None)
|
||
|
|
if load_collection is not None:
|
||
|
|
await load_collection(collection_name=self._collection)
|
||
|
|
rows: list[dict[str, Any]] = []
|
||
|
|
for source in sources:
|
||
|
|
vector = await self._embed(source["content"])
|
||
|
|
if len(vector) != VECTOR_DIM:
|
||
|
|
raise RecoverableAgentError("画像向量维度不一致")
|
||
|
|
existing = await self._client.query(
|
||
|
|
collection_name=self._collection,
|
||
|
|
filter=f'memory_uuid == "{source["memory_uuid"]}"',
|
||
|
|
output_fields=["memory_uuid", "version", "customer_id"],
|
||
|
|
)
|
||
|
|
if existing and int(existing[0].get("version", 0)) > source["version"]:
|
||
|
|
continue
|
||
|
|
rows.append({
|
||
|
|
"memory_uuid": source["memory_uuid"],
|
||
|
|
"customer_id": customer_id,
|
||
|
|
"content": source["content"],
|
||
|
|
"embedding": vector,
|
||
|
|
"memory_type": source["memory_type"],
|
||
|
|
"memory_key": source["memory_key"],
|
||
|
|
"confidence": source["confidence"],
|
||
|
|
"version": source["version"],
|
||
|
|
"status": "active",
|
||
|
|
"valid_until_ts": source["valid_until_ts"],
|
||
|
|
"updated_at_ts": source["updated_at_ts"],
|
||
|
|
})
|
||
|
|
if rows:
|
||
|
|
await self._client.upsert(collection_name=self._collection, data=rows)
|
||
|
|
|
||
|
|
@staticmethod
|
||
|
|
def _normalize(
|
||
|
|
payload: dict[str, Any],
|
||
|
|
) -> tuple[int, int, list[dict[str, Any]]]:
|
||
|
|
customer_id = payload.get("customer_id")
|
||
|
|
profile_version = payload.get("profile_version")
|
||
|
|
sources = payload.get("memory_sources")
|
||
|
|
if not isinstance(customer_id, int) or customer_id <= 0:
|
||
|
|
raise ValueError("customer_id is invalid")
|
||
|
|
if not isinstance(profile_version, int) or profile_version <= 0:
|
||
|
|
raise ValueError("profile_version is invalid")
|
||
|
|
if not isinstance(sources, list):
|
||
|
|
raise ValueError("memory_sources is invalid")
|
||
|
|
normalized: list[dict[str, Any]] = []
|
||
|
|
now = int(datetime.now(UTC).timestamp())
|
||
|
|
for source in sources:
|
||
|
|
if not isinstance(source, dict):
|
||
|
|
raise ValueError("memory source is invalid")
|
||
|
|
required = [source.get(name) for name in (
|
||
|
|
"memory_uuid", "memory_key", "content", "memory_type"
|
||
|
|
)]
|
||
|
|
if not all(isinstance(value, str) and value.strip() for value in required):
|
||
|
|
raise ValueError("memory source fields are invalid")
|
||
|
|
try:
|
||
|
|
memory_uuid = str(UUID(str(source["memory_uuid"])))
|
||
|
|
except ValueError as exc:
|
||
|
|
raise ValueError("memory_uuid is invalid") from exc
|
||
|
|
memory_key = str(source["memory_key"]).strip()
|
||
|
|
if not (memory_key.startswith("preference:") or memory_key.startswith("goal:")):
|
||
|
|
raise ValueError("memory key is not projectable")
|
||
|
|
confidence = source.get("confidence")
|
||
|
|
version = source.get("version")
|
||
|
|
if not isinstance(confidence, (int, float)) or not 0 <= confidence <= 1:
|
||
|
|
raise ValueError("memory confidence is invalid")
|
||
|
|
if not isinstance(version, int) or version <= 0:
|
||
|
|
raise ValueError("memory version is invalid")
|
||
|
|
valid_until = source.get("valid_until")
|
||
|
|
valid_until_ts = None
|
||
|
|
if isinstance(valid_until, str) and valid_until:
|
||
|
|
try:
|
||
|
|
valid_until_ts = int(datetime.fromisoformat(valid_until).timestamp())
|
||
|
|
except ValueError as exc:
|
||
|
|
raise ValueError("memory valid_until is invalid") from exc
|
||
|
|
normalized.append({
|
||
|
|
"memory_uuid": memory_uuid,
|
||
|
|
"memory_key": memory_key,
|
||
|
|
"content": sanitize_customer_service_message(str(source["content"])).strip(),
|
||
|
|
"memory_type": str(source["memory_type"]).strip(),
|
||
|
|
"confidence": float(confidence),
|
||
|
|
"version": version,
|
||
|
|
"valid_until_ts": valid_until_ts,
|
||
|
|
"updated_at_ts": now,
|
||
|
|
})
|
||
|
|
return customer_id, profile_version, normalized
|