diff --git a/app/infrastructure/neo4j_profile_projection.py b/app/infrastructure/neo4j_profile_projection.py new file mode 100644 index 0000000..50f2ddd --- /dev/null +++ b/app/infrastructure/neo4j_profile_projection.py @@ -0,0 +1,152 @@ +"""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, query: str, **parameters: 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 diff --git a/tests/unit/infrastructure/test_neo4j_profile_projection.py b/tests/unit/infrastructure/test_neo4j_profile_projection.py new file mode 100644 index 0000000..ac5df3d --- /dev/null +++ b/tests/unit/infrastructure/test_neo4j_profile_projection.py @@ -0,0 +1,92 @@ +from types import SimpleNamespace + +import pytest + +from app.infrastructure.neo4j_profile_projection import Neo4jProfileProjection + + +class FakeDriver: + def __init__(self, *, applied: bool = True) -> None: + self.applied = applied + self.calls: list[tuple[str, dict[str, object]]] = [] + + async def execute_query(self, query: str, **parameters: object) -> SimpleNamespace: + self.calls.append((query, parameters)) + return SimpleNamespace(records=[{"applied": True}] if self.applied else []) + + +def payload() -> dict[str, object]: + return { + "customer_id": 7, + "profile_uuid": "profile-7-v1", + "profile_version": 1, + "memory_sources": [ + { + "memory_uuid": "memory-1", + "memory_key": "preference:risk_level", + "content": "稳健型", + "memory_type": "preference", + "confidence": 0.9, + "version": 2, + "valid_until": None, + }, + { + "memory_uuid": "memory-2", + "memory_key": "goal:liquidity", + "content": "保持流动性", + "memory_type": "goal", + "confidence": 0.8, + "version": 1, + "valid_until": None, + }, + ], + } + + +@pytest.mark.asyncio +async def test_projects_only_fixed_preference_and_goal_queries() -> None: + driver = FakeDriver() + result = await Neo4jProfileProjection(driver).upsert(payload()) + + assert result.applied is True + assert len(driver.calls) == 3 + assert "MERGE (c:Customer" in driver.calls[0][0] + assert "PREFERS" in driver.calls[1][0] + assert "HAS_GOAL" in driver.calls[2][0] + assert driver.calls[1][1]["items"][0]["memory_uuid"] == "memory-1" + + +@pytest.mark.asyncio +async def test_lower_profile_version_is_skipped_without_writes() -> None: + driver = FakeDriver(applied=False) + result = await Neo4jProfileProjection(driver).upsert(payload()) + + assert result.applied is False + assert result.reason == "newer_profile_version_exists" + assert len(driver.calls) == 1 + + +@pytest.mark.asyncio +async def test_sensitive_memory_content_is_redacted_before_projection() -> None: + data = payload() + source = data["memory_sources"][0] + assert isinstance(source, dict) + source["content"] = "我的密码是123456,手机号13800138000" + driver = FakeDriver() + + await Neo4jProfileProjection(driver).upsert(data) + + projected = driver.calls[1][1]["items"][0]["content"] + assert "123456" not in projected + assert "13800138000" not in projected + + +@pytest.mark.asyncio +async def test_unknown_memory_key_is_rejected() -> None: + data = payload() + source = data["memory_sources"][0] + assert isinstance(source, dict) + source["memory_key"] = "account:balance" + + with pytest.raises(ValueError, match="not projectable"): + await Neo4jProfileProjection(FakeDriver()).upsert(data)