""" 从 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