fix:修复了数据库模块的信息泄露风险等问题。
This commit is contained in:
@@ -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,
|
||||
):
|
||||
"""
|
||||
一键全量导入:
|
||||
|
||||
Reference in New Issue
Block a user