tradingSystem/app/db/session.py

200 lines
7.9 KiB
Python
Raw Permalink Normal View History

# -*- coding: utf-8 -*-
"""
数据库连接与严格单表访问守卫
==================================
开发约定 (README): 153 代理侧数据库**严格单表访问** 代理不支持多表联查,
写错了要到线上才报错本模块在 SQL 执行入口做静态检查, 把这条纪律钉死在代码里
三个库 (键名对齐 bionic_trader, 同一 .env 可共用):
proxy PROXY_DB_URL 153 代理: pms_* 全部表 + trading_* + strategy_daily_results
factor SOURCE_DB_EXT_DSN 因子分表 (自算参考位/MA/ATR)
index DB_MYSQL_URL 大盘指数 zs_day_data (页面区制提示)
连接失败一律抛 DBUnavailable, 由上层 (API/任务) 捕获并降级 页面照常打开
调度任务照常守成不产生新指令 (设计 §13 三条铁律之故障即守成)
"""
from __future__ import annotations
import logging
import re
import time
import threading
from sqlalchemy import create_engine, text
from sqlalchemy.exc import SQLAlchemyError
from config.settings import settings
logger = logging.getLogger("pms.db")
_engines = {}
_lock = threading.Lock()
# 短路冷却: 库刚连挂过就直接快速失败, 不再逐次等 TCP 超时。
# 目的 —— 库不可达时页面仍能秒开 (每个接口只吃一次超时), 调度任务也不会被拖死。
FAIL_COOLDOWN_SEC = 10.0
_last_fail = {}
_DSN = {
"proxy": lambda: settings.PROXY_DB_URL,
"factor": lambda: settings.SOURCE_DB_EXT_DSN,
"index": lambda: settings.DB_MYSQL_URL,
}
class DBUnavailable(RuntimeError):
"""连库/执行失败。上层必须捕获并降级, 不得让异常冒泡成 500。"""
class MultiTableSQL(ValueError):
"""违反严格单表访问纪律。"""
# ---------------------------------------------------------------- 单表守卫
_TABLE_RE = re.compile(r"\b(?:from|join|into|update)\s+`?([a-zA-Z_][a-zA-Z0-9_]*)`?", re.I)
_JOIN_RE = re.compile(r"\bjoin\b", re.I)
# FROM 子句 (到下一个子句关键字为止) —— 用于识别 "FROM a, b" 这种逗号连表
_FROM_CLAUSE_RE = re.compile(
r"\bfrom\s+(.+?)(?=\bwhere\b|\bgroup\b|\border\b|\blimit\b|\bhaving\b|\bunion\b|\bon\b|\)|;|$)",
re.I | re.S)
2026-07-28 15:48:57 +08:00
# `ON DUPLICATE KEY UPDATE` 里的 UPDATE 后面跟的是**列名**不是表名。不先摘掉这一段,
# 下面的 _TABLE_RE 会把第一个被赋值的列当成第二张表, 于是所有 upsert 一律误判成多表。
# 2026-07-28 修: 四处 upsert 全部命中该误判 —— set_param (页面改参数)、ensure_position
# (建仓落账本第一步)、upsert_report (日报落库)、upsert_industry (行业映射导入),
# 一执行就抛 MultiTableSQL。单测走内存桩不经过本函数, 所以一直没暴露。
# 见 scripts/test_batch6_units.py 的「单表守卫」一组。
_ODKU_RE = re.compile(r"\bon\s+duplicate\s+key\s+update\b", re.I)
def assert_single_table(sql: str) -> str:
"""静态检查: 一条 SQL 只能碰一张表, 且不得出现 JOIN / 逗号连表 / 跨表子查询。"""
s = re.sub(r"--[^\n]*", " ", str(sql))
s = re.sub(r"/\*.*?\*/", " ", s, flags=re.S)
2026-07-28 15:48:57 +08:00
s = _ODKU_RE.sub(" ", s) # 见上: 摘掉 upsert 尾巴, 它不引入新表
if _JOIN_RE.search(s):
raise MultiTableSQL(f"严格单表访问: SQL 含 JOIN —— {s[:120]}")
for clause in _FROM_CLAUSE_RE.findall(s):
parts = [x.strip() for x in clause.split(",") if x.strip()]
if len(parts) > 1:
raise MultiTableSQL(f"严格单表访问: FROM 子句逗号连表 {parts} —— {s[:120]}")
tables = {m.lower() for m in _TABLE_RE.findall(s)}
if len(tables) > 1:
raise MultiTableSQL(f"严格单表访问: SQL 涉及多表 {sorted(tables)} —— {s[:120]}")
return sql
# ---------------------------------------------------------------- 引擎
def get_engine(name: str = "proxy"):
if name in _engines:
return _engines[name]
with _lock:
if name in _engines:
return _engines[name]
dsn_fn = _DSN.get(name)
if not dsn_fn:
raise DBUnavailable(f"未知数据源 {name}")
try:
eng = create_engine(dsn_fn(), pool_pre_ping=True, pool_size=5, max_overflow=5,
pool_recycle=3600, future=True,
# 连接超时兜底: 库不可达时快速失败, 不拖死页面与调度任务
connect_args={"connect_timeout": 3})
except Exception as e: # DSN 非法等
raise DBUnavailable(f"数据源 {name} 初始化失败: {e}") from e
_engines[name] = eng
return eng
def _check_cooldown(name: str):
ts = _last_fail.get(name)
if ts and (time.monotonic() - ts) < FAIL_COOLDOWN_SEC:
raise DBUnavailable(
f"数据源 {name} 刚刚连接失败, {FAIL_COOLDOWN_SEC:.0f} 秒内快速失败 "
f"(避免逐次等待超时); 恢复后自动重连")
def _mark(name: str, ok: bool):
if ok:
_last_fail.pop(name, None)
else:
_last_fail[name] = time.monotonic()
def ping(name: str = "proxy") -> dict:
try:
with get_engine(name).connect() as c:
c.execute(text("SELECT 1"))
_mark(name, True)
return {"ok": True, "source": name}
except Exception as e:
_mark(name, False)
return {"ok": False, "source": name, "error": f"{type(e).__name__}: {e}"}
# ---------------------------------------------------------------- 执行助手
def fetch_all(sql: str, params=None, source: str = "proxy") -> list:
assert_single_table(sql)
_check_cooldown(source)
try:
with get_engine(source).connect() as c:
rows = c.execute(text(sql), params or {}).mappings().fetchall()
_mark(source, True)
return [dict(r) for r in rows]
except MultiTableSQL:
raise
except SQLAlchemyError as e:
_mark(source, _is_sql_error(e))
raise DBUnavailable(f"查询失败({source}): {type(e).__name__}: {e}") from e
except Exception as e:
_mark(source, False)
raise DBUnavailable(f"查询异常({source}): {type(e).__name__}: {e}") from e
def _is_sql_error(e) -> bool:
"""区分「SQL 本身错」(库是通的, 不进冷却) 与「连不上库」(进冷却)。"""
msg = str(e).lower()
return not any(k in msg for k in ("can't connect", "connection refused", "timed out",
"lost connection", "gone away", "no route to host",
"name or service not known"))
def fetch_one(sql: str, params=None, source: str = "proxy"):
rows = fetch_all(sql, params, source)
return rows[0] if rows else None
def execute(sql: str, params=None, source: str = "proxy") -> int:
"""写操作 (INSERT/UPDATE/DELETE)。返回受影响行数。"""
assert_single_table(sql)
_check_cooldown(source)
try:
with get_engine(source).begin() as c:
r = c.execute(text(sql), params or {})
_mark(source, True)
return int(r.rowcount or 0)
except MultiTableSQL:
raise
except SQLAlchemyError as e:
_mark(source, _is_sql_error(e))
raise DBUnavailable(f"写入失败({source}): {type(e).__name__}: {e}") from e
except Exception as e:
_mark(source, False)
raise DBUnavailable(f"写入异常({source}): {type(e).__name__}: {e}") from e
def execute_many(sql: str, seq_params: list, source: str = "proxy") -> int:
assert_single_table(sql)
if not seq_params:
return 0
_check_cooldown(source)
try:
with get_engine(source).begin() as c:
r = c.execute(text(sql), seq_params)
_mark(source, True)
return int(r.rowcount or 0)
except SQLAlchemyError as e:
_mark(source, _is_sql_error(e))
raise DBUnavailable(f"批量写入失败({source}): {type(e).__name__}: {e}") from e
except Exception as e:
_mark(source, False)
raise DBUnavailable(f"批量写入异常({source}): {type(e).__name__}: {e}") from e