Files
SuperBizAgent-java/scripts/query_mysql.py

118 lines
3.6 KiB
Python

#!/usr/bin/env python3
"""Read-only MySQL query helper used by local verification."""
import os
import sys
if sys.stdout.encoding and sys.stdout.encoding.lower() != "utf-8":
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
try:
import pymysql
import pymysql.cursors
except ImportError:
print("缺少依赖,请先执行: pip install pymysql")
sys.exit(1)
def required_env(name: str) -> str:
value = os.getenv(name)
if value is None or not value.strip():
print(f"缺少必需环境变量: {name}")
sys.exit(2)
return value.strip()
DB_CONFIG = {
"host": required_env("SUPERBIZ_MYSQL_HOST"),
"port": int(required_env("SUPERBIZ_MYSQL_PORT")),
"user": required_env("SUPERBIZ_MYSQL_USERNAME"),
"password": required_env("SUPERBIZ_MYSQL_PASSWORD"),
"database": required_env("SUPERBIZ_MYSQL_DATABASE"),
"charset": "utf8mb4",
"cursorclass": pymysql.cursors.DictCursor,
"read_timeout": 10,
"write_timeout": 10,
}
def validate_read_only_sql(sql: str) -> str:
normalized = sql.strip()
if not normalized:
raise ValueError("SQL 不能为空")
statements = [part.strip() for part in normalized.split(";") if part.strip()]
if len(statements) != 1:
raise ValueError("只允许单条 SELECT")
statement = statements[0]
upper = statement.upper()
if upper != "SELECT" and not upper.startswith("SELECT "):
raise ValueError("只允许 SELECT 查询")
padded = f" {upper} "
forbidden = (
" INTO ",
" FOR UPDATE",
" SHOW ",
" DESCRIBE ",
" INFORMATION_SCHEMA",
" WITH ",
" UNION ",
" CALL ",
)
if any(token in padded for token in forbidden):
raise ValueError("查询包含禁止的只读或元数据语法")
return statement
def run_query(sql: str) -> None:
statement = validate_read_only_sql(sql)
connection = pymysql.connect(**DB_CONFIG)
try:
with connection.cursor() as cursor:
cursor.execute("START TRANSACTION READ ONLY")
cursor.execute(statement)
rows = cursor.fetchall()
if not rows:
print("(空结果)")
return
columns = list(rows[0].keys())
widths = {
column: max(len(column), max(len(str(row[column])) for row in rows))
for column in columns
}
header = " | ".join(column.ljust(widths[column]) for column in columns)
print(header)
print("-" * len(header))
for row in rows:
print(" | ".join(str(row[column]).ljust(widths[column]) for column in columns))
print(f"\n({len(rows)} 行)")
finally:
connection.rollback()
connection.close()
if __name__ == "__main__":
if len(sys.argv) > 1:
try:
run_query(" ".join(sys.argv[1:]))
except ValueError as exc:
print(f"拒绝执行: {exc}")
sys.exit(3)
else:
print("MySQL 只读交互模式(输入 exit 退出)")
print(
f"连接: {DB_CONFIG['user']}@{DB_CONFIG['host']}:"
f"{DB_CONFIG['port']}/{DB_CONFIG['database']}"
)
print("-" * 50)
while True:
try:
sql = input("sql> ").strip()
if sql.lower() in ("exit", "quit", "q"):
break
if sql:
run_query(sql)
except KeyboardInterrupt:
break
except Exception as exc:
print(f"错误: {exc}")