From 50c7acb02ce8419ce4f20fe9b1b91e7672b61f60 Mon Sep 17 00:00:00 2001 From: shellway-pc <413209390@qq.com> Date: Sat, 1 Aug 2026 10:34:57 +0800 Subject: [PATCH] =?UTF-8?q?fix:=E4=BF=AE=E5=A4=8D=E4=BA=86=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E5=BA=93=E6=A8=A1=E5=9D=97=E7=9A=84=E4=BF=A1=E6=81=AF?= =?UTF-8?q?=E6=B3=84=E9=9C=B2=E9=A3=8E=E9=99=A9=E7=AD=89=E9=97=AE=E9=A2=98?= =?UTF-8?q?=E3=80=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- alpha/.env.example | 11 ++++ alpha/.gitignore | 1 + alpha/config.py | 33 +++++++++-- quantitative_data/config.py | 21 ++++++- quantitative_data/importer.py | 92 +++++++++++++++++++----------- quantitative_data/requirements.txt | 1 + 6 files changed, 121 insertions(+), 38 deletions(-) create mode 100644 alpha/.env.example create mode 100644 alpha/.gitignore diff --git a/alpha/.env.example b/alpha/.env.example new file mode 100644 index 0000000..d363cfe --- /dev/null +++ b/alpha/.env.example @@ -0,0 +1,11 @@ +# ======================================== +# 阿尔法研究 — 环境变量配置模板 +# 复制此文件为 .env 并填入真实值 +# ======================================== + +# PostgreSQL 数据库连接 +DB_HOST=localhost +DB_PORT=12345 +DB_NAME=quant_db +DB_USER=postgres +DB_PASSWORD=your_password_here \ No newline at end of file diff --git a/alpha/.gitignore b/alpha/.gitignore new file mode 100644 index 0000000..2eea525 --- /dev/null +++ b/alpha/.gitignore @@ -0,0 +1 @@ +.env \ No newline at end of file diff --git a/alpha/config.py b/alpha/config.py index 2c09779..c366b35 100644 --- a/alpha/config.py +++ b/alpha/config.py @@ -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 # 初始资金 diff --git a/quantitative_data/config.py b/quantitative_data/config.py index 745f6e8..8f10fca 100644 --- a/quantitative_data/config.py +++ b/quantitative_data/config.py @@ -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") # 数据起始日期 diff --git a/quantitative_data/importer.py b/quantitative_data/importer.py index 16dd189..c82f734 100644 --- a/quantitative_data/importer.py +++ b/quantitative_data/importer.py @@ -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, ): """ 一键全量导入: diff --git a/quantitative_data/requirements.txt b/quantitative_data/requirements.txt index 54d195d..a00d0b3 100644 --- a/quantitative_data/requirements.txt +++ b/quantitative_data/requirements.txt @@ -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