"""Local entity-reader compatibility adapter for the legacy Zep consumer shape.""" from __future__ import annotations from dataclasses import dataclass, field from typing import Any, Optional from .memory_repository import SqlAlchemyMemoryRepository @dataclass class LocalEntityNode: uuid: str name: str labels: list[str] summary: str attributes: dict[str, Any] related_edges: list[dict[str, Any]] = field(default_factory=list) related_nodes: list[dict[str, Any]] = field(default_factory=list) def to_dict(self) -> dict[str, Any]: return { "uuid": self.uuid, "name": self.name, "labels": self.labels, "summary": self.summary, "attributes": self.attributes, "related_edges": self.related_edges, "related_nodes": self.related_nodes, } def get_entity_type(self) -> Optional[str]: return next((label for label in self.labels if label not in {"Entity", "Node"}), None) @dataclass class LocalFilteredEntities: entities: list[LocalEntityNode] entity_types: set[str] total_count: int filtered_count: int def to_dict(self) -> dict[str, Any]: return { "entities": [entity.to_dict() for entity in self.entities], "entity_types": sorted(self.entity_types), "total_count": self.total_count, "filtered_count": self.filtered_count, } class LocalEntityReader: def __init__( self, session_or_repository, *, organization_id: str | None = None, graph_id: str | None = None, owns_session: bool = False, ): self.owns_session = False if isinstance(session_or_repository, SqlAlchemyMemoryRepository): self.repository = session_or_repository else: if organization_id is None or graph_id is None: raise ValueError("memory_reader_scope_required") self.repository = SqlAlchemyMemoryRepository( session_or_repository, organization_id=organization_id, graph_id=graph_id, ) self.owns_session = owns_session def _validate_graph_id(self, graph_id: str | None = None) -> None: if graph_id is not None and graph_id != self.repository.graph_id: raise ValueError("memory_graph_scope_conflict") def close(self) -> None: if self.owns_session: self.repository.session.close() self.owns_session = False @staticmethod def _node_dict(node) -> dict[str, Any]: return { "uuid": node.id, "name": node.canonical_name, "labels": list(node.labels or []), "summary": node.summary or "", "attributes": dict(node.attributes or {}), } def _entity(self, node, *, enrich_with_edges: bool = True) -> LocalEntityNode: related_edges: list[dict[str, Any]] = [] related_nodes: list[dict[str, Any]] = [] if enrich_with_edges: all_nodes = {item.id: item for item in self.repository.list_nodes()} for edge in self.repository.get_node_edges(node.id): source = all_nodes.get(edge.source_node_id) target = all_nodes.get(edge.target_node_id) if edge.source_node_id == node.id: related_edges.append( { "direction": "outgoing", "edge_name": edge.relation, "fact": edge.fact, "target_node_uuid": edge.target_node_id, } ) related_node = target else: related_edges.append( { "direction": "incoming", "edge_name": edge.relation, "fact": edge.fact, "source_node_uuid": edge.source_node_id, } ) related_node = source if related_node is not None: related_nodes.append( { "uuid": related_node.id, "name": related_node.canonical_name, "labels": list(related_node.labels or []), "summary": related_node.summary or "", } ) return LocalEntityNode( uuid=node.id, name=node.canonical_name, labels=list(node.labels or []), summary=node.summary or "", attributes=dict(node.attributes or {}), related_edges=related_edges, related_nodes=related_nodes, ) def get_all_nodes(self, graph_id: str | None = None) -> list[dict[str, Any]]: self._validate_graph_id(graph_id) return [self._node_dict(node) for node in self.repository.list_nodes()] def get_all_edges(self, graph_id: str | None = None) -> list[dict[str, Any]]: self._validate_graph_id(graph_id) return [ { "uuid": edge.id, "name": edge.relation, "fact": edge.fact, "source_node_uuid": edge.source_node_id, "target_node_uuid": edge.target_node_id, "attributes": dict(edge.attributes or {}), } for edge in self.repository.list_edges() ] def filter_defined_entities( self, graph_id: str | None = None, *, defined_entity_types: Optional[list[str]] = None, enrich_with_edges: bool = True, ) -> LocalFilteredEntities: self._validate_graph_id(graph_id) nodes = self.repository.list_nodes() allowed = set(defined_entity_types or []) entity_types: set[str] = set() entities: list[LocalEntityNode] = [] for node in nodes: custom_labels = [label for label in node.labels or [] if label not in {"Entity", "Node"}] if not custom_labels: continue if defined_entity_types: matching_labels = [label for label in custom_labels if label in allowed] if not matching_labels: continue entity_type = matching_labels[0] else: entity_type = custom_labels[0] entity_types.add(entity_type) entities.append(self._entity(node, enrich_with_edges=enrich_with_edges)) return LocalFilteredEntities( entities=entities, entity_types=entity_types, total_count=len(nodes), filtered_count=len(entities), ) def get_entity_with_context( self, graph_id_or_entity_uuid: str | None = None, entity_uuid: str | None = None, *, graph_id: str | None = None, ) -> Optional[LocalEntityNode]: requested_graph_id = graph_id if entity_uuid is not None and requested_graph_id is None: requested_graph_id = graph_id_or_entity_uuid self._validate_graph_id(requested_graph_id) node_id = entity_uuid or graph_id_or_entity_uuid if not node_id: raise ValueError("memory_entity_id_required") node = self.repository.get_node(node_id) return self._entity(node, enrich_with_edges=True) if node is not None else None def get_entities_by_type( self, graph_id_or_entity_type: str | None = None, entity_type: str | None = None, *, graph_id: str | None = None, enrich_with_edges: bool = True, ) -> list[LocalEntityNode]: requested_graph_id = graph_id if entity_type is not None and requested_graph_id is None: requested_graph_id = graph_id_or_entity_type self._validate_graph_id(requested_graph_id) selected_type = entity_type or graph_id_or_entity_type if not selected_type: raise ValueError("memory_entity_type_required") return [ self._entity(node, enrich_with_edges=enrich_with_edges) for node in self.repository.list_nodes() if selected_type in (node.labels or []) ] def make_local_entity_reader_factory(session_factory, *, organization_id: str): """Create per-worker readers; each reader owns and closes its own session.""" if not isinstance(organization_id, str) or not organization_id.strip(): raise ValueError("memory_reader_organization_required") def factory(graph_id: str) -> LocalEntityReader: return LocalEntityReader( session_factory(), organization_id=organization_id, graph_id=graph_id, owns_session=True, ) return factory