Compare commits
2
Commits
f45d130a58
...
7a31c29b44
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7a31c29b44 | ||
|
|
50c7acb02c |
@@ -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 dataclasses import dataclass, field
|
||||||
from typing import List, Optional
|
from typing import List, Optional
|
||||||
|
|
||||||
@@ -10,11 +23,21 @@ class AlphaConfig:
|
|||||||
"""阿尔法研究全局配置"""
|
"""阿尔法研究全局配置"""
|
||||||
|
|
||||||
# ---- 数据库 ----
|
# ---- 数据库 ----
|
||||||
db_host: str = "192.168.27.15"
|
db_host: str = field(
|
||||||
db_port: int = 12345
|
default_factory=lambda: os.getenv("DB_HOST", "localhost")
|
||||||
db_name: str = "quant_db"
|
)
|
||||||
db_user: str = "postgres"
|
db_port: int = field(
|
||||||
db_password: str = "postgres"
|
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 # 初始资金
|
initial_cash: float = 1_000_000.0 # 初始资金
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ try:
|
|||||||
env_path = cwd_path
|
env_path = cwd_path
|
||||||
|
|
||||||
if env_path.exists():
|
if env_path.exists():
|
||||||
load_dotenv(dotenv_path=env_path, override=True)
|
load_dotenv(dotenv_path=env_path, override=False)
|
||||||
print(f"✓ 已加载环境变量文件: {env_path}")
|
print(f"✓ 已加载环境变量文件: {env_path}")
|
||||||
_LOADED = True
|
_LOADED = True
|
||||||
else:
|
else:
|
||||||
@@ -53,6 +53,25 @@ PASSWORD_ENCODED = quote_plus(_PASSWORD) if _PASSWORD else ""
|
|||||||
# Tushare API Token
|
# Tushare API Token
|
||||||
TUSHARE_TOKEN = os.environ.get("TUSHARE_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")) # 每批次插入行数
|
BATCH_SIZE = int(os.environ.get("QUANT_BATCH_SIZE", "5000")) # 每批次插入行数
|
||||||
START_DATE = os.environ.get("QUANT_START_DATE", "2010-01-01") # 数据起始日期
|
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}: 空数据,跳过")
|
logger.warning(f" {table_name}: 空数据,跳过")
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
|
# pd.NaT / pd.NaT / numpy NaN 等无法被 psycopg2 识别,统一替换为 Python None
|
||||||
|
df = df.where(pd.notna(df), None)
|
||||||
|
|
||||||
columns = list(df.columns)
|
columns = list(df.columns)
|
||||||
rows = [tuple(row) for row in df.itertuples(index=False)]
|
rows = [tuple(row) for row in df.itertuples(index=False)]
|
||||||
|
|
||||||
# 构建 ON CONFLICT 子句
|
# 构建 ON CONFLICT 子句(使用 sql.Identifier 防止注入/语法错误)
|
||||||
conflict_str = ", ".join(conflict_columns)
|
conflict_identifiers = sql.SQL(", ").join(map(sql.Identifier, conflict_columns))
|
||||||
# 构建 UPDATE SET 子句 (排除冲突列)
|
# 构建 UPDATE SET 子句 (排除冲突列)
|
||||||
update_cols = [c for c in columns if c not in conflict_columns]
|
update_cols = [c for c in columns if c not in conflict_columns]
|
||||||
if not update_cols:
|
if not update_cols:
|
||||||
@@ -125,7 +128,7 @@ def batch_insert(table_name: str, df: pd.DataFrame, conn, conflict_columns: List
|
|||||||
).format(
|
).format(
|
||||||
table=sql.Identifier(table_name),
|
table=sql.Identifier(table_name),
|
||||||
cols=sql.SQL(", ").join(map(sql.Identifier, columns)),
|
cols=sql.SQL(", ").join(map(sql.Identifier, columns)),
|
||||||
conflict=sql.SQL(conflict_str),
|
conflict=conflict_identifiers,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
update_set = sql.SQL(", ").join(
|
update_set = sql.SQL(", ").join(
|
||||||
@@ -138,7 +141,7 @@ def batch_insert(table_name: str, df: pd.DataFrame, conn, conflict_columns: List
|
|||||||
).format(
|
).format(
|
||||||
table=sql.Identifier(table_name),
|
table=sql.Identifier(table_name),
|
||||||
cols=sql.SQL(", ").join(map(sql.Identifier, columns)),
|
cols=sql.SQL(", ").join(map(sql.Identifier, columns)),
|
||||||
conflict=sql.SQL(conflict_str),
|
conflict=conflict_identifiers,
|
||||||
update_set=update_set,
|
update_set=update_set,
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -246,7 +249,7 @@ def import_stock_basic():
|
|||||||
# 2. 导入交易日历
|
# 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)
|
导入交易日历 (trade_cal)
|
||||||
Tushare: 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(
|
def import_daily_batch(
|
||||||
stock_list: List[str],
|
stock_list: List[str],
|
||||||
start_date: str = None,
|
start_date: Optional[str] = None,
|
||||||
end_date: str = None,
|
end_date: Optional[str] = None,
|
||||||
sleep_interval: float = 0.3,
|
sleep_interval: float = 0.3,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -430,10 +433,10 @@ def import_daily_by_year(
|
|||||||
# ============================================================
|
# ============================================================
|
||||||
|
|
||||||
def import_daily_basic(
|
def import_daily_basic(
|
||||||
ts_code: str = None,
|
ts_code: Optional[str] = None,
|
||||||
trade_date: str = None,
|
trade_date: Optional[str] = None,
|
||||||
start_date: str = None,
|
start_date: Optional[str] = None,
|
||||||
end_date: str = None,
|
end_date: Optional[str] = None,
|
||||||
conn=None,
|
conn=None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -488,8 +491,8 @@ def import_daily_basic(
|
|||||||
|
|
||||||
|
|
||||||
def import_daily_basic_by_date(
|
def import_daily_basic_by_date(
|
||||||
start_date: str = None,
|
start_date: Optional[str] = None,
|
||||||
end_date: 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):
|
for i, td in enumerate(trade_dates, 1):
|
||||||
try:
|
try:
|
||||||
pro = get_ts_pro()
|
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:
|
if df is not None and not df.empty:
|
||||||
df = normalize_columns(df)
|
df = normalize_columns(df)
|
||||||
if "trade_date" in df.columns:
|
if "trade_date" in df.columns:
|
||||||
@@ -552,9 +559,9 @@ def import_daily_basic_by_date(
|
|||||||
# ============================================================
|
# ============================================================
|
||||||
|
|
||||||
def import_adj_factor(
|
def import_adj_factor(
|
||||||
ts_code: str = None,
|
ts_code: Optional[str] = None,
|
||||||
start_date: str = None,
|
start_date: Optional[str] = None,
|
||||||
end_date: str = None,
|
end_date: Optional[str] = None,
|
||||||
conn=None,
|
conn=None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
@@ -601,8 +608,8 @@ def import_adj_factor(
|
|||||||
|
|
||||||
def import_adj_factor_batch(
|
def import_adj_factor_batch(
|
||||||
stock_list: List[str],
|
stock_list: List[str],
|
||||||
start_date: str = None,
|
start_date: Optional[str] = None,
|
||||||
end_date: str = None,
|
end_date: Optional[str] = None,
|
||||||
):
|
):
|
||||||
"""批量导入复权因子"""
|
"""批量导入复权因子"""
|
||||||
if start_date is None:
|
if start_date is None:
|
||||||
@@ -636,8 +643,8 @@ def import_adj_factor_batch(
|
|||||||
|
|
||||||
def import_financial_statements(
|
def import_financial_statements(
|
||||||
stock_list: List[str],
|
stock_list: List[str],
|
||||||
start_date: str = None,
|
start_date: Optional[str] = None,
|
||||||
end_date: str = None,
|
end_date: Optional[str] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
按股票批量导入三大报表 + 财务指标
|
按股票批量导入三大报表 + 财务指标
|
||||||
@@ -703,7 +710,7 @@ def import_financial_statements(
|
|||||||
batch_insert(table_name, df, conn, conflict_cols)
|
batch_insert(table_name, df, conn, conflict_cols)
|
||||||
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug(f" [{ts_code}] {table_name}: {e}")
|
logger.warning(f" [{ts_code}] {table_name}: {e}")
|
||||||
conn.rollback()
|
conn.rollback()
|
||||||
|
|
||||||
if i % 50 == 0 or i == total:
|
if i % 50 == 0 or i == total:
|
||||||
@@ -719,9 +726,9 @@ def import_financial_statements(
|
|||||||
# ============================================================
|
# ============================================================
|
||||||
|
|
||||||
def import_index_daily(
|
def import_index_daily(
|
||||||
index_codes: List[str] = None,
|
index_codes: Optional[List[str]] = None,
|
||||||
start_date: str = None,
|
start_date: Optional[str] = None,
|
||||||
end_date: str = None,
|
end_date: Optional[str] = None,
|
||||||
):
|
):
|
||||||
"""
|
"""
|
||||||
导入指数日线行情 (index_daily)
|
导入指数日线行情 (index_daily)
|
||||||
@@ -824,22 +831,79 @@ def init_database():
|
|||||||
conn = get_pg_connection()
|
conn = get_pg_connection()
|
||||||
try:
|
try:
|
||||||
import os as _os
|
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")
|
schema_path = _os.path.join(_os.path.dirname(_os.path.abspath(__file__)), "schema.sql")
|
||||||
with open(schema_path, "r", encoding="utf-8") as f:
|
with open(schema_path, "r", encoding="utf-8") as f:
|
||||||
ddl_sql = f.read()
|
ddl_sql = f.read()
|
||||||
|
|
||||||
# 按分号分割,逐条执行 (忽略被注释掉的分区表DDL)
|
if _USE_SQLPARSE:
|
||||||
statements = [s.strip() for s in ddl_sql.split(";") if s.strip()
|
raw_statements = _sqlparse.split(ddl_sql)
|
||||||
and not s.strip().startswith("--")]
|
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()
|
cursor = conn.cursor()
|
||||||
for stmt in statements:
|
for stmt in statements:
|
||||||
if stmt and not stmt.startswith("--"):
|
if stmt and not stmt.startswith("--"):
|
||||||
try:
|
try:
|
||||||
cursor.execute(stmt)
|
cursor.execute(stmt)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.debug(f" SQL 跳过: {str(e)[:100]}")
|
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
|
||||||
|
|
||||||
conn.commit()
|
|
||||||
cursor.close()
|
cursor.close()
|
||||||
logger.info(" Schema 初始化完成")
|
logger.info(" Schema 初始化完成")
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -903,10 +967,10 @@ def get_stock_codes_from_db(conn=None) -> List[str]:
|
|||||||
# ============================================================
|
# ============================================================
|
||||||
|
|
||||||
def full_import(
|
def full_import(
|
||||||
start_date: str = None,
|
start_date: Optional[str] = None,
|
||||||
end_date: str = None,
|
end_date: Optional[str] = None,
|
||||||
import_financials: bool = True,
|
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
|
psycopg2-binary>=2.9.0
|
||||||
sqlalchemy>=2.0.0
|
sqlalchemy>=2.0.0
|
||||||
python-dotenv>=1.0.0
|
python-dotenv>=1.0.0
|
||||||
|
sqlparse>=0.4.0
|
||||||
|
|||||||
Reference in New Issue
Block a user