From 97b1c0300be85555cfe4559c73527d19e477e68f Mon Sep 17 00:00:00 2001 From: Ariel Vernaza Date: Thu, 16 Jul 2026 21:25:41 +0200 Subject: [PATCH] feat(nodes, server): migrate neo4j onto the graph base class Renames db_neo4j to graph_neo4j and derives it from GraphGlobalBase / GraphInstanceBase, the base extracted in #1584. The node keeps only what is specific to Neo4J -- the Bolt driver, READ access-mode sessions, and schema reflection -- while the lane handling, the natural-language-to-Cypher loop and the shared tools come from the base. IInstance drops from 431 to 129 lines and the local utils.py copy of the safety helpers is gone. The rename follows the graph_ convention and moves the node out of the Databases category into Graph Databases, so classType, protocol, path and the form field prefix change with it. get_schema is overridden to keep the label filter, and get_cypher stays as a deprecated alias for get_query: it mirrors the statement under the old `cypher` key, since get_query returns it under `query` and the alias would otherwise keep the name while silently changing the payload. Fixes carried in the migration: - _run_query refused unsafe Cypher and opened sessions in READ access mode. Bolt has no server-side read-only mode for a standalone instance, so is_cypher_safe is the primary defence here, not just defence-in-depth. - _run_query honours the base row contract: limit+1 rows for truncation detection, and read_row_cap as the ceiling when the pipeline lane calls with no limit. - affected_rows is counted unconditionally, since a Cypher write commonly also returns rows (CREATE (n) RETURN n). Pipelines are upgraded to services version 2 so saved configs still resolve. This covers both graph renames: db_neo4j -> graph_neo4j, and the earlier tool_falkordb -> graph_falkordb from #1584, which shipped without an upgrade path. Only the provider (and config type) changes -- the stored field names are untouched by either rename, and the new falkordb fields have defaults. Fixes #1600 --- docs/README-nodes.md | 8 +- nodes/src/nodes/db_neo4j/IGlobal.py | 384 --------------- nodes/src/nodes/db_neo4j/IInstance.py | 431 ---------------- nodes/src/nodes/db_neo4j/utils.py | 84 ---- nodes/src/nodes/graph_falkordb/README.md | 4 +- nodes/src/nodes/graph_neo4j/IGlobal.py | 291 +++++++++++ nodes/src/nodes/graph_neo4j/IInstance.py | 137 ++++++ .../nodes/{db_neo4j => graph_neo4j}/README.md | 32 +- .../{db_neo4j => graph_neo4j}/__init__.py | 0 .../nodes/{db_neo4j => graph_neo4j}/neo4j.svg | 0 .../requirements.txt | 0 .../{db_neo4j => graph_neo4j}/services.json | 52 +- nodes/test/test_graph_neo4j.py | 461 ++++++++++++++++++ packages/ai/src/ai/common/graph/__init__.py | 8 +- packages/docs/redirects.ts | 3 +- .../engLib/store/headers/services.hpp | 2 +- .../engLib/store/pipeline/pipeline_config.cpp | 12 + .../test/store/pipeline/pipeline_config.cpp | 23 + 18 files changed, 980 insertions(+), 952 deletions(-) delete mode 100644 nodes/src/nodes/db_neo4j/IGlobal.py delete mode 100644 nodes/src/nodes/db_neo4j/IInstance.py delete mode 100644 nodes/src/nodes/db_neo4j/utils.py create mode 100644 nodes/src/nodes/graph_neo4j/IGlobal.py create mode 100644 nodes/src/nodes/graph_neo4j/IInstance.py rename nodes/src/nodes/{db_neo4j => graph_neo4j}/README.md (77%) rename nodes/src/nodes/{db_neo4j => graph_neo4j}/__init__.py (100%) rename nodes/src/nodes/{db_neo4j => graph_neo4j}/neo4j.svg (100%) rename nodes/src/nodes/{db_neo4j => graph_neo4j}/requirements.txt (100%) rename nodes/src/nodes/{db_neo4j => graph_neo4j}/services.json (85%) create mode 100644 nodes/test/test_graph_neo4j.py diff --git a/docs/README-nodes.md b/docs/README-nodes.md index f989c27e1..96fc75789 100644 --- a/docs/README-nodes.md +++ b/docs/README-nodes.md @@ -150,13 +150,13 @@ channel; they have no data lanes and **bind to an agent** (see | `db_postgres` | answers, questions → answers, table, text | PostgreSQL and Supabase (insert + NL-to-SQL) | | `db_mysql` | answers, questions → answers, table, text | MySQL | | `db_clickhouse` | questions → answers, table, text | ClickHouse (NL-to-SQL) | -| `db_neo4j` | questions → answers, table, text | Neo4j graph database | ### Graph Databases -| Service | Data flow (in → out) | Description | -| ---------------- | -------------------------------- | --------------------------------------------------- | -| `graph_falkordb` | questions → answers, table, text | FalkorDB graph database (NL-to-Cypher, multi-graph) | +| Service | Data flow (in → out) | Description | +| ---------------- | -------------------------------- | ---------------------------------------------------- | +| `graph_falkordb` | questions → answers, table, text | FalkorDB graph database (NL-to-Cypher, multi-graph) | +| `graph_neo4j` | questions → answers, table, text | Neo4j graph database (NL-to-Cypher, READ access) | ### Document Processing diff --git a/nodes/src/nodes/db_neo4j/IGlobal.py b/nodes/src/nodes/db_neo4j/IGlobal.py deleted file mode 100644 index ad2c8512a..000000000 --- a/nodes/src/nodes/db_neo4j/IGlobal.py +++ /dev/null @@ -1,384 +0,0 @@ -# ============================================================================= -# MIT License -# Copyright (c) 2026 Aparavi Software AG -# -# Permission is hereby granted, free of charge, to any person obtaining a copy -# of this software and associated documentation files (the "Software"), to deal -# in the Software without restriction, including without limitation the rights -# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell -# copies of the Software, and to permit persons to whom the Software is -# furnished to do so, subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in -# all copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, -# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE -# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER -# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, -# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE -# SOFTWARE. -# ============================================================================= - -""" -Global (connection-level) state for the Neo4J database node. - -Manages the neo4j driver lifecycle, graph schema reflection, query execution, -and query validation. All Neo4J-specific knowledge lives here — IInstance -calls these methods without knowing the underlying driver. -""" - -from __future__ import annotations - -from typing import Any, Dict, List, Optional, Tuple - -import neo4j -from neo4j.exceptions import Neo4jError, ServiceUnavailable - -from rocketlib import IGlobalBase, debug, error, warning -from ai.common.config import Config - -from .utils import _is_cypher_safe - -DEFAULT_MAX_EXECUTE_ROWS = 25000 - - -class IGlobal(IGlobalBase): - """Neo4J-specific global connection state.""" - - QUERY_TIMEOUT: float = 30.0 # Maximum seconds a Cypher query may run before being aborted. - - # neo4j.Driver instance — opened in beginGlobal, closed in endGlobal. - driver: Optional[neo4j.Driver] = None - - # Cached graph schema: {'nodes': {label: [(prop, type), ...]}, - # 'relationships': [{type, start, end}, ...]} - graph_schema: Dict[str, Any] - - # Unprefixed config values set during beginGlobal. - uri: str = '' - database: str = 'neo4j' - # label: str = 'Row' - db_description: str = '' - max_validation_attempts: int = 5 - allow_execute: bool = False - max_execute_rows: int = DEFAULT_MAX_EXECUTE_ROWS - - # ------------------------------------------------------------------ - # Lifecycle - # ------------------------------------------------------------------ - - def beginGlobal(self) -> None: - """Open the Neo4J driver, verify connectivity, and cache the graph schema.""" - self.graph_schema = {} - config = Config.getNodeConfig(self.glb.logicalType, self.glb.connConfig) - - self.uri = config.get('uri', 'neo4j://localhost:7687').strip() - self.database = config.get('database', 'neo4j').strip() or 'neo4j' - # self.label = config.get('label', 'Row').strip() or 'Row' - self.db_description = config.get('db_description', '') - - try: - self.max_validation_attempts = int(config.get('max_attempts', 5)) - except (ValueError, TypeError): - self.max_validation_attempts = 5 - - # EXECUTE path is opt-in: a caller passing QuestionType.EXECUTE bypasses - # the LLM translation + _is_cypher_safe gate, so the node owner must - # explicitly enable the capability. Strings like 'false' / '0' must - # not be truthy here, so don't use bool() directly. - allow_execute = config.get('allow_execute', False) - if isinstance(allow_execute, str): - self.allow_execute = allow_execute.strip().lower() in {'1', 'true', 'yes', 'on'} - else: - self.allow_execute = bool(allow_execute) - - try: - self.max_execute_rows = max(1, int(config.get('max_execute_rows', DEFAULT_MAX_EXECUTE_ROWS))) - except (TypeError, ValueError): - self.max_execute_rows = DEFAULT_MAX_EXECUTE_ROWS - - auth = self._build_auth(config) - - try: - self.driver = neo4j.GraphDatabase.driver(self.uri, auth=auth) - # Verify DBMS connectivity and authentication. - self.driver.verify_connectivity() - # verify_connectivity() uses the home database — probe the configured - # database explicitly so a wrong name or missing permissions fail fast. - with self.driver.session(database=self.database) as session: - session.run('RETURN 1').consume() - except (ServiceUnavailable, Neo4jError) as e: - error(f'Neo4J connection failed: {e}') - raise - - self.graph_schema = self._reflect_schema() - debug(f'Neo4J connected: {self.uri}, database={self.database}') - - def endGlobal(self) -> None: - """Close the Neo4J driver and release the connection.""" - if self.driver is not None: - try: - self.driver.close() - except Exception as e: - warning(f'Error closing Neo4J driver: {e}') - finally: - self.driver = None - - # ------------------------------------------------------------------ - # Public helpers used by IInstance and the driver - # ------------------------------------------------------------------ - - def _run_query(self, cypher: str, params: Optional[Dict] = None, *, timeout: float = QUERY_TIMEOUT) -> List[Dict]: - """Execute a Cypher query and return rows as a list of plain dicts. - - Args: - cypher (str): The Cypher query to execute. - params (Optional[Dict]): Query parameters to bind into the Cypher statement. - timeout (float): Maximum seconds the query may run before being aborted. - Defaults to ``QUERY_TIMEOUT``. - - Returns: - List[Dict]: Result rows, each serialised to a plain Python dict. - - Raises: - neo4j.exceptions.Neo4jError: If the driver reports a query or connection error. - """ - if params is None: - params = {} - - # Defence-in-depth: block writes even if the caller skips the safety check. - if not _is_cypher_safe(cypher): - raise ValueError('Refusing to execute unsafe (write/admin) Cypher statement') - - with self.driver.session(database=self.database) as session: - result = session.run(neo4j.Query(cypher, timeout=timeout), params) - return [_record_to_dict(record) for record in result] - - def _run_query_raw(self, cypher: str, *, timeout: float = QUERY_TIMEOUT) -> Dict[str, Any]: - """Execute a raw Cypher statement without the ``_is_cypher_safe`` gate. - - Used by the EXECUTE path where the caller has accepted the risk of running - write/admin Cypher directly. Returns ``{'rows': [...], 'affected_rows': N}`` - to mirror the SQL ``_executeRawQuery`` shape — ``affected_rows`` is derived - from the result summary counters when no rows are returned (e.g. CREATE - without RETURN, DELETE). - - Raises: - neo4j.exceptions.Neo4jError: Caught at the IInstance handler per precedent. - """ - max_rows = self.max_execute_rows - with self.driver.session(database=self.database) as session: - result = session.run(neo4j.Query(cypher, timeout=timeout)) - rows = [_record_to_dict(record) for _, record in zip(range(max_rows + 1), result)] - if len(rows) > max_rows: - raise ValueError(f'EXECUTE query exceeded max_execute_rows={max_rows}') - counters = result.consume().counters - affected = ( - counters.nodes_created - + counters.nodes_deleted - + counters.relationships_created - + counters.relationships_deleted - + counters.properties_set - ) - return {'rows': rows, 'affected_rows': 0 if rows else affected} - - def _validate_query(self, cypher: str) -> Tuple[bool, str]: - """Run EXPLAIN on a Cypher statement to check syntax without executing it. - - Returns (True, '') on success or (False, error_message) on failure. - """ - try: - with self.driver.session(database=self.database) as session: - session.run(f'EXPLAIN {cypher}').consume() - return True, '' - except Neo4jError as e: - return False, str(e.message or e) - except Exception as e: - return False, str(e) - - def validateConfig(self) -> None: - """Test connectivity with a trivial read query; safe to call at save-time. - - Surfaces driver error messages via ``warning()`` so the user sees exactly - what went wrong (wrong password, unreachable host, etc.). - """ - config = Config.getNodeConfig(self.glb.logicalType, self.glb.connConfig) - uri = config.get('uri', 'neo4j://localhost:7687').strip() - database = config.get('database', 'neo4j').strip() or 'neo4j' - auth = self._build_auth(config) - - tmp_driver = None - try: - tmp_driver = neo4j.GraphDatabase.driver(uri, auth=auth, connection_timeout=5) - tmp_driver.verify_connectivity() - with tmp_driver.session(database=database) as session: - session.run('RETURN 1').consume() - except neo4j.exceptions.AuthError as e: - warning(f'Authentication failed: {e}') - return - except ServiceUnavailable as e: - warning(f'Could not connect to Neo4J at {uri}: {e}') - return - except Neo4jError as e: - warning(str(e.message or e)) - return - except Exception as e: - warning(str(e)) - return - finally: - if tmp_driver is not None: - try: - tmp_driver.close() - except Exception: - pass - - # ------------------------------------------------------------------ - # Schema reflection - # ------------------------------------------------------------------ - - def _reflect_schema(self) -> Dict[str, Any]: - """Reflect node labels, properties, and relationship types from Neo4J.""" - schema: Dict[str, Any] = {'nodes': {}, 'relationships': []} - - try: - with self.driver.session(database=self.database) as session: - # Node labels and their properties. - node_props: Dict[str, List[Tuple[str, str]]] = {} - try: - result = session.run('CALL db.schema.nodeTypeProperties()') - for record in result: - raw_label = record.get('nodeType', '') - label = raw_label.lstrip(':') if raw_label else '' - prop = record.get('propertyName') or '' - types = record.get('propertyTypes') or [] - prop_type = types[0] if types else 'ANY' - if label: - if label not in node_props: - node_props[label] = [] - if prop: - node_props[label].append((prop, prop_type)) - except Neo4jError: - # Fall back to just listing labels without properties. - try: - result = session.run('CALL db.labels()') - for record in result: - label = record.get('label') or '' - if label: - node_props[label] = [] - except Neo4jError: - pass - - schema['nodes'] = node_props - - # Relationship types with start/end labels from schema visualization. - try: - result = session.run('CALL db.schema.visualization()') - rels = [] - for record in result: - for rel in record.get('relationships') or []: - # neo4j.graph.Relationship exposes metadata as attributes, - # not dict keys — rel.get('type') returns None. - rel_type = getattr(rel, 'type', None) or '' - start_node = getattr(rel, 'start_node', None) - end_node = getattr(rel, 'end_node', None) - start_label = _first_label(start_node) - end_label = _first_label(end_node) - if rel_type: - rels.append( - { - 'type': rel_type, - 'start': start_label, - 'end': end_label, - } - ) - schema['relationships'] = rels - except Neo4jError: - # Fall back to listing relationship types without endpoints. - try: - result = session.run('CALL db.relationshipTypes()') - schema['relationships'] = [ - {'type': r.get('relationshipType', ''), 'start': '', 'end': ''} - for r in result - if r.get('relationshipType') - ] - except Neo4jError: - pass - - except Exception as e: - warning(f'Neo4J schema reflection failed: {e}') - - return schema - - # ------------------------------------------------------------------ - # Internal helpers - # ------------------------------------------------------------------ - - @staticmethod - def _build_auth(config: Dict[str, Any]) -> Any: - """Build a neo4j auth tuple or bearer-token auth from config. - - Respects the ``auth_method`` field: ``'token'`` → bearer auth, - anything else (default ``'userpass'``) → username/password. - - Args: - config (Dict[str, Any]): Unprefixed node config dict. - - Returns: - Any: A ``neo4j.bearer_auth`` token or a ``(user, password)`` tuple. - """ - auth_method = config.get('auth_method', 'userpass').strip() - if auth_method == 'token': - return neo4j.bearer_auth(config.get('token', '')) - user = config.get('user', 'neo4j').strip() or 'neo4j' - password = config.get('password', '') - return (user, password) - - -# --------------------------------------------------------------------------- -# Module-level helpers -# --------------------------------------------------------------------------- - - -def _record_to_dict(record: neo4j.Record) -> Dict[str, Any]: - """Convert a neo4j Record to a plain dict, serialising graph objects.""" - return {key: _serialize_value(record[key]) for key in record.keys()} - - -def _serialize_value(value: Any) -> Any: - """Recursively convert Neo4J-specific types to JSON-serializable values.""" - if value is None or isinstance(value, (str, int, float, bool)): - return value - if isinstance(value, (list, tuple)): - return [_serialize_value(v) for v in value] - if isinstance(value, dict): - return {k: _serialize_value(v) for k, v in value.items()} - # neo4j.graph.Node - if hasattr(value, 'labels') and hasattr(value, 'items'): - return {'_labels': list(value.labels), **{k: _serialize_value(v) for k, v in value.items()}} - # neo4j.graph.Relationship - if hasattr(value, 'type') and hasattr(value, 'items') and hasattr(value, 'start_node'): - return {'_type': value.type, **{k: _serialize_value(v) for k, v in value.items()}} - # neo4j.graph.Path - if hasattr(value, 'nodes') and hasattr(value, 'relationships'): - return { - 'nodes': [_serialize_value(n) for n in value.nodes], - 'relationships': [_serialize_value(r) for r in value.relationships], - } - # neo4j temporal types (DateTime, Date, Time, Duration) - if hasattr(value, 'isoformat'): - return value.isoformat() - if hasattr(value, '__str__'): - return str(value) - return value - - -def _first_label(node: Any) -> str: - """Extract the first label from a neo4j Node object, or '' if unavailable.""" - if node is None: - return '' - if hasattr(node, 'labels'): - labels = list(node.labels) - return labels[0] if labels else '' - return '' diff --git a/nodes/src/nodes/db_neo4j/IInstance.py b/nodes/src/nodes/db_neo4j/IInstance.py deleted file mode 100644 index 315ca9f05..000000000 --- a/nodes/src/nodes/db_neo4j/IInstance.py +++ /dev/null @@ -1,431 +0,0 @@ -# ============================================================================= -# MIT License -# Copyright (c) 2026 Aparavi Software AG -# ============================================================================= - -""" -Instance-level state for the Neo4J database node. - -Handles pipeline lane traffic (questions, table, answers), translates -natural-language questions to Cypher via the connected LLM, executes -queries, and inserts data as graph nodes. -""" - -from __future__ import annotations - -import json -from typing import Any, Dict, Optional - -from rocketlib import IInstanceBase, tool_function, error, warning -from ai.common.schema import Answer, Question, QuestionType -from ai.common.table import Table -from rocketlib.types import IInvokeLLM - -from .IGlobal import DEFAULT_MAX_EXECUTE_ROWS, IGlobal -from .utils import _is_cypher_safe, _parse_is_valid - - -class IInstance(IInstanceBase): - """Neo4J-specific instance state.""" - - IGlobal: IGlobal - - # ------------------------------------------------------------------ - # Tool methods - # ------------------------------------------------------------------ - - @tool_function( - input_schema={ - 'type': 'object', - 'required': ['question'], - 'properties': { - 'question': { - 'type': 'string', - 'description': 'Natural-language description of the graph data you want to retrieve', - }, - 'limit': { - 'type': 'integer', - 'description': f'Maximum number of rows to return (default 250, max {DEFAULT_MAX_EXECUTE_ROWS}).', - }, - }, - }, - output_schema={ - 'type': 'object', - 'properties': { - 'rows': { - 'type': 'array', - 'description': 'Result rows returned by the Cypher query.', - 'items': {'type': 'object'}, - }, - 'cypher': {'type': 'string', 'description': 'The generated Cypher query that was executed.'}, - 'row_limit': {'type': 'integer', 'description': 'The row cap applied to this query.'}, - 'error': {'type': 'string', 'description': 'Error message if query generation or execution failed.'}, - 'answer': { - 'type': 'string', - 'description': 'LLM text response when the question is not a graph query.', - }, - }, - }, - description=( - 'Accepts a natural-language description of the graph data you want, ' - 'converts it to a safe Cypher MATCH query, executes it against the Neo4J ' - 'graph database, and returns the result rows. ' - 'No schema lookup or Cypher knowledge required — just describe what you need. ' - 'Results may be large — consider using peek or store.' - ), - ) - def get_data(self, args): - """Translate natural language to Cypher and execute.""" - if not isinstance(args, dict): - raise ValueError('Tool input must be a JSON object') - question = args.get('question') - if not question or not isinstance(question, str) or not question.strip(): - raise ValueError('"question" is required and must be a non-empty string') - - limit = _clamp_limit(args.get('limit')) - cypher_result = self.get_cypher({'question': question.strip(), 'limit': limit}) - if not cypher_result.get('valid'): - return cypher_result - - cypher = cypher_result['cypher'] - try: - rows = self.IGlobal._run_query(cypher) - except Exception as e: - return {'error': str(e), 'cypher': cypher, 'rows': []} - return {'rows': rows, 'cypher': cypher, 'row_limit': limit} - - @tool_function( - input_schema={ - 'type': 'object', - 'properties': { - 'label': { - 'type': 'string', - 'description': 'Optional node label to filter schema to a single node type.', - }, - }, - }, - output_schema={ - 'type': 'object', - 'properties': { - 'nodes': {'type': 'object', 'description': 'Map of node label to list of {property, type} objects.'}, - 'relationships': { - 'type': 'array', - 'description': 'List of {type, start, end} relationship descriptors.', - }, - 'database': {'type': 'string'}, - }, - }, - description=( - 'Returns the Neo4J graph schema: node labels with their properties and types, and relationship types with their start and end node labels. Do NOT call this preemptively — only use when get_data fails or returns unexpected results.' - ), - ) - def get_schema(self, args): - """Return the cached graph schema.""" - if args is not None and not isinstance(args, dict): - raise ValueError('Tool input must be a JSON object or empty') - if not args: - args = {} - - schema = self.IGlobal.graph_schema - label_filter = args.get('label') - - nodes = schema.get('nodes', {}) - rels = schema.get('relationships', []) - - if label_filter: - filtered = nodes.get(label_filter) - if filtered is None: - return {'error': f'Node label :{label_filter} not found'} - nodes = {label_filter: filtered} - rels = [r for r in rels if r.get('start') == label_filter or r.get('end') == label_filter] - - return { - 'database': self.IGlobal.database, - 'nodes': {label: [{'property': p, 'type': t} for p, t in props] for label, props in nodes.items()}, - 'relationships': rels, - } - - @tool_function( - input_schema={ - 'type': 'object', - 'required': ['question'], - 'properties': { - 'question': { - 'type': 'string', - 'description': 'Natural-language question to convert into a Cypher query', - }, - }, - }, - output_schema={ - 'type': 'object', - 'properties': { - 'cypher': {'type': 'string', 'description': 'The generated Cypher MATCH statement.'}, - 'valid': {'type': 'boolean', 'description': 'Whether a valid, safe Cypher query was generated.'}, - 'error': {'type': 'string', 'description': 'Error message if the generated Cypher was unsafe.'}, - 'answer': { - 'type': 'string', - 'description': 'LLM text response when the question is not a graph query.', - }, - }, - }, - description=( - 'Accepts a natural-language description and returns the equivalent Cypher MATCH statement without executing it. Only use when the user explicitly asks to see the Cypher — for actual data retrieval, use get_data instead.' - ), - ) - def get_cypher(self, args): - """Translate natural language to Cypher without executing.""" - if not isinstance(args, dict): - raise ValueError('Tool input must be a JSON object') - question = args.get('question') - if not question or not isinstance(question, str) or not question.strip(): - raise ValueError('"question" is required and must be a non-empty string') - - limit = _clamp_limit(args.get('limit')) - result = self._buildCypherQuery(question.strip(), limit=limit) - is_valid = _parse_is_valid(result.get('isValid', False)) - cypher = result.get('query', '') - - if is_valid and cypher and _is_cypher_safe(cypher): - return {'cypher': cypher, 'valid': True} - elif is_valid and cypher: - return {'error': 'Generated query contains unsafe Cypher', 'cypher': cypher, 'valid': False} - else: - return {'answer': cypher, 'valid': False} - - # ------------------------------------------------------------------ - # Pipeline lane handlers - # ------------------------------------------------------------------ - - def writeQuestions(self, question: Question) -> None: - """Translate a natural-language question to Cypher, execute it, emit results.""" - question_text = question.questions[0].text if question.questions else None - - if not question_text: - warning('No question text provided.') - return - - lanes = self.instance.getListeners() - - # DIALECT: dialect-discovery request — emit {'dialect': 'neo4j'} on the - # answers lane so SDK callers can tell they're talking to a graph DB - # rather than a relational one. - if question.type == QuestionType.DIALECT: - if 'answers' in lanes: - answer = Answer() - answer.setAnswer(json.dumps({'dialect': 'neo4j'})) - self.instance.writeAnswers(answer) - return - - # EXECUTE: caller passes raw Cypher; bypass LLM translation + safety check. - if question.type == QuestionType.EXECUTE: - if not self.IGlobal.allow_execute: - warning('QuestionType.EXECUTE is disabled for this node (set allow_execute=true to enable).') - return - try: - execute_result = self.IGlobal._run_query_raw(question_text) - rows = execute_result['rows'] - affected = execute_result['affected_rows'] - markdown = self._formatResultAsMarkdown(rows) if rows else None - - if 'text' in lanes: - self.instance.writeText(markdown if markdown else f'{affected} rows affected') - - if 'table' in lanes and rows: - self.instance.writeTable(markdown) - - if 'answers' in lanes: - answer = Answer() - answer.setAnswer(json.dumps(execute_result, default=str)) - self.instance.writeAnswers(answer) - except Exception as e: - error(f'Error handling execute question: {e}') - return - - try: - query_json = self._buildCypherQuery(question_text) - is_valid = _parse_is_valid(query_json.get('isValid', False)) - cypher = query_json.get('query', '') - - executed = is_valid and bool(cypher) and _is_cypher_safe(cypher) - - if executed: - result = self.IGlobal._run_query(cypher) - else: - result = cypher - - if 'text' in lanes: - self.instance.writeText(str(result)) - - if 'table' in lanes and executed and result: - self.instance.writeTable(self._formatResultAsMarkdown(result)) - - if 'answers' in lanes: - answer = Answer() - if executed and result: - answer.setAnswer(self._formatResultAsMarkdown(result)) - else: - answer.setAnswer(str(result)) - self.instance.writeAnswers(answer) - - except Exception as e: - error(f'Error handling question: {e}') - - # ------------------------------------------------------------------ - # Cypher query building - # ------------------------------------------------------------------ - - def _buildCypherQuery(self, question_text: str, *, limit: int = 250) -> Dict: - """Generate a Cypher query, validate with EXPLAIN, retry on failure.""" - previous_cypher: Optional[str] = None - last_error: Optional[str] = None - result: Dict = {} - - for attempt in range(self.IGlobal.max_validation_attempts): - result = self._buildCypherQueryOnce( - question_text, - limit=limit, - previous_cypher=previous_cypher, - error_message=last_error, - ) - - is_valid = _parse_is_valid(result.get('isValid', False)) - cypher = result.get('query', '') - - if not is_valid or not cypher or not _is_cypher_safe(cypher): - return result - - ok, explain_error = self.IGlobal._validate_query(cypher) - if ok: - return result - - warning( - f'Cypher validation attempt {attempt + 1}/{self.IGlobal.max_validation_attempts} failed: {explain_error}' - ) - previous_cypher = cypher - last_error = explain_error - - warning( - f'Cypher validation failed after {self.IGlobal.max_validation_attempts} attempt(s); returning last result.' - ) - return result - - def _buildCypherQueryOnce( - self, - question_text: str, - *, - limit: int = 250, - previous_cypher: Optional[str] = None, - error_message: Optional[str] = None, - ) -> Dict: - """Single LLM call: translate a natural-language question into Cypher.""" - - def describe_schema(schema: Dict) -> str: - lines = [] - nodes = schema.get('nodes', {}) - for label, props in nodes.items(): - lines.append(f'Node :{label}') - for prop_name, prop_type in props: - lines.append(f' {prop_name}: {prop_type}') - lines.append('') - for rel in schema.get('relationships', []): - rel_type = rel.get('type', '') - start = rel.get('start', '') - end = rel.get('end', '') - if start and end: - lines.append(f'Relationship :{rel_type}') - lines.append(f' (:{start})-[:{rel_type}]->(:{end})') - elif rel_type: - lines.append(f'Relationship :{rel_type}') - lines.append('') - return '\n'.join(lines).strip() - - schema_description = describe_schema(self.IGlobal.graph_schema) - - question: Question = Question(type=QuestionType.QUESTION, role='You are a technical assistant.') - question.addQuestion(question_text) - - if self.IGlobal.db_description: - question.addContext(f'Graph description: {self.IGlobal.db_description}') - - if schema_description: - question.addContext(schema_description) - - question.expectJson = True - - question.addInstruction( - 'Cypher Query Generation Guidelines', - 'Generate a Cypher query based only on the node labels and relationship types provided in context.', - ) - question.addInstruction( - 'LIMIT', f'Limit the results to {limit} rows using LIMIT {limit} at the end of the query.' - ) - question.addInstruction( - 'Formatting', - 'Do not wrap the Cypher query in markdown (e.g., no triple backticks) and abide by formatting in the provided examples.', - ) - question.addInstruction( - 'Commands', - 'You are only permitted to use MATCH, OPTIONAL MATCH, WITH, WHERE, RETURN, ORDER BY, SKIP, and LIMIT. Avoid any write operations (CREATE, MERGE, DELETE, DETACH DELETE, SET, REMOVE, DROP).', - ) - question.addInstruction( - 'Ambiguity', - "If the user's question is ambiguous, make reasonable assumptions and attempt to craft a query. If you infer that the user's question is entirely unrelated to querying the graph, attempt to answer the question in a manner similar to the provided examples.", - ) - - question.addExample( - "Who are Alice's colleagues?", - { - 'isValid': 'true', - 'query': f"MATCH (alice:Person {{name: 'Alice'}})-[:WORKS_WITH]->(colleague:Person)\nRETURN colleague.name AS name, colleague.role AS role\nLIMIT {limit}", - }, - ) - question.addExample( - 'When did the Visigoths sack Rome?', - { - 'isValid': 'false', - 'query': 'The Visigoths sacked Rome in 410 AD, under the leadership of their king, Alaric I.', - }, - ) - - if previous_cypher and error_message: - question.addContext( - f'Your previous attempt produced the following Cypher:\n\n{previous_cypher}\n\nNeo4J rejected it with this error:\n\n{error_message}\n\nPlease fix the query and try again.' - ) - - result = self.instance.invoke(IInvokeLLM.Ask(question=question)) - - if not result or not result.answer: - raise ValueError('LLM failed to return a Cypher query.') - - return result.answer - - # ------------------------------------------------------------------ - # Result formatting - # ------------------------------------------------------------------ - - def _formatResultAsMarkdown(self, result: Any) -> str: - """Convert a query result (list of dicts) to a markdown table string.""" - headers = None - data = [] - - if isinstance(result, list) and result: - first = result[0] - if isinstance(first, dict): - headers = list(first.keys()) - data = [[str(row.get(key, '')) for key in headers] for row in result] - elif isinstance(first, (list, tuple)): - data = [[str(cell) for cell in row] for row in result] - else: - data = [[str(row)] for row in result] - else: - data = [[str(result)]] - - return Table.generate_markdown_table(data, headers) - - -def _clamp_limit(raw_limit) -> int: - """Clamp a user-supplied limit to [1, DEFAULT_MAX_EXECUTE_ROWS].""" - try: - return max(1, min(int(raw_limit), DEFAULT_MAX_EXECUTE_ROWS)) if raw_limit is not None else 250 - except (ValueError, TypeError): - return 250 diff --git a/nodes/src/nodes/db_neo4j/utils.py b/nodes/src/nodes/db_neo4j/utils.py deleted file mode 100644 index 0884378bb..000000000 --- a/nodes/src/nodes/db_neo4j/utils.py +++ /dev/null @@ -1,84 +0,0 @@ -# ============================================================================= -# MIT License -# Copyright (c) 2026 Aparavi Software AG -# -# Permission is hereby granted, free of charge, to any person obtaining a copy -# of this software and associated documentation files (the "Software"), to deal -# in the Software without restriction, including without limitation the rights -# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell -# copies of the Software, and to permit persons to whom the Software is -# furnished to do so, subject to the following conditions: -# -# The above copyright notice and this permission notice shall be included in -# all copies or substantial portions of the Software. -# -# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR -# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, -# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE -# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER -# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, -# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE -# SOFTWARE. -# ============================================================================= - -"""Shared utilities for the Neo4J database node.""" - -from __future__ import annotations - -import re - - -# --------------------------------------------------------------------------- -# Cypher safety check — read-only queries only -# --------------------------------------------------------------------------- - -_UNSAFE_CYPHER = re.compile( - r'\b(?:CREATE|MERGE|DELETE|DETACH\s+DELETE|SET|REMOVE|DROP|FOREACH|LOAD\s+CSV|' - r'CALL\s+apoc\.(?:create|merge|delete|periodic\.commit|refactor|load))\b', - re.IGNORECASE, -) - - -def _parse_is_valid(value: object) -> bool: - """Normalise an ``isValid`` value from LLM JSON output to a Python bool. - - Args: - value (object): Raw value from the LLM response dict — may be a - ``bool`` (``True``/``False``) or a ``str`` (``'true'``/``'false'``). - - Returns: - bool: ``True`` only when the value is the boolean ``True`` or the - case-insensitive string ``'true'``. - """ - if isinstance(value, bool): - return value - return str(value).lower() == 'true' - - -def _is_cypher_safe(cypher: str) -> bool: - """Return True when the Cypher statement is read-only (MATCH/RETURN/CALL schema only). - - Args: - cypher (str): The Cypher statement to inspect. - - Returns: - bool: ``True`` if the statement contains no write or admin clauses, - ``False`` otherwise. - """ - # Strip both single-line and block comments before checking. - stripped = re.sub(r'//[^\n]*', '', cypher) - stripped = re.sub(r'/\*.*?\*/', '', stripped, flags=re.DOTALL) - return not bool(_UNSAFE_CYPHER.search(stripped)) - - -def _strip_ns(tool_name: str) -> str: - """Strip the ``'neo4j.'`` namespace prefix from a tool name. - - Args: - tool_name (str): Fully-qualified tool name (e.g. ``'neo4j.get_data'``). - - Returns: - str: Bare tool name without the namespace prefix. - """ - prefix = 'neo4j.' - return tool_name[len(prefix) :] if tool_name.startswith(prefix) else tool_name diff --git a/nodes/src/nodes/graph_falkordb/README.md b/nodes/src/nodes/graph_falkordb/README.md index 0cab5d170..97578ab3d 100644 --- a/nodes/src/nodes/graph_falkordb/README.md +++ b/nodes/src/nodes/graph_falkordb/README.md @@ -16,8 +16,8 @@ Queries a FalkorDB graph in two ways, and you can use either or both: It derives from `ai.common.graph.GraphInstanceBase`, the base class shared by every graph database node: the lane handling, the natural-language-to-Cypher loop and the common tools live there, so this node only implements what is specific to FalkorDB — the Redis-protocol client, -multi-graph selection, and server-side read-only execution. (`db_neo4j` predates the base class -and is migrated to it separately.) +multi-graph selection, and server-side read-only execution. (`graph_neo4j` derives from the same +base.) Queries are **read-only by default**: they run through `GRAPH.RO_QUERY`, so the FalkorDB server itself rejects any write clause (`CREATE`/`MERGE`/`SET`/`DELETE`) — the restriction is enforced diff --git a/nodes/src/nodes/graph_neo4j/IGlobal.py b/nodes/src/nodes/graph_neo4j/IGlobal.py new file mode 100644 index 000000000..ddfb35947 --- /dev/null +++ b/nodes/src/nodes/graph_neo4j/IGlobal.py @@ -0,0 +1,291 @@ +# ============================================================================= +# RocketRide Engine +# ============================================================================= +# MIT License +# Copyright (c) 2026 Aparavi Software AG +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +# ============================================================================= + +"""Neo4J graph node — global (shared) connection state. + +Thin driver over ``GraphGlobalBase``: it owns the Bolt driver and the Neo4J +schema procedures, and inherits the lifecycle, config parsing, schema caching +and LLM query loop from the base class. +""" + +from __future__ import annotations + +from typing import Any, Dict, List, Optional, Tuple + +import neo4j +from neo4j.exceptions import Neo4jError, ServiceUnavailable + +from rocketlib import error, warning + +from ai.common.graph import GraphGlobalBase, is_cypher_safe + +DEFAULT_URI = 'neo4j://localhost:7687' +DEFAULT_DATABASE = 'neo4j' + + +class IGlobal(GraphGlobalBase): + """Neo4J-specific global connection state.""" + + QUERY_TIMEOUT: float = 30.0 # Maximum seconds a Cypher query may run before being aborted. + + driver: Optional[neo4j.Driver] = None + + uri: str = '' + database: str = DEFAULT_DATABASE + + # ------------------------------------------------------------------ + # Driver hooks + # ------------------------------------------------------------------ + + def _open_driver(self, config: Dict[str, Any]) -> None: + """Open the Bolt driver and verify both the DBMS and the target database.""" + self.uri = str(config.get('uri') or DEFAULT_URI).strip() or DEFAULT_URI + self.database = str(config.get('database') or DEFAULT_DATABASE).strip() or DEFAULT_DATABASE + + try: + self.driver = neo4j.GraphDatabase.driver(self.uri, auth=self._build_auth(config)) + self.driver.verify_connectivity() + # verify_connectivity() only checks the home database — probe the + # configured one so a wrong name or missing permission fails fast. + with self.driver.session(database=self.database) as session: + session.run('RETURN 1').consume() + except (ServiceUnavailable, Neo4jError) as e: + error(f'Neo4J connection failed: {e}') + raise + + def _close_driver(self) -> None: + if self.driver is not None: + try: + self.driver.close() + finally: + self.driver = None + + def _probe_connection(self, config: Dict[str, Any]) -> None: + """Save-time probe: connect and run RETURN 1, reporting failures as warnings.""" + uri = str(config.get('uri') or DEFAULT_URI).strip() or DEFAULT_URI + database = str(config.get('database') or DEFAULT_DATABASE).strip() or DEFAULT_DATABASE + + tmp_driver = None + try: + tmp_driver = neo4j.GraphDatabase.driver(uri, auth=self._build_auth(config), connection_timeout=5) + tmp_driver.verify_connectivity() + with tmp_driver.session(database=database) as session: + session.run('RETURN 1').consume() + except neo4j.exceptions.AuthError as e: + warning(f'Authentication failed: {e}') + except ServiceUnavailable as e: + warning(f'Could not connect to Neo4J at {uri}: {e}') + except Neo4jError as e: + warning(str(e.message or e)) + except Exception as e: + warning(str(e)) + finally: + if tmp_driver is not None: + try: + tmp_driver.close() + except Exception: + pass + + def _run_query(self, query: str, params: Optional[Dict] = None, limit: Optional[int] = None) -> List[Dict]: + """Execute a read-only Cypher query and return rows as plain dicts. + + Unlike FalkorDB's GRAPH.RO_QUERY, Bolt has no server-side read-only + execution mode for a standalone instance: READ access only guarantees + routing to a replica in a cluster. The is_cypher_safe gate is therefore + the primary defence here, not just defence-in-depth. + + With ``limit`` set the caller wants truncation detection, so return one + row past it (still bounded by ``read_row_cap``); with no limit (the lane + path) fall back to the ``read_row_cap`` ceiling. + """ + if not is_cypher_safe(query): + raise ValueError('Refusing to execute unsafe (write/admin) Cypher statement') + + cap = self.read_row_cap if limit is None else min(int(limit), self.read_row_cap) + 1 + with self.driver.session( + database=self.database, + default_access_mode=neo4j.READ_ACCESS, + ) as session: + result = session.run(neo4j.Query(query, timeout=self.QUERY_TIMEOUT), params or {}) + return [_record_to_dict(record) for _, record in zip(range(cap), result)] + + def _run_query_raw(self, query: str) -> Dict[str, Any]: + """Execute a raw statement with no safety gate (the EXECUTE path).""" + max_rows = self.max_execute_rows + with self.driver.session(database=self.database) as session: + result = session.run(neo4j.Query(query, timeout=self.QUERY_TIMEOUT)) + rows = [_record_to_dict(record) for _, record in zip(range(max_rows + 1), result)] + if len(rows) > max_rows: + raise ValueError(f'EXECUTE query exceeded max_execute_rows={max_rows}') + + counters = result.consume().counters + # Count writes unconditionally: a Cypher write commonly also returns rows + # (CREATE (n) RETURN n), so a non-empty result set does not mean nothing + # changed. The counters are 0 for a pure read anyway. + affected = ( + counters.nodes_created + + counters.nodes_deleted + + counters.relationships_created + + counters.relationships_deleted + + counters.properties_set + + counters.labels_added + + counters.labels_removed + ) + return {'rows': rows, 'affected_rows': affected} + + def _validate_query(self, query: str) -> Tuple[bool, str]: + """Run EXPLAIN to check syntax without executing the query.""" + try: + with self.driver.session(database=self.database) as session: + session.run(f'EXPLAIN {query}').consume() + return True, '' + except Neo4jError as e: + return False, str(e.message or e) + except Exception as e: + return False, str(e) + + def _reflect_schema(self) -> Dict[str, Any]: + """Reflect node labels, their properties, and relationship types.""" + schema: Dict[str, Any] = {'nodes': {}, 'relationships': []} + + try: + with self.driver.session(database=self.database) as session: + schema['nodes'] = self._reflect_nodes(session) + schema['relationships'] = self._reflect_relationships(session) + except Exception as e: + warning(f'Neo4J schema reflection failed: {e}') + + return schema + + # ------------------------------------------------------------------ + # Neo4J helpers + # ------------------------------------------------------------------ + + @staticmethod + def _reflect_nodes(session) -> Dict[str, List[Tuple[str, str]]]: + """Read labels and their property types, falling back to bare labels.""" + node_props: Dict[str, List[Tuple[str, str]]] = {} + try: + for record in session.run('CALL db.schema.nodeTypeProperties()'): + raw_label = record.get('nodeType', '') + label = raw_label.lstrip(':') if raw_label else '' + if not label: + continue + node_props.setdefault(label, []) + prop = record.get('propertyName') or '' + if prop: + types = record.get('propertyTypes') or [] + node_props[label].append((prop, types[0] if types else 'ANY')) + except Neo4jError: + try: + for record in session.run('CALL db.labels()'): + label = record.get('label') or '' + if label: + node_props[label] = [] + except Neo4jError: + pass + return node_props + + @staticmethod + def _reflect_relationships(session) -> List[Dict[str, str]]: + """Read relationship types with endpoints, falling back to bare types.""" + try: + rels = [] + for record in session.run('CALL db.schema.visualization()'): + for rel in record.get('relationships') or []: + # neo4j.graph.Relationship exposes metadata as attributes, + # not dict keys — rel.get('type') returns None. + rel_type = getattr(rel, 'type', None) or '' + if rel_type: + rels.append( + { + 'type': rel_type, + 'start': _first_label(getattr(rel, 'start_node', None)), + 'end': _first_label(getattr(rel, 'end_node', None)), + } + ) + return rels + except Neo4jError: + try: + return [ + {'type': r.get('relationshipType', ''), 'start': '', 'end': ''} + for r in session.run('CALL db.relationshipTypes()') + if r.get('relationshipType') + ] + except Neo4jError: + return [] + + @staticmethod + def _build_auth(config: Dict[str, Any]) -> Any: + """Build a bearer token or (user, password) tuple from the auth_method field.""" + auth_method = str(config.get('auth_method') or 'userpass').strip() + if auth_method == 'token': + return neo4j.bearer_auth(config.get('token', '')) + user = str(config.get('user') or 'neo4j').strip() or 'neo4j' + return (user, config.get('password', '')) + + +# --------------------------------------------------------------------------- +# Module-level helpers +# --------------------------------------------------------------------------- + + +def _record_to_dict(record: neo4j.Record) -> Dict[str, Any]: + """Convert a neo4j Record to a plain dict, serialising graph objects.""" + return {key: _serialize_value(record[key]) for key in record.keys()} + + +def _serialize_value(value: Any) -> Any: + """Recursively convert Neo4J-specific types to JSON-serializable values.""" + if value is None or isinstance(value, (str, int, float, bool)): + return value + if isinstance(value, (list, tuple)): + return [_serialize_value(v) for v in value] + if isinstance(value, dict): + return {k: _serialize_value(v) for k, v in value.items()} + # neo4j.graph.Node + if hasattr(value, 'labels') and hasattr(value, 'items'): + return {'_labels': list(value.labels), **{k: _serialize_value(v) for k, v in value.items()}} + # neo4j.graph.Relationship + if hasattr(value, 'type') and hasattr(value, 'items') and hasattr(value, 'start_node'): + return {'_type': value.type, **{k: _serialize_value(v) for k, v in value.items()}} + # neo4j.graph.Path + if hasattr(value, 'nodes') and hasattr(value, 'relationships'): + return { + 'nodes': [_serialize_value(n) for n in value.nodes], + 'relationships': [_serialize_value(r) for r in value.relationships], + } + # neo4j temporal types (DateTime, Date, Time, Duration) + if hasattr(value, 'isoformat'): + return value.isoformat() + return str(value) + + +def _first_label(node: Any) -> str: + """Extract the first label from a neo4j Node object, or '' if unavailable.""" + if node is None or not hasattr(node, 'labels'): + return '' + labels = list(node.labels) + return labels[0] if labels else '' diff --git a/nodes/src/nodes/graph_neo4j/IInstance.py b/nodes/src/nodes/graph_neo4j/IInstance.py new file mode 100644 index 000000000..f3528a6d2 --- /dev/null +++ b/nodes/src/nodes/graph_neo4j/IInstance.py @@ -0,0 +1,137 @@ +# ============================================================================= +# RocketRide Engine +# ============================================================================= +# MIT License +# Copyright (c) 2026 Aparavi Software AG +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +# ============================================================================= + +"""Neo4J graph node instance. + +Inherits the graph tool surface (``get_data``, ``get_schema``, ``get_query``, +``execute``, ``dialect``) and the questions lane from ``GraphInstanceBase``. +""" + +from __future__ import annotations + +from rocketlib import tool_function + +from ai.common.graph import GraphInstanceBase +from ai.common.utils import normalize_tool_input, optional_str + +from .IGlobal import IGlobal + + +class IInstance(GraphInstanceBase): + IGlobal: IGlobal + + def _db_display_name(self) -> str: + return 'Neo4J' + + def _db_dialect(self) -> str: + return 'neo4j' + + @tool_function( + input_schema={ + 'type': 'object', + 'properties': { + 'label': { + 'type': 'string', + 'description': 'Optional node label to filter the schema to a single node type.', + }, + }, + }, + output_schema={ + 'type': 'object', + 'properties': { + 'labels': {'type': 'array', 'items': {'type': 'string'}}, + 'nodes': {'type': 'object', 'description': 'Map of node label to its properties and types.'}, + 'relationships': {'type': 'array', 'items': {'type': 'object'}}, + 'database': {'type': 'string', 'description': 'The Neo4J database the schema was read from.'}, + 'error': {'type': 'string'}, + }, + }, + description=lambda self: ( + f'Returns the {self._db_display_name()} graph schema: node labels with their properties, and the ' + f'relationship types connecting them. Optionally filter to a single node label. Do NOT call this ' + f'preemptively -- only when get_data fails or returns unexpected results.' + ), + ) + def get_schema(self, args): + """Return the reflected schema, optionally filtered to one node label. + + Overrides the base ``get_schema`` to keep Neo4J's ``label`` filter and + ``database`` key, which the previous tool surface exposed. + """ + args = normalize_tool_input(args, tool_name='get_schema') + label_filter = optional_str(args, 'label', tool_name='get_schema') + + schema = self.IGlobal.graph_schema or {} + nodes = schema.get('nodes', {}) + rels = schema.get('relationships', []) + + if label_filter: + filtered = nodes.get(label_filter) + if filtered is None: + return {'error': f'Node label :{label_filter} not found'} + nodes = {label_filter: filtered} + rels = [r for r in rels if r.get('start') == label_filter or r.get('end') == label_filter] + + return { + 'database': self.IGlobal.database, + 'labels': list(nodes.keys()), + 'nodes': {label: [{'property': p, 'type': t} for p, t in props] for label, props in nodes.items()}, + 'relationships': rels, + } + + @tool_function( + input_schema={ + 'type': 'object', + 'required': ['question'], + 'properties': { + 'question': {'type': 'string', 'description': 'Natural-language question to convert into Cypher'}, + 'limit': {'type': 'integer'}, + }, + }, + output_schema={ + 'type': 'object', + 'properties': { + 'cypher': {'type': 'string', 'description': 'The generated Cypher MATCH statement.'}, + 'query': {'type': 'string'}, + 'valid': {'type': 'boolean'}, + 'error': {'type': 'string'}, + 'answer': {'type': 'string'}, + }, + }, + description=( + 'Deprecated alias for get_query, kept so agents written against the previous Neo4J tool ' + 'surface keep working. Prefer get_query.' + ), + ) + def get_cypher(self, args): + """Deprecated alias for the get_query tool. + + The old tool returned the statement under ``cypher``; get_query returns it + under ``query``. Mirror it so agents reading ``result['cypher']`` still work. + """ + out = self.get_query(args) + if isinstance(out, dict) and 'query' in out: + out = {**out, 'cypher': out['query']} + return out diff --git a/nodes/src/nodes/db_neo4j/README.md b/nodes/src/nodes/graph_neo4j/README.md similarity index 77% rename from nodes/src/nodes/db_neo4j/README.md rename to nodes/src/nodes/graph_neo4j/README.md index 991e1b9a6..67cd31c39 100644 --- a/nodes/src/nodes/db_neo4j/README.md +++ b/nodes/src/nodes/graph_neo4j/README.md @@ -1,10 +1,12 @@ -# db_neo4j +# graph_neo4j -A RocketRide database and tool node that answers natural-language questions against a Neo4J graph database by translating them to Cypher with a connected LLM. +A RocketRide graph and tool node that answers natural-language questions against a Neo4J graph database by translating them to Cypher with a connected LLM. ## What it does -Connects to a Neo4J instance over the Bolt protocol using the official **neo4j Python driver** and plays two roles. As a **pipeline node**, it receives natural-language questions on the `questions` lane, asks the connected LLM to generate a Cypher query, executes it against the graph, and emits results downstream on `table`, `text`, and `answers`. As a **tool node**, it exposes `get_data`, `get_schema`, and `get_cypher` directly to an agent. Designed for knowledge graph retrieval, entity linking, and graph-based RAG workflows. +Connects to a Neo4J instance over the Bolt protocol using the official **neo4j Python driver** and plays two roles. As a **pipeline node**, it receives natural-language questions on the `questions` lane, asks the connected LLM to generate a Cypher query, executes it against the graph, and emits results downstream on `table`, `text`, and `answers`. As a **tool node**, it exposes `get_data`, `get_query`, `get_schema`, `execute`, and `dialect` to an agent, plus a deprecated `get_cypher` alias for `get_query` kept for agents written against the previous tool surface. Designed for knowledge graph retrieval, entity linking, and graph-based RAG workflows. + +It derives from `ai.common.graph.GraphInstanceBase`, the base class shared by every graph database node (also used by `graph_falkordb`): the lane handling, the natural-language-to-Cypher loop and the common tools live there, so this node only implements what is specific to Neo4J — the Bolt driver, READ access-mode sessions, and schema reflection. The graph schema (node labels with property types, and relationship types with their start/end labels) is reflected once at pipeline start using `db.schema.nodeTypeProperties()` and `db.schema.visualization()` (falling back to `db.labels()` and `db.relationshipTypes()` on older servers) and included in every LLM prompt so Cypher is generated against the real graph structure. @@ -28,7 +30,7 @@ The node is **read-only by design**: every generated or supplied Cypher statemen |---------|-----------|-------------| | `questions` | `table`, `text`, `answers` | Translate question to Cypher, execute, emit results on each connected lane | -For a normal question, results are emitted as a markdown table on `table` and `answers`, and as plain text on `text`. If the LLM judges the question unrelated to the graph, its text reply is forwarded in place of a query result. +For a normal question, results are emitted as a markdown table on `table` and `answers`, and as plain text on `text`. If the LLM judges the question unrelated to the graph, its text reply is forwarded in place of a query result. Rows read on this lane are capped at 25,000 so a broad query cannot exhaust worker memory. Two special question types are handled on the `questions` lane: @@ -115,16 +117,16 @@ Connectivity and authentication are verified at pipeline start with `verify_conn | Field | Type | Description | Default | |---|---|---|---| -| `neo4jdb.allow_execute` | `boolean` | **Allow direct query execution**
Permit QuestionType.EXECUTE callers to run raw Cypher without LLM translation or safety checks. Leave OFF unless a trusted application explicitly needs to issue Cypher directly. | `false` | -| `neo4jdb.auth_method` | `string` | **Authentication** | `"userpass"` | -| `neo4jdb.database` | `string` | **Database name**
Name of the Neo4J database to connect to. Use 'neo4j' for the default database. | `"neo4j"` | -| `neo4jdb.db_description` | `string` | **Graph description**
What is this graph used for? Describe its content and domain, this helps the LLM generate more accurate Cypher queries. | `""` | -| `neo4jdb.max_attempts` | `integer` | **Max validation attempts**
Maximum number of times to re-ask the LLM if EXPLAIN rejects the generated Cypher query | `5` | -| `neo4jdb.password` | `string` | **Password**
Password to authenticate with the Neo4J instance. | | -| `neo4jdb.profile` | `string` | | `"default"` | -| `neo4jdb.token` | `string` | **Bearer token**
Bearer token for token-based authentication (e.g. Neo4J Aura cloud). | | -| `neo4jdb.uri` | `string` | **Connection URI**
Bolt URI for the Neo4J instance. Use neo4j:// or bolt:// for plaintext, neo4j+s:// or bolt+s:// for TLS (e.g. Neo4J Aura cloud) | `"neo4j://localhost:7687"` | -| `neo4jdb.user` | `string` | **User**
Username to authenticate with the Neo4J instance. | `"neo4j"` | +| `graph_neo4j.allow_execute` | `boolean` | **Allow direct query execution**
Permit QuestionType.EXECUTE callers to run raw Cypher without LLM translation or safety checks. Leave OFF unless a trusted application explicitly needs to issue Cypher directly. | `false` | +| `graph_neo4j.auth_method` | `string` | **Authentication** | `"userpass"` | +| `graph_neo4j.database` | `string` | **Database name**
Name of the Neo4J database to connect to. Use 'neo4j' for the default database. | `"neo4j"` | +| `graph_neo4j.db_description` | `string` | **Graph description**
What is this graph used for? Describe its content and domain, this helps the LLM generate more accurate Cypher queries. | `""` | +| `graph_neo4j.max_attempts` | `integer` | **Max validation attempts**
Maximum number of times to re-ask the LLM if EXPLAIN rejects the generated Cypher query | `5` | +| `graph_neo4j.password` | `string` | **Password**
Password to authenticate with the Neo4J instance. | | +| `graph_neo4j.profile` | `string` | | `"default"` | +| `graph_neo4j.token` | `string` | **Bearer token**
Bearer token for token-based authentication (e.g. Neo4J Aura cloud). | | +| `graph_neo4j.uri` | `string` | **Connection URI**
Bolt URI for the Neo4J instance. Use neo4j:// or bolt:// for plaintext, neo4j+s:// or bolt+s:// for TLS (e.g. Neo4J Aura cloud) | `"neo4j://localhost:7687"` | +| `graph_neo4j.user` | `string` | **User**
Username to authenticate with the Neo4J instance. | `"neo4j"` | ## Dependencies @@ -132,5 +134,5 @@ Connectivity and authentication are verified at pipeline start with `verify_conn ## Source -[ View source](https://github.com/rocketride-org/rocketride-server/tree/develop/nodes/src/nodes/db_neo4j) +[ View source](https://github.com/rocketride-org/rocketride-server/tree/develop/nodes/src/nodes/graph_neo4j) diff --git a/nodes/src/nodes/db_neo4j/__init__.py b/nodes/src/nodes/graph_neo4j/__init__.py similarity index 100% rename from nodes/src/nodes/db_neo4j/__init__.py rename to nodes/src/nodes/graph_neo4j/__init__.py diff --git a/nodes/src/nodes/db_neo4j/neo4j.svg b/nodes/src/nodes/graph_neo4j/neo4j.svg similarity index 100% rename from nodes/src/nodes/db_neo4j/neo4j.svg rename to nodes/src/nodes/graph_neo4j/neo4j.svg diff --git a/nodes/src/nodes/db_neo4j/requirements.txt b/nodes/src/nodes/graph_neo4j/requirements.txt similarity index 100% rename from nodes/src/nodes/db_neo4j/requirements.txt rename to nodes/src/nodes/graph_neo4j/requirements.txt diff --git a/nodes/src/nodes/db_neo4j/services.json b/nodes/src/nodes/graph_neo4j/services.json similarity index 85% rename from nodes/src/nodes/db_neo4j/services.json rename to nodes/src/nodes/graph_neo4j/services.json index a7f7963d6..03e31d80e 100644 --- a/nodes/src/nodes/db_neo4j/services.json +++ b/nodes/src/nodes/graph_neo4j/services.json @@ -8,12 +8,12 @@ // Required: // The protocol is the endpoint protocol // - "protocol": "db_neo4j://", + "protocol": "graph_neo4j://", // // Required: // Class type of the node - what it does // - "classType": ["database", "tool"], + "classType": ["graph", "tool"], // // Required: // Capabilities are flags that change the behavior of the underlying @@ -37,7 +37,7 @@ // The path is the executable/script code - it is node dependent // and is optional for most node // - "path": "nodes.db_neo4j", + "path": "nodes.graph_neo4j", // // Required: // The prefix map when added/removed when converting URLs <=> paths @@ -97,13 +97,13 @@ // in the shape // "fields": { - "neo4jdb.uri": { + "graph_neo4j.uri": { "type": "string", "title": "Connection URI", "default": "neo4j://localhost:7687", "description": "Bolt URI for the Neo4J instance. Use neo4j:// or bolt:// for plaintext, neo4j+s:// or bolt+s:// for TLS (e.g. Neo4J Aura cloud)" }, - "neo4jdb.auth_method": { + "graph_neo4j.auth_method": { "type": "string", "title": "Authentication", "default": "userpass", @@ -114,21 +114,21 @@ "conditional": [ { "value": "userpass", - "properties": ["neo4jdb.user", "neo4jdb.password"] + "properties": ["graph_neo4j.user", "graph_neo4j.password"] }, { "value": "token", - "properties": ["neo4jdb.token"] + "properties": ["graph_neo4j.token"] } ] }, - "neo4jdb.user": { + "graph_neo4j.user": { "type": "string", "title": "User", "default": "neo4j", "description": "Username to authenticate with the Neo4J instance." }, - "neo4jdb.password": { + "graph_neo4j.password": { "type": "string", "title": "Password", "description": "Password to authenticate with the Neo4J instance.", @@ -137,7 +137,7 @@ "ui:widget": "password" } }, - "neo4jdb.token": { + "graph_neo4j.token": { "type": "string", "title": "Bearer token", "description": "Bearer token for token-based authentication (e.g. Neo4J Aura cloud).", @@ -146,19 +146,19 @@ "ui:widget": "password" } }, - "neo4jdb.database": { + "graph_neo4j.database": { "type": "string", "title": "Database name", "default": "neo4j", "description": "Name of the Neo4J database to connect to. Use 'neo4j' for the default database." }, - // "neo4jdb.label": { + // "graph_neo4j.label": { // "type": "string", // "title": "Node label", // "default": "Row", // "description": "Label applied to graph nodes created when inserting data from the pipeline (e.g. Person, Event, Row)." // }, - "neo4jdb.db_description": { + "graph_neo4j.db_description": { "type": "string", "title": "Graph description", "default": "", @@ -167,7 +167,7 @@ "ui:widget": "textarea" } }, - "neo4jdb.max_attempts": { + "graph_neo4j.max_attempts": { "type": "integer", "title": "Max validation attempts", "default": 5, @@ -175,25 +175,25 @@ "maximum": 20, "description": "Maximum number of times to re-ask the LLM if EXPLAIN rejects the generated Cypher query" }, - "neo4jdb.allow_execute": { + "graph_neo4j.allow_execute": { "type": "boolean", "title": "Allow direct query execution", "default": false, "description": "Permit QuestionType.EXECUTE callers to run raw Cypher without LLM translation or safety checks. Leave OFF unless a trusted application explicitly needs to issue Cypher directly." }, - "neo4jdb.default": { + "graph_neo4j.default": { "object": "default", "properties": [ - "neo4jdb.db_description", - "neo4jdb.uri", - "neo4jdb.auth_method", - "neo4jdb.database", - // "neo4jdb.label", - "neo4jdb.max_attempts", - "neo4jdb.allow_execute" + "graph_neo4j.db_description", + "graph_neo4j.uri", + "graph_neo4j.auth_method", + "graph_neo4j.database", + // "graph_neo4j.label", + "graph_neo4j.max_attempts", + "graph_neo4j.allow_execute" ] }, - "neo4jdb.profile": { + "graph_neo4j.profile": { "hidden": true, "type": "string", "default": "default", @@ -201,7 +201,7 @@ "conditional": [ { "value": "default", - "properties": ["neo4jdb.default"] + "properties": ["graph_neo4j.default"] } ] } @@ -215,7 +215,7 @@ { "section": "Pipe", "title": "Neo4J", - "properties": ["neo4jdb.profile"] + "properties": ["graph_neo4j.profile"] } ] } diff --git a/nodes/test/test_graph_neo4j.py b/nodes/test/test_graph_neo4j.py new file mode 100644 index 000000000..e1d0f38b6 --- /dev/null +++ b/nodes/test/test_graph_neo4j.py @@ -0,0 +1,461 @@ +# ============================================================================= +# RocketRide Engine +# ============================================================================= +# MIT License +# Copyright (c) 2026 Aparavi Software AG +# +# Permission is hereby granted, free of charge, to any person obtaining a copy +# of this software and associated documentation files (the "Software"), to deal +# in the Software without restriction, including without limitation the rights +# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +# copies of the Software, and to permit persons to whom the Software is +# furnished to do so, subject to the following conditions: +# +# The above copyright notice and this permission notice shall be included in +# all copies or substantial portions of the Software. +# +# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +# SOFTWARE. +# ============================================================================= + +"""Tests for the Neo4J graph node. + +The node derives from ``ai.common.graph``, so the real base classes are loaded +from disk with ``rocketlib`` and the ``neo4j`` driver stubbed, mirroring +``test_graph_falkordb.py``. +""" + +from __future__ import annotations + +import importlib.util +import sys +import types + +from contextlib import contextmanager +from pathlib import Path +from typing import Iterator + +import pytest + +_REPO = Path(__file__).resolve().parents[2] +_GRAPH_PKG = _REPO / 'packages' / 'ai' / 'src' / 'ai' / 'common' / 'graph' +_UTILS_DIR = _REPO / 'packages' / 'ai' / 'src' / 'ai' / 'common' / 'utils' +_CONFIG_UTILS = _UTILS_DIR / 'config_utils.py' +_TOOL_ARGS = _UTILS_DIR / 'tool_args.py' +_NODE_DIR = _REPO / 'nodes' / 'src' / 'nodes' / 'graph_neo4j' + + +class _StubNeo4jError(Exception): + def __init__(self, message=''): + super().__init__(message) + self.message = message + + +class _StubServiceUnavailable(_StubNeo4jError): + pass + + +class _StubAuthError(_StubNeo4jError): + pass + + +class _StubBase: + """Stand-in for IInstanceBase / IGlobalBase — the engine supplies the real one.""" + + def __init__(self, *args, **kwargs): + pass + + +class _StubTable: + @staticmethod + def generate_markdown_table(data, headers=None): + return '\n'.join([' | '.join(map(str, row)) for row in data]) + + +_STUB_NAMES = ( + 'rocketlib', + 'rocketlib.types', + 'ai', + 'ai.common', + 'ai.common.schema', + 'ai.common.table', + 'ai.common.config', + 'ai.common.utils', + 'neo4j', + 'neo4j.exceptions', +) + + +def _install_stubs() -> None: + rocketlib = types.ModuleType('rocketlib') + rocketlib.IInstanceBase = _StubBase + rocketlib.IGlobalBase = _StubBase + rocketlib.tool_function = lambda **kwargs: lambda f: f + rocketlib.debug = lambda *a, **kw: None + rocketlib.error = lambda *a, **kw: None + rocketlib.warning = lambda *a, **kw: None + rocketlib.OPEN_MODE = types.SimpleNamespace(CONFIG=object()) + rocketlib_types = types.ModuleType('rocketlib.types') + rocketlib_types.IInvokeLLM = types.SimpleNamespace(Ask=lambda **kw: kw) + rocketlib.types = rocketlib_types + + ai_pkg = types.ModuleType('ai') + ai_pkg.__path__ = [] + common_pkg = types.ModuleType('ai.common') + common_pkg.__path__ = [] + + schema = types.ModuleType('ai.common.schema') + schema.Answer = type('Answer', (), {'setAnswer': lambda self, v: setattr(self, 'value', v)}) + schema.Question = type('Question', (), {}) + schema.QuestionType = types.SimpleNamespace(QUESTION=1, DIALECT=2, EXECUTE=3) + + table = types.ModuleType('ai.common.table') + table.Table = _StubTable + + config = types.ModuleType('ai.common.config') + config.Config = types.SimpleNamespace(getNodeConfig=lambda *a, **kw: {}) + + config_utils = _load_from_path('ai.common.config_utils', _CONFIG_UTILS) + utils = types.ModuleType('ai.common.utils') + utils.parse_bool = config_utils.parse_bool + utils.config_int = config_utils.config_int + + neo4j = types.ModuleType('neo4j') + neo4j.READ_ACCESS = 'READ' + neo4j.Record = object + neo4j.Query = lambda q, timeout=None: q + neo4j.bearer_auth = lambda token: ('bearer', token) + neo4j.GraphDatabase = types.SimpleNamespace(driver=lambda *a, **kw: None) + neo4j.graph = types.SimpleNamespace(Node=object, Relationship=object) + neo4j_exceptions = types.ModuleType('neo4j.exceptions') + neo4j_exceptions.Neo4jError = _StubNeo4jError + neo4j_exceptions.ServiceUnavailable = _StubServiceUnavailable + neo4j_exceptions.AuthError = _StubAuthError + neo4j.exceptions = neo4j_exceptions + + sys.modules.update( + { + 'rocketlib': rocketlib, + 'rocketlib.types': rocketlib_types, + 'ai': ai_pkg, + 'ai.common': common_pkg, + 'ai.common.schema': schema, + 'ai.common.table': table, + 'ai.common.config': config, + 'ai.common.utils': utils, + 'neo4j': neo4j, + 'neo4j.exceptions': neo4j_exceptions, + } + ) + + tool_args = _load_from_path('ai.common.tool_args', _TOOL_ARGS) + for _fn in ('normalize_tool_input', 'require_str', 'require_dict', 'require_int', 'optional_int', 'optional_str'): + setattr(utils, _fn, getattr(tool_args, _fn)) + + +@contextmanager +def _scoped_stubs() -> Iterator[None]: + saved = {name: sys.modules.get(name) for name in _STUB_NAMES} + _install_stubs() + try: + yield + finally: + for name, module in saved.items(): + if module is None: + sys.modules.pop(name, None) + else: + sys.modules[name] = module + + +def _load_from_path(name: str, path: Path, *, is_package: bool = False): + search = [str(path.parent)] if is_package else None + spec = importlib.util.spec_from_file_location(name, path, submodule_search_locations=search) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + sys.modules[name] = module + spec.loader.exec_module(module) + return module + + +def _load_node(): + with _scoped_stubs(): + _load_from_path('ai.common.graph', _GRAPH_PKG / '__init__.py', is_package=True) + iglobal = _load_from_path('nodes.graph_neo4j.IGlobal', _NODE_DIR / 'IGlobal.py') + pkg = types.ModuleType('nodes.graph_neo4j') + pkg.__path__ = [str(_NODE_DIR)] + pkg.IGlobal = iglobal + sys.modules['nodes.graph_neo4j'] = pkg + iinstance = _load_from_path('nodes.graph_neo4j.IInstance', _NODE_DIR / 'IInstance.py') + return iglobal, iinstance + + +_glb_mod, mod = _load_node() +graph_base = sys.modules['ai.common.graph'] + + +# --------------------------------------------------------------------------- +# Fakes: Bolt driver / session / result +# --------------------------------------------------------------------------- + + +class _FakeCounters: + _ATTRS = ( + 'nodes_created', + 'nodes_deleted', + 'relationships_created', + 'relationships_deleted', + 'properties_set', + 'labels_added', + 'labels_removed', + ) + + def __init__(self, **kw): + for a in self._ATTRS: + setattr(self, a, kw.get(a, 0)) + + +class _FakeRecord: + def __init__(self, data): + self._data = data + + def keys(self): + return list(self._data.keys()) + + def __getitem__(self, key): + return self._data[key] + + +class _FakeResult: + def __init__(self, records=None, counters=None, raise_error=None): + self._records = [_FakeRecord(r) for r in (records or [])] + self._counters = counters or _FakeCounters() + self._raise = raise_error + + def __iter__(self): + if self._raise: + raise self._raise + return iter(self._records) + + def consume(self): + if self._raise: + raise self._raise + return types.SimpleNamespace(counters=self._counters) + + +class _FakeSession: + def __init__(self, driver): + self._driver = driver + + def __enter__(self): + return self + + def __exit__(self, *a): + return False + + def run(self, query, params=None): + self._driver.calls.append((str(query), params)) + if self._driver.raise_error: + raise self._driver.raise_error + return self._driver.result + + +class _FakeDriver: + def __init__(self, result=None, raise_error=None): + self.result = result or _FakeResult() + self.raise_error = raise_error + self.calls = [] + self.session_kwargs = [] + + def session(self, **kwargs): + self.session_kwargs.append(kwargs) + return _FakeSession(self) + + def verify_connectivity(self): + pass + + def close(self): + pass + + +class _FakeGlobal(_glb_mod.IGlobal): + """The real Neo4J IGlobal with a fake Bolt driver.""" + + def __init__(self, driver, *, database='neo4j'): + self.driver = driver + self.database = database + self.max_execute_rows = 25000 + self.max_validation_attempts = 5 + self.allow_execute = False + self.db_description = '' + self.graph_schema = {'nodes': {}, 'relationships': []} + + +def _instance(global_state): + inst = mod.IInstance() + inst.IGlobal = global_state + return inst + + +# --------------------------------------------------------------------------- +# _run_query: read-only gate, READ access, limit/truncation +# --------------------------------------------------------------------------- + + +def test_run_query_refuses_write_cypher(): + glb = _FakeGlobal(_FakeDriver()) + with pytest.raises(ValueError): + glb._run_query('MATCH (n) DELETE n') + # The unsafe query must never reach the driver. + assert glb.driver.calls == [] + + +def test_run_query_uses_read_access_mode(): + driver = _FakeDriver(_FakeResult(records=[{'name': 'Alice'}])) + glb = _FakeGlobal(driver) + rows = glb._run_query('MATCH (n) RETURN n.name AS name') + assert rows == [{'name': 'Alice'}] + assert driver.session_kwargs[0].get('default_access_mode') == 'READ' + + +def test_run_query_returns_one_past_limit_for_truncation(): + records = [{'n': i} for i in range(10)] + glb = _FakeGlobal(_FakeDriver(_FakeResult(records=records))) + rows = glb._run_query('MATCH (n) RETURN n', limit=3) + # limit+1 rows so the caller can detect truncation. + assert len(rows) == 4 + + +def test_run_query_without_limit_falls_back_to_read_row_cap(): + records = [{'n': i} for i in range(10)] + glb = _FakeGlobal(_FakeDriver(_FakeResult(records=records))) + glb.max_execute_rows = 4 # read_row_cap defaults to this + # The pipeline `questions` lane calls with no limit; it must not stream + # the whole result set into worker memory. + assert len(glb._run_query('MATCH (n) RETURN n')) == 4 + + +def test_run_query_clamps_limit_to_read_row_cap(): + records = [{'n': i} for i in range(10)] + glb = _FakeGlobal(_FakeDriver(_FakeResult(records=records))) + glb.max_execute_rows = 2 + # A caller limit above the cap cannot lift it: cap+1 rows at most. + assert len(glb._run_query('MATCH (n) RETURN n', limit=9)) == 3 + + +# --------------------------------------------------------------------------- +# _run_query_raw: EXECUTE path, affected_rows counted unconditionally +# --------------------------------------------------------------------------- + + +def test_affected_rows_counted_even_when_write_returns_rows(): + # CREATE (n) RETURN n writes AND returns a row — affected_rows must not be 0. + result = _FakeResult(records=[{'n': 'Alice'}], counters=_FakeCounters(nodes_created=1, properties_set=1)) + glb = _FakeGlobal(_FakeDriver(result)) + out = glb._run_query_raw('CREATE (n:Person {name: "Alice"}) RETURN n') + assert out['rows'] + assert out['affected_rows'] == 2 # would have been 0 before the fix + + +def test_execute_raw_caps_at_max_execute_rows(): + glb = _FakeGlobal(_FakeDriver(_FakeResult(records=[{'n': i} for i in range(10)]))) + glb.max_execute_rows = 3 + with pytest.raises(ValueError, match='max_execute_rows'): + glb._run_query_raw('MATCH (n) RETURN n') + + +# --------------------------------------------------------------------------- +# _validate_query: EXPLAIN +# --------------------------------------------------------------------------- + + +def test_validate_query_runs_explain(): + driver = _FakeDriver(_FakeResult()) + ok, err = _FakeGlobal(driver)._validate_query('MATCH (n) RETURN n') + assert ok is True and err == '' + assert driver.calls[0][0].startswith('EXPLAIN ') + + +def test_validate_query_reports_error(): + driver = _FakeDriver(raise_error=_StubNeo4jError('syntax error')) + ok, err = _FakeGlobal(driver)._validate_query('MATCH (n RETURN n') + assert ok is False and 'syntax error' in err + + +# --------------------------------------------------------------------------- +# Tool surface: dialect, get_cypher alias, get_schema label filter +# --------------------------------------------------------------------------- + + +def test_dialect_reports_neo4j(): + inst = _instance(_FakeGlobal(_FakeDriver())) + assert inst.dialect({}) == {'dialect': 'neo4j'} + + +def test_get_cypher_is_alias_for_get_query(): + inst = _instance(_FakeGlobal(_FakeDriver())) + seen = {} + + def _fake_get_query(args): + seen['args'] = args + return {'query': 'MATCH (n) RETURN n', 'valid': True} + + inst.get_query = _fake_get_query + out = inst.get_cypher({'question': 'all nodes'}) + # The old tool returned the statement under `cypher`; keep that key working + # for agents written against it, without dropping get_query's `query`. + assert out == {'query': 'MATCH (n) RETURN n', 'cypher': 'MATCH (n) RETURN n', 'valid': True} + assert seen['args'] == {'question': 'all nodes'} + + +def test_get_cypher_passes_through_non_query_replies(): + inst = _instance(_FakeGlobal(_FakeDriver())) + # The not-a-graph-question path has no `query` key, so there is nothing to + # mirror and no empty `cypher` should be invented. + inst.get_query = lambda args: {'answer': 'I am a teapot', 'valid': False} + assert inst.get_cypher({'question': 'hi'}) == {'answer': 'I am a teapot', 'valid': False} + + +def test_get_schema_filters_by_label(): + glb = _FakeGlobal(_FakeDriver()) + glb.graph_schema = { + 'nodes': {'Person': [('name', 'STRING')], 'City': [('name', 'STRING')]}, + 'relationships': [{'type': 'LIVES_IN', 'start': 'Person', 'end': 'City'}], + } + out = _instance(glb).get_schema({'label': 'Person'}) + assert out['labels'] == ['Person'] + assert 'City' not in out['nodes'] + assert out['database'] == 'neo4j' + + +def test_get_schema_unknown_label_returns_error(): + glb = _FakeGlobal(_FakeDriver()) + glb.graph_schema = {'nodes': {'Person': []}, 'relationships': []} + out = _instance(glb).get_schema({'label': 'Ghost'}) + assert 'error' in out + + +def test_get_schema_unwraps_input_envelope(): + glb = _FakeGlobal(_FakeDriver()) + glb.graph_schema = {'nodes': {'Person': [('name', 'STRING')]}, 'relationships': []} + out = _instance(glb).get_schema({'input': {'label': 'Person'}}) + assert out['labels'] == ['Person'] + + +# --------------------------------------------------------------------------- +# get_data: limit clamped, truncation honest +# --------------------------------------------------------------------------- + + +def test_get_data_enforces_limit_and_flags_truncation(): + glb = _FakeGlobal(_FakeDriver(_FakeResult(records=[{'n': i} for i in range(10)]))) + inst = _instance(glb) + inst.get_query = lambda args: {'query': 'MATCH (n) RETURN n', 'valid': True} + out = inst.get_data({'question': 'all nodes', 'limit': 3}) + assert out['valid'] is True + assert len(out['rows']) == 3 + assert out['truncated'] is True diff --git a/packages/ai/src/ai/common/graph/__init__.py b/packages/ai/src/ai/common/graph/__init__.py index 1a8c7ddba..358ba801d 100644 --- a/packages/ai/src/ai/common/graph/__init__.py +++ b/packages/ai/src/ai/common/graph/__init__.py @@ -25,10 +25,10 @@ """Base classes for graph database nodes. -The one node that derives from this base today is ``graph_falkordb``; it was the -reference the abstraction was extracted from. ``db_neo4j`` is graph-capable but -still ships its own copy of this logic — migrating it onto this base is the -intended next step, not a promise this module already keeps. +``graph_falkordb`` was the reference the abstraction was extracted from, and +``graph_neo4j`` derives from it too. Those are the graph-query nodes the repo +has: ``db_arango`` is multi-model and ``db_hydradb`` retrieves natively, so +neither follows the natural-language-to-query loop this base implements. """ from .cypher_safety import is_cypher_safe diff --git a/packages/docs/redirects.ts b/packages/docs/redirects.ts index f3a67f5fa..598bb6a8d 100644 --- a/packages/docs/redirects.ts +++ b/packages/docs/redirects.ts @@ -222,8 +222,9 @@ const redirects: RedirectItem[] = [ ] }, { - "to": "/nodes/db_neo4j", + "to": "/nodes/graph_neo4j", "from": [ + "/nodes/db_neo4j", "/database/neo4j", "/database/neo4j/neo4j" ] diff --git a/packages/server/engine-lib/engLib/store/headers/services.hpp b/packages/server/engine-lib/engLib/store/headers/services.hpp index 947a89e98..03f502bbe 100644 --- a/packages/server/engine-lib/engLib/store/headers/services.hpp +++ b/packages/server/engine-lib/engLib/store/headers/services.hpp @@ -30,7 +30,7 @@ namespace engine::store { //------------------------------------------------------------------------- class IServices { public: - _const int VERSION = 1; + _const int VERSION = 2; //----------------------------------------------------------------- /// @details diff --git a/packages/server/engine-lib/engLib/store/pipeline/pipeline_config.cpp b/packages/server/engine-lib/engLib/store/pipeline/pipeline_config.cpp index 0a784f983..8b6666623 100644 --- a/packages/server/engine-lib/engLib/store/pipeline/pipeline_config.cpp +++ b/packages/server/engine-lib/engLib/store/pipeline/pipeline_config.cpp @@ -536,6 +536,18 @@ Error PipelineConfig::upgradeComponent(json::Value &component, } } + // Graph nodes were renamed when they moved onto the graph base class, so + // pipelines saved with the old provider still resolve to the new service. + if (version < 2) { + if (provider == "tool_falkordb") { + component["provider"] = "graph_falkordb"; + if (config.isMember("type")) config["type"] = "graph_falkordb"; + } else if (provider == "db_neo4j") { + component["provider"] = "graph_neo4j"; + if (config.isMember("type")) config["type"] = "graph_neo4j"; + } + } + return {}; } diff --git a/packages/server/engine-lib/test/store/pipeline/pipeline_config.cpp b/packages/server/engine-lib/test/store/pipeline/pipeline_config.cpp index b88f1bd85..84df8295d 100644 --- a/packages/server/engine-lib/test/store/pipeline/pipeline_config.cpp +++ b/packages/server/engine-lib/test/store/pipeline/pipeline_config.cpp @@ -126,6 +126,29 @@ TEST_CASE("PipelineConfig") { "'pipeline.version' is unsupported"); } + SECTION("version:upgrade:graph") { + auto &component = pipeline["components"][4]; + + SECTION("falkordb") { + component["provider"] = "tool_falkordb"; + component["config"]["type"] = "tool_falkordb"; + + REQUIRE_NO_ERROR(config.validate()); + REQUIRE(component["provider"] == "graph_falkordb"); + REQUIRE(component["config"]["type"] == "graph_falkordb"); + } + + SECTION("neo4j") { + component["provider"] = "db_neo4j"; + + REQUIRE_NO_ERROR(config.validate()); + REQUIRE(component["provider"] == "graph_neo4j"); + REQUIRE_FALSE(component["config"].isMember("type")); + } + + REQUIRE(pipeline["version"] == IServices::VERSION); + } + SECTION("source:missing:optional") { pipeline.removeMember("source");