0.1.1 暂存
This commit is contained in:
@@ -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.
Binary file not shown.
Binary file not shown.
@@ -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]
|
||||
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user