"""Neo4j 客户画像最小投影适配器。 该模块只接受已审核画像快照的结构化来源,不接受模型生成的 Cypher 或关系名称。 """ from dataclasses import dataclass from datetime import UTC, datetime from typing import Any, Protocol from app.core.conversation_privacy import sanitize_customer_service_message class Neo4jQueryDriver(Protocol): async def execute_query(self, *args: Any, **kwargs: Any) -> Any: ... @dataclass(frozen=True) class ProjectionResult: """一次画像投影结果;`applied=False` 表示版本已被更新版本覆盖。""" applied: bool reason: str = "" _CUSTOMER_QUERY = """ MERGE (c:Customer {customer_id: $customer_id}) WITH c, coalesce(c.profile_version, 0) AS current_version WHERE current_version < $profile_version SET c.profile_version = $profile_version, c.updated_at = $updated_at RETURN true AS applied """ _PREFERENCE_QUERY = """ UNWIND $items AS item MERGE (p:Preference {customer_id: $customer_id, key: item.memory_key}) WITH p, item WHERE coalesce(p.version, 0) <= $profile_version SET p.value = item.content, p.memory_uuid = item.memory_uuid, p.version = item.version, p.confidence = item.confidence WITH p, item MATCH (c:Customer {customer_id: $customer_id}) MERGE (c)-[r:PREFERS {memory_uuid: item.memory_uuid}]->(p) SET r.confidence = item.confidence, r.version = item.version, r.valid_from = item.valid_from, r.valid_until = item.valid_until RETURN count(p) AS projected """ _GOAL_QUERY = """ UNWIND $items AS item MERGE (g:Goal {customer_id: $customer_id, key: item.memory_key}) WITH g, item WHERE coalesce(g.version, 0) <= $profile_version SET g.value = item.content, g.memory_uuid = item.memory_uuid, g.version = item.version, g.confidence = item.confidence WITH g, item MATCH (c:Customer {customer_id: $customer_id}) MERGE (c)-[r:HAS_GOAL {memory_uuid: item.memory_uuid}]->(g) SET r.confidence = item.confidence, r.version = item.version, r.valid_from = item.valid_from, r.valid_until = item.valid_until RETURN count(g) AS projected """ class Neo4jProfileProjection: """把已审核画像来源投影为受控 Neo4j 节点和关系。""" def __init__(self, driver: Neo4jQueryDriver) -> None: self._driver = driver async def upsert(self, payload: dict[str, Any]) -> ProjectionResult: customer_id, profile_version, updated_at, sources = self._normalize(payload) customer_result = await self._driver.execute_query( _CUSTOMER_QUERY, customer_id=customer_id, profile_version=profile_version, updated_at=updated_at, ) if not getattr(customer_result, "records", None): return ProjectionResult(False, "newer_profile_version_exists") grouped = { "preference": [item for item in sources if item["kind"] == "preference"], "goal": [item for item in sources if item["kind"] == "goal"], } for kind, items in grouped.items(): if not items: continue query = _PREFERENCE_QUERY if kind == "preference" else _GOAL_QUERY await self._driver.execute_query( query, customer_id=customer_id, profile_version=profile_version, items=items, ) return ProjectionResult(True, "applied") @staticmethod def _normalize( payload: dict[str, Any], ) -> tuple[int, int, str, list[dict[str, Any]]]: customer_id = payload.get("customer_id") profile_version = payload.get("profile_version") profile_uuid = payload.get("profile_uuid") 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(profile_uuid, str) or not profile_uuid.strip(): raise ValueError("profile_uuid is invalid") if not isinstance(sources, list): raise ValueError("memory_sources is invalid") normalized: list[dict[str, Any]] = [] for source in sources: if not isinstance(source, dict): raise ValueError("memory source is invalid") memory_uuid = source.get("memory_uuid") memory_key = source.get("memory_key") content = source.get("content") memory_type = source.get("memory_type") if not isinstance(memory_uuid, str) or not memory_uuid.strip(): raise ValueError("memory source fields are invalid") if not isinstance(memory_key, str) or not memory_key.strip(): raise ValueError("memory source fields are invalid") if not isinstance(content, str) or not content.strip(): raise ValueError("memory source fields are invalid") if not isinstance(memory_type, str) or not memory_type.strip(): raise ValueError("memory source fields are invalid") if memory_key.startswith("preference:"): kind = "preference" elif memory_key.startswith("goal:"): kind = "goal" else: 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") normalized.append({ "kind": kind, "memory_uuid": memory_uuid.strip(), "memory_key": memory_key.strip(), "content": sanitize_customer_service_message(content).strip(), "memory_type": memory_type.strip(), "confidence": float(confidence), "version": version, "valid_until": source.get("valid_until"), "valid_from": source.get("valid_from"), }) updated_at = str(payload.get("updated_at") or datetime.now(UTC).isoformat()) return customer_id, profile_version, updated_at, normalized