196 lines
5.1 KiB
Python
196 lines
5.1 KiB
Python
#!/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
|