Refactor api_server.py to import environment setup and schema loading from bootstrap.py, enhancing modularity. Introduce a new function in Text2SQLOrchestrator to prioritize VCUserAccessibleFunction in table selection, improving SQL generation accuracy. Update validation logic to enforce restrictions on CJK characters in SQL string literals, ensuring compliance with business rules. Enhance prompts to clarify SQL generation constraints regarding date conditions and CJK usage.
This commit is contained in:
@@ -4,7 +4,7 @@ SQL 验证工具集
|
||||
|
||||
import re
|
||||
import logging
|
||||
from typing import Tuple, List, Dict
|
||||
from typing import Tuple, List, Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -14,11 +14,56 @@ _CJK_IN_STRING_RE = re.compile(
|
||||
)
|
||||
|
||||
|
||||
def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
扫描 SQL 中单引号字符串(含 T-SQL N'…'),若字面量内出现 CJK 则判失败。
|
||||
def _cjk_text_in_string_literal(text: str) -> bool:
|
||||
return bool(_CJK_IN_STRING_RE.search(text))
|
||||
|
||||
跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本,避免误报。
|
||||
|
||||
def _node_in_subtree(root, target) -> bool:
|
||||
if root is None:
|
||||
return False
|
||||
for n in root.walk():
|
||||
if n is target:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _cjk_string_in_forbidden_context(node, exp) -> bool:
|
||||
"""
|
||||
禁止含 CJK 的字面量出现在「比对/过滤」语境:WHERE、HAVING、JOIN ON、
|
||||
以及 CASE 分支的 WHEN 条件(含简单 CASE 的 WHEN 值),
|
||||
但允许出现在 SELECT 投影、CASE 的 THEN/ELSE 结果等纯展示位置。
|
||||
"""
|
||||
if node.find_ancestor(exp.Where):
|
||||
return True
|
||||
if node.find_ancestor(exp.Having):
|
||||
return True
|
||||
join = node.find_ancestor(exp.Join)
|
||||
if join is not None:
|
||||
on = join.args.get("on")
|
||||
if on is not None and _node_in_subtree(on, node):
|
||||
return True
|
||||
|
||||
case = node.find_ancestor(exp.Case)
|
||||
while case is not None:
|
||||
default = case.args.get("default")
|
||||
if default is not None and _node_in_subtree(default, node):
|
||||
return False
|
||||
for br in case.args.get("ifs") or []:
|
||||
then_expr = br.args.get("true")
|
||||
cond = br.this
|
||||
if then_expr is not None and _node_in_subtree(then_expr, node):
|
||||
return False
|
||||
if cond is not None and _node_in_subtree(cond, node):
|
||||
return True
|
||||
case = case.find_ancestor(exp.Case)
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def _check_no_cjk_legacy_text_scan(sql: str) -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
解析失败时的回退:扫描单引号字符串(含 N'…'),字面量内出现 CJK 即失败。
|
||||
跳过 ``--`` 行注释与 ``/* */`` 块注释内的文本。
|
||||
"""
|
||||
errors: List[str] = []
|
||||
i = 0
|
||||
@@ -27,7 +72,6 @@ def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
|
||||
in_block_comment = False
|
||||
|
||||
def _read_single_quoted_string(start: int) -> Tuple[str, int]:
|
||||
"""从 start 指向的 opening `'` 之后开始读,返回 (内容, 闭合引号后下标)。"""
|
||||
j = start
|
||||
parts: List[str] = []
|
||||
while j < n:
|
||||
@@ -66,10 +110,9 @@ def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
|
||||
i += 2
|
||||
continue
|
||||
|
||||
# N' 或 n' 前缀的 Unicode 字面量
|
||||
if i + 1 < n and sql[i] in "Nn" and sql[i + 1] == "'":
|
||||
body, i = _read_single_quoted_string(i + 2)
|
||||
if _CJK_IN_STRING_RE.search(body):
|
||||
if _cjk_text_in_string_literal(body):
|
||||
prev = body[:48] + ("…" if len(body) > 48 else "")
|
||||
errors.append(
|
||||
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
|
||||
@@ -79,7 +122,7 @@ def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
|
||||
|
||||
if sql[i] == "'":
|
||||
body, i = _read_single_quoted_string(i + 1)
|
||||
if _CJK_IN_STRING_RE.search(body):
|
||||
if _cjk_text_in_string_literal(body):
|
||||
prev = body[:48] + ("…" if len(body) > 48 else "")
|
||||
errors.append(
|
||||
"SQL 字符串字面量中含中文或与业务中文直接作为比对值(禁止)。"
|
||||
@@ -92,6 +135,41 @@ def check_no_cjk_in_sql_string_literals(sql: str) -> Tuple[bool, List[str]]:
|
||||
return len(errors) == 0, errors
|
||||
|
||||
|
||||
def check_no_cjk_in_sql_string_literals(sql: str, dialect: str = "tsql") -> Tuple[bool, List[str]]:
|
||||
"""
|
||||
禁止在「过滤/比对」语境使用含中日韩字符的字符串字面量(含 T-SQL ``N'…'``)。
|
||||
|
||||
允许在 SELECT 投影、CASE 的 THEN/ELSE 结果等展示用字面量中使用中文标签。
|
||||
解析失败时回退为全文扫描(与旧版一致,偏严)。
|
||||
"""
|
||||
from sqlglot import parse_one, exp
|
||||
|
||||
errors: List[str] = []
|
||||
|
||||
try:
|
||||
parsed = parse_one(sql, dialect=dialect)
|
||||
except Exception as e:
|
||||
logger.debug("CJK 校验回退为全文扫描(SQL 解析失败): %s", e)
|
||||
return _check_no_cjk_legacy_text_scan(sql)
|
||||
|
||||
for node in parsed.walk():
|
||||
text: Optional[str] = None
|
||||
if isinstance(node, exp.Literal) and node.is_string:
|
||||
text = str(node.this)
|
||||
elif isinstance(node, exp.National):
|
||||
text = str(node.this)
|
||||
if text is None or not _cjk_text_in_string_literal(text):
|
||||
continue
|
||||
if _cjk_string_in_forbidden_context(node, exp):
|
||||
prev = text[:48] + ("…" if len(text) > 48 else "")
|
||||
errors.append(
|
||||
"SQL 在 WHERE/HAVING/JOIN ON 或 CASE/WHEN 条件中出现含中文的字符串字面量(禁止)。"
|
||||
f"片段近似: …'{prev}'… — 请改用 Schema 注释中的代码/枚举,或通过维表 JOIN 用键列过滤。"
|
||||
)
|
||||
|
||||
return len(errors) == 0, errors
|
||||
|
||||
|
||||
# 危险操作关键词(除非明确允许)
|
||||
DANGEROUS_KEYWORDS = [
|
||||
"DROP", "DELETE", "UPDATE", "INSERT", "ALTER", "TRUNCATE",
|
||||
|
||||
Reference in New Issue
Block a user