fix:修复了数据库模块的信息泄露风险等问题。

This commit is contained in:
2026-08-01 10:34:57 +08:00
parent f45d130a58
commit 50c7acb02c
6 changed files with 121 additions and 38 deletions
+11
View File
@@ -0,0 +1,11 @@
# ========================================
# 阿尔法研究 — 环境变量配置模板
# 复制此文件为 .env 并填入真实值
# ========================================
# PostgreSQL 数据库连接
DB_HOST=localhost
DB_PORT=12345
DB_NAME=quant_db
DB_USER=postgres
DB_PASSWORD=your_password_here
+1
View File
@@ -0,0 +1 @@
.env
+28 -5
View File
@@ -1,6 +1,19 @@
"""
阿尔法模块配置
敏感信息(数据库密码等)通过环境变量管理:
- DB_HOST : 数据库主机地址(默认 localhost)
- DB_PORT : 数据库端口(默认 12345)
- DB_NAME : 数据库名称(默认 quant_db)
- DB_USER : 数据库用户(默认 postgres)
- DB_PASSWORD : 数据库密码
设置方式:
Windows: set DB_PASSWORD=your_password
Linux: export DB_PASSWORD=your_password
也可创建 .env 文件(参考 .env.example
"""
import os
from dataclasses import dataclass, field
from typing import List, Optional
@@ -10,11 +23,21 @@ class AlphaConfig:
"""阿尔法研究全局配置"""
# ---- 数据库 ----
db_host: str = "192.168.27.15"
db_port: int = 12345
db_name: str = "quant_db"
db_user: str = "postgres"
db_password: str = "postgres"
db_host: str = field(
default_factory=lambda: os.getenv("DB_HOST", "localhost")
)
db_port: int = field(
default_factory=lambda: int(os.getenv("DB_PORT", "12345"))
)
db_name: str = field(
default_factory=lambda: os.getenv("DB_NAME", "quant_db")
)
db_user: str = field(
default_factory=lambda: os.getenv("DB_USER", "postgres")
)
db_password: str = field(
default_factory=lambda: os.getenv("DB_PASSWORD", "")
)
# ---- 回测基础参数 ----
initial_cash: float = 1_000_000.0 # 初始资金
+20 -1
View File
@@ -28,7 +28,7 @@ try:
env_path = cwd_path
if env_path.exists():
load_dotenv(dotenv_path=env_path, override=True)
load_dotenv(dotenv_path=env_path, override=False)
print(f"✓ 已加载环境变量文件: {env_path}")
_LOADED = True
else:
@@ -53,6 +53,25 @@ PASSWORD_ENCODED = quote_plus(_PASSWORD) if _PASSWORD else ""
# Tushare API Token
TUSHARE_TOKEN = os.environ.get("TUSHARE_TOKEN", "")
# ---- 启动期强校验:关键凭证未设置时直接报错,避免失败点后置 ----
if not _PASSWORD:
raise ValueError(
"QUANT_DB_PASSWORD 环境变量未设置。"
"请设置环境变量或创建 .env 文件(参考 .env.example)。\n"
" Windows: set QUANT_DB_PASSWORD=your_password\n"
" Linux: export QUANT_DB_PASSWORD=your_password"
)
_TUSHARE = os.environ.get("TUSHARE_TOKEN", "")
if not _TUSHARE:
raise ValueError(
"TUSHARE_TOKEN 环境变量未设置。"
"请设置环境变量或创建 .env 文件(参考 .env.example)。\n"
" Windows: set TUSHARE_TOKEN=your_token\n"
" Linux: export TUSHARE_TOKEN=your_token\n"
" 注册地址: https://tushare.pro"
)
# 批量导入参数
BATCH_SIZE = int(os.environ.get("QUANT_BATCH_SIZE", "5000")) # 每批次插入行数
START_DATE = os.environ.get("QUANT_START_DATE", "2010-01-01") # 数据起始日期
+60 -32
View File
@@ -110,11 +110,14 @@ def batch_insert(table_name: str, df: pd.DataFrame, conn, conflict_columns: List
logger.warning(f" {table_name}: 空数据,跳过")
return 0
# pd.NaT / pd.NaT / numpy NaN 等无法被 psycopg2 识别,统一替换为 Python None
df = df.where(pd.notna(df), None)
columns = list(df.columns)
rows = [tuple(row) for row in df.itertuples(index=False)]
# 构建 ON CONFLICT 子句
conflict_str = ", ".join(conflict_columns)
# 构建 ON CONFLICT 子句(使用 sql.Identifier 防止注入/语法错误)
conflict_identifiers = sql.SQL(", ").join(map(sql.Identifier, conflict_columns))
# 构建 UPDATE SET 子句 (排除冲突列)
update_cols = [c for c in columns if c not in conflict_columns]
if not update_cols:
@@ -125,7 +128,7 @@ def batch_insert(table_name: str, df: pd.DataFrame, conn, conflict_columns: List
).format(
table=sql.Identifier(table_name),
cols=sql.SQL(", ").join(map(sql.Identifier, columns)),
conflict=sql.SQL(conflict_str),
conflict=conflict_identifiers,
)
else:
update_set = sql.SQL(", ").join(
@@ -138,7 +141,7 @@ def batch_insert(table_name: str, df: pd.DataFrame, conn, conflict_columns: List
).format(
table=sql.Identifier(table_name),
cols=sql.SQL(", ").join(map(sql.Identifier, columns)),
conflict=sql.SQL(conflict_str),
conflict=conflict_identifiers,
update_set=update_set,
)
@@ -246,7 +249,7 @@ def import_stock_basic():
# 2. 导入交易日历
# ============================================================
def import_trade_cal(start_date: str = None, end_date: str = None):
def import_trade_cal(start_date: Optional[str] = None, end_date: Optional[str] = None):
"""
导入交易日历 (trade_cal)
Tushare: trade_cal
@@ -346,8 +349,8 @@ def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) -
def import_daily_batch(
stock_list: List[str],
start_date: str = None,
end_date: str = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
sleep_interval: float = 0.3,
):
"""
@@ -430,10 +433,10 @@ def import_daily_by_year(
# ============================================================
def import_daily_basic(
ts_code: str = None,
trade_date: str = None,
start_date: str = None,
end_date: str = None,
ts_code: Optional[str] = None,
trade_date: Optional[str] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
conn=None,
):
"""
@@ -488,8 +491,8 @@ def import_daily_basic(
def import_daily_basic_by_date(
start_date: str = None,
end_date: str = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
):
"""
按日期批量导入每日指标 (全市场)
@@ -528,7 +531,11 @@ def import_daily_basic_by_date(
for i, td in enumerate(trade_dates, 1):
try:
pro = get_ts_pro()
df = pro.daily_basic(trade_date=td)
def fetch_daily_basic():
return pro.daily_basic(trade_date=td)
df = fetch_with_retry(fetch_daily_basic, max_retries=3)
if df is not None and not df.empty:
df = normalize_columns(df)
if "trade_date" in df.columns:
@@ -552,9 +559,9 @@ def import_daily_basic_by_date(
# ============================================================
def import_adj_factor(
ts_code: str = None,
start_date: str = None,
end_date: str = None,
ts_code: Optional[str] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
conn=None,
):
"""
@@ -601,8 +608,8 @@ def import_adj_factor(
def import_adj_factor_batch(
stock_list: List[str],
start_date: str = None,
end_date: str = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
):
"""批量导入复权因子"""
if start_date is None:
@@ -636,8 +643,8 @@ def import_adj_factor_batch(
def import_financial_statements(
stock_list: List[str],
start_date: str = None,
end_date: str = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
):
"""
按股票批量导入三大报表 + 财务指标
@@ -703,7 +710,7 @@ def import_financial_statements(
batch_insert(table_name, df, conn, conflict_cols)
except Exception as e:
logger.debug(f" [{ts_code}] {table_name}: {e}")
logger.warning(f" [{ts_code}] {table_name}: {e}")
conn.rollback()
if i % 50 == 0 or i == total:
@@ -719,9 +726,9 @@ def import_financial_statements(
# ============================================================
def import_index_daily(
index_codes: List[str] = None,
start_date: str = None,
end_date: str = None,
index_codes: Optional[List[str]] = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
):
"""
导入指数日线行情 (index_daily)
@@ -824,20 +831,41 @@ def init_database():
conn = get_pg_connection()
try:
import os as _os
# 尝试使用 sqlparse 按语句拆分 (处理含分号的字符串字面量、函数体等)
try:
import sqlparse as _sqlparse
_USE_SQLPARSE = True
except ImportError:
_USE_SQLPARSE = False
logger.warning(" sqlparse 未安装,回退为简单分号拆分")
schema_path = _os.path.join(_os.path.dirname(_os.path.abspath(__file__)), "schema.sql")
with open(schema_path, "r", encoding="utf-8") as f:
ddl_sql = f.read()
# 按分号分割,逐条执行 (忽略被注释掉的分区表DDL)
statements = [s.strip() for s in ddl_sql.split(";") if s.strip()
and not s.strip().startswith("--")]
if _USE_SQLPARSE:
raw_statements = _sqlparse.split(ddl_sql)
statements = [
s.strip() for s in raw_statements
if s.strip() and not s.strip().startswith("--")
]
else:
statements = [
s.strip() for s in ddl_sql.split(";")
if s.strip() and not s.strip().startswith("--")
]
cursor = conn.cursor()
for stmt in statements:
if stmt and not stmt.startswith("--"):
try:
cursor.execute(stmt)
except Exception as e:
logger.debug(f" SQL 跳过: {str(e)[:100]}")
logger.error(f" DDL 执行失败: {str(e)[:200]}\n SQL: {stmt[:300]}")
conn.rollback()
cursor.close()
raise
conn.commit()
cursor.close()
@@ -903,10 +931,10 @@ def get_stock_codes_from_db(conn=None) -> List[str]:
# ============================================================
def full_import(
start_date: str = None,
end_date: str = None,
start_date: Optional[str] = None,
end_date: Optional[str] = None,
import_financials: bool = True,
stock_codes: List[str] = None,
stock_codes: Optional[List[str]] = None,
):
"""
一键全量导入:
+1
View File
@@ -3,3 +3,4 @@ pandas>=1.5.0
psycopg2-binary>=2.9.0
sqlalchemy>=2.0.0
python-dotenv>=1.0.0
sqlparse>=0.4.0