Files
ai-agent-admin/backend-fastapi/core/database_manager/sql_utils.py
T
2026-06-08 18:14:59 +08:00

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