Files
ai-g3sb-backman2.0/tools/dbhub_sql_parser.py
T
2026-04-14 10:28:22 +08:00

269 lines
7.5 KiB
Python

"""
从 DBHub sql-parser.ts 等价移植:按方言剥离注释/字符串、切分语句。
参见 dbhub/src/utils/sql-parser.ts
"""
from __future__ import annotations
import re
from typing import Callable, Literal, TypedDict
ConnectorType = Literal["postgres", "mysql", "mariadb", "sqlite", "sqlserver"]
class _Token(TypedDict):
type: int # 0 Plain, 1 Comment, 2 QuotedBlock
end: int
_TOKEN_PLAIN = 0
_TOKEN_COMMENT = 1
_TOKEN_QUOTED = 2
def _plain_token(i: int) -> _Token:
return {"type": _TOKEN_PLAIN, "end": i + 1}
def _scan_single_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "-" or sql[i + 1] != "-":
return None
j = i
while j < len(sql) and sql[j] != "\n":
j += 1
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_multi_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
j = i + 2
while j + 1 < len(sql) and not (sql[j] == "*" and sql[j + 1] == "/"):
j += 1
if j + 1 < len(sql):
j += 2
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_multi_line_comment_mysql(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
nxt = sql[i + 2] if i + 2 < len(sql) else ""
nxt2 = sql[i + 3] if i + 3 < len(sql) else ""
if nxt == "!" or (nxt == "M" and nxt2 == "!"):
return None
return _scan_multi_line_comment(sql, i)
def _scan_nested_multi_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
j = i + 2
depth = 1
while j < len(sql) and depth > 0:
if j + 1 < len(sql) and sql[j] == "/" and sql[j + 1] == "*":
depth += 1
j += 2
elif j + 1 < len(sql) and sql[j] == "*" and sql[j + 1] == "/":
depth -= 1
j += 2
else:
j += 1
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_single_quoted_string(sql: str, i: int) -> _Token | None:
if sql[i] != "'":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "'" and sql[j + 1] == "'":
j += 2
elif sql[j] == "'":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_double_quoted_string(sql: str, i: int) -> _Token | None:
if sql[i] != '"':
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == '"' and sql[j + 1] == '"':
j += 2
elif sql[j] == '"':
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
_dollar_quote_open_regex = re.compile(r"^\$([a-zA-Z_]\w*)?\$")
def _scan_dollar_quoted_block(sql: str, i: int) -> _Token | None:
if sql[i] != "$":
return None
nxt = sql[i + 1] if i + 1 < len(sql) else ""
if nxt.isdigit():
return None
remaining = sql[i:]
m = _dollar_quote_open_regex.match(remaining)
if not m:
return None
tag = m.group(0)
body_start = i + len(tag)
close_idx = sql.find(tag, body_start)
end = close_idx + len(tag) if close_idx != -1 else len(sql)
return {"type": _TOKEN_QUOTED, "end": end}
def _scan_backtick_quoted_identifier(sql: str, i: int) -> _Token | None:
if sql[i] != "`":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "`" and sql[j + 1] == "`":
j += 2
elif sql[j] == "`":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_bracket_quoted_identifier(sql: str, i: int) -> _Token | None:
if sql[i] != "[":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "]" and sql[j + 1] == "]":
j += 2
elif sql[j] == "]":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_token_ansi(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _plain_token(i)
)
def _scan_token_postgres(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_nested_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_dollar_quoted_block(sql, i)
or _plain_token(i)
)
def _scan_token_mysql(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment_mysql(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_backtick_quoted_identifier(sql, i)
or _plain_token(i)
)
def _scan_token_sqlite(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_backtick_quoted_identifier(sql, i)
or _scan_bracket_quoted_identifier(sql, i)
or _plain_token(i)
)
def _scan_token_sqlserver(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_bracket_quoted_identifier(sql, i)
or _plain_token(i)
)
_DIALECT_SCANNERS: dict[ConnectorType, Callable[[str, int], _Token]] = {
"postgres": _scan_token_postgres,
"mysql": _scan_token_mysql,
"mariadb": _scan_token_mysql,
"sqlite": _scan_token_sqlite,
"sqlserver": _scan_token_sqlserver,
}
def _get_scanner(dialect: ConnectorType | None) -> Callable[[str, int], _Token]:
if dialect and dialect in _DIALECT_SCANNERS:
return _DIALECT_SCANNERS[dialect]
return _scan_token_ansi
def strip_comments_and_strings(sql: str, dialect: ConnectorType | None = None) -> str:
"""Replace comments, string literals, and dialect-specific quoted blocks with a single space each."""
scan_token = _get_scanner(dialect)
parts: list[str] = []
plain_start = -1
i = 0
n = len(sql)
while i < n:
token = scan_token(sql, i)
if token["type"] == _TOKEN_PLAIN:
if plain_start == -1:
plain_start = i
else:
if plain_start != -1:
parts.append(sql[plain_start:i])
plain_start = -1
parts.append(" ")
i = token["end"]
if plain_start != -1:
parts.append(sql[plain_start:])
return "".join(parts)
def split_sql_statements(sql: str, dialect: ConnectorType | None = None) -> list[str]:
"""Split SQL into individual statements, handling semicolons inside quoted contexts."""
scan_token = _get_scanner(dialect)
statements: list[str] = []
stmt_start = 0
i = 0
n = len(sql)
while i < n:
if sql[i] == ";":
trimmed = sql[stmt_start:i].strip()
if trimmed:
statements.append(trimmed)
stmt_start = i + 1
i += 1
continue
token = scan_token(sql, i)
i = token["end"]
trimmed = sql[stmt_start:].strip()
if trimmed:
statements.append(trimmed)
return statements