Python Sqlparse如何判断SQL查询格式合法性并返回布尔状态
首先要明确:你提到的两种SQL案例,后者的“不合法”属于语义层面错误(SELECT字段包含未在GROUP BY中声明的location),而sqlparse本身仅负责SQL的语法格式解析(比如检查缺失括号、关键字错误这类语法问题),不做语义校验。下面分两种场景给出解决方案:
一、用sqlparse检查SQL语法格式合法性
sqlparse没有直接返回True/False的校验函数,但可以通过尝试解析SQL并捕获异常来判断语法是否合法。如果解析过程抛出异常,说明语法格式不合法;反之则语法格式合法。
示例代码:
import sqlparse from sqlparse.exceptions import ParseError def is_sql_syntactically_valid(sql): try: parsed = sqlparse.parse(sql) # 排除空解析的情况 return len(parsed) > 0 except ParseError: return False # 测试语法合法的SQL(即使语义有问题,语法合法也会返回True) valid_syntax_sql = """select count(users), department, location from usertable group by department""" print(is_sql_syntactically_valid(valid_syntax_sql)) # 输出True # 测试语法错误的SQL(比如缺少闭合括号) invalid_syntax_sql = "select count(users from usertable" print(is_sql_syntactically_valid(invalid_syntax_sql)) # 输出False
二、校验你示例中的语义合法性(GROUP BY与SELECT字段匹配)
sqlparse无法处理这类语义校验,因为它不理解SQL的业务逻辑和数据库表结构。常见的解决方式有两种:
- 方式1:连接数据库做预编译检查
利用数据库的预编译功能,不需要执行SQL,只让数据库校验语义合法性。示例(以PostgreSQL为例,用psycopg2):
import psycopg2
from psycopg2.errors import GroupByError
def is_sql_semantically_valid(sql, db_config):
conn = None
try:
conn = psycopg2.connect(**db_config)
cur = conn.cursor()
# 尝试预编译SQL,不执行也不提交
cur.execute(f"PREPARE test_plan AS {sql}")
cur.execute("DEALLOCATE test_plan")
return True
except GroupByError:
return False
except Exception:
# 捕获其他语义类错误
return False
finally:
if conn:
conn.close()
- **方式2:解析AST手动校验** 用sqlparse解析出SQL的抽象语法树(AST),手动提取SELECT字段和GROUP BY字段,对比是否符合规则。这种方式仅适合简单场景,复杂SQL需要更完善的逻辑: ```python import sqlparse from sqlparse.sql import IdentifierList, Identifier from sqlparse.tokens import Keyword def extract_select_fields(parsed): select_fields = [] for stmt in parsed: if stmt.get_type() == 'SELECT': for token in stmt.tokens: if isinstance(token, IdentifierList): for item in token.get_identifiers(): # 提取非聚合函数的字段名 if hasattr(item, 'get_real_name'): field_name = item.get_real_name() if not field_name.startswith('count('): select_fields.append(field_name) return select_fields def extract_group_by_fields(parsed): group_by_fields = [] for stmt in parsed: if stmt.get_type() == 'SELECT': for idx, token in enumerate(stmt.tokens): if token.ttype == Keyword and token.value.upper() == 'GROUP BY': # 提取GROUP BY后的字段 for item in stmt.tokens[idx+1:]: if isinstance(item, IdentifierList): group_by_fields.extend([sub.get_real_name() for sub in item.get_identifiers()]) elif isinstance(item, Identifier): group_by_fields.append(item.get_real_name()) break return group_by_fields def is_group_by_matching_select(sql): parsed = sqlparse.parse(sql) select_fields = extract_select_fields(parsed) group_by_fields = extract_group_by_fields(parsed) return all(field in group_by_fields for field in select_fields) # 测试你的示例 valid_semantic_sql = """select count(users), department from usertable group by department""" print(is_group_by_matching_select(valid_semantic_sql)) # 输出True invalid_semantic_sql = """select count(users), department, location from usertable group by department""" print(is_group_by_matching_select(invalid_semantic_sql)) # 输出False
内容的提问来源于stack exchange,提问作者Rahul Neekhra

