"""Local compatibility tools for the legacy ZepToolsService result contract.""" from __future__ import annotations from dataclasses import dataclass, field from typing import Any, Optional from .memory_repository import SqlAlchemyMemoryRepository @dataclass class LocalSearchResult: facts: list[str] edges: list[dict[str, Any]] nodes: list[dict[str, Any]] query: str total_count: int def to_dict(self) -> dict[str, Any]: return { "facts": self.facts, "edges": self.edges, "nodes": self.nodes, "query": self.query, "total_count": self.total_count, } def to_text(self) -> str: parts = [f"Search query: {self.query}", f"Found {self.total_count} relevant items"] if self.facts: parts.append("\n### Relevant facts:") parts.extend(f"{index}. {fact}" for index, fact in enumerate(self.facts, 1)) return "\n".join(parts) @dataclass class LocalNodeInfo: uuid: str name: str labels: list[str] summary: str attributes: dict[str, Any] def to_dict(self) -> dict[str, Any]: return { "uuid": self.uuid, "name": self.name, "labels": self.labels, "summary": self.summary, "attributes": self.attributes, } @dataclass class LocalInsightForgeResult: query: str simulation_requirement: str sub_queries: list[str] semantic_facts: list[str] = field(default_factory=list) entity_insights: list[dict[str, Any]] = field(default_factory=list) relationship_chains: list[str] = field(default_factory=list) total_facts: int = 0 total_entities: int = 0 total_relationships: int = 0 def to_text(self) -> str: parts = [ "## Local Memory Deep Analysis", f"Analysis question: {self.query}", f"Prediction scenario: {self.simulation_requirement}", f"\n### Statistics\n- Facts: {self.total_facts}\n- Entities: {self.total_entities}\n- Relationships: {self.total_relationships}", ] if self.semantic_facts: parts.append("\n### Key facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.semantic_facts, 1))) if self.entity_insights: parts.append( "\n### Entities\n" + "\n".join( f"- {item.get('name', 'Unknown')} ({item.get('type', 'Entity')}): {item.get('summary', '')}" for item in self.entity_insights ) ) if self.relationship_chains: parts.append("\n### Relationships\n" + "\n".join(f"- {chain}" for chain in self.relationship_chains)) return "\n".join(parts) @dataclass class LocalPanoramaResult: query: str all_nodes: list[LocalNodeInfo] = field(default_factory=list) all_edges: list[LocalEdgeInfo] = field(default_factory=list) active_facts: list[str] = field(default_factory=list) historical_facts: list[str] = field(default_factory=list) def to_text(self) -> str: parts = [ "## Local Memory Panorama", f"Query: {self.query}", f"\n### Statistics\n- Total nodes: {len(self.all_nodes)}\n- Total edges: {len(self.all_edges)}\n- Active facts: {len(self.active_facts)}\n- Historical facts: {len(self.historical_facts)}", ] if self.active_facts: parts.append("\n### Active facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.active_facts, 1))) if self.historical_facts: parts.append("\n### Historical facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.historical_facts, 1))) if self.all_nodes: parts.append("\n### Entities\n" + "\n".join(f"- {node.name}: {node.summary}" for node in self.all_nodes)) return "\n".join(parts) @dataclass class LocalInterviewResult: interview_topic: str summary: str def to_text(self) -> str: return ( "## Local Memory Interview\n" f"Topic: {self.interview_topic}\n\n" f"{self.summary}" ) @dataclass class LocalEdgeInfo: uuid: str name: str fact: str source_node_uuid: str target_node_uuid: str source_node_name: Optional[str] = None target_node_name: Optional[str] = None created_at: Optional[str] = None valid_at: Optional[str] = None invalid_at: Optional[str] = None expired_at: Optional[str] = None def to_dict(self) -> dict[str, Any]: return { "uuid": self.uuid, "name": self.name, "fact": self.fact, "source_node_uuid": self.source_node_uuid, "target_node_uuid": self.target_node_uuid, "source_node_name": self.source_node_name, "target_node_name": self.target_node_name, "created_at": self.created_at, "valid_at": self.valid_at, "invalid_at": self.invalid_at, "expired_at": self.expired_at, } @property def is_expired(self) -> bool: return self.expired_at is not None @property def is_invalid(self) -> bool: return self.invalid_at is not None class LocalMemoryTools: def __init__(self, session_or_repository, *, organization_id: str | None = None, graph_id: str | None = None): 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_tools_scope_required") self.repository = SqlAlchemyMemoryRepository( session_or_repository, organization_id=organization_id, graph_id=graph_id, ) 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") @staticmethod def _node_info(node) -> LocalNodeInfo: return LocalNodeInfo( uuid=node.id, name=node.canonical_name, labels=list(node.labels or []), summary=node.summary or "", attributes=dict(node.attributes or {}), ) def _edge_info(self, edge, nodes_by_id: dict[str, Any]) -> LocalEdgeInfo: source = nodes_by_id.get(edge.source_node_id) target = nodes_by_id.get(edge.target_node_id) return LocalEdgeInfo( uuid=edge.id, name=edge.relation, fact=edge.fact, source_node_uuid=edge.source_node_id, target_node_uuid=edge.target_node_id, source_node_name=source.canonical_name if source else None, target_node_name=target.canonical_name if target else None, created_at=edge.created_at.isoformat() if edge.created_at else None, valid_at=edge.valid_at.isoformat() if edge.valid_at else None, invalid_at=edge.invalid_at.isoformat() if edge.invalid_at else None, expired_at=edge.expired_at.isoformat() if edge.expired_at else None, ) def get_all_nodes(self, *, limit: int = 10_000) -> list[LocalNodeInfo]: return [self._node_info(node) for node in self.repository.list_nodes(limit=limit)] def get_all_edges(self, *, limit: int = 20_000) -> list[LocalEdgeInfo]: nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)} return [self._edge_info(edge, nodes) for edge in self.repository.list_edges(limit=limit)] @staticmethod def _score(query: str, *values: str) -> int: normalized_query = query.casefold().strip() haystack = " ".join(value or "" for value in values).casefold() if not normalized_query or not haystack: return 0 if normalized_query in haystack: return 100 return sum(10 for token in normalized_query.split() if token in haystack) def search_graph( self, query: str, *, limit: int = 10, scope: str = "edges", graph_id: str | None = None, ) -> LocalSearchResult: self._validate_graph_id(graph_id) if scope not in {"edges", "nodes", "both"}: raise ValueError("invalid_memory_search_scope") safe_limit = min(max(int(limit), 1), 100) facts: list[str] = [] edges: list[dict[str, Any]] = [] nodes: list[dict[str, Any]] = [] all_nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)} if scope in {"edges", "both"}: scored_edges = [ (self._score(query, edge.fact, edge.relation), edge) for edge in self.repository.list_edges(limit=20_000) ] for score, edge in sorted(scored_edges, key=lambda item: (-item[0], item[1].id))[:safe_limit]: if score <= 0: continue edge_info = self._edge_info(edge, all_nodes).to_dict() edges.append(edge_info) if edge.fact: facts.append(edge.fact) if scope in {"nodes", "both"}: scored_nodes = [ (self._score(query, node.canonical_name, node.summary), node) for node in all_nodes.values() ] for score, node in sorted(scored_nodes, key=lambda item: (-item[0], item[1].id))[:safe_limit]: if score <= 0: continue node_info = self._node_info(node).to_dict() node_info.pop("attributes", None) nodes.append(node_info) if node.summary: facts.append(f"[{node.canonical_name}]: {node.summary}") return LocalSearchResult( facts=facts, edges=edges, nodes=nodes, query=query, total_count=len(facts), ) def quick_search( self, query: str, *, limit: int = 10, graph_id: str | None = None, ) -> LocalSearchResult: return self.search_graph(query, limit=limit, scope="edges", graph_id=graph_id) def get_node_detail(self, node_uuid: str) -> LocalNodeInfo | None: node = self.repository.get_node(node_uuid) return self._node_info(node) if node is not None else None def get_node_edges(self, node_uuid: str, *, limit: int = 500) -> list[LocalEdgeInfo]: nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)} return [self._edge_info(edge, nodes) for edge in self.repository.get_node_edges(node_uuid, limit=limit)] def insight_forge( self, *, graph_id: str, query: str, simulation_requirement: str = "", report_context: str = "", ) -> LocalInsightForgeResult: self._validate_graph_id(graph_id) result = self.search_graph(query, limit=100, scope="edges", graph_id=graph_id) facts = list(dict.fromkeys(result.facts)) nodes_by_id = { node.id: node for node in self.repository.list_nodes(limit=10_000) } related_node_ids = list( dict.fromkeys( node_id for edge in result.edges for node_id in (edge.get("source_node_uuid"), edge.get("target_node_uuid")) if node_id ) ) entity_insights = [] for node_id in related_node_ids: node = nodes_by_id.get(node_id) if node is None: continue entity_insights.append( { "uuid": node.id, "name": node.canonical_name, "type": next( (label for label in (node.labels or []) if label not in {"Entity", "Node"}), "Entity", ), "summary": node.summary or "", "related_facts": [ fact for fact in facts if node.canonical_name.casefold() in fact.casefold() ], } ) relationship_chains = [] seen_chains = set() for edge in result.edges: chain = ( f"{edge.get('source_node_name') or edge.get('source_node_uuid')} " f"--[{edge.get('name', '')}]--> " f"{edge.get('target_node_name') or edge.get('target_node_uuid')}" ) if chain not in seen_chains: seen_chains.add(chain) relationship_chains.append(chain) return LocalInsightForgeResult( query=query, simulation_requirement=simulation_requirement, sub_queries=[query] if query else [], semantic_facts=facts, entity_insights=entity_insights, relationship_chains=relationship_chains, total_facts=len(facts), total_entities=len(entity_insights), total_relationships=len(relationship_chains), ) def panorama_search( self, *, graph_id: str, query: str, include_expired: bool = True, limit: int = 50, ) -> LocalPanoramaResult: self._validate_graph_id(graph_id) safe_limit = min(max(int(limit), 1), 100) nodes = self.get_all_nodes(limit=10_000) edges = self.get_all_edges(limit=20_000) active_facts = [] historical_facts = [] for edge in edges: if not edge.fact: continue if edge.is_invalid or edge.is_expired: valid_at = edge.valid_at or "Unknown" invalid_at = edge.invalid_at or edge.expired_at or "Unknown" historical_facts.append(f"[{valid_at} - {invalid_at}] {edge.fact}") else: active_facts.append(edge.fact) query_lower = query.casefold() keywords = [ word.strip() for word in query_lower.replace(",", " ").replace(",", " ").split() if len(word.strip()) > 1 ] def relevance_score(fact: str) -> int: fact_lower = fact.casefold() score = 100 if query_lower in fact_lower else 0 return score + sum(10 for keyword in keywords if keyword in fact_lower) active_facts.sort(key=relevance_score, reverse=True) historical_facts.sort(key=relevance_score, reverse=True) return LocalPanoramaResult( query=query, all_nodes=nodes, all_edges=edges, active_facts=active_facts[:safe_limit], historical_facts=historical_facts[:safe_limit] if include_expired else [], ) def get_entity_summary(self, *, graph_id: str, entity_name: str) -> dict[str, Any]: self._validate_graph_id(graph_id) normalized = " ".join((entity_name or "").casefold().split()) node = next( ( item for item in self.repository.list_nodes(limit=10_000) if item.normalized_name == normalized ), None, ) if node is None: return {"entity_name": entity_name, "found": False, "summary": "", "related_facts": []} edges = self.get_node_edges(node.id, limit=500) return { "entity_name": node.canonical_name, "found": True, "uuid": node.id, "labels": list(node.labels or []), "summary": node.summary or "", "attributes": dict(node.attributes or {}), "related_facts": [edge.fact for edge in edges], } def get_entities_by_type(self, *, graph_id: str, entity_type: str) -> list[LocalNodeInfo]: self._validate_graph_id(graph_id) target = (entity_type or "").casefold().strip() return [ self._node_info(node) for node in self.repository.list_nodes(limit=10_000) if target in {str(label).casefold() for label in (node.labels or [])} ] def interview_agents( self, *, simulation_id: str, interview_requirement: str, simulation_requirement: str, max_agents: int = 5, ) -> LocalInterviewResult: del simulation_id, max_agents return LocalInterviewResult( interview_topic=interview_requirement, summary=( "Local memory stores graph facts, not simulation transcripts. " "Use the retrieved entities and relationships as evidence; no synthetic interview response was generated." ), ) def get_simulation_context( self, *, graph_id: str, simulation_requirement: str, limit: int = 30, ) -> dict[str, Any]: self._validate_graph_id(graph_id) result = self.search_graph( simulation_requirement, limit=limit, scope="both", graph_id=graph_id, ) nodes = self.repository.list_nodes(limit=10_000) entities = [ { "name": node.canonical_name, "type": next((label for label in (node.labels or []) if label not in {"Entity", "Node"}), "Entity"), "summary": node.summary or "", } for node in nodes if any(label not in {"Entity", "Node"} for label in (node.labels or [])) ] return { "simulation_requirement": simulation_requirement, "related_facts": result.facts, "graph_statistics": self.get_graph_statistics(graph_id), "entities": entities[: max(int(limit), 1)], "total_entities": len(entities), } def get_graph_statistics(self, graph_id: str | None = None) -> dict[str, Any]: self._validate_graph_id(graph_id) nodes = self.repository.list_nodes(limit=10_000) edges = self.repository.list_edges(limit=20_000) entity_types: dict[str, int] = {} for node in nodes: for label in node.labels or []: if label not in {"Entity", "Node"}: entity_types[label] = entity_types.get(label, 0) + 1 relation_types: dict[str, int] = {} for edge in edges: relation_types[edge.relation] = relation_types.get(edge.relation, 0) + 1 return { "graph_id": self.repository.graph_id, "node_count": len(nodes), "edge_count": len(edges), "total_nodes": len(nodes), "total_edges": len(edges), "entity_types": entity_types, "relation_types": relation_types, }