269 lines
7.5 KiB
Python
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
|