fix:修复了数据库模块的信息泄露风险等问题。
This commit is contained in:
@@ -0,0 +1,11 @@
|
||||
# ========================================
|
||||
# 阿尔法研究 — 环境变量配置模板
|
||||
# 复制此文件为 .env 并填入真实值
|
||||
# ========================================
|
||||
|
||||
# PostgreSQL 数据库连接
|
||||
DB_HOST=localhost
|
||||
DB_PORT=12345
|
||||
DB_NAME=quant_db
|
||||
DB_USER=postgres
|
||||
DB_PASSWORD=your_password_here
|
||||
@@ -0,0 +1 @@
|
||||
.env
|
||||
+28
-5
@@ -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 # 初始资金
|
||||
|
||||
@@ -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") # 数据起始日期
|
||||
|
||||
@@ -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,
|
||||
):
|
||||
"""
|
||||
一键全量导入:
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user