#!/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}")