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:
陈辅元
2026-04-17 15:15:01 +08:00
parent bcb9a205fa
commit 3891aae7fe
19 changed files with 707 additions and 306 deletions
Binary file not shown.
+87 -9
View File
@@ -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",