feat(harness): add readonly mysql tool
This commit is contained in:
+75
-47
@@ -1,19 +1,11 @@
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
"""
|
||||
通用 MySQL 查询脚本
|
||||
用法:
|
||||
python scripts/query_mysql.py "SELECT * FROM diagnosis_session ORDER BY created_at DESC LIMIT 5"
|
||||
python scripts/query_mysql.py # 交互模式
|
||||
依赖:pip install pymysql
|
||||
"""
|
||||
"""Read-only MySQL query helper used by local verification."""
|
||||
|
||||
import sys
|
||||
import os
|
||||
import sys
|
||||
|
||||
# Windows 控制台 UTF-8 输出
|
||||
if sys.stdout.encoding and sys.stdout.encoding.lower() != 'utf-8':
|
||||
sys.stdout.reconfigure(encoding='utf-8', errors='replace')
|
||||
if sys.stdout.encoding and sys.stdout.encoding.lower() != "utf-8":
|
||||
sys.stdout.reconfigure(encoding="utf-8", errors="replace")
|
||||
|
||||
try:
|
||||
import pymysql
|
||||
@@ -22,6 +14,7 @@ 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():
|
||||
@@ -31,59 +24,94 @@ def required_env(name: str) -> str:
|
||||
|
||||
|
||||
DB_CONFIG = {
|
||||
"host": os.getenv("SUPERBIZ_MYSQL_HOST", "119.29.78.52"),
|
||||
"port": int(os.getenv("SUPERBIZ_MYSQL_PORT", "33306")),
|
||||
"user": os.getenv("SUPERBIZ_MYSQL_USERNAME", "root"),
|
||||
"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": os.getenv("SUPERBIZ_MYSQL_DATABASE", "superbiz_agent"),
|
||||
"database": required_env("SUPERBIZ_MYSQL_DATABASE"),
|
||||
"charset": "utf8mb4",
|
||||
"cursorclass": pymysql.cursors.DictCursor,
|
||||
"read_timeout": 10,
|
||||
"write_timeout": 10,
|
||||
}
|
||||
|
||||
|
||||
def run_query(sql: str):
|
||||
conn = pymysql.connect(**DB_CONFIG)
|
||||
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 conn.cursor() as cur:
|
||||
cur.execute(sql)
|
||||
if sql.strip().upper().startswith("SELECT") or sql.strip().upper().startswith("SHOW"):
|
||||
rows = cur.fetchall()
|
||||
if not rows:
|
||||
print("(空结果)")
|
||||
return
|
||||
# 打印列头
|
||||
cols = list(rows[0].keys())
|
||||
col_widths = {c: max(len(c), max(len(str(r[c])) for r in rows)) for c in cols}
|
||||
header = " | ".join(c.ljust(col_widths[c]) for c in cols)
|
||||
print(header)
|
||||
print("-" * len(header))
|
||||
for row in rows:
|
||||
print(" | ".join(str(row[c]).ljust(col_widths[c]) for c in cols))
|
||||
print(f"\n({len(rows)} 行)")
|
||||
else:
|
||||
conn.commit()
|
||||
print(f"OK,影响行数: {cur.rowcount}")
|
||||
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:
|
||||
conn.close()
|
||||
connection.rollback()
|
||||
connection.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if len(sys.argv) > 1:
|
||||
sql = " ".join(sys.argv[1:])
|
||||
run_query(sql)
|
||||
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']}:{DB_CONFIG['port']}/{DB_CONFIG['database']}")
|
||||
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 not sql:
|
||||
continue
|
||||
run_query(sql)
|
||||
if sql:
|
||||
run_query(sql)
|
||||
except KeyboardInterrupt:
|
||||
break
|
||||
except Exception as e:
|
||||
print(f"错误: {e}")
|
||||
except Exception as exc:
|
||||
print(f"错误: {exc}")
|
||||
|
||||
Reference in New Issue
Block a user