Files

1253 lines
49 KiB
Python
Raw Permalink Normal View History

2026-04-14 10:28:22 +08:00
"""
与 DBHub 对齐的数据库工具:search_objects(元数据探索)、execute_sql(SQL 执行,对齐 execute-sql.ts)、
execute_sql_all(同 execute_sql 但不截断行数)、execute_sql_count_only(仅统计行数,不返回明细)。
Text2SQL Validator 使用 ``probe_sql_execution_status`` / ``probe_sql_execution_status_ex`` 做库上探针。
- 模块级 ``execute_sql`` / ``execute_sql_all`` / ``execute_sql_count_only`` / ``search_objects``:便捷入口,委托给 ``default_dbhub_tools``。
- ``DbHubTools``:可实例化的工具类,方法语义与上述模块函数相同。
- ``default_dbhub_tools``:默认单例实例,与模块级函数共用。
"""
from __future__ import annotations
import json
import logging
import os
import re
from decimal import Decimal
from typing import Any, Literal
from sqlalchemy import inspect, text
from sqlalchemy.engine import Engine
from sqlalchemy.exc import SQLAlchemyError
from db.dbhub_allowed_keywords import allowed_keywords_list, is_read_only_sql
from db.dbhub_sql_parser import (
ConnectorType,
_get_scanner,
_TOKEN_PLAIN,
split_sql_statements,
strip_comments_and_strings,
)
from db.engine import get_engine
log = logging.getLogger(__name__)
def _get_optional_str(key: str) -> str | None:
"""从环境变量读取(与根目录 .env / load_dotenv 一致);键名大小写不敏感。"""
v = os.getenv(key)
if v is not None and str(v).strip():
return str(v).strip()
v2 = os.getenv(key.upper())
if v2 is not None and str(v2).strip():
return str(v2).strip()
return None
ObjectType = Literal["schema", "table", "column", "procedure", "function", "index"]
DetailLevel = Literal["names", "summary", "full"]
_SKIP_SCHEMAS = frozenset(
{
"information_schema",
"pg_catalog",
"pg_toast",
"mysql",
"performance_schema",
"sys",
},
)
def _like_pattern_to_regex(pattern: str) -> re.Pattern[str]:
"""SQL LIKE → regex,与 DBHub likePatternToRegex 一致(% / _,其余字符 re.escape)。"""
parts: list[str] = []
for c in pattern:
if c == "%":
parts.append(".*")
elif c == "_":
parts.append(".")
else:
parts.append(re.escape(c))
return re.compile("^" + "".join(parts) + "$", re.IGNORECASE)
def _bare_table_name(name: str) -> str:
return name.strip().rsplit(".", 1)[-1].strip()
def _normalize_user_object_name(name: str) -> str:
"""去掉首尾空白;若整段为 [Name] 形式则去括号(单层)。"""
t = name.strip()
if len(t) >= 2 and t[0] == "[" and t[-1] == "]":
return t[1:-1].replace("]]", "]")
return t
def _resolve_table_or_view_bare_name(insp: Any, schema: str | None, user_table: str) -> str:
"""
将用户输入的表/视图名解析为与 SQLAlchemy inspect 一致、可在 get_columns/get_indexes 中使用的名称。
SQL Server 等库对标识符大小写不敏感,但反射时传入与系统目录不一致的大小写可能导致 get_columns 返回空;
视图不在 get_table_names 中,仅传表名会漏匹配。本函数在指定 schema 下对表名与视图名做不区分大小写对齐。
"""
want = _normalize_user_object_name(user_table)
if not want:
return user_table.strip()
key = want.lower()
for lister in (insp.get_table_names, insp.get_view_names):
try:
for raw in lister(schema=schema):
bare = _bare_table_name(str(raw))
if bare.lower() == key:
return bare
except Exception: # noqa: BLE001
continue
return want
def _filter_schemas(raw: list[str]) -> list[str]:
return [s for s in raw if s and s.lower() not in _SKIP_SCHEMAS]
def _qualified_table_sql(engine: Any, schema: str | None, bare_table: str) -> str:
prep = engine.dialect.identifier_preparer
tq = prep.quote(bare_table)
if schema:
return f"{prep.quote_schema(schema)}.{tq}"
return tq
def _table_row_count(engine: Any, schema: str | None, bare_table: str) -> int | None:
q = _qualified_table_sql(engine, schema, bare_table)
try:
with engine.connect() as conn:
r = conn.execute(text(f"SELECT COUNT(*) AS c FROM {q}"))
row = r.fetchone()
if row is None:
return None
return int(row[0])
except Exception: # noqa: BLE001
return None
def _pk_columns(insp: Any, schema: str | None, bare: str) -> tuple[str, ...]:
try:
pk = insp.get_pk_constraint(bare, schema=schema)
cols = pk.get("constrained_columns") or []
return tuple(str(c) for c in cols)
except Exception: # noqa: BLE001
return ()
def _index_dicts(insp: Any, schema: str | None, bare: str) -> list[dict[str, Any]]:
pk_cols = _pk_columns(insp, schema, bare)
try:
raw_idx = list(insp.get_indexes(bare, schema=schema))
except Exception: # noqa: BLE001
return []
out: list[dict[str, Any]] = []
for idx in raw_idx:
cols = list(idx.get("column_names") or [])
unique = bool(idx.get("unique"))
primary = bool(pk_cols) and tuple(cols) == pk_cols
out.append(
{
"name": str(idx.get("name") or ""),
"columns": cols,
"unique": unique,
"primary": primary,
},
)
return out
def _table_comment(insp: Any, schema: str | None, bare: str) -> str | None:
"""获取表注释,若不支持或无注释则返回 None。"""
try:
tc = insp.get_table_comment(bare, schema=schema)
if isinstance(tc, dict):
text = (tc.get("text") or "").strip()
return text if text else None
except (NotImplementedError, AttributeError, TypeError):
pass
except Exception: # noqa: BLE001
pass
return None
def _column_dicts(insp: Any, schema: str | None, bare: str) -> list[dict[str, Any]]:
try:
cols = list(insp.get_columns(bare, schema=schema))
except Exception: # noqa: BLE001
return []
out: list[dict[str, Any]] = []
for c in cols:
if not isinstance(c, dict):
continue
fname = str(c.get("name") or "").strip()
if not fname:
continue
ftype = c.get("type")
dtype = str(ftype).strip() if ftype is not None else ""
nullable = c.get("nullable")
null_b = nullable is True or (isinstance(nullable, str) and nullable.upper() == "YES")
entry: dict[str, Any] = {
"column_name": fname,
"data_type": dtype,
"is_nullable": "YES" if null_b else "NO",
"column_default": c.get("default"),
}
# 添加列描述(如果存在)
comment = c.get("comment")
if comment is not None:
comment_str = str(comment).strip()
if comment_str:
entry["description"] = comment_str
out.append(entry)
return out
def _fetch_routines(
engine: Any,
dialect: str,
*,
schema_filter: str | None,
routine_sql_type: str | None,
) -> list[tuple[str, str, str, str | None, str | None, str | None]]:
"""
返回 (schema, name, routine_type, definition_or_none, language_or_none, return_type_or_none)。
routine_sql_type: 'PROCEDURE' | 'FUNCTION' | None(两者都要)
注意:INFORMATION_SCHEMA.ROUTINES 在不同数据库中包含的字段不同:
- SQL Server: 有 EXTERNAL_LANGUAGE(但通常为 NULL),无 RETURN_TYPE
- MySQL/MariaDB: 有 EXTERNAL_LANGUAGE、DTD_IDENTIFIER(返回类型)
- PostgreSQL: INFORMATION_SCHEMA 中存储过程支持有限,通常需要通过 pg_proc 查询
"""
rows: list[tuple[str, str, str, str | None, str | None, str | None]] = []
if dialect not in ("mssql", "postgresql", "mysql", "mariadb"):
return rows
# 基础查询:所有数据库都支持的字段
sql = """
SELECT ROUTINE_SCHEMA, ROUTINE_NAME, ROUTINE_TYPE, ROUTINE_DEFINITION
FROM INFORMATION_SCHEMA.ROUTINES
WHERE 1=1
"""
params: dict[str, Any] = {}
if schema_filter:
sql += " AND ROUTINE_SCHEMA = :sch"
params["sch"] = schema_filter
if routine_sql_type:
sql += " AND ROUTINE_TYPE = :rt"
params["rt"] = routine_sql_type
with engine.connect() as conn:
for r in conn.execute(text(sql), params):
defn = r[3]
rows.append(
(
str(r[0]),
str(r[1]),
str(r[2]),
None if defn is None else str(defn),
None, # language - INFORMATION_SCHEMA 中通常不可用
None, # return_type - 需要额外查询
),
)
# 对于 MySQL/MariaDB,尝试获取更多信息
if dialect in ("mysql", "mariadb") and rows:
try:
extra_sql = """
SELECT ROUTINE_SCHEMA, ROUTINE_NAME, EXTERNAL_LANGUAGE, DTD_IDENTIFIER
FROM INFORMATION_SCHEMA.ROUTINES
WHERE 1=1
"""
extra_params: dict[str, Any] = {}
if schema_filter:
extra_sql += " AND ROUTINE_SCHEMA = :sch"
extra_params["sch"] = schema_filter
if routine_sql_type:
extra_sql += " AND ROUTINE_TYPE = :rt"
extra_params["rt"] = routine_sql_type
extra_map: dict[tuple[str, str], tuple[str | None, str | None]] = {}
for r in conn.execute(text(extra_sql), extra_params):
key = (str(r[0]), str(r[1]))
lang = str(r[2]).strip() if r[2] else None
ret_type = str(r[3]).strip() if r[3] else None
extra_map[key] = (lang, ret_type)
# 更新已有行
updated_rows: list[tuple[str, str, str, str | None, str | None, str | None]] = []
for sch, name, rtype, defn, _, _ in rows:
key = (sch, name)
if key in extra_map:
lang, ret_type = extra_map[key]
updated_rows.append((sch, name, rtype, defn, lang, ret_type))
else:
updated_rows.append((sch, name, rtype, defn, None, None))
rows = updated_rows
except Exception: # noqa: BLE001
# 如果额外查询失败,保持原有数据
pass
return rows
def sqlalchemy_dialect_to_connector(dialect_name: str) -> ConnectorType:
"""
将 SQLAlchemy dialect.name 映射到 DBHub ConnectorType。
未知方言按 postgres 规则做只读校验(与 DBHub 移植约定一致)。
"""
m: dict[str, ConnectorType] = {
"postgresql": "postgres",
"mysql": "mysql",
"mariadb": "mariadb",
"sqlite": "sqlite",
"mssql": "sqlserver",
}
return m.get(dialect_name, "postgres")
_COUNT_WRAP_FIRST_WORDS = frozenset({"select", "with", "explain"})
_AGGREGATE_SCALAR_PREFIX = re.compile(
r"(?is)^(?:distinct\s+)?(?:count|sum|avg|min|max|stdev|stdevp|string_agg|group_concat|"
r"variance|var_pop|var_samp|stddev|stddev_pop|stddev_samp)\s*\(",
)
def _parse_top_level_select_list(sql: str, connector: ConnectorType) -> tuple[int, str] | None:
"""
定位最外层 ``SELECT`` 与深度 0 上首个 ``FROM`` 之间的选择列表,返回 (列数, 列表原文断片)。
``WITH`` 开头、``SELECT *``、或无法配对到 ``FROM`` 时返回 None。
"""
if _first_sql_keyword(sql, connector) == "with":
return None
scan = _get_scanner(connector)
n = len(sql)
i = 0
list_paren = 0
col_commas = 0
state: Literal["LEADING", "IN_LIST"] = "LEADING"
list_start = 0
while i < n:
tok = scan(sql, i)
if tok["type"] != _TOKEN_PLAIN:
i = tok["end"]
continue
if state == "LEADING":
m = re.match(r"(?is)select\s+", sql[i:n])
if m:
state = "IN_LIST"
list_start = i + m.end()
i += m.end()
continue
i += 1
continue
if list_paren == 0:
mf = re.match(r"(?is)\bfrom\b", sql[i:n])
if mf:
frag = sql[list_start:i]
if re.match(r"(?is)^\s*\*\s*$", frag.strip()):
return None
return (col_commas + 1, frag)
c = sql[i]
if c == "(":
list_paren += 1
elif c == ")":
list_paren = max(0, list_paren - 1)
elif c == "," and list_paren == 0:
col_commas += 1
i += 1
return None
def _has_top_level_group_by(sql: str, connector: ConnectorType) -> bool:
scan = _get_scanner(connector)
n = len(sql)
i = 0
depth = 0
while i < n:
tok = scan(sql, i)
if tok["type"] != _TOKEN_PLAIN:
i = tok["end"]
continue
if depth == 0:
m = re.match(r"(?is)\bgroup\s+by\b", sql[i:n])
if m:
return True
c = sql[i]
if c == "(":
depth += 1
elif c == ")":
depth = max(0, depth - 1)
i += 1
return False
def _is_aggregate_scalar_fast_path(inner: str, connector: ConnectorType) -> bool:
"""单列、无 GROUP BY、且选择列表以常见聚合函数开头时,结果集为标量,宜直接执行内层而非 COUNT(*) 包裹。"""
parsed = _parse_top_level_select_list(inner, connector)
if parsed is None:
return False
ncols, frag = parsed
if ncols != 1:
return False
if _has_top_level_group_by(inner, connector):
return False
fs = frag.strip()
if not fs or re.match(r"(?is)^\*\s*$", fs):
return False
return bool(_AGGREGATE_SCALAR_PREFIX.match(fs))
def _strip_outer_trailing_order_by(sql: str, connector: ConnectorType) -> str:
"""
SQL Server / MySQL / MariaDB:派生表内 ``ORDER BY`` 若无 TOP/OFFSET/LIMIT 会报语法错误。
在 COUNT 子查询包装前去掉最外层(括号深度为 0)最后一次出现的 ``ORDER BY`` 及其后内容;
不改变行数统计语义。注释与字符串内的括号/关键字不参与解析。
"""
if connector not in ("sqlserver", "mysql", "mariadb"):
return sql
scan = _get_scanner(connector)
n = len(sql)
i = 0
depth = 0
last_order_by_start = -1
while i < n:
tok = scan(sql, i)
if tok["type"] != _TOKEN_PLAIN:
i = tok["end"]
continue
# 方言扫描器对「普通字符」常一次只前进一个字符;ORDER BY 必须从当前位置看到串尾才匹配得到
if depth == 0:
m = re.match(r"(?is)order\s+by\b", sql[i:n])
if m:
last_order_by_start = i
i += m.end()
continue
c = sql[i]
if c == "(":
depth += 1
elif c == ")":
depth = max(0, depth - 1)
i += 1
if last_order_by_start < 0:
return sql
return sql[:last_order_by_start].rstrip().rstrip(";").rstrip()
def _first_sql_keyword(sql: str, connector: ConnectorType) -> str:
"""剥离注释与字符串字面量后,取首个非空白词(小写)。"""
cleaned = strip_comments_and_strings(sql, connector)
cleaned = cleaned.strip().lower()
m = re.search(r"\S+", cleaned)
return m.group(0) if m else ""
def _sql_wrapped_count(
inner_sql: str,
connector: ConnectorType,
ncols: int | None,
) -> str:
inner = inner_sql.strip()
if not inner:
raise ValueError("SQL 不能为空")
inner = _strip_outer_trailing_order_by(inner, connector)
core = (
"SELECT COUNT(*) AS __dbhub_cnt FROM (\n"
f"{inner}\n"
") AS __dbhub_subq"
)
if ncols is not None and connector in ("sqlserver", "mysql", "mariadb"):
names = ", ".join(f"__dbhub_c{k}" for k in range(ncols))
return f"{core} ({names})"
return core
def _execute_sql_statements_count_only(
engine: Engine,
statements: list[str],
connector: ConnectorType,
) -> int:
"""
顺序执行多条语句;对最后一条:优先用 COUNT 子查询只取标量(select/with/explain),
否则单次执行后在 Python 侧逐行计数(不组装 rows 明细)。
"""
with engine.connect() as conn:
for i, stmt in enumerate(statements):
s = stmt.strip()
if i < len(statements) - 1:
r = conn.execute(text(s))
if r.returns_rows:
r.fetchall()
continue
if not s:
raise ValueError("SQL 为空")
first = _first_sql_keyword(s, connector)
if first in _COUNT_WRAP_FIRST_WORDS:
inner = _strip_outer_trailing_order_by(s.strip(), connector)
if _is_aggregate_scalar_fast_path(inner, connector):
r = conn.execute(text(inner))
row = r.fetchone()
if row is None or row[0] is None:
return 0
cell = row[0]
if isinstance(cell, Decimal):
return int(cell)
return int(cell)
parsed = _parse_top_level_select_list(inner, connector)
if parsed is None:
r = conn.execute(text(inner))
if not r.returns_rows:
return 0
n = 0
for _ in r:
n += 1
return n
ncols, _frag = parsed
wrapped = _sql_wrapped_count(inner, connector, ncols)
r = conn.execute(text(wrapped))
row = r.fetchone()
if row is None:
return 0
return int(row[0])
r = conn.execute(text(s))
if not r.returns_rows:
return 0
n = 0
for _ in r:
n += 1
return n
def _serialize_sql_cell(val: object) -> object:
if isinstance(val, Decimal):
return str(val)
if isinstance(val, (bytes, bytearray)):
return bytes(val).decode("utf-8", errors="replace")
if hasattr(val, "isoformat"):
try:
return val.isoformat()
except Exception: # noqa: BLE001
return val
return val
def _execute_sql_statements(
engine: Engine,
statements: list[str],
*,
max_rows: int | None,
) -> tuple[list[str], list[dict[str, object]], bool]:
"""顺序执行多条语句,仅返回最后一条有结果集语句的列与行。``max_rows`` 为 None 时对最后一条 ``fetchall``,不截断。"""
last_cols: list[str] = []
last_rows: list[dict[str, object]] = []
truncated = False
with engine.connect() as conn:
result = None
for i, stmt in enumerate(statements):
s = stmt.strip()
result = conn.execute(text(s))
if i < len(statements) - 1:
if result.returns_rows:
result.fetchall()
continue
if not result.returns_rows:
last_cols = []
last_rows = []
truncated = False
break
last_cols = list(result.keys())
if max_rows is None:
rows_data = result.fetchall()
truncated = False
else:
mr = max(1, int(max_rows))
fetched = result.fetchmany(mr + 1)
truncated = len(fetched) > mr
rows_data = fetched[:mr]
last_rows = []
for row in rows_data:
row_map: dict[str, object] = {}
for j, col in enumerate(last_cols):
row_map[col] = _serialize_sql_cell(row[j])
last_rows.append(row_map)
return last_cols, last_rows, truncated
def _execute_sql(
sql: str,
*,
readonly: bool = True,
max_rows: int | None = None,
) -> dict[str, object]:
"""
在已配置的业务库上执行 SQL(可多语句,分号分隔)。
:param sql: 待执行 SQL 字符串。多条语句用分号分隔;
:param readonly: 是否启用只读校验,默认 True。为 True 时,任一条语句不符合 G3SB 允许的首关键字
(如 select、with、explain 等,随方言而异)或含变更类关键字则抛出 ValueError。
为 False 时不做上述校验,可执行 DML/DDL(风险自负,勿用于不可信输入)。
:param max_rows: 对「最后一条返回结果集的语句」最多取多少行;None 时使用配置项 sql_max_rows,
并在实现侧夹紧到 1~10000。前面的语句若产生结果集会被消费掉但不返回。
:return: 字典,含 rows、count、source_id(固定 ``default``)、columns、truncated。
最后一条语句无结果集时 rows、columns 为空,count 为 0。未配置 database_url、SQL 为空或执行失败时
抛出 ValueError(或包装后的底层异常信息)。
"""
raw = (sql or "").strip()
if not raw:
raise ValueError("SQL 不能为空")
url = (_get_optional_str("database_url") or "").strip()
if not url:
raise ValueError("database_url 未配置")
engine = get_engine()
connector = sqlalchemy_dialect_to_connector(engine.dialect.name)
statements = split_sql_statements(raw, connector)
if not statements:
raise ValueError("SQL 为空")
if readonly:
for st in statements:
if not is_read_only_sql(st, connector):
kw = ", ".join(allowed_keywords_list(connector)) or "none"
raise ValueError(
f"Read-only mode is enabled. Only the following SQL operations are allowed: {kw}",
)
mr = max_rows if max_rows is not None else int(_get_optional_str("sql_max_rows") or 10000)
mr = max(1, min(10000, int(mr)))
log.info(f"_execute_sql 执行语句数={len(statements)} readonly={readonly} max_rows={mr}")
try:
cols, rows, truncated = _execute_sql_statements(engine, statements, max_rows=mr)
except Exception as e: # noqa: BLE001
log.error(f"_execute_sql 执行失败: {e}")
raise ValueError(str(e)) from e
count = len(rows)
return {
"rows": rows,
"count": count,
"source_id": "default",
"columns": cols,
"truncated": truncated,
}
def _execute_sql_all(
sql: str,
*,
readonly: bool = True,
) -> dict[str, object]:
"""
与 ``_execute_sql`` 相同校验与返回结构,但对最后一条有结果集的语句不做 ``max_rows`` 截断,全部 ``fetchall``。
超大结果集会占用大量内存,仅用于可信 SQL / 已自行 LIMIT 的场景。
"""
raw = (sql or "").strip()
if not raw:
raise ValueError("SQL 不能为空")
url = (_get_optional_str("database_url") or "").strip()
if not url:
raise ValueError("database_url 未配置")
engine = get_engine()
connector = sqlalchemy_dialect_to_connector(engine.dialect.name)
statements = split_sql_statements(raw, connector)
if not statements:
raise ValueError("SQL 为空")
if readonly:
for st in statements:
if not is_read_only_sql(st, connector):
kw = ", ".join(allowed_keywords_list(connector)) or "none"
raise ValueError(
f"Read-only mode is enabled. Only the following SQL operations are allowed: {kw}",
)
log.info(f"_execute_sql_all 执行语句数={len(statements)} readonly={readonly}")
try:
cols, rows, truncated = _execute_sql_statements(engine, statements, max_rows=None)
except Exception as e: # noqa: BLE001
log.error(f"_execute_sql_all 执行失败: {e}")
raise ValueError(str(e)) from e
count = len(rows)
return {
"rows": rows,
"count": count,
"source_id": "default",
"columns": cols,
"truncated": truncated,
}
def _execute_sql_count_only(
sql: str,
*,
readonly: bool = True,
) -> dict[str, int]:
"""
在已配置的业务库上执行 SQL,仅返回「最后一条有结果集语句」的行数,不返回单元格明细。
对以 select / with / explain 开头(注释剥离后)的最后一条语句:若已为单列常见聚合(如 ``count(1)``)且无 ``GROUP BY``,
则直接执行该语句并取标量,避免 ``COUNT(*)`` 外包一层只得到 1 行。否则在库侧用 ``SELECT COUNT(*) FROM ( ... )``;
对 SQL Server / MySQL / MariaDB 会为派生表补上 ``(__dbhub_c0, ...)`` 列名(避免匿名列 8155 等),并在包装前去掉
最外层无意义的 ``ORDER BY``。无法解析选择列表时(如 ``WITH``、``SELECT *``)改为拉全结果在 Python 侧逐行计数。
其它只读语句(如部分 show/pragma)仍单次执行后逐行计数,不组装为 dict。
多条语句用分号分隔时,前面的语句照常执行并消费结果集;仅对最后一条给出计数。
最后一条无结果集时 count 为 0。进入执行阶段后若任一步失败(含语法、权限、连接中断等),返回 ``{"count": -1}``;
SQL 为空、未配置 ``database_url``、只读校验不通过仍抛 ``ValueError``。
"""
raw = (sql or "").strip()
if not raw:
raise ValueError("SQL 不能为空")
url = (_get_optional_str("database_url") or "").strip()
if not url:
raise ValueError("database_url 未配置")
engine = get_engine()
connector = sqlalchemy_dialect_to_connector(engine.dialect.name)
statements = split_sql_statements(raw, connector)
if not statements:
raise ValueError("SQL 为空")
if readonly:
for st in statements:
if not is_read_only_sql(st, connector):
kw = ", ".join(allowed_keywords_list(connector)) or "none"
raise ValueError(
f"Read-only mode is enabled. Only the following SQL operations are allowed: {kw}",
)
log.info(f"_execute_sql_count_only 执行语句数={len(statements)} readonly={readonly}")
try:
n = _execute_sql_statements_count_only(engine, statements, connector)
except Exception as e: # noqa: BLE001
log.error(f"_execute_sql_count_only 执行失败: {e}")
return {"count": -1}
return {"count": n}
def probe_sql_execution_status_ex(
sql: str, *, max_rows: int = 1
) -> tuple[int | None, str | None]:
"""
在已配置的业务库上对 SQL 做只读试执行,返回状态码与失败时的错误摘要。
Returns:
二元组 ``(status, error_message)``:
- ``(None, None)``:未配置 ``database_url``,跳过探针。
- ``(1, None)``:执行成功,且**至少返回一行数据**(与「有列无行」的 SELECT 区分)。
- ``(0, None)``:执行成功,但**数据行数为 0**(可有列名而无行,或无非空结果集)。
- ``(-1, msg)``:执行失败;``msg`` 为异常信息摘要(便于生成阶段重试)。
"""
url = (_get_optional_str("database_url") or "").strip()
if not url:
return None, None
try:
out = _execute_sql(sql, readonly=True, max_rows=max_rows)
except Exception as e: # noqa: BLE001
log.warning(f"probe_sql_execution_status 执行失败: {e}")
detail = str(e).strip() or "unknown error"
return -1, detail
rows = list(out.get("rows") or [])
if len(rows) > 0:
return 1, None
return 0, None
def probe_sql_execution_status(sql: str, *, max_rows: int = 1) -> int | None:
"""
在已配置的业务库上对 SQL 做只读试执行,返回紧凑状态码(用于 Validator 探针)。
语义与 :func:`probe_sql_execution_status_ex` 的首个返回值一致;失败细节请用 ``_ex``。
"""
status, _ = probe_sql_execution_status_ex(sql, max_rows=max_rows)
return status
def _search_objects(
object_type: ObjectType,
pattern: str = "%",
schema: str | None = None,
table: str | None = None,
detail_level: DetailLevel = "names",
limit: int = 100,
) -> dict[str, Any]:
"""
按对象类型在已配置的业务库中探索 schema、表、列、索引、存储过程/函数等元数据(对齐 DBHub search_objects)。
会过滤系统 schema(如information_schema、sys、pg_catalog 等),再按 object_type 与 pattern 做 SQL LIKE 风格匹配;
:param object_type: 要探索的对象类别。schema-模式名;table-数据表;column-列(可与 table、schema 联用
限定单表);index-索引(同上);procedure-存储过程;function-函数(标量/表值等,视库而定)。
:param pattern: SQL LIKE 模式,默认 "%" 表示不过滤名称。"%" 匹配任意长度子串,"_" 匹配单个字符;对表名、
列名、索引名、例程名等做大小写不敏感匹配。
:param schema: 限定在某个 schema 内查找;None 表示在多个非系统 schema 上依次查找。若给出具体名称,
必须是库中已存在的 schema,否则抛 ValueError。使用参数 table 时必须同时指定 schema。
:param table: 仅在 object_type 为 column 或 index 时允许传入,与 schema 共同限定「只查这一张表或视图」上的列
或索引;用于其它 object_type 时会抛 ValueError。名称会在该 schema 下与反射得到的表名、视图名做
不区分大小写匹配(并支持 [Name] 写法),以兼容 SQL Server 等对目录大小写不敏感但 API 需真实写法的情况。
:param detail_level: 返回粒度。names-仅对象名及定位字段(如 schema、table);summary-增加简要元数据
(如表的 column_count、row_count,列的类型、可空等);full-表级返回列列表、索引列表等完整结构。
:param limit: 最多返回的结果条数,默认 100,有效范围 1~1000(传入值会被夹紧到该区间)。
:return: 包含 object_type、pattern、schema、table、detail_level、count、results、truncated 的字典。
未配置 database_url、连接失败或内省失败时抛出 ValueError(或其它底层异常)。
"""
if table and not schema:
raise ValueError("The 'table' parameter requires 'schema' to be specified")
if table and object_type not in ("column", "index"):
raise ValueError(
f"The 'table' parameter only applies to object_type 'column' or 'index', not '{object_type}'",
)
lim = max(1, min(1000, int(limit)))
pat = pattern if pattern is not None else "%"
rx = _like_pattern_to_regex(pat)
url = (_get_optional_str("database_url") or "").strip()
if not url:
raise ValueError("database_url 未配置")
try:
engine = get_engine()
except ValueError as e:
raise ValueError(str(e)) from e
try:
insp = inspect(engine)
except SQLAlchemyError as e:
log.error(f"_search_objects inspect 失败: {e}")
raise ValueError(f"数据库内省失败: {e}") from e
try:
all_schema_names = list(insp.get_schema_names())
except Exception: # noqa: BLE001
all_schema_names = []
filtered_schemas = _filter_schemas(all_schema_names)
if schema:
if schema not in all_schema_names:
avail = ", ".join(all_schema_names) or "(none)"
raise ValueError(f"Schema '{schema}' does not exist. Available schemas: {avail}")
schemas_to_search = [schema]
else:
schemas_to_search = filtered_schemas if filtered_schemas else [None]
dialect = engine.dialect.name
results: list[Any] = []
if object_type == "schema":
base_schemas = filtered_schemas if filtered_schemas else all_schema_names
candidates = [s for s in base_schemas if rx.match(s)]
for schema_name in candidates:
if len(results) >= lim:
break
if detail_level == "names":
results.append({"name": schema_name})
else:
try:
tbls = list(insp.get_table_names(schema=schema_name))
except Exception: # noqa: BLE001
tbls = []
results.append({"name": schema_name, "table_count": len(tbls)})
elif object_type == "table":
for schema_name in schemas_to_search:
if len(results) >= lim:
break
try:
tbls = list(insp.get_table_names(schema=schema_name))
except Exception: # noqa: BLE001
continue
for table_name in tbls:
if len(results) >= lim:
break
bare = _bare_table_name(str(table_name))
if not rx.match(bare):
continue
if detail_level == "names":
results.append({"name": bare, "schema": schema_name})
elif detail_level == "summary":
cols = _column_dicts(insp, schema_name, bare)
rc = _table_row_count(engine, schema_name, bare)
comment = _table_comment(insp, schema_name, bare)
result_dict: dict[str, Any] = {
"name": bare,
"schema": schema_name,
"column_count": len(cols),
"row_count": rc,
}
if comment:
result_dict["comment"] = comment
results.append(result_dict)
else:
cols = _column_dicts(insp, schema_name, bare)
idxs = _index_dicts(insp, schema_name, bare)
rc = _table_row_count(engine, schema_name, bare)
comment = _table_comment(insp, schema_name, bare)
result_dict_full: dict[str, Any] = {
"name": bare,
"schema": schema_name,
"column_count": len(cols),
"row_count": rc,
"columns": [
{
"name": c["column_name"],
"type": c["data_type"],
"nullable": c["is_nullable"] == "YES",
"default": c["column_default"],
**({"description": c["description"]} if "description" in c else {}),
}
for c in cols
],
"indexes": idxs,
}
if comment:
result_dict_full["comment"] = comment
results.append(result_dict_full)
elif object_type in ("column", "index"):
for schema_name in schemas_to_search:
if len(results) >= lim:
break
try:
if table:
tables_to_search = [_resolve_table_or_view_bare_name(insp, schema_name, table)]
else:
tables_to_search = [_bare_table_name(str(t)) for t in insp.get_table_names(schema=schema_name)]
except Exception: # noqa: BLE001
continue
for table_name in tables_to_search:
if len(results) >= lim:
break
bare = _bare_table_name(str(table_name))
if object_type == "column":
cols = _column_dicts(insp, schema_name, bare)
for c in cols:
if len(results) >= lim:
break
if not rx.match(c["column_name"]):
continue
if detail_level == "names":
results.append({"name": c["column_name"], "table": bare, "schema": schema_name})
else:
result_dict_col: dict[str, Any] = {
"name": c["column_name"],
"table": bare,
"schema": schema_name,
"type": c["data_type"],
"nullable": c["is_nullable"] == "YES",
"default": c["column_default"],
}
if "description" in c:
result_dict_col["description"] = c["description"]
results.append(result_dict_col)
else:
idxs = _index_dicts(insp, schema_name, bare)
for idx in idxs:
if len(results) >= lim:
break
iname = str(idx.get("name") or "")
if not rx.match(iname):
continue
if detail_level == "names":
results.append({"name": iname, "table": bare, "schema": schema_name})
else:
results.append(
{
"name": iname,
"table": bare,
"schema": schema_name,
"columns": idx.get("columns"),
"unique": idx.get("unique"),
"primary": idx.get("primary"),
},
)
elif object_type in ("procedure", "function"):
rt = "PROCEDURE" if object_type == "procedure" else "FUNCTION"
try:
routine_rows = _fetch_routines(engine, dialect, schema_filter=schema, routine_sql_type=rt)
except Exception as e: # noqa: BLE001
log.warning(f"_search_objects 例程列表失败 dialect={dialect}: {e}")
routine_rows = []
for sch, name, rtype, defn in routine_rows:
if len(results) >= lim:
break
if not rx.match(name):
continue
if detail_level == "names":
results.append({"name": name, "schema": sch})
elif detail_level == "summary":
results.append(
{
"name": name,
"schema": sch,
"type": rtype,
"language": None,
"return_type": None,
},
)
else:
results.append(
{
"name": name,
"schema": sch,
"type": rtype,
"language": None,
"parameters": None,
"return_type": None,
"definition": defn,
},
)
else:
raise ValueError(f"Unsupported object_type: {object_type}")
return {
"object_type": object_type,
"pattern": pat,
"schema": schema,
"table": table,
"detail_level": detail_level,
"count": len(results),
"results": results,
"truncated": len(results) == lim,
}
class DbHubTools:
"""
与 DBHub 对齐的数据库工具封装:SQL 执行与元数据探索。
方法 execute_sql 行为对齐 dbhub execute-sql.ts;
方法 search_objects 对齐 DBHub search_objects。
execute_sql_all 为不截断行的全量拉取;
execute_sql_count_only 为仅统计行数;
实现委托至模块内 ``_execute_sql`` / ``_execute_sql_all`` / ``_execute_sql_count_only`` / ``_search_objects``。
"""
def execute_sql(
self,
sql: str,
*,
readonly: bool = True,
max_rows: int | None = None,
) -> dict[str, object]:
"""
在已配置的业务库上执行 SQL(可多语句,分号分隔)。
:param sql: 待执行 SQL 字符串。多条语句用分号分隔;
:param readonly: 是否启用只读校验,默认 True。为 True 时,任一条语句不符合 G3SB 允许的首关键字
(如 select、with、explain 等,随方言而异)或含变更类关键字则抛出 ValueError
为 False 时不做上述校验,可执行 DML/DDL(风险自负,勿用于不可信输入)。
:param max_rows: 对「最后一条返回结果集的语句」最多取多少行;None 时使用配置项 sql_max_rows,
并在实现侧夹紧到 1~10000。前面的语句若产生结果集会被消费掉但不返回。
:return: 字典,含 rows(每行一个 dict,列名到值的映射)、count(等于 rows 长度,等同 DBHub rowCount)、
columns(列名列表)、truncated(是否因超过 max_rows 而截断)。
最后一条语句无结果集时 rows、columns 为空,count 为 0。未配置 database_url、SQL 为空或执行失败时
抛出 ValueError(或包装后的底层异常信息)。
"""
return _execute_sql(sql, readonly=readonly, max_rows=max_rows)
def search_objects(
self,
object_type: ObjectType,
pattern: str = "%",
schema: str | None = None,
table: str | None = None,
detail_level: DetailLevel = "names",
limit: int = 100,
) -> dict[str, Any]:
"""
按对象类型在已配置的业务库中探索 schema、表、列、索引、存储过程/函数等元数据(对齐 DBHub search_objects)。
会过滤系统 schema(如information_schema、sys、pg_catalog 等),再按 object_type 与 pattern 做 SQL LIKE 风格匹配;
:param object_type: 要探索的对象类别。schema-模式名;table-数据表;column-列(可与 table、schema 联用
限定单表);index-索引(同上);procedure-存储过程;function-函数(标量/表值等,视库而定)。
:param pattern: SQL LIKE 模式,默认 "%" 表示不过滤名称。"%" 匹配任意长度子串,"_" 匹配单个字符;对表名、
列名、索引名、例程名等做大小写不敏感匹配。
:param schema: 限定在某个 schema 内查找;None 表示在多个非系统 schema 上依次查找。若给出具体名称,
必须是库中已存在的 schema,否则抛 ValueError。使用参数 table 时必须同时指定 schema。
:param table: 仅在 object_type 为 column 或 index 时允许传入,与 schema 共同限定「只查这一张表或视图」上的列
或索引;用于其它 object_type 时会抛 ValueError。名称会在该 schema 下与反射得到的表名、视图名做
不区分大小写匹配(并支持 [Name] 写法),以兼容 SQL Server 等对目录大小写不敏感但 API 需真实写法的情况。
:param detail_level: 返回粒度。names-仅对象名及定位字段(如 schema、table);summary-增加简要元数据
(如表的 column_count、row_count,列的类型、可空等);full-表级返回列列表、索引列表等完整结构。
:param limit: 最多返回的结果条数,默认 100,有效范围 1~1000(传入值会被夹紧到该区间)。
:return: 包含 object_type、pattern、schema、table、detail_level、count、results、truncated 的字典。
未配置 database_url、连接失败或内省失败时抛出 ValueError(或其它底层异常)。
"""
return _search_objects(
object_type,
pattern=pattern,
schema=schema,
table=table,
detail_level=detail_level,
limit=limit,
)
def execute_sql_all(
self,
sql: str,
*,
readonly: bool = True,
) -> dict[str, object]:
"""
在已配置的业务库上执行 SQL,对最后一条产生结果集的语句 fetchall 全量取行,不做行数上限截断。(4.SQL 执行接口)
:param sql: 待执行 SQL。多条语句用分号分隔;仅返回最后一条有结果集语句的 columns/rows,
前面语句若产生结果集会被执行并消费掉但不返回。
:param readonly: 是否启用只读校验,默认 True。为 True 时任一条不符合 G3SB 允许的首关键字(如 select、with、explain 等,
随方言而异)或含变更类关键字则抛出 ValueError;为 False 时不做该校验(风险自负,勿用于不可信输入)。
:return: 字典含 rows、count(等于 rows 长度)、columns、truncated。本路径下 truncated 恒为 False。
最后一条语句无结果集时 rows、columns 为空,count 为 0。未配置 database_url、SQL 为空、校验失败或执行失败时
抛出 ValueError(或带底层信息的包装异常)。
"""
return _execute_sql_all(sql, readonly=readonly)
def execute_sql_count_only(self, sql: str) -> dict[str, int]:
"""
仅返回结果集语句的行数,不返回 rows/columns 明细;(3.SQL 检验接口)
:param sql: 待执行 SQL 字符串。多条语句用分号分隔;
:return: {count: int}
"""
return _execute_sql_count_only(sql, readonly=True)
# 默认工具实例;下方模块级 execute_sql / execute_sql_all / execute_sql_count_only / search_objects 均委托至此。
dbhub_tools = DbHubTools()
def search_objects(**kwargs: Any) -> dict[str, Any]:
return dbhub_tools.search_objects(**kwargs)
def execute_sql(sql: str, **kwargs: Any) -> dict[str, object]:
return dbhub_tools.execute_sql(sql, **kwargs)
def execute_sql_all(sql: str, **kwargs: Any) -> dict[str, object]:
return dbhub_tools.execute_sql_all(sql, **kwargs)
def execute_sql_count_only(sql: str, **kwargs: Any) -> dict[str, int]:
return dbhub_tools.execute_sql_count_only(sql, **kwargs)
def main() -> None:
"""
命令行演示:需已配置 database_url。SQL Server 示例默认 schema 为 dbo;
PostgreSQL 可将下面示例中的 schema 改为 public。
"""
demos: list[tuple[str, dict[str, Any]]] = [
(
"1) 列出 schema(detail_level=names)",
{"object_type": "schema", "detail_level": "names", "limit": 5},
),
(
"2) 列出某 schema 下的表名",
{"object_type": "table", "schema": "dbo", "detail_level": "names", "limit": 5},
),
(
"3) 表名 LIKE 模糊匹配(%Account%)",
{
"object_type": "table",
"schema": "dbo",
"pattern": "%Account%",
"detail_level": "names",
"limit": 20,
},
),
(
"4) 表级摘要(列数、行数 COUNT)",
{"object_type": "table", "schema": "dbo", "detail_level": "summary", "limit": 5},
),
(
"6) 表级返回列列表、索引列表等完整结构",
{"object_type": "table", "schema": "dbo", "detail_level": "full", "limit": 5},
),
(
"5) 某表或视图的列(支持大小写/视图;无则改 table 为库中真实对象名)",
{
"object_type": "column",
"schema": "dbo",
"table": "BCAccountAccruedCustodianFee",
"pattern": "%",
"detail_level": "names",
"limit": 5,
},
),
]
# for title, kwargs in demos:
# print(f"\n=== {title} ===")
# try:
# out = dbhub_tools.search_objects(**kwargs)
# print(json.dumps(out, ensure_ascii=False, indent=2, default=str))
# except Exception as e: # noqa: BLE001
# log.warning(f"示例跳过或失败: {title} err={e}")
# print(f"(失败) {e}")
sql = """
SELECT
b.BrokerID,
m.Name AS BrokerName,
b.CurrencyID,
b.Settled AS OwedToUsAmount,
ABS(b.Settled) AS OutstandingAmount
FROM BCBrokerCash b
INNER JOIN MCBroker m ON b.BrokerID = m.BrokerID
WHERE b.Settled < 0
ORDER BY ABS(b.Settled) DESC;
"""
out = dbhub_tools.execute_sql(sql)
print(json.dumps(out, ensure_ascii=False, indent=2, default=str))
sql = """
SELECT
b.BrokerID,
m.Name AS BrokerName,
b.CurrencyID,
b.Settled AS OwedToUsAmount,
ABS(b.Settled) AS OutstandingAmount
FROM BCBrokerCash b
INNER JOIN MCBroker m ON b.BrokerID = m.BrokerID
WHERE b.Settled < 0
ORDER BY ABS(b.Settled) DESC;
"""
out = dbhub_tools.execute_sql_count_only(sql)
print(json.dumps(out, ensure_ascii=False, indent=2, default=str))
sql = """
SELECT
b.BrokerID,
m.Name AS BrokerName,
b.CurrencyID,
b.Settled AS OwedToUsAmount,
ABS(b.Settled) AS OutstandingAmount
FROM BCBrokerCash b
INNER JOIN MCBroker m ON b.BrokerID = m.BrokerID
WHERE b.Settled < 0
ORDER BY ABS(b.Settled) DESC;
"""
out = dbhub_tools.execute_sql_all(sql)
print(json.dumps(out, ensure_ascii=False, indent=2, default=str))
if __name__ == "__main__":
main()