Cai
2026-08-31 3fdf9063cb61ff10d9f54aeed6eaf0eca1777844
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
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