From 1d1ae398c4e3600a37076ee563c67202331a8f46 Mon Sep 17 00:00:00 2001 From: ZohaibHassan16 Date: Sun, 15 Mar 2026 21:50:55 +0500 Subject: [PATCH] fix: add thread safety and pagination to ContextGraph --- semantica/context/context_graph.py | 720 ++++++++++++++++------------- semantica/explorer/session.py | 52 ++- 2 files changed, 417 insertions(+), 355 deletions(-) diff --git a/semantica/context/context_graph.py b/semantica/context/context_graph.py index 6cae5ff9..d9bdadff 100644 --- a/semantica/context/context_graph.py +++ b/semantica/context/context_graph.py @@ -108,6 +108,7 @@ Production Use Cases: from collections import defaultdict, deque from dataclasses import dataclass, field from datetime import datetime, timezone +import threading from typing import Any, Dict, List, Optional, Set, Tuple, Union import uuid @@ -270,6 +271,9 @@ class ContextGraph: self.entity_linker = self.config.get("entity_linker") or EntityLinker() + # Thread safety lock + self._lock = threading.RLock() + # Stable identifier so this graph can be referenced after save/load self.graph_id: str = str(uuid.uuid4()) @@ -331,37 +335,38 @@ class ContextGraph: Number of nodes added """ count = 0 - for node in nodes: - # Extract content from properties if not explicit - node_props = node.get("properties", {}) - content = node_props.get("content", node.get("id")) - # Restore validity windows from properties (written there by ContextNode.to_dict) - # or from top-level keys on the node dict - valid_from = ( - node.get("valid_from") - or node_props.get("valid_from") - ) - valid_until = ( - node.get("valid_until") - or node_props.get("valid_until") - ) - metadata = { - k: v for k, v in node_props.items() - if k not in ("content", "valid_from", "valid_until") - } + with self._lock: + for node in nodes: + # Extract content from properties if not explicit + node_props = node.get("properties", {}) + content = node_props.get("content", node.get("id")) + # Restore validity windows from properties (written there by ContextNode.to_dict) + # or from top-level keys on the node dict + valid_from = ( + node.get("valid_from") + or node_props.get("valid_from") + ) + valid_until = ( + node.get("valid_until") + or node_props.get("valid_until") + ) + metadata = { + k: v for k, v in node_props.items() + if k not in ("content", "valid_from", "valid_until") + } - internal_node = ContextNode( - node_id=node.get("id"), - node_type=node.get("type", "entity"), - content=content, - metadata=metadata, - properties=node_props, - valid_from=valid_from, - valid_until=valid_until, - ) + internal_node = ContextNode( + node_id=node.get("id"), + node_type=node.get("type", "entity"), + content=content, + metadata=metadata, + properties=node_props, + valid_from=valid_from, + valid_until=valid_until, + ) - if self._add_internal_node(internal_node): - count += 1 + if self._add_internal_node(internal_node): + count += 1 return count def add_edges(self, edges: List[Dict[str, Any]]) -> int: @@ -376,32 +381,35 @@ class ContextGraph: Number of edges added """ count = 0 - for edge in edges: - edge_props = edge.get("properties", {}) - # Restore validity windows — ContextEdge.to_dict() writes them at top level - valid_from = edge.get("valid_from") or edge_props.get("valid_from") - valid_until = edge.get("valid_until") or edge_props.get("valid_until") - internal_edge = ContextEdge( - source_id=edge.get("source_id"), - target_id=edge.get("target_id"), - edge_type=edge.get("type", "related_to"), - weight=edge.get("weight", 1.0), - metadata=edge_props, - valid_from=valid_from, - valid_until=valid_until, - ) + with self._lock: + for edge in edges: + edge_props = edge.get("properties", {}) + # Restore validity windows — ContextEdge.to_dict() writes them at top level + valid_from = edge.get("valid_from") or edge_props.get("valid_from") + valid_until = edge.get("valid_until") or edge_props.get("valid_until") + internal_edge = ContextEdge( + source_id=edge.get("source_id"), + target_id=edge.get("target_id"), + edge_type=edge.get("type", "related_to"), + weight=edge.get("weight", 1.0), + metadata=edge_props, + valid_from=valid_from, + valid_until=valid_until, + ) - if self._add_internal_edge(internal_edge): - count += 1 + if self._add_internal_edge(internal_edge): + count += 1 return count def __contains__(self, node_id: object) -> bool: if not isinstance(node_id, str): return False - return node_id in self.nodes + with self._lock: + return node_id in self.nodes def has_node(self, node_id: str) -> bool: - return node_id in self.nodes + with self._lock: + return node_id in self.nodes def neighbors(self, node_id: str) -> List[Dict[str, Any]]: return self.get_neighbors(node_id, hops=1) @@ -411,55 +419,61 @@ class ContextGraph: node_id: str, relationship_types: Optional[List[str]] = None, ) -> List[str]: - if node_id not in self.nodes: - return [] + with self._lock: + if node_id not in self.nodes: + return [] - rel_filter = set(relationship_types) if relationship_types else None - neighbor_ids: List[str] = [] - for edge in self._adjacency.get(node_id, []): - if rel_filter is None or edge.edge_type in rel_filter: - neighbor_ids.append(edge.target_id) - return neighbor_ids + rel_filter = set(relationship_types) if relationship_types else None + neighbor_ids: List[str] = [] + for edge in self._adjacency.get(node_id, []): + if rel_filter is None or edge.edge_type in rel_filter: + neighbor_ids.append(edge.target_id) + return neighbor_ids def get_nodes_by_label(self, label: str) -> List[Dict[str, Any]]: result = [] - for nid in self.node_type_index.get(label, set()): - node = self.nodes.get(nid) - if node: - result.append({ - "id": node.node_id, - "content": node.content, - "type": node.node_type, - "metadata": (node.properties or {}).copy(), - }) + with self._lock: + for nid in self.node_type_index.get(label, set()): + node = self.nodes.get(nid) + if node: + result.append({ + "id": node.node_id, + "content": node.content, + "type": node.node_type, + "metadata": (node.properties or {}).copy(), + }) return result def get_node_property(self, node_id: str, property_name: str) -> Any: - node = self.nodes.get(node_id) - if not node: - return None - return node.properties.get(property_name) + with self._lock: + node = self.nodes.get(node_id) + if not node: + return None + return node.properties.get(property_name) def get_node_attributes(self, node_id: str) -> Dict[str, Any]: - node = self.nodes.get(node_id) - if not node: - return {} - return node.properties.copy() + with self._lock: + node = self.nodes.get(node_id) + if not node: + return {} + return node.properties.copy() def add_node_attribute(self, node_id: str, attributes: Dict[str, Any]) -> None: - node = self.nodes.get(node_id) - if not node: - return - node.properties.update(attributes) - node.metadata.update(attributes) + with self._lock: + node = self.nodes.get(node_id) + if not node: + return + node.properties.update(attributes) + node.metadata.update(attributes) def get_edge_data(self, source_id: str, target_id: str) -> Dict[str, Any]: - for edge in self._adjacency.get(source_id, []): - if edge.target_id == target_id: - data = edge.metadata.copy() - data["type"] = edge.edge_type - data["weight"] = edge.weight - return data + with self._lock: + for edge in self._adjacency.get(source_id, []): + if edge.target_id == target_id: + data = edge.metadata.copy() + data["type"] = edge.edge_type + data["weight"] = edge.weight + return data return {} def get_neighbors( @@ -468,6 +482,8 @@ class ContextGraph: hops: int = 1, relationship_types: Optional[List[str]] = None, min_weight: float = 0.0, + skip: int = 0, + limit: Optional[int] = None, ) -> List[Dict[str, Any]]: """ Get neighbors of a node. @@ -479,57 +495,66 @@ class ContextGraph: min_weight: Minimum edge weight required to traverse an edge (default 0.0 means all edges pass). Use e.g. ``min_weight=0.5`` to follow only strong/high-confidence relationships. + skip: Number of items to skip for pagination. + limit: Maximum items to return. Returns: List of dicts with neighbor info (id, type, content, relationship, weight, hop). """ - if node_id not in self.nodes: - return [] + with self._lock: + if node_id not in self.nodes: + return [] - neighbors: List[Dict[str, Any]] = [] - visited = {node_id} - queue = deque([(node_id, 0)]) - rel_filter = set(relationship_types) if relationship_types else None + neighbors: List[Dict[str, Any]] = [] + visited = {node_id} + queue = deque([(node_id, 0)]) + rel_filter = set(relationship_types) if relationship_types else None - while queue: - current_id, current_hop = queue.popleft() - if current_hop >= hops: - continue - - outgoing_edges = self._adjacency.get(current_id, []) - for edge in outgoing_edges: - if rel_filter is not None and edge.edge_type not in rel_filter: + while queue: + current_id, current_hop = queue.popleft() + if current_hop >= hops: continue - if edge.weight < min_weight: - continue - neighbor_id = edge.target_id - if neighbor_id in visited: - continue - visited.add(neighbor_id) - queue.append((neighbor_id, current_hop + 1)) - node = self.nodes.get(neighbor_id) - if not node: - continue - neighbors.append( - { - "id": node.node_id, - "type": node.node_type, - "content": node.content, - "relationship": edge.edge_type, - "weight": edge.weight, - "hop": current_hop + 1, - } - ) + outgoing_edges = self._adjacency.get(current_id, []) + for edge in outgoing_edges: + if rel_filter is not None and edge.edge_type not in rel_filter: + continue + if edge.weight < min_weight: + continue + neighbor_id = edge.target_id + if neighbor_id in visited: + continue + visited.add(neighbor_id) + queue.append((neighbor_id, current_hop + 1)) - return neighbors + node = self.nodes.get(neighbor_id) + if not node: + continue + neighbors.append( + { + "id": node.node_id, + "type": node.node_type, + "content": node.content, + "relationship": edge.edge_type, + "weight": edge.weight, + "hop": current_hop + 1, + } + ) - def query(self, query: str) -> List[Dict[str, Any]]: + if limit is not None: + return neighbors[skip: skip + limit] + return neighbors[skip:] + + def query( + self, query: str, skip: int = 0, limit: Optional[int] = None + ) -> List[Dict[str, Any]]: """ Execute a simple keyword search query on the graph nodes. Args: query: Keyword query string + skip: Items to skip + limit: Max items to return Returns: List of matching node dicts @@ -537,22 +562,26 @@ class ContextGraph: results = [] query_lower = query.lower().split() - for node in self.nodes.values(): - content_lower = node.content.lower() - if any(word in content_lower for word in query_lower): - # Calculate simple score - overlap = sum(1 for word in query_lower if word in content_lower) - score = overlap / len(query_lower) if query_lower else 0.0 + with self._lock: + for node in self.nodes.values(): + content_lower = node.content.lower() + if any(word in content_lower for word in query_lower): + # Calculate simple score + overlap = sum(1 for word in query_lower if word in content_lower) + score = overlap / len(query_lower) if query_lower else 0.0 - results.append( - { - "node": node.to_dict(), - "score": score, - "content": node.content, - } - ) + results.append( + { + "node": node.to_dict(), + "score": score, + "content": node.content, + } + ) - return sorted(results, key=lambda x: x["score"], reverse=True) + sorted_res = sorted(results, key=lambda x: x["score"], reverse=True) + if limit is not None: + return sorted_res[skip: skip + limit] + return sorted_res[skip:] def add_node( self, @@ -574,17 +603,18 @@ class ContextGraph: content = content or node_id valid_from = properties.pop("valid_from", None) valid_until = properties.pop("valid_until", None) - return self._add_internal_node( - ContextNode( - node_id=node_id, - node_type=node_type, - content=content, - metadata=properties, - properties=properties, - valid_from=valid_from, - valid_until=valid_until, + with self._lock: + return self._add_internal_node( + ContextNode( + node_id=node_id, + node_type=node_type, + content=content, + metadata=properties, + properties=properties, + valid_from=valid_from, + valid_until=valid_until, + ) ) - ) def add_edge( self, @@ -607,17 +637,18 @@ class ContextGraph: """ valid_from = properties.pop("valid_from", None) valid_until = properties.pop("valid_until", None) - return self._add_internal_edge( - ContextEdge( - source_id=source_id, - target_id=target_id, - edge_type=edge_type, - weight=weight, - metadata=properties, - valid_from=valid_from, - valid_until=valid_until, + with self._lock: + return self._add_internal_edge( + ContextEdge( + source_id=source_id, + target_id=target_id, + edge_type=edge_type, + weight=weight, + metadata=properties, + valid_from=valid_from, + valid_until=valid_until, + ) ) - ) def save_to_file(self, path: str) -> None: """ @@ -628,25 +659,26 @@ class ContextGraph: """ import json - # Serialise cross-graph link metadata (object references are not serialisable, - # so we store other_graph_id; callers can reconnect with resolve_links()). - links_data = [] - for link_id, (other_graph, source_node_id, target_node_id) in self._linked_graphs.items(): - links_data.append( - { - "link_id": link_id, - "source_node_id": source_node_id, - "target_node_id": target_node_id, - "other_graph_id": other_graph.graph_id, - } - ) + with self._lock: + # Serialise cross-graph link metadata (object references are not serialisable, + # so we store other_graph_id; callers can reconnect with resolve_links()). + links_data = [] + for link_id, (other_graph, source_node_id, target_node_id) in self._linked_graphs.items(): + links_data.append( + { + "link_id": link_id, + "source_node_id": source_node_id, + "target_node_id": target_node_id, + "other_graph_id": other_graph.graph_id, + } + ) - data = { - "graph_id": self.graph_id, - "nodes": [node.to_dict() for node in self.nodes.values()], - "edges": [edge.to_dict() for edge in self.edges], - "links": links_data, - } + data = { + "graph_id": self.graph_id, + "nodes": [node.to_dict() for node in self.nodes.values()], + "edges": [edge.to_dict() for edge in self.edges], + "links": links_data, + } with open(path, "w", encoding="utf-8") as f: json.dump(data, f, indent=2, ensure_ascii=False) @@ -670,72 +702,80 @@ class ContextGraph: with open(path, "r", encoding="utf-8") as f: data = json.load(f) - # Clear existing - self.nodes.clear() - self.edges.clear() - self._adjacency.clear() - self.node_type_index.clear() - self.edge_type_index.clear() - self._linked_graphs.clear() - self._unresolved_links.clear() + with self._lock: + # Clear existing + self.nodes.clear() + self.edges.clear() + self._adjacency.clear() + self.node_type_index.clear() + self.edge_type_index.clear() + self._linked_graphs.clear() + self._unresolved_links.clear() - # Restore stable graph identity - if "graph_id" in data: - self.graph_id = data["graph_id"] + if "graph_id" in data: + self.graph_id = data["graph_id"] - # Load nodes - nodes = data.get("nodes", []) - self.add_nodes(nodes) + + nodes = data.get("nodes", []) + self.add_nodes(nodes) - # Load edges - edges = data.get("edges", []) - self.add_edges(edges) + edges = data.get("edges", []) + self.add_edges(edges) - # Restore link metadata — object references require resolve_links() to reconnect - for link_meta in data.get("links", []): - link_id = link_meta.get("link_id") - if link_id: - self._unresolved_links[link_id] = link_meta + + for link_meta in data.get("links", []): + link_id = link_meta.get("link_id") + if link_id: + self._unresolved_links[link_id] = link_meta self.logger.info(f"Loaded context graph from {path}") def find_node(self, node_id: str) -> Optional[Dict[str, Any]]: """Find a node by ID.""" - node = self.nodes.get(node_id) - if node: - merged_metadata = {} - merged_metadata.update(getattr(node, "metadata", {}) or {}) - merged_metadata.update(getattr(node, "properties", {}) or {}) - return { - "id": node.node_id, - "type": node.node_type, - "content": node.content, - "metadata": merged_metadata, - } - return None + with self._lock: + node = self.nodes.get(node_id) + if node: + merged_metadata = {} + merged_metadata.update(getattr(node, "metadata", {}) or {}) + merged_metadata.update(getattr(node, "properties", {}) or {}) + return { + "id": node.node_id, + "type": node.node_type, + "content": node.content, + "metadata": merged_metadata, + } + return None - def find_nodes(self, node_type: Optional[str] = None) -> List[Dict[str, Any]]: + def find_nodes( + self, node_type: Optional[str] = None, skip: int = 0, limit: Optional[int] = None + ) -> List[Dict[str, Any]]: """Find nodes, optionally filtered by type.""" - if node_type: - node_ids = self.node_type_index.get(node_type, set()) - nodes = [self.nodes[nid] for nid in node_ids] - else: - nodes = self.nodes.values() + with self._lock: + if node_type: + node_ids = self.node_type_index.get(node_type, set()) + nodes = [self.nodes[nid] for nid in node_ids] + else: + nodes = list(self.nodes.values()) - return [ - { - "id": n.node_id, - "type": n.node_type, - "content": n.content, - "metadata": {**(getattr(n, "metadata", {}) or {}), **(getattr(n, "properties", {}) or {})}, - } - for n in nodes - ] + results = [ + { + "id": n.node_id, + "type": n.node_type, + "content": n.content, + "metadata": {**(getattr(n, "metadata", {}) or {}), **(getattr(n, "properties", {}) or {})}, + } + for n in nodes + ] + if limit is not None: + return results[skip: skip + limit] + return results[skip:] def find_active_nodes( self, node_type: Optional[str] = None, at_time: Optional[datetime] = None, + skip: int = 0, + limit: Optional[int] = None, ) -> List[Dict[str, Any]]: """ Find nodes that are currently active within their validity window. @@ -745,32 +785,38 @@ class ContextGraph: Args: node_type: Optional node type filter. at_time: Point in time to evaluate validity (defaults to ``datetime.utcnow()``). + skip: Items to skip + limit: Max items to return Returns: List of active node dicts (same format as :meth:`find_nodes`). """ now = at_time or datetime.utcnow() - if node_type: - node_ids = self.node_type_index.get(node_type, set()) - nodes_iter = [self.nodes[nid] for nid in node_ids if nid in self.nodes] - else: - nodes_iter = list(self.nodes.values()) + with self._lock: + if node_type: + node_ids = self.node_type_index.get(node_type, set()) + nodes_iter = [self.nodes[nid] for nid in node_ids if nid in self.nodes] + else: + nodes_iter = list(self.nodes.values()) - result = [] - for node in nodes_iter: - if node.is_active(now): - result.append( - { - "id": node.node_id, - "type": node.node_type, - "content": node.content, - "metadata": { - **(getattr(node, "metadata", {}) or {}), - **(getattr(node, "properties", {}) or {}), - }, - } - ) - return result + result = [] + for node in nodes_iter: + if node.is_active(now): + result.append( + { + "id": node.node_id, + "type": node.node_type, + "content": node.content, + "metadata": { + **(getattr(node, "metadata", {}) or {}), + **(getattr(node, "properties", {}) or {}), + }, + } + ) + + if limit is not None: + return result[skip: skip + limit] + return result[skip:] def link_graph( self, @@ -799,38 +845,36 @@ class ContextGraph: KeyError: If source_node_id is not in this graph or target_node_id is not in other_graph. """ - if source_node_id not in self.nodes: - raise KeyError(f"Source node '{source_node_id}' not found in this graph") - if target_node_id not in other_graph.nodes: - raise KeyError(f"Target node '{target_node_id}' not found in other_graph") + with self._lock: + if source_node_id not in self.nodes: + raise KeyError(f"Source node '{source_node_id}' not found in this graph") + if target_node_id not in other_graph.nodes: + raise KeyError(f"Target node '{target_node_id}' not found in other_graph") - link_id = str(uuid.uuid4()) - self._linked_graphs[link_id] = (other_graph, source_node_id, target_node_id) + link_id = str(uuid.uuid4()) + self._linked_graphs[link_id] = (other_graph, source_node_id, target_node_id) - # Create a dedicated marker node so it is clearly typed and does not pollute - # the entity namespace. _add_internal_edge auto-creates missing targets as - # "entity" nodes — by pre-inserting a "cross_graph_link" node we prevent that. - marker_node_id = f"__cross_graph_{link_id}" - self._add_internal_node( - ContextNode( - node_id=marker_node_id, - node_type="cross_graph_link", - content=f"Cross-graph link → {target_node_id}", - metadata={"cross_graph": True, "link_id": link_id, "target_node_id": target_node_id}, - properties={}, + marker_node_id = f"__cross_graph_{link_id}" + self._add_internal_node( + ContextNode( + node_id=marker_node_id, + node_type="cross_graph_link", + content=f"Cross-graph link → {target_node_id}", + metadata={"cross_graph": True, "link_id": link_id, "target_node_id": target_node_id}, + properties={}, + ) ) - ) - # Record a marker edge so the link shows up in graph traversal - self._add_internal_edge( - ContextEdge( - source_id=source_node_id, - target_id=marker_node_id, - edge_type=link_type, - weight=1.0, - metadata={"cross_graph": True, "link_id": link_id}, + + self._add_internal_edge( + ContextEdge( + source_id=source_node_id, + target_id=marker_node_id, + edge_type=link_type, + weight=1.0, + metadata={"cross_graph": True, "link_id": link_id}, + ) ) - ) - return link_id + return link_id def navigate_to(self, link_id: str) -> Tuple["ContextGraph", str]: """ @@ -885,59 +929,69 @@ class ContextGraph: resolved = g1b.resolve_links({g2b.graph_id: g2b}) """ resolved = 0 - for link_id, meta in list(self._unresolved_links.items()): - other_graph_id = meta.get("other_graph_id") - if other_graph_id in graphs: - other_graph = graphs[other_graph_id] - source_node_id = meta["source_node_id"] - target_node_id = meta["target_node_id"] - # Validate target node still exists in the restored graph - if target_node_id in other_graph.nodes: - self._linked_graphs[link_id] = (other_graph, source_node_id, target_node_id) - del self._unresolved_links[link_id] - resolved += 1 - else: - self.logger.warning( - f"resolve_links: target node '{target_node_id}' not found in " - f"graph '{other_graph_id}' for link '{link_id}'" - ) + with self._lock: + for link_id, meta in list(self._unresolved_links.items()): + other_graph_id = meta.get("other_graph_id") + if other_graph_id in graphs: + other_graph = graphs[other_graph_id] + source_node_id = meta["source_node_id"] + target_node_id = meta["target_node_id"] + # Validate target node still exists in the restored graph + if target_node_id in other_graph.nodes: + self._linked_graphs[link_id] = (other_graph, source_node_id, target_node_id) + del self._unresolved_links[link_id] + resolved += 1 + else: + self.logger.warning( + f"resolve_links: target node '{target_node_id}' not found in " + f"graph '{other_graph_id}' for link '{link_id}'" + ) return resolved - def find_edges(self, edge_type: Optional[str] = None) -> List[Dict[str, Any]]: + def find_edges( + self, edge_type: Optional[str] = None, skip: int = 0, limit: Optional[int] = None + ) -> List[Dict[str, Any]]: """Find edges, optionally filtered by type.""" - if edge_type: - edges = self.edge_type_index.get(edge_type, []) - else: - edges = self.edges + with self._lock: + if edge_type: + edges = self.edge_type_index.get(edge_type, []) + else: + edges = self.edges - return [ - { - "source": e.source_id, - "target": e.target_id, - "type": e.edge_type, - "weight": e.weight, - "metadata": e.metadata, - } - for e in edges - ] + results = [ + { + "source": e.source_id, + "target": e.target_id, + "type": e.edge_type, + "weight": e.weight, + "metadata": e.metadata, + } + for e in edges + ] + + if limit is not None: + return results[skip: skip + limit] + return results[skip:] def stats(self) -> Dict[str, Any]: """Get graph statistics.""" - return { - "node_count": len(self.nodes), - "edge_count": len(self.edges), - "node_types": {k: len(v) for k, v in self.node_type_index.items()}, - "edge_types": {k: len(v) for k, v in self.edge_type_index.items()}, - "density": self.density(), - } + with self._lock: + return { + "node_count": len(self.nodes), + "edge_count": len(self.edges), + "node_types": {k: len(v) for k, v in self.node_type_index.items()}, + "edge_types": {k: len(v) for k, v in self.edge_type_index.items()}, + "density": self.density(), + } def density(self) -> float: """Calculate graph density.""" - n = len(self.nodes) - if n < 2: - return 0.0 - max_edges = n * (n - 1) # Directed graph - return len(self.edges) / max_edges + with self._lock: + n = len(self.nodes) + if n < 2: + return 0.0 + max_edges = n * (n - 1) + return len(self.edges) / max_edges # --- Internal Helpers --- @@ -977,31 +1031,33 @@ class ContextGraph: def _add_internal_node(self, node: ContextNode) -> bool: """Internal method to add a node.""" - self.nodes[node.node_id] = node - # Handle edge case where node_type might be None or not a string - if hasattr(node, 'node_type') and isinstance(node.node_type, str): - self.node_type_index[node.node_type].add(node.node_id) - else: - # Use 'unknown' as fallback for invalid node_type - self.node_type_index['unknown'].add(node.node_id) - return True + with self._lock: + self.nodes[node.node_id] = node + # Handle edge case where node_type might be None or not a string + if hasattr(node, 'node_type') and isinstance(node.node_type, str): + self.node_type_index[node.node_type].add(node.node_id) + else: + # Use 'unknown' as fallback for invalid node_type + self.node_type_index['unknown'].add(node.node_id) + return True def _add_internal_edge(self, edge: ContextEdge) -> bool: """Internal method to add an edge.""" - # Ensure nodes exist - if edge.source_id not in self.nodes: - self._add_internal_node( - ContextNode(edge.source_id, "entity", edge.source_id) - ) - if edge.target_id not in self.nodes: - self._add_internal_node( - ContextNode(edge.target_id, "entity", edge.target_id) - ) + with self._lock: + # Ensure nodes exist + if edge.source_id not in self.nodes: + self._add_internal_node( + ContextNode(edge.source_id, "entity", edge.source_id) + ) + if edge.target_id not in self.nodes: + self._add_internal_node( + ContextNode(edge.target_id, "entity", edge.target_id) + ) - self.edges.append(edge) - self.edge_type_index[edge.edge_type].append(edge) - self._adjacency[edge.source_id].append(edge) - return True + self.edges.append(edge) + self.edge_type_index[edge.edge_type].append(edge) + self._adjacency[edge.source_id].append(edge) + return True # --- Builder Methods (Legacy/Utility) --- diff --git a/semantica/explorer/session.py b/semantica/explorer/session.py index de0ff3b2..5c61984b 100644 --- a/semantica/explorer/session.py +++ b/semantica/explorer/session.py @@ -128,20 +128,23 @@ class GraphSession: ) -> tuple[list[dict[str, Any]], int]: """Return a paginated slice of nodes and the total count.""" with self._lock: - all_nodes = self.graph.find_nodes(node_type=node_type) - - # Optional keyword filter - if search: - search_lower = search.lower() - all_nodes = [ - n for n in all_nodes - if search_lower in n.get("id", "").lower() - or search_lower in n.get("content", "").lower() - or search_lower in str(n.get("metadata", {})).lower() - ] - - total = len(all_nodes) - page = all_nodes[skip: skip + limit] + # We fetch all nodes if there is a search filter as we need to filter them in memory, + # otherwise we let the context graph handle the pagination natively. + if search: + all_nodes = self.graph.find_nodes(node_type=node_type) + search_lower = search.lower() + all_nodes = [ + n for n in all_nodes + if search_lower in n.get("id", "").lower() + or search_lower in n.get("content", "").lower() + or search_lower in str(n.get("metadata", {})).lower() + ] + total = len(all_nodes) + page = all_nodes[skip: skip + limit] + else: + total = self.graph.stats().get("node_types", {}).get(node_type, 0) if node_type else self.graph.stats().get("node_count", 0) + page = self.graph.find_nodes(node_type=node_type, skip=skip, limit=limit) + return page, total def get_edges( @@ -154,15 +157,18 @@ class GraphSession: ) -> tuple[list[dict[str, Any]], int]: """Return a paginated slice of edges and the total count.""" with self._lock: - all_edges = self.graph.find_edges(edge_type=edge_type) - - if source: - all_edges = [e for e in all_edges if e.get("source") == source] - if target: - all_edges = [e for e in all_edges if e.get("target") == target] - - total = len(all_edges) - page = all_edges[skip: skip + limit] + if source or target: + all_edges = self.graph.find_edges(edge_type=edge_type) + if source: + all_edges = [e for e in all_edges if e.get("source") == source] + if target: + all_edges = [e for e in all_edges if e.get("target") == target] + total = len(all_edges) + page = all_edges[skip: skip + limit] + else: + total = self.graph.stats().get("edge_types", {}).get(edge_type, 0) if edge_type else self.graph.stats().get("edge_count", 0) + page = self.graph.find_edges(edge_type=edge_type, skip=skip, limit=limit) + return page, total def get_neighbors(self, node_id: str, depth: int = 1) -> List[Dict[str, Any]]: