Build lightweight AI agent admin
This commit is contained in:
@@ -0,0 +1,195 @@
|
||||
#!/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
|
||||
Reference in New Issue
Block a user