0.1.1 暂存

This commit is contained in:
陈辅元
2026-04-14 10:28:22 +08:00
parent 4a0638ba2e
commit cc83fe963a
83 changed files with 4005 additions and 179 deletions
+31
View File
@@ -0,0 +1,31 @@
"""
数据库工具:对齐 DBHub 的 ``execute_sql`` / ``search_objects``。
使用前请将 ``backend`` 目录加入 ``sys.path``(与 ``api_server.py`` / ``backend/main.py`` 一致),并已在进程内加载根目录 ``.env``。
"""
from __future__ import annotations
from db.dbhub_tools import (
DbHubTools,
dbhub_tools,
execute_sql,
execute_sql_all,
execute_sql_count_only,
probe_sql_execution_status,
probe_sql_execution_status_ex,
search_objects,
)
from db.engine import get_engine
__all__ = [
"DbHubTools",
"dbhub_tools",
"execute_sql",
"execute_sql_all",
"execute_sql_count_only",
"get_engine",
"probe_sql_execution_status",
"probe_sql_execution_status_ex",
"search_objects",
]
Binary file not shown.
Binary file not shown.
Binary file not shown.
+87
View File
@@ -0,0 +1,87 @@
"""
从 DBHub allowed-keywords.ts 等价移植:只读 SQL 判定。
参见 dbhub/src/utils/allowed-keywords.ts
"""
from __future__ import annotations
import re
from typing import Literal
from db.dbhub_sql_parser import ConnectorType, strip_comments_and_strings
ALLOWED_KEYWORDS: dict[ConnectorType, list[str]] = {
"postgres": ["select", "with", "explain", "show"],
"mysql": ["select", "with", "explain", "show", "describe", "desc"],
"mariadb": ["select", "with", "explain", "show", "describe", "desc"],
"sqlite": ["select", "with", "explain", "pragma"],
"sqlserver": ["select", "with", "explain", "showplan"],
}
_MUTATING = [
"insert",
"update",
"delete",
"drop",
"alter",
"create",
"truncate",
"merge",
"grant",
"revoke",
"rename",
]
_mutating_pattern = re.compile(rf"\b(?:{'|'.join(_MUTATING)})\b", re.IGNORECASE)
_mutating_pattern_with_replace = re.compile(
rf"\b(?:{'|'.join(_MUTATING)}|replace\s+(?:(?:low_priority|delayed)\s+)?into)\b",
re.IGNORECASE,
)
_MUTATING_PATTERNS: dict[ConnectorType, re.Pattern[str]] = {
"postgres": _mutating_pattern,
"mysql": _mutating_pattern_with_replace,
"mariadb": _mutating_pattern_with_replace,
"sqlite": _mutating_pattern_with_replace,
"sqlserver": _mutating_pattern,
}
_SELECT_INTO_PATTERN = re.compile(r"\bselect\b[\s\S]+\binto\b", re.IGNORECASE)
_EXPLAIN_ANALYZE_PATTERN = re.compile(
r"^explain\s+(?:\([^)]*\banalyze\b(?!\s*(?:=\s*)?(?:false|off|0)\b)[^)]*\)|\banalyze\b(?!\s*(?:=\s*)?(?:false|off|0)\b)(?:\s+verbose\b)?)",
re.IGNORECASE,
)
def _check_read_only(cleaned_sql: str, connector_type: ConnectorType | str) -> bool:
if not cleaned_sql:
return False
m = re.search(r"\S+", cleaned_sql)
first_word = m.group(0) if m else ""
keyword_list = ALLOWED_KEYWORDS.get(connector_type, []) # type: ignore[arg-type]
if first_word not in keyword_list:
return False
if first_word == "with":
pat = _MUTATING_PATTERNS.get(connector_type, _mutating_pattern) # type: ignore[arg-type]
if pat.search(cleaned_sql):
return False
if first_word in ("select", "with") and _SELECT_INTO_PATTERN.search(cleaned_sql):
return False
if first_word == "explain":
em = _EXPLAIN_ANALYZE_PATTERN.match(cleaned_sql)
if em:
after_explain = cleaned_sql[em.end() :].strip()
if after_explain and not _check_read_only(after_explain, connector_type):
return False
return True
def is_read_only_sql(sql: str, connector_type: ConnectorType | str) -> bool:
"""Check if a SQL query is read-only (DBHub-compatible)."""
cleaned = strip_comments_and_strings(sql, connector_type if connector_type in ALLOWED_KEYWORDS else None)
cleaned = cleaned.strip().lower()
return _check_read_only(cleaned, connector_type)
def allowed_keywords_list(connector_type: ConnectorType | str) -> list[str]:
return ALLOWED_KEYWORDS.get(connector_type, []) # type: ignore[arg-type]
+268
View File
@@ -0,0 +1,268 @@
"""
从 DBHub sql-parser.ts 等价移植:按方言剥离注释/字符串、切分语句。
参见 dbhub/src/utils/sql-parser.ts
"""
from __future__ import annotations
import re
from typing import Callable, Literal, TypedDict
ConnectorType = Literal["postgres", "mysql", "mariadb", "sqlite", "sqlserver"]
class _Token(TypedDict):
type: int # 0 Plain, 1 Comment, 2 QuotedBlock
end: int
_TOKEN_PLAIN = 0
_TOKEN_COMMENT = 1
_TOKEN_QUOTED = 2
def _plain_token(i: int) -> _Token:
return {"type": _TOKEN_PLAIN, "end": i + 1}
def _scan_single_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "-" or sql[i + 1] != "-":
return None
j = i
while j < len(sql) and sql[j] != "\n":
j += 1
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_multi_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
j = i + 2
while j + 1 < len(sql) and not (sql[j] == "*" and sql[j + 1] == "/"):
j += 1
if j + 1 < len(sql):
j += 2
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_multi_line_comment_mysql(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
nxt = sql[i + 2] if i + 2 < len(sql) else ""
nxt2 = sql[i + 3] if i + 3 < len(sql) else ""
if nxt == "!" or (nxt == "M" and nxt2 == "!"):
return None
return _scan_multi_line_comment(sql, i)
def _scan_nested_multi_line_comment(sql: str, i: int) -> _Token | None:
if i + 1 >= len(sql) or sql[i] != "/" or sql[i + 1] != "*":
return None
j = i + 2
depth = 1
while j < len(sql) and depth > 0:
if j + 1 < len(sql) and sql[j] == "/" and sql[j + 1] == "*":
depth += 1
j += 2
elif j + 1 < len(sql) and sql[j] == "*" and sql[j + 1] == "/":
depth -= 1
j += 2
else:
j += 1
return {"type": _TOKEN_COMMENT, "end": j}
def _scan_single_quoted_string(sql: str, i: int) -> _Token | None:
if sql[i] != "'":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "'" and sql[j + 1] == "'":
j += 2
elif sql[j] == "'":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_double_quoted_string(sql: str, i: int) -> _Token | None:
if sql[i] != '"':
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == '"' and sql[j + 1] == '"':
j += 2
elif sql[j] == '"':
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
_dollar_quote_open_regex = re.compile(r"^\$([a-zA-Z_]\w*)?\$")
def _scan_dollar_quoted_block(sql: str, i: int) -> _Token | None:
if sql[i] != "$":
return None
nxt = sql[i + 1] if i + 1 < len(sql) else ""
if nxt.isdigit():
return None
remaining = sql[i:]
m = _dollar_quote_open_regex.match(remaining)
if not m:
return None
tag = m.group(0)
body_start = i + len(tag)
close_idx = sql.find(tag, body_start)
end = close_idx + len(tag) if close_idx != -1 else len(sql)
return {"type": _TOKEN_QUOTED, "end": end}
def _scan_backtick_quoted_identifier(sql: str, i: int) -> _Token | None:
if sql[i] != "`":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "`" and sql[j + 1] == "`":
j += 2
elif sql[j] == "`":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_bracket_quoted_identifier(sql: str, i: int) -> _Token | None:
if sql[i] != "[":
return None
j = i + 1
while j < len(sql):
if j + 1 < len(sql) and sql[j] == "]" and sql[j + 1] == "]":
j += 2
elif sql[j] == "]":
j += 1
break
else:
j += 1
return {"type": _TOKEN_QUOTED, "end": j}
def _scan_token_ansi(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _plain_token(i)
)
def _scan_token_postgres(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_nested_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_dollar_quoted_block(sql, i)
or _plain_token(i)
)
def _scan_token_mysql(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment_mysql(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_backtick_quoted_identifier(sql, i)
or _plain_token(i)
)
def _scan_token_sqlite(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_backtick_quoted_identifier(sql, i)
or _scan_bracket_quoted_identifier(sql, i)
or _plain_token(i)
)
def _scan_token_sqlserver(sql: str, i: int) -> _Token:
return (
_scan_single_line_comment(sql, i)
or _scan_multi_line_comment(sql, i)
or _scan_single_quoted_string(sql, i)
or _scan_double_quoted_string(sql, i)
or _scan_bracket_quoted_identifier(sql, i)
or _plain_token(i)
)
_DIALECT_SCANNERS: dict[ConnectorType, Callable[[str, int], _Token]] = {
"postgres": _scan_token_postgres,
"mysql": _scan_token_mysql,
"mariadb": _scan_token_mysql,
"sqlite": _scan_token_sqlite,
"sqlserver": _scan_token_sqlserver,
}
def _get_scanner(dialect: ConnectorType | None) -> Callable[[str, int], _Token]:
if dialect and dialect in _DIALECT_SCANNERS:
return _DIALECT_SCANNERS[dialect]
return _scan_token_ansi
def strip_comments_and_strings(sql: str, dialect: ConnectorType | None = None) -> str:
"""Replace comments, string literals, and dialect-specific quoted blocks with a single space each."""
scan_token = _get_scanner(dialect)
parts: list[str] = []
plain_start = -1
i = 0
n = len(sql)
while i < n:
token = scan_token(sql, i)
if token["type"] == _TOKEN_PLAIN:
if plain_start == -1:
plain_start = i
else:
if plain_start != -1:
parts.append(sql[plain_start:i])
plain_start = -1
parts.append(" ")
i = token["end"]
if plain_start != -1:
parts.append(sql[plain_start:])
return "".join(parts)
def split_sql_statements(sql: str, dialect: ConnectorType | None = None) -> list[str]:
"""Split SQL into individual statements, handling semicolons inside quoted contexts."""
scan_token = _get_scanner(dialect)
statements: list[str] = []
stmt_start = 0
i = 0
n = len(sql)
while i < n:
if sql[i] == ";":
trimmed = sql[stmt_start:i].strip()
if trimmed:
statements.append(trimmed)
stmt_start = i + 1
i += 1
continue
token = scan_token(sql, i)
i = token["end"]
trimmed = sql[stmt_start:].strip()
if trimmed:
statements.append(trimmed)
return statements
File diff suppressed because it is too large Load Diff
+46
View File
@@ -0,0 +1,46 @@
"""
SQLAlchemy 引擎:使用环境变量 ``database_url`` / ``DATABASE_URL``(与项目根目录 ``.env`` 一致)。
"""
from __future__ import annotations
import logging
import os
from typing import Optional
from sqlalchemy import create_engine
from sqlalchemy.engine import Engine
logger = logging.getLogger(__name__)
_engine: Engine | None = None
def _database_url_from_env() -> str:
for key in ("database_url", "DATABASE_URL"):
v = os.getenv(key)
if v is not None and str(v).strip():
return str(v).strip()
return ""
def get_engine(*, url: Optional[str] = None, reset: bool = False) -> Engine:
"""
返回默认业务库引擎(单例)。未配置 ``database_url`` 时抛出 ``ValueError``。
:param url: 若传入,则忽略单例并为此 URL 新建引擎(便于测试)。
:param reset: 为 True 时丢弃已缓存的单例,下次再按环境变量创建。
"""
global _engine
if reset:
_engine = None
if url is not None:
return create_engine(url, pool_pre_ping=True)
if _engine is not None:
return _engine
u = _database_url_from_env()
if not u:
raise ValueError("database_url 未配置")
_engine = create_engine(u, pool_pre_ping=True)
logger.info("SQLAlchemy engine initialized from database_url")
return _engine