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
+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,
):
"""
一键全量导入: