""" 量化数据导入核心模块 连接 Docker PostgreSQL (192.168.27.11:5438) 从 Tushare 拉取数据并批量导入 """ import time import logging from datetime import datetime, timedelta from typing import Optional, List, Dict import tushare as ts import pandas as pd import psycopg2 from psycopg2 import sql from psycopg2.extras import execute_values from sqlalchemy import create_engine from config import DB_CONFIG, TUSHARE_TOKEN, BATCH_SIZE, START_DATE, END_DATE, PASSWORD_ENCODED # ============================================================ # 日志配置 # ============================================================ logging.basicConfig( level=logging.INFO, format="%(asctime)s [%(levelname)s] %(message)s", handlers=[ logging.FileHandler("import_data.log", encoding="utf-8"), logging.StreamHandler(), ], ) logger = logging.getLogger(__name__) # ============================================================ # 初始化连接 # ============================================================ # psycopg2 原生连接 (用于执行 DDL / 精细控制) def get_pg_connection(): """获取 psycopg2 原生连接""" return psycopg2.connect( host=DB_CONFIG["host"], port=DB_CONFIG["port"], database=DB_CONFIG["database"], user=DB_CONFIG["user"], password=DB_CONFIG["password"], ) # SQLAlchemy 引擎 (用于 DataFrame.to_sql) _sqlalchemy_engine = None def get_sqlalchemy_engine(): """获取 SQLAlchemy 引擎 (单例)""" global _sqlalchemy_engine if _sqlalchemy_engine is None: db_url = ( f"postgresql://{DB_CONFIG['user']}:{PASSWORD_ENCODED}" f"@{DB_CONFIG['host']}:{DB_CONFIG['port']}/{DB_CONFIG['database']}" ) _sqlalchemy_engine = create_engine(db_url, pool_size=5, max_overflow=10) return _sqlalchemy_engine # Tushare Pro API (单例) _ts_pro = None def get_ts_pro(): """获取 Tushare Pro API 实例 (单例)""" global _ts_pro if _ts_pro is None: ts.set_token(TUSHARE_TOKEN) _ts_pro = ts.pro_api() logger.info("Tushare Pro API 初始化完成") return _ts_pro # ============================================================ # 通用导入工具函数 # ============================================================ def normalize_columns(df: pd.DataFrame) -> pd.DataFrame: """ 标准化列名:Tushare 返回的列名可能有大小写差异,统一转小写 """ df.columns = [c.lower() for c in df.columns] return df def safe_float(val): """安全转换为浮点数,NaN -> None""" try: if pd.isna(val): return None return float(val) except (ValueError, TypeError): return None def batch_insert(table_name: str, df: pd.DataFrame, conn, conflict_columns: List[str]): """ 使用 execute_values 批量 UPSERT (INSERT ... ON CONFLICT) - table_name: 目标表名 - df: 待导入 DataFrame - conn: psycopg2 连接 - conflict_columns: 冲突列 (唯一约束列),冲突时更新其他列 """ if df.empty: 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 子句(使用 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: # 没有需要更新的列,使用 DO NOTHING upsert_sql = sql.SQL( "INSERT INTO {table} ({cols}) VALUES %s " "ON CONFLICT ({conflict}) DO NOTHING" ).format( table=sql.Identifier(table_name), cols=sql.SQL(", ").join(map(sql.Identifier, columns)), conflict=conflict_identifiers, ) else: update_set = sql.SQL(", ").join( sql.SQL("{col} = EXCLUDED.{col}").format(col=sql.Identifier(c)) for c in update_cols ) upsert_sql = sql.SQL( "INSERT INTO {table} ({cols}) VALUES %s " "ON CONFLICT ({conflict}) DO UPDATE SET {update_set}" ).format( table=sql.Identifier(table_name), cols=sql.SQL(", ").join(map(sql.Identifier, columns)), conflict=conflict_identifiers, update_set=update_set, ) cursor = conn.cursor() try: execute_values(cursor, upsert_sql.as_string(cursor), rows, page_size=BATCH_SIZE) conn.commit() logger.info(f" {table_name}: 成功导入 {len(rows)} 条记录") return len(rows) except Exception as e: conn.rollback() logger.error(f" {table_name}: 批量导入失败 - {e}") raise finally: cursor.close() def fetch_with_retry(fetch_func, max_retries: int = 3, delay: float = 2.0): """ 带重试的数据获取装饰器 - fetch_func: 数据获取函数 (返回 DataFrame) - max_retries: 最大重试次数 - delay: 重试间隔(秒) """ for attempt in range(max_retries): try: result = fetch_func() if result is not None and not result.empty: return result logger.warning(f" 第 {attempt+1} 次获取返回空数据,重试...") except Exception as e: logger.warning(f" 第 {attempt+1} 次获取失败: {e}") if attempt < max_retries - 1: time.sleep(delay * (attempt + 1)) # 递增延迟 return pd.DataFrame() # ============================================================ # 1. 导入股票基本信息 # ============================================================ def import_stock_basic(): """ 导入股票基本信息 (stock_basic) Tushare: stock_basic """ # 确保数据库表结构已初始化 init_database() logger.info("=" * 60) logger.info("[1/7] 导入股票基本信息 (stock_basic) ...") pro = get_ts_pro() conn = get_pg_connection() try: # 获取全量股票基本信息 df = pro.stock_basic( exchange="", list_status="L", fields="ts_code,symbol,name,area,industry,market,list_status,list_date,is_hs,act_name,act_ent_type", ) if df is None or df.empty: logger.warning("未获取到股票基本信息") return df = normalize_columns(df) # 转换日期格式 if "list_date" in df.columns: df["list_date"] = pd.to_datetime(df["list_date"], format="%Y%m%d", errors="coerce") logger.info(f" 获取到 {len(df)} 条股票基本信息") # 批量导入 conflict_cols = ["ts_code"] batch_insert("stock_basic", df, conn, conflict_cols) # 也尝试获取退市的股票 try: df_d = pro.stock_basic( exchange="", list_status="D", fields="ts_code,symbol,name,area,industry,market,list_status,list_date,is_hs", ) if df_d is not None and not df_d.empty: df_d = normalize_columns(df_d) if "list_date" in df_d.columns: df_d["list_date"] = pd.to_datetime( df_d["list_date"], format="%Y%m%d", errors="coerce" ) batch_insert("stock_basic", df_d, conn, conflict_cols) logger.info(f" 额外导入 {len(df_d)} 条退市股票信息") except Exception as e: logger.warning(f" 获取退市股票信息失败: {e}") except Exception as e: logger.error(f" 导入股票基本信息失败: {e}") raise finally: conn.close() # ============================================================ # 2. 导入交易日历 # ============================================================ def import_trade_cal(start_date: Optional[str] = None, end_date: Optional[str] = None): """ 导入交易日历 (trade_cal) Tushare: trade_cal """ if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE logger.info("=" * 60) logger.info(f"[2/7] 导入交易日历 (trade_cal): {start_date} ~ {end_date}") pro = get_ts_pro() conn = get_pg_connection() try: # 上交所 for exchange, ex_name in [("SSE", "上交所"), ("SZSE", "深交所")]: df = pro.trade_cal( exchange=exchange, start_date=start_date.replace("-", ""), end_date=end_date.replace("-", ""), ) if df is None or df.empty: logger.warning(f" {ex_name} 交易日历为空") continue df = normalize_columns(df) # 转换日期 for col in ["cal_date", "pretrade_date"]: if col in df.columns: df[col] = pd.to_datetime(df[col], format="%Y%m%d", errors="coerce") # 重命名 is_open (Tushare 返回 0/1 整数) if "is_open" in df.columns: df["is_open"] = df["is_open"].astype(int) conflict_cols = ["exchange", "cal_date"] batch_insert("trade_cal", df, conn, conflict_cols) logger.info(f" {ex_name}: {len(df)} 条交易日历") except Exception as e: logger.error(f" 导入交易日历失败: {e}") raise finally: conn.close() time.sleep(0.5) # API 频率限制 # ============================================================ # 3. 导入日线行情 (核心表) # ============================================================ def import_daily_for_stock(ts_code: str, start_date: str, end_date: str, conn) -> int: """ 导入单只股票的日线行情 返回导入的记录数 """ pro = get_ts_pro() def fetch(): return pro.daily( ts_code=ts_code, start_date=start_date.replace("-", ""), end_date=end_date.replace("-", ""), ) df = fetch_with_retry(fetch, max_retries=2) if df is None or df.empty: return 0 df = normalize_columns(df) # 转换日期 if "trade_date" in df.columns: df["trade_date"] = pd.to_datetime(df["trade_date"], format="%Y%m%d", errors="coerce") # 数值列处理 NaN numeric_cols = [ "open", "high", "low", "close", "pre_close", "change", "pct_chg", "vol", "amount", ] for col in numeric_cols: if col in df.columns: df[col] = pd.to_numeric(df[col], errors="coerce") # 额外列 (Tushare Pro 不同版本返回字段可能不同) for col in ["turnover_rate", "volume_ratio", "ma5", "ma10", "ma20", "ma_v_5", "ma_v_10", "ma_v_20"]: if col not in df.columns: df[col] = None conflict_cols = ["ts_code", "trade_date"] return batch_insert("daily", df, conn, conflict_cols) def import_daily_batch( stock_list: List[str], start_date: Optional[str] = None, end_date: Optional[str] = None, sleep_interval: float = 0.3, ): """ 批量导入多只股票的日线行情 - stock_list: 股票代码列表 - sleep_interval: API 调用间隔 (避免频率限制) """ if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE total = len(stock_list) logger.info("=" * 60) logger.info( f"[3/7] 导入日线行情 (daily): {start_date} ~ {end_date}, " f"共 {total} 只股票" ) conn = get_pg_connection() success_count = 0 fail_list = [] for i, ts_code in enumerate(stock_list, 1): try: n = import_daily_for_stock(ts_code, start_date, end_date, conn) if n > 0: success_count += 1 if i % 50 == 0 or i == total: logger.info(f" 进度: {i}/{total} 成功={success_count} 失败={len(fail_list)}") except Exception as e: logger.error(f" [{ts_code}] 导入失败: {e}") fail_list.append(ts_code) conn.rollback() time.sleep(sleep_interval) # API 频率控制 conn.close() logger.info(f" 日线行情导入完成: 成功 {success_count}/{total}") if fail_list: logger.warning(f" 失败列表({len(fail_list)}): {fail_list[:20]}...") return fail_list def import_daily_by_year( stock_list: List[str], start_year: int = 2010, end_year: int = 2025, ): """ 按年份逐批导入日线行情 (断点续传友好) 适合大数据量导入,每年每只股票可单独重试 """ logger.info("=" * 60) logger.info( f"[3/7] 按年导入日线行情: {start_year} ~ {end_year}, " f"共 {len(stock_list)} 只股票" ) total_imported = 0 for year in range(start_year, end_year + 1): year_start = f"{year}-01-01" year_end = f"{year}-12-31" logger.info(f"--- 导入 {year} 年日线行情 ---") fail_list = import_daily_batch( stock_list, start_date=year_start, end_date=year_end, sleep_interval=0.2, ) total_imported += 1 logger.info(f" {year} 年完成\n") logger.info(f" 所有年份日线行情导入完成!") # ============================================================ # 4. 导入每日指标 (daily_basic) - 估值/基本面 # ============================================================ def import_daily_basic( ts_code: Optional[str] = None, trade_date: Optional[str] = None, start_date: Optional[str] = None, end_date: Optional[str] = None, conn=None, ): """ 导入每日指标 (daily_basic) Tushare: daily_basic - 支持按单只股票导入 - 支持按日期范围全市场导入 """ if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE pro = get_ts_pro() own_conn = conn is None if own_conn: conn = get_pg_connection() try: # 构建参数 kwargs = { "start_date": start_date.replace("-", ""), "end_date": end_date.replace("-", ""), } if ts_code: kwargs["ts_code"] = ts_code if trade_date: kwargs["trade_date"] = trade_date.replace("-", "") def fetch(): return pro.daily_basic(**kwargs) df = fetch_with_retry(fetch, max_retries=2) if df is None or df.empty: return 0 df = normalize_columns(df) if "trade_date" in df.columns: df["trade_date"] = pd.to_datetime(df["trade_date"], format="%Y%m%d", errors="coerce") numeric_cols = df.select_dtypes(include=["number"]).columns.tolist() for col in numeric_cols: df[col] = pd.to_numeric(df[col], errors="coerce") conflict_cols = ["ts_code", "trade_date"] return batch_insert("daily_basic", df, conn, conflict_cols) finally: if own_conn: conn.close() def import_daily_basic_by_date( start_date: Optional[str] = None, end_date: Optional[str] = None, ): """ 按日期批量导入每日指标 (全市场) Tushare daily_basic 接口可按交易日获取全市场数据,比较高效 """ if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE logger.info("=" * 60) logger.info(f"[4/7] 导入每日指标 (daily_basic): {start_date} ~ {end_date}") # 获取交易日列表 conn = get_pg_connection() try: cursor = conn.cursor() cursor.execute( """ SELECT DISTINCT cal_date FROM trade_cal WHERE is_open = 1 AND cal_date >= %s AND cal_date <= %s ORDER BY cal_date """, (start_date, end_date), ) trade_dates = [row[0].strftime("%Y%m%d") for row in cursor.fetchall()] cursor.close() finally: conn.close() total = len(trade_dates) logger.info(f" 共 {total} 个交易日") conn = get_pg_connection() for i, td in enumerate(trade_dates, 1): try: pro = get_ts_pro() 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: df["trade_date"] = pd.to_datetime(df["trade_date"], format="%Y%m%d", errors="coerce") conflict_cols = ["ts_code", "trade_date"] batch_insert("daily_basic", df, conn, conflict_cols) except Exception as e: logger.warning(f" [{td}] 导入失败: {e}") conn.rollback() if i % 20 == 0 or i == total: logger.info(f" 进度: {i}/{total}") time.sleep(0.3) conn.close() logger.info(" 每日指标导入完成") # ============================================================ # 5. 导入复权因子 # ============================================================ def import_adj_factor( ts_code: Optional[str] = None, start_date: Optional[str] = None, end_date: Optional[str] = None, conn=None, ): """ 导入复权因子 (adj_factor) Tushare: adj_factor """ if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE pro = get_ts_pro() own_conn = conn is None if own_conn: conn = get_pg_connection() try: kwargs = { "start_date": start_date.replace("-", ""), "end_date": end_date.replace("-", ""), } if ts_code: kwargs["ts_code"] = ts_code def fetch(): return pro.adj_factor(**kwargs) df = fetch_with_retry(fetch, max_retries=2) if df is None or df.empty: return 0 df = normalize_columns(df) if "trade_date" in df.columns: df["trade_date"] = pd.to_datetime(df["trade_date"], format="%Y%m%d", errors="coerce") conflict_cols = ["ts_code", "trade_date"] return batch_insert("adj_factor", df, conn, conflict_cols) finally: if own_conn: conn.close() def import_adj_factor_batch( stock_list: List[str], start_date: Optional[str] = None, end_date: Optional[str] = None, ): """批量导入复权因子""" if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE total = len(stock_list) logger.info("=" * 60) logger.info(f"[5/7] 导入复权因子 (adj_factor): {total} 只股票") conn = get_pg_connection() for i, ts_code in enumerate(stock_list, 1): try: import_adj_factor(ts_code=ts_code, start_date=start_date, end_date=end_date, conn=conn) except Exception as e: logger.warning(f" [{ts_code}] 复权因子导入失败: {e}") conn.rollback() if i % 100 == 0 or i == total: logger.info(f" 进度: {i}/{total}") time.sleep(0.25) conn.close() logger.info(" 复权因子导入完成") # ============================================================ # 6. 导入财务数据 (利润表、资产负债表、现金流量表、财务指标) # ============================================================ def import_financial_statements( stock_list: List[str], start_date: Optional[str] = None, end_date: Optional[str] = None, ): """ 按股票批量导入三大报表 + 财务指标 - income: 利润表 - balancesheet: 资产负债表 - cashflow: 现金流量表 - fina_indicator: 财务指标 """ if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE total = len(stock_list) logger.info("=" * 60) logger.info( f"[6/7] 导入财务数据: {start_date} ~ {end_date}, " f"共 {total} 只股票" ) period_start = start_date.replace("-", "") period_end = end_date.replace("-", "") pro = get_ts_pro() conn = get_pg_connection() for i, ts_code in enumerate(stock_list, 1): for table_name, fetch_method in [ ("income", pro.income), ("balancesheet", pro.balancesheet), ("cashflow", pro.cashflow), ("fina_indicator", pro.fina_indicator), ]: try: if table_name == "fina_indicator": # fina_indicator 参数略有不同 df = fetch_method( ts_code=ts_code, start_date=period_start, end_date=period_end, ) else: df = fetch_method( ts_code=ts_code, start_date=period_start, end_date=period_end, ) if df is None or df.empty: continue df = normalize_columns(df) # 转换日期列 for col in ["ann_date", "f_ann_date", "end_date"]: if col in df.columns: df[col] = pd.to_datetime(df[col], format="%Y%m%d", errors="coerce") if table_name == "fina_indicator": conflict_cols = ["ts_code", "end_date"] else: conflict_cols = ["ts_code", "end_date", "report_type"] batch_insert(table_name, df, conn, conflict_cols) except Exception as e: logger.warning(f" [{ts_code}] {table_name}: {e}") conn.rollback() if i % 50 == 0 or i == total: logger.info(f" 财务数据进度: {i}/{total}") time.sleep(0.3) conn.close() logger.info(" 财务数据导入完成") # ============================================================ # 7. 导入指数日线行情 # ============================================================ def import_index_daily( index_codes: Optional[List[str]] = None, start_date: Optional[str] = None, end_date: Optional[str] = None, ): """ 导入指数日线行情 (index_daily) 默认导入主要指数:上证指数、深证成指、沪深300、中证500、创业板指、科创50 """ if index_codes is None: index_codes = [ "000001.SH", # 上证指数 "399001.SZ", # 深证成指 "000300.SH", # 沪深300 "000905.SH", # 中证500 "399006.SZ", # 创业板指 "000688.SH", # 科创50 ] if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE logger.info("=" * 60) logger.info(f"[7/7] 导入指数日线行情: {len(index_codes)} 个指数") pro = get_ts_pro() conn = get_pg_connection() for idx_code in index_codes: try: def fetch(): return pro.index_daily( ts_code=idx_code, start_date=start_date.replace("-", ""), end_date=end_date.replace("-", ""), ) df = fetch_with_retry(fetch, max_retries=2) if df is not None and not df.empty: df = normalize_columns(df) if "trade_date" in df.columns: df["trade_date"] = pd.to_datetime( df["trade_date"], format="%Y%m%d", errors="coerce" ) conflict_cols = ["ts_code", "trade_date"] n = batch_insert("index_daily", df, conn, conflict_cols) logger.info(f" {idx_code}: {n} 条") else: logger.warning(f" {idx_code}: 无数据") except Exception as e: logger.error(f" {idx_code}: {e}") conn.rollback() time.sleep(0.3) conn.close() logger.info(" 指数日线行情导入完成") # ============================================================ # 8. 初始化数据库 Schema # ============================================================ def init_database(): """ 执行 DDL,创建所有表结构 """ logger.info("=" * 60) logger.info("初始化数据库 Schema ...") # 尝试创建数据库 (如果不存在) try: admin_conn = psycopg2.connect( host=DB_CONFIG["host"], port=DB_CONFIG["port"], database="postgres", # 连接默认 postgres 库来创建新库 user=DB_CONFIG["user"], password=DB_CONFIG["password"], ) admin_conn.autocommit = True cursor = admin_conn.cursor() # 检查数据库是否存在 cursor.execute( "SELECT 1 FROM pg_database WHERE datname = %s", (DB_CONFIG["database"],), ) if cursor.fetchone() is None: cursor.execute( sql.SQL("CREATE DATABASE {} ENCODING 'UTF8'").format( sql.Identifier(DB_CONFIG["database"]) ) ) logger.info(f" 数据库 {DB_CONFIG['database']} 创建成功") else: logger.info(f" 数据库 {DB_CONFIG['database']} 已存在") cursor.close() admin_conn.close() except Exception as e: logger.warning(f" 创建数据库步骤跳过 (可能无权限): {e}") # 执行 DDL 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() 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("--") ] # 按依赖关系排序:CREATE TABLE → CREATE INDEX → COMMENT ON → CREATE VIEW # 避免 sqlparse 拆分后语句乱序导致 UndefinedTable 错误 _priority = { "CREATE TABLE": 1, "CREATE INDEX": 2, "COMMENT ON": 3, "CREATE OR REPLACE VIEW": 4, "CREATE VIEW": 4, } def _stmt_priority(s): upper = s.upper() for keyword, prio in _priority.items(): if upper.startswith(keyword): return prio return 99 # 兜底:最后执行 statements.sort(key=_stmt_priority) # 设置 autocommit 模式:每条 DDL 独立事务,互不影响 # 否则 rollback 会撤销之前已成功执行的 CREATE TABLE conn.autocommit = True cursor = conn.cursor() for stmt in statements: if stmt and not stmt.startswith("--"): try: cursor.execute(stmt) except Exception as e: stmt_upper = stmt.upper() # 以下类型的语句失败视为可恢复的 warning,不中断整个初始化流程: # 1. 视图创建(依赖的基础表可能尚未创建) # 2. 索引创建(依赖的表可能尚未导入数据) # 3. COMMENT ON(依赖的表/视图可能尚未创建) is_recoverable = ( "CREATE OR REPLACE VIEW" in stmt_upper or "CREATE VIEW" in stmt_upper or "CREATE INDEX" in stmt_upper or "COMMENT ON" in stmt_upper ) if is_recoverable: logger.warning(f" DDL 暂跳过(依赖尚未就绪): {str(e)[:150]}\n SQL: {stmt[:200]}") else: logger.error(f" DDL 执行失败: {str(e)[:200]}\n SQL: {stmt[:300]}") cursor.close() conn.close() raise cursor.close() logger.info(" Schema 初始化完成") except Exception as e: conn.rollback() logger.error(f" Schema 初始化失败: {e}") raise finally: conn.close() # ============================================================ # 9. 获取所有股票列表 (辅助) # ============================================================ def get_all_stock_codes(include_delisted: bool = False) -> List[str]: """ 从 Tushare 获取所有 A 股股票代码列表 """ pro = get_ts_pro() codes = [] for status, label in [("L", "上市"), ("D", "退市"), ("P", "暂停")]: if status != "L" and not include_delisted: continue try: df = pro.stock_basic( exchange="", list_status=status, fields="ts_code", ) if df is not None and not df.empty: codes.extend(df["ts_code"].tolist()) except Exception as e: logger.warning(f" 获取 {label} 股票列表失败: {e}") logger.info(f" 获取到 {len(codes)} 只股票代码") return codes def get_stock_codes_from_db(conn=None) -> List[str]: """ 从已导入的 stock_basic 表获取股票代码列表 """ own_conn = conn is None if own_conn: conn = get_pg_connection() try: cursor = conn.cursor() cursor.execute("SELECT ts_code FROM stock_basic WHERE list_status = 'L' ORDER BY ts_code") codes = [row[0] for row in cursor.fetchall()] cursor.close() return codes finally: if own_conn: conn.close() # ============================================================ # 10. 一键全量导入 # ============================================================ def full_import( start_date: Optional[str] = None, end_date: Optional[str] = None, import_financials: bool = True, stock_codes: Optional[List[str]] = None, ): """ 一键全量导入: 1. 初始化 Schema 2. 股票基本信息 3. 交易日历 4. 日线行情 5. 每日指标(估值) 6. 复权因子 7. 财务数据 (可选) 8. 指数日线行情 参数: - start_date, end_date: 数据范围 - import_financials: 是否导入财务数据 (耗时较长) - stock_codes: 指定股票列表,不传则全量导入 """ if start_date is None: start_date = START_DATE if end_date is None: end_date = END_DATE start_time = datetime.now() logger.info("=" * 70) logger.info(f" 开始全量数据导入: {start_date} ~ {end_date}") logger.info(f" PostgreSQL: {DB_CONFIG['host']}:{DB_CONFIG['port']}/{DB_CONFIG['database']}") logger.info("=" * 70) # Step 0: 初始化 Schema init_database() # Step 1: 股票基本信息 import_stock_basic() # Step 2: 交易日历 import_trade_cal(start_date, end_date) # 获取股票列表 if stock_codes is None: stock_codes = get_stock_codes_from_db() if not stock_codes: logger.error("无法获取股票列表,请先导入 stock_basic") return # Step 3: 日线行情 (按年导入) import_daily_by_year( stock_codes, start_year=int(start_date[:4]), end_year=int(end_date[:4]), ) # Step 4: 每日指标 (按日期导入) import_daily_basic_by_date(start_date, end_date) # Step 5: 复权因子 import_adj_factor_batch(stock_codes, start_date, end_date) # Step 6: 财务数据 if import_financials: import_financial_statements(stock_codes, start_date, end_date) # Step 7: 指数日线行情 import_index_daily(start_date=start_date, end_date=end_date) elapsed = datetime.now() - start_time logger.info("=" * 70) logger.info(f" 全量数据导入完成! 总耗时: {elapsed}") logger.info("=" * 70) if __name__ == "__main__": # 测试连接 try: conn = get_pg_connection() logger.info(f"成功连接到 PostgreSQL: {DB_CONFIG['host']}:{DB_CONFIG['port']}") conn.close() except Exception as e: logger.error(f"无法连接到 PostgreSQL: {e}") logger.error("请确认 Docker 容器已启动,且 config.py 中的连接参数正确")