#!/usr/bin/env python # -*- coding: utf-8 -*- """SQL 工具:标识符转义、语句拆分、系统对象黑名单""" from typing import List PROTECTED_DATABASES = frozenset({ "postgres", "template0", "template1", "mysql", "information_schema", "performance_schema", "sys", "master", "tempdb", "model", "msdb", }) PROTECTED_SCHEMAS = frozenset({ "pg_catalog", "information_schema", "pg_toast", "sys", "guest", "db_owner", "db_accessadmin", "db_securityadmin", "db_ddladmin", "db_backupoperator", "db_datareader", "db_datawriter", "db_denydatareader", "db_denydatawriter", }) ORACLE_PROTECTED_SCHEMAS = frozenset({ "sys", "system", "outln", "xdb", "ctxsys", "mdsys", "orddata", "ordsys", "lbacsys", "dbsnmp", "appqossys", "audsys", "gsmadmin_internal", "oav", "ordplugins", "si_informtn_schema", "wmsys", }) def is_protected_database(name: str, db_type: str = "") -> bool: lower = name.lower() if lower in PROTECTED_DATABASES: return True if (db_type or "").lower() == "oracle": return False return False def is_protected_schema(name: str, db_type: str = "") -> bool: lower = name.lower() if lower in PROTECTED_SCHEMAS: return True if lower.startswith("db_"): return True if (db_type or "").lower() == "oracle" and lower in ORACLE_PROTECTED_SCHEMAS: return True return False def _normalize_db_type(db_type: str) -> str: db = (db_type or "postgresql").lower() if db in ("mssql", "sql server"): return "sqlserver" if db in ("postgres", "psql"): return "postgresql" return db def quote_identifier(name: str, db_type: str) -> str: """按数据库类型转义标识符""" db = _normalize_db_type(db_type) if db == "mysql": return f"`{name.replace('`', '``')}`" if db in ("sqlserver", "mssql"): return f"[{name.replace(']', ']]')}]" if db == "oracle": return f'"{name.replace(chr(34), chr(34) + chr(34))}"' return f'"{name.replace(chr(34), chr(34) + chr(34))}"' def quote_table(schema: str | None, table: str, db_type: str) -> str: if schema: return f"{quote_identifier(schema, db_type)}.{quote_identifier(table, db_type)}" return quote_identifier(table, db_type) def split_sql_statements(sql: str) -> List[str]: """按分号拆分 SQL,忽略字符串与 dollar-quote 内的分号""" if not sql or not sql.strip(): return [] statements: List[str] = [] current: List[str] = [] i = 0 n = len(sql) in_single = False in_double = False in_backtick = False dollar_tag: str | None = None while i < n: if dollar_tag is None and not in_single and not in_double and not in_backtick: if sql[i] == "$": j = i + 1 while j < n and sql[j] != "$" and (sql[j].isalnum() or sql[j] == "_"): j += 1 if j < n and sql[j] == "$": dollar_tag = sql[i : j + 1] current.append(dollar_tag) i = j + 1 continue if sql[i] == "'": in_single = True current.append(sql[i]) i += 1 continue if sql[i] == '"': in_double = True current.append(sql[i]) i += 1 continue if sql[i] == "`": in_backtick = True current.append(sql[i]) i += 1 continue if sql[i] == ";": stmt = "".join(current).strip() if stmt: statements.append(stmt) current = [] i += 1 continue elif dollar_tag is not None: if sql.startswith(dollar_tag, i): current.append(dollar_tag) i += len(dollar_tag) dollar_tag = None continue elif in_single: if sql[i] == "'" and i + 1 < n and sql[i + 1] == "'": current.append("''") i += 2 continue if sql[i] == "'": in_single = False current.append(sql[i]) i += 1 continue elif in_double: if sql[i] == '"' and i + 1 < n and sql[i + 1] == '"': current.append('""') i += 2 continue if sql[i] == '"': in_double = False current.append(sql[i]) i += 1 continue elif in_backtick: if sql[i] == "`" and i + 1 < n and sql[i + 1] == "`": current.append("``") i += 2 continue if sql[i] == "`": in_backtick = False current.append(sql[i]) i += 1 continue current.append(sql[i]) i += 1 stmt = "".join(current).strip() if stmt: statements.append(stmt) return statements