from datetime import UTC, datetime, timedelta from typing import Any from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.model.memory import MemorySyncOutbox from app.service.graph_model import merge_node_clause, node_id, node_spec, relation_allowed from app.service.relationship_service import RelationshipService class GraphProjectionWorker: """Projects approved domain events to Neo4j; callers provide durable dedup storage.""" def __init__( self, relationships: RelationshipService, session: AsyncSession | None = None ) -> None: self.relationships = relationships self.session = session self.processed_event_ids: set[str] = set() async def project(self, event_id: str, payload: dict[str, Any]) -> bool: if event_id in self.processed_event_ids: return False if self.session is not None: existing = await self.session.scalar( select(MemorySyncOutbox).where(MemorySyncOutbox.event_uuid == event_id) ) if existing is not None and existing.status == "processed": return False relation = payload.get("relationship") if relation not in RelationshipService.ALLOWED_RELATIONSHIPS: raise ValueError("relationship is not allowed") # 节点按**类型**写标签与主属性:原先统一写 `:Entity {entity_id}`,而读服务查的是 # `:Customer {customer_id}`,两侧从未对齐,写进去的关系永远读不出来。 # 标签与属性名来自 `graph_model` 的受控常量,调用方传不进任意字符串。 source_type = str(payload.get("source_type", "")).strip().lower() target_type = str(payload.get("target_type", "")).strip().lower() source = node_spec(source_type) target = node_spec(target_type) if not relation_allowed(relation, source_type, target_type): raise ValueError( f"relationship {relation} cannot connect {source_type} -> {target_type}" ) query = ( f"{merge_node_clause('a', source, 'source_id')} " f"{merge_node_clause('b', target, 'target_id')} " f"MERGE (a)-[r:{relation}]->(b) " "SET r.trace_id = $trace_id, r.confidence = $confidence, r.updated_at = $updated_at" ) try: await self.relationships.driver.execute_query( query, source_id=node_id(source, payload["source_id"]), target_id=node_id(target, payload["target_id"]), trace_id=str(payload.get("trace_id", "")), confidence=float(payload.get("confidence", 0.0)), updated_at=datetime.now(UTC).replace(tzinfo=None).isoformat(), ) except Exception as exc: if self.session is not None: record = await self.session.scalar( select(MemorySyncOutbox).where(MemorySyncOutbox.event_uuid == event_id) ) if record is not None: record.retry_count += 1 record.last_error = str(exc)[:500] record.status = "dead" if record.retry_count >= 5 else "pending" record.next_retry_at = datetime.now(UTC).replace(tzinfo=None) + timedelta( seconds=min(300, 2 ** record.retry_count) ) await self.session.commit() raise self.processed_event_ids.add(event_id) if self.session is not None: record = await self.session.scalar( select(MemorySyncOutbox).where(MemorySyncOutbox.event_uuid == event_id) ) if record is not None: record.status = "processed" await self.session.commit() return True