from __future__ import annotations
|
|
import json
|
import os
|
from dataclasses import dataclass
|
from pathlib import Path
|
from typing import Any, Mapping
|
|
import mysql.connector
|
|
|
TARGET_DATABASE = "stock_valuation"
|
MARKET_DATABASE = "trading_xuntou"
|
|
|
class DatabaseConfigError(RuntimeError):
|
pass
|
|
|
@dataclass(frozen=True)
|
class MySQLSettings:
|
host: str
|
port: int
|
user: str
|
password: str
|
database: str = TARGET_DATABASE
|
market_database: str = MARKET_DATABASE
|
|
def connection_args(self, database: str) -> dict[str, Any]:
|
return {
|
"host": self.host,
|
"port": self.port,
|
"user": self.user,
|
"password": self.password,
|
"database": database,
|
"charset": "utf8mb4",
|
"collation": "utf8mb4_0900_ai_ci",
|
"autocommit": False,
|
"connection_timeout": 10,
|
}
|
|
|
def _read_config(path: Path | None) -> Mapping[str, Any]:
|
if path is None:
|
raw = os.environ.get("STOCK_VALUATION_MYSQL_CONFIG")
|
path = Path(raw).expanduser() if raw else None
|
if path is None:
|
return {}
|
try:
|
value = json.loads(path.read_text(encoding="utf-8"))
|
except (OSError, UnicodeDecodeError, json.JSONDecodeError) as exc:
|
raise DatabaseConfigError("无法读取 MySQL 本机配置") from exc
|
if not isinstance(value, dict):
|
raise DatabaseConfigError("MySQL 本机配置根节点必须是对象")
|
return value
|
|
|
def load_settings(
|
*,
|
database: str | None = None,
|
config_path: Path | None = None,
|
) -> MySQLSettings:
|
local = _read_config(config_path)
|
|
def choose(env_name: str, key: str, default: Any = None) -> Any:
|
value = os.environ.get(env_name)
|
return value if value is not None else local.get(key, default)
|
|
selected_database = database or choose(
|
"STOCK_VALUATION_MYSQL_DATABASE", "database", TARGET_DATABASE
|
)
|
market_database = choose(
|
"STOCK_VALUATION_MARKET_DATABASE", "market_database", MARKET_DATABASE
|
)
|
allow_test = os.environ.get("STOCK_VALUATION_ALLOW_TEST_DATABASE") == "1"
|
valid_target = selected_database == TARGET_DATABASE or (
|
allow_test and str(selected_database).startswith("stock_valuation_test_")
|
)
|
if not valid_target:
|
raise DatabaseConfigError("目标库必须是 stock_valuation 或显式授权的隔离测试库")
|
if market_database != MARKET_DATABASE:
|
raise DatabaseConfigError("行情只读库必须是 trading_xuntou")
|
user = choose("STOCK_VALUATION_MYSQL_USER", "user")
|
password = choose("STOCK_VALUATION_MYSQL_PASSWORD", "password")
|
if not isinstance(user, str) or not user:
|
raise DatabaseConfigError("缺少 MySQL 用户配置")
|
if not isinstance(password, str):
|
raise DatabaseConfigError("缺少 MySQL 密码配置")
|
try:
|
port = int(choose("STOCK_VALUATION_MYSQL_PORT", "port", 3306))
|
except (TypeError, ValueError) as exc:
|
raise DatabaseConfigError("MySQL 端口配置无效") from exc
|
return MySQLSettings(
|
host=str(choose("STOCK_VALUATION_MYSQL_HOST", "host", "127.0.0.1")),
|
port=port,
|
user=user,
|
password=password,
|
database=str(selected_database),
|
market_database=str(market_database),
|
)
|
|
|
def connect_target(settings: MySQLSettings):
|
return mysql.connector.connect(**settings.connection_args(settings.database))
|
|
|
def connect_market_readonly(settings: MySQLSettings):
|
connection = mysql.connector.connect(
|
**settings.connection_args(settings.market_database)
|
)
|
connection.start_transaction(readonly=True)
|
return connection
|
|
|
def safe_mysql_error(exc: BaseException) -> str:
|
errno = getattr(exc, "errno", None)
|
sqlstate = getattr(exc, "sqlstate", None)
|
details = []
|
if errno is not None:
|
details.append(f"errno={errno}")
|
if sqlstate:
|
details.append(f"sqlstate={sqlstate}")
|
suffix = f"({', '.join(details)})" if details else ""
|
return "MySQL 操作失败" + suffix
|