kgrag/graph_search/cypher_validator.py
2026-06-30 13:35:52 +08:00

351 lines
16 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""
Cypher 验证模块
功能:验证 Cypher 语句的语法、Schema 和语义正确性
"""
import re
import json
import logging
from typing import Optional, List, Dict, Tuple, Any, Set
from neo4j.exceptions import CypherSyntaxError
logger = logging.getLogger(__name__)
class CypherValidator:
"""Cypher 验证器类"""
def __init__(
self,
driver,
neo4j_schema: str,
allowed_node_types: Optional[List[str]] = None,
allowed_relationship_types: Optional[List[str]] = None
):
"""
初始化验证器
Args:
driver: Neo4j driver 实例
neo4j_schema: Neo4j schema 字符串
allowed_node_types: 允许的节点类型列表从graph_service.py传递
allowed_relationship_types: 允许的关系类型列表从graph_service.py传递
"""
self.driver = driver
self.neo4j_schema = neo4j_schema
# 使用传入的节点类型和关系类型,如果没有传入则使用空集合(不进行验证)
self.allowed_node_types = set(allowed_node_types) if allowed_node_types else set()
self.allowed_relationship_types = set(allowed_relationship_types) if allowed_relationship_types else set()
def validate_syntax(self, cypher: str) -> Tuple[bool, Optional[str]]:
"""
第一层验证:语法验证(快速检查)
"""
if not cypher or not cypher.strip():
return False, "Cypher语句为空"
cypher_upper = cypher.upper()
# ========== 必须包含的关键字检查 ==========
if "MATCH" not in cypher_upper or "RETURN" not in cypher_upper:
return False, "缺少必要的Cypher关键字(MATCH/RETURN)"
# ========== 危险操作检查(防止注入) ==========
dangerous = ["DELETE", "DROP", "REMOVE", "SET", "CREATE", "MERGE"]
for kw in dangerous:
# 使用正则表达式匹配完整的单词,避免匹配到属性名中的子串
# 例如:避免将 created_at 中的 CREATE 误判为危险操作
pattern = r'\b' + re.escape(kw) + r'\b'
matches = list(re.finditer(pattern, cypher_upper))
if matches:
return_pos = cypher_upper.find("RETURN")
# 检查是否有危险关键字在RETURN之前作为独立单词
for match in matches:
kw_pos = match.start()
if return_pos == -1 or kw_pos < return_pos:
return False, f"检测到危险操作: {kw}"
# ========== 括号匹配检查(使用栈方法,力扣经典题) ==========
# 使用栈来检查括号是否匹配
bracket_stack = [] # 栈:存储 (位置, 括号类型, 上下文信息)
bracket_pairs = {')': '(', ']': '[', '}': '{'}
i = 0
while i < len(cypher):
char = cypher[i]
# 遇到左括号,入栈
if char in ['(', '[', '{']:
# 获取括号前的上下文(用于判断是否合法)
before_start = max(0, i - 30)
before_context = cypher[before_start:i]
bracket_stack.append((i, char, before_context))
# 遇到右括号,检查是否匹配
elif char in [')', ']', '}']:
if not bracket_stack:
return False, f"括号不匹配:位置 {i} 处有多余的右括号 '{char}'"
last_pos, last_char, last_context = bracket_stack.pop()
expected_left = bracket_pairs[char]
if last_char != expected_left:
return False, f"括号类型不匹配:位置 {last_pos}'{last_char}' 与位置 {i}'{char}' 不匹配"
i += 1
# 检查栈是否为空(所有括号是否都匹配)
if bracket_stack:
unmatched_pos, unmatched_char, _ = bracket_stack[0]
return False, f"括号不匹配:位置 {unmatched_pos} 处有未匹配的左括号 '{unmatched_char}'"
# ========== 检查多余的括号嵌套(使用栈方法分析上下文) ==========
# 检查连续的左括号 `( (` 是否是合法的
i = 0
while i < len(cypher) - 1:
# 查找 `( (` 模式(外层括号)
if cypher[i] == '(':
# 跳过空白字符,查找内层括号
inner_pos = i + 1
while inner_pos < len(cypher) and cypher[inner_pos] in ' \t\n\r':
inner_pos += 1
# 如果找到内层括号 `(`
if inner_pos < len(cypher) and cypher[inner_pos] == '(':
# 获取上下文
before_start = max(0, i - 30)
before_context = cypher[before_start:i]
after_end = min(len(cypher), inner_pos + 30)
after_context = cypher[inner_pos:after_end]
# 检查是否是合法的括号嵌套
is_valid = self._is_valid_bracket_nesting(before_context, after_context)
if not is_valid:
error_context = cypher[max(0, i-30):min(len(cypher), inner_pos+30)]
return False, f"语法错误:检测到多余的括号。位置:...{error_context}..."
i += 1
return True, None
def _is_valid_bracket_nesting(self, before_context: str, after_context: str) -> bool:
"""
判断括号嵌套是否合法
"""
# 检查是否是节点定义(变量名后跟冒号)
# 支持中文字符,允许开头有空白字符
is_node_def = re.match(r'\s*[a-zA-Z_\u4e00-\u9fa5][a-zA-Z0-9_\u4e00-\u9fa5]*\s*:', after_context)
# 检查是否是属性访问(变量名后跟点号)
# 支持中文字符,允许开头有空白字符
is_property_access = re.match(r'\s*[a-zA-Z_\u4e00-\u9fa5][a-zA-Z0-9_\u4e00-\u9fa5]*\s*\.', after_context)
# 检查是否在WHERE/CASE/WHEN上下文中这些上下文中的括号嵌套通常是合法的
before_upper = before_context.upper()
is_in_where_context = (
'WHERE' in before_upper or
'AND' in before_upper or
'OR' in before_upper or
'CASE' in before_upper or
'WHEN' in before_upper or
'THEN' in before_upper or
'RETURN' in before_upper # RETURN子句中的CASE表达式也可能有括号嵌套
)
# 检查是否在函数调用中(如 count(, sum(, max( 等)
is_in_function = re.search(r'[a-zA-Z_][a-zA-Z0-9_]*\s*\(\s*$', before_context)
# 如果匹配任何一种合法情况返回True
return is_node_def or is_property_access or is_in_where_context or is_in_function
def dry_run(self, cypher: str, parse_error_func=None) -> Tuple[bool, Optional[str], Optional[Dict[str, Any]]]:
"""
第三层验证执行前验证Dry Run使用EXPLAIN检查
Args:
cypher: 待验证的Cypher语句
parse_error_func: 错误解析函数(可选)
Returns:
Tuple[bool, Optional[str], Optional[Dict]]: (是否有效, 错误信息, 错误详情)
错误详情包含error_type ('syntax_error' | 'schema_error'), error_message, error_position等
"""
try:
explain_cypher = f"EXPLAIN {cypher}"
result = self.driver.execute_query(explain_cypher)
summary = None
notifications = None
# 方式1: 直接访问 result.summary
if hasattr(result, 'summary'):
summary = result.summary
# 方式2: 如果是元组,取第二个元素
elif isinstance(result, tuple) and len(result) >= 2:
summary = result[1]
# 方式3: 尝试通过索引访问
elif hasattr(result, '__getitem__'):
try:
if len(result) > 1:
summary = result[1]
except:
pass
# 如果找到了 summary尝试获取 notifications
if summary:
notifications = getattr(summary, 'notifications', None)
# 如果 notifications 是 None尝试其他属性名
if notifications is None:
for attr_name in ['notifications', '_notifications', 'warnings', '_warnings']:
if hasattr(summary, attr_name):
notifications = getattr(summary, attr_name)
break
# 如果仍然没有找到,尝试直接从 result 访问
if notifications is None:
for attr_name in ['notifications', '_notifications', 'summary', '_summary']:
if hasattr(result, attr_name):
obj = getattr(result, attr_name)
if hasattr(obj, 'notifications'):
notifications = obj.notifications
break
elif isinstance(obj, list):
notifications = obj
break
return True, None, None
except CypherSyntaxError as e:
error_info = None
if parse_error_func:
error_info = parse_error_func(str(e), cypher)
# 确保 error_info 包含 error_type
if error_info:
error_info['error_type'] = 'syntax_error'
else:
error_info = {'error_type': 'syntax_error', 'error_message': str(e)}
return False, f"语法错误: {str(e)}", error_info
except Exception as e:
error_msg = str(e)
error_info = None
if parse_error_func:
error_info = parse_error_func(error_msg, cypher)
if "SyntaxError" in error_msg or "Syntax" in error_msg:
if error_info:
error_info['error_type'] = 'syntax_error'
else:
error_info = {'error_type': 'syntax_error', 'error_message': error_msg}
return False, f"语法错误: {error_msg}", error_info
# 其他异常,默认为语法错误
if error_info:
error_info['error_type'] = 'syntax_error'
else:
error_info = {'error_type': 'syntax_error', 'error_message': error_msg}
return False, f"执行计划失败: {error_msg}", error_info
def validate_schema(self, cypher: str) -> Tuple[bool, Optional[str]]:
"""
第二层验证Schema 验证(验证节点类型和关系类型)
简单逻辑:在 MATCH 后面找到 `-`,向前查找 `()` 或 `[]`,提取 `:` 后的内容进行验证。
Args:
cypher: 待验证的Cypher语句
Returns:
Tuple[bool, Optional[str]]: (是否有效, 错误信息)
"""
if not cypher or not cypher.strip():
return False, "Cypher语句为空"
# 如果节点类型和关系类型列表都为空,跳过验证
if not self.allowed_node_types and not self.allowed_relationship_types:
return True, None
cypher_upper = cypher.upper()
invalid_labels = []
invalid_types = []
# 找到所有 MATCH 关键字的位置
i = 0
while i < len(cypher):
match_pos = cypher_upper.find("MATCH", i)
if match_pos == -1:
break
# 从 MATCH 后开始查找,找到子句结束位置
start_pos = match_pos + 5
end_pos = len(cypher)
for keyword in ['WHERE', 'WITH', 'RETURN', 'OPTIONAL', 'MATCH']:
kw_pos = cypher_upper.find(keyword, start_pos)
if kw_pos != -1 and kw_pos < end_pos:
end_pos = kw_pos
# 在 MATCH 子句范围内查找所有 `-`
pos = start_pos
while pos < end_pos:
dash_pos = cypher.find('-', pos, end_pos)
if dash_pos == -1:
break
# 向前查找最近的 `(` 或 `[`
bracket_start = -1
bracket_type = None
for j in range(dash_pos - 1, max(0, dash_pos - 200), -1):
if cypher[j] == '(':
bracket_start = j
bracket_type = '('
break
elif cypher[j] == '[':
bracket_start = j
bracket_type = '['
break
if bracket_start != -1:
# 找到对应的右括号
right_bracket = ')' if bracket_type == '(' else ']'
bracket_end = cypher.find(right_bracket, bracket_start + 1, dash_pos + 100)
if bracket_end != -1:
# 提取括号内容
content = cypher[bracket_start + 1:bracket_end]
# 查找 `:` 后的内容
colon_pos = content.find(':')
if colon_pos != -1:
after_colon = content[colon_pos + 1:]
# 去除 `{}` 及其内容
after_colon = re.sub(r'\{[^}]*\}', '', after_colon).strip()
if bracket_type == '(':
# 节点标签,用 `:` 分隔
for label in after_colon.split(':'):
label = label.strip()
# 只取第一个单词(去除空格后的内容)
label = re.split(r'[\s\{]', label)[0]
if label and self.allowed_node_types and label not in self.allowed_node_types:
invalid_labels.append(label)
else:
# 关系类型,用 `|` 分隔
for rel_type in after_colon.split('|'):
rel_type = rel_type.strip()
# 只取关系类型本体,剥离可变长度/属性等修饰
# 例如:`:包含*0..8` 应识别为关系类型 `包含`
# 例如:`:REL*` / `:REL*1..3` / `:REL {a:1}` / `:REL*0..8 {a:1}`
rel_type = re.split(r'[\s\{]', rel_type)[0]
rel_type = re.sub(r'\*.*$', '', rel_type).strip()
if rel_type and self.allowed_relationship_types and rel_type not in self.allowed_relationship_types:
invalid_types.append(rel_type)
pos = dash_pos + 1
i = end_pos
# 返回错误信息
if invalid_labels:
return False, f"检测到不允许的节点类型: {', '.join(sorted(set(invalid_labels)))}"
if invalid_types:
return False, f"检测到不允许的关系类型: {', '.join(sorted(set(invalid_types)))}"
return True, None