SHA256
init
This commit is contained in:
@@ -39,6 +39,11 @@ def bs_query(query_fn, *args, **kwargs):
|
||||
row = rs.get_row_data()
|
||||
"""
|
||||
bs_login()
|
||||
short_name = query_fn.__name__.replace("query_", "")
|
||||
sig_parts = [repr(a) for a in args]
|
||||
sig_parts += [f"{k}={v!r}" for k, v in kwargs.items()]
|
||||
sig = ", ".join(sig_parts)
|
||||
print(f" [BS] {short_name}({sig})", flush=True)
|
||||
with _lock:
|
||||
rs = query_fn(*args, **kwargs)
|
||||
yield rs
|
||||
|
||||
@@ -14,7 +14,7 @@
|
||||
|
||||
from sqlalchemy import (
|
||||
Column, String, Date, DateTime, Float, Integer, Text,
|
||||
UniqueConstraint, Index, create_engine, MetaData, func, text,
|
||||
UniqueConstraint, Index, create_engine, MetaData, func, select, text,
|
||||
)
|
||||
from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker
|
||||
from sqlalchemy.dialects.mysql import insert as mysql_insert
|
||||
@@ -97,6 +97,26 @@ class FinancialCashflow(Base):
|
||||
data = Column(Text, comment="JSON格式现金流数据")
|
||||
|
||||
|
||||
# ── 指数日线行情 ──
|
||||
class IndexDaily(Base):
|
||||
__tablename__ = "index_daily"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("code", "date", name="uq_index_code_date"),
|
||||
Index("ix_index_date", "date"),
|
||||
)
|
||||
|
||||
id = Column(Integer, primary_key=True, autoincrement=True)
|
||||
code = Column(String(10), nullable=False, comment="指数代码")
|
||||
date = Column(Date, nullable=False, comment="交易日期")
|
||||
open = Column(Float, comment="开盘价")
|
||||
high = Column(Float, comment="最高价")
|
||||
low = Column(Float, comment="最低价")
|
||||
close = Column(Float, comment="收盘价")
|
||||
volume = Column(Float, comment="成交量")
|
||||
amount = Column(Float, comment="成交额")
|
||||
pct_change = Column(Float, comment="涨跌幅%")
|
||||
|
||||
|
||||
# ── 分红送转 ──
|
||||
class StockDividend(Base):
|
||||
__tablename__ = "stock_dividend"
|
||||
@@ -245,6 +265,20 @@ class StockSector(Base):
|
||||
# ── 数据库连接管理 ──
|
||||
_engine = None
|
||||
_SessionFactory = None
|
||||
_stock_codes_cache: list[str] | None = None
|
||||
|
||||
|
||||
def get_stock_codes() -> list[str]:
|
||||
"""获取全部股票代码(进程内缓存,避免重复查询)"""
|
||||
global _stock_codes_cache
|
||||
if _stock_codes_cache is None:
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(select(StockInfo.code))
|
||||
_stock_codes_cache = [row[0] for row in result]
|
||||
finally:
|
||||
session.close()
|
||||
return _stock_codes_cache
|
||||
|
||||
|
||||
def get_engine():
|
||||
|
||||
+113
-37
@@ -1,25 +1,22 @@
|
||||
"""日线行情抓取模块 — 使用 BaoStock
|
||||
|
||||
单数据源架构,代码大幅简化。
|
||||
跳过策略(停牌天按天记录):
|
||||
跳过策略:
|
||||
- 数据完整 = 行情记录数 + 已标记停牌天数 >= 交易日总数
|
||||
- 未上市股票(ipo_date > 查询结束日期)
|
||||
- 抓取成功后自动识别缺失交易日并标记为停牌
|
||||
- 增量抓取:只请求缺失的日期段,不重复抓已有数据
|
||||
"""
|
||||
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
import baostock as bs
|
||||
from src.baostock_conn import bs_query, code_to_bs, bs_login
|
||||
from src.config import get_fetch_config
|
||||
from src.db import StockInfo, StockDaily, StockNoData, batch_upsert, get_session
|
||||
from src.db import StockInfo, StockDaily, StockNoData, batch_upsert, get_session, get_stock_codes
|
||||
from sqlalchemy import select, func
|
||||
|
||||
|
||||
def _clean(val):
|
||||
"""将空字符串、无效值转为 None"""
|
||||
if val is None:
|
||||
return None
|
||||
if isinstance(val, str) and val.strip() == "":
|
||||
@@ -27,17 +24,7 @@ def _clean(val):
|
||||
return val
|
||||
|
||||
|
||||
def _get_stock_codes() -> list[str]:
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(select(StockInfo.code))
|
||||
return [row[0] for row in result]
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _get_not_listed(end_date: str) -> set[str]:
|
||||
"""查询在end_date之后上市的股票(未上市,需跳过)"""
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(
|
||||
@@ -80,12 +67,85 @@ def _get_complete_codes(start_date: str, end_date: str, trading_days: list[str])
|
||||
session.close()
|
||||
|
||||
|
||||
def _record_nodata_days(code: str, days: list[str]):
|
||||
"""记录个股的停牌/无数据日期(按天粒度)"""
|
||||
if not days:
|
||||
def _get_gaps(codes: list[str], start_date: str, end_date: str,
|
||||
trading_days: list[str]) -> dict[str, list[str]]:
|
||||
"""查询每只股票缺失的日期段,返回 {code: [gap_start, gap_end]}"""
|
||||
sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"
|
||||
ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
|
||||
td_set = set(trading_days)
|
||||
session = get_session()
|
||||
try:
|
||||
# 已有行情日期
|
||||
result = session.execute(
|
||||
select(StockDaily.code, StockDaily.date)
|
||||
.where(StockDaily.code.in_(codes))
|
||||
.where(StockDaily.date >= sd)
|
||||
.where(StockDaily.date <= ed)
|
||||
)
|
||||
code_dates: dict[str, set[str]] = {}
|
||||
for code, date in result:
|
||||
code_dates.setdefault(code, set()).add(str(date))
|
||||
|
||||
# 已标记停牌日期
|
||||
nodata_result = session.execute(
|
||||
select(StockNoData.code, StockNoData.date)
|
||||
.where(StockNoData.code.in_(codes))
|
||||
.where(StockNoData.date >= sd)
|
||||
.where(StockNoData.date <= ed)
|
||||
)
|
||||
for code, date in nodata_result:
|
||||
code_dates.setdefault(code, set()).add(str(date))
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
gaps: dict[str, list[str]] = {}
|
||||
for code in codes:
|
||||
existing = code_dates.get(code, set())
|
||||
missing = [d for d in trading_days if d not in existing]
|
||||
if not missing:
|
||||
continue
|
||||
gaps[code] = [missing[0], missing[-1]]
|
||||
return gaps
|
||||
|
||||
|
||||
def _mark_suspensions_local(codes: list[str], start_date: str, end_date: str,
|
||||
trading_days: list[str]):
|
||||
"""纯本地数据库批量标记停牌天:比较 stock_daily 与 trading_day"""
|
||||
if not codes or not trading_days:
|
||||
return
|
||||
rows = [{"code": code, "date": d} for d in days]
|
||||
batch_upsert(StockNoData, rows, ["code", "date"])
|
||||
sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"
|
||||
ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(
|
||||
select(StockDaily.code, StockDaily.date)
|
||||
.where(StockDaily.code.in_(codes))
|
||||
.where(StockDaily.date >= sd)
|
||||
.where(StockDaily.date <= ed)
|
||||
)
|
||||
code_dates: dict[str, set[str]] = {}
|
||||
for code, date in result:
|
||||
code_dates.setdefault(code, set()).add(str(date))
|
||||
|
||||
nodata_result = session.execute(
|
||||
select(StockNoData.code, StockNoData.date)
|
||||
.where(StockNoData.code.in_(codes))
|
||||
.where(StockNoData.date >= sd)
|
||||
.where(StockNoData.date <= ed)
|
||||
)
|
||||
nodata_existing = {(code, str(date)) for code, date in nodata_result}
|
||||
|
||||
rows = []
|
||||
for code in codes:
|
||||
existing = code_dates.get(code, set())
|
||||
for td in trading_days:
|
||||
if td not in existing and (code, td) not in nodata_existing:
|
||||
rows.append({"code": code, "date": td})
|
||||
if rows:
|
||||
batch_upsert(StockNoData, rows, ["code", "date"])
|
||||
print(f" 本地批量标记停牌: {len(rows)} 条", flush=True)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _fetch_baostock(code: str, start_date: str, end_date: str) -> list[dict] | None:
|
||||
@@ -136,7 +196,7 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None):
|
||||
cfg = get_fetch_config()
|
||||
delay = cfg.get("delay", 0.1)
|
||||
|
||||
codes = _get_stock_codes()
|
||||
codes = get_stock_codes()
|
||||
if not codes:
|
||||
print(" 无股票列表,请先运行 --stock-info", flush=True)
|
||||
return
|
||||
@@ -167,39 +227,55 @@ def fetch_daily(start_date: str | None = None, end_date: str | None = None):
|
||||
print(" 所有股票数据已完整,无需抓取", flush=True)
|
||||
return
|
||||
|
||||
# 查询每只股票实际缺失的日期范围
|
||||
print(f" 正在分析 {total} 只股票的数据缺口...", flush=True)
|
||||
gaps = _get_gaps(codes, start_date, end_date, trading_days)
|
||||
|
||||
print(f"正在抓取日线行情 {start_date} ~ {end_date},{len(gaps)} 只需更新...", flush=True)
|
||||
|
||||
success = 0
|
||||
fail = 0
|
||||
nodata_count = 0
|
||||
skipped = len(skip_set)
|
||||
fetched_codes = []
|
||||
t_start = time.time()
|
||||
|
||||
print(f"正在抓取日线行情 {start_date} ~ {end_date},需抓取 {total} 只(跳过 {skipped} 只)...", flush=True)
|
||||
|
||||
for i, code in enumerate(codes):
|
||||
rows = _fetch_baostock(code, start_date, end_date)
|
||||
g = gaps.get(code)
|
||||
if not g:
|
||||
continue
|
||||
# 增量:只请求缺失的日期段
|
||||
gap_start = g[0].replace("-", "")
|
||||
gap_end = g[-1].replace("-", "")
|
||||
|
||||
t0 = time.time()
|
||||
rows = _fetch_baostock(code, gap_start, gap_end)
|
||||
t_fetch = time.time() - t0
|
||||
t1 = time.time()
|
||||
if rows is not None:
|
||||
try:
|
||||
batch_upsert(StockDaily, rows, ["code", "date"])
|
||||
success += 1
|
||||
returned_dates = {row["date"] for row in rows}
|
||||
missing_days = [d for d in trading_days if d not in returned_dates]
|
||||
if missing_days:
|
||||
_record_nodata_days(code, missing_days)
|
||||
except Exception as e:
|
||||
print(f" {code} 写入失败: {e}", flush=True)
|
||||
fail += 1
|
||||
continue
|
||||
else:
|
||||
# 无数据:标记所有交易日为停牌
|
||||
_record_nodata_days(code, trading_days)
|
||||
nodata_count += 1
|
||||
fetched_codes.append(code)
|
||||
t_write = time.time() - t1
|
||||
|
||||
total_elapsed = time.time() - t_start
|
||||
avg = total_elapsed / (i + 1)
|
||||
elapsed = time.time() - t_start
|
||||
avg = elapsed / (i + 1)
|
||||
eta = avg * (total - i - 1)
|
||||
print(f" [{i+1}/{total}] {code} 成功:{success} 失败:{fail} 停牌:{nodata_count} "
|
||||
f"已用时:{total_elapsed:.0f}s 预计剩余:{eta:.0f}s", flush=True)
|
||||
|
||||
print(f" [{i+1}/{total}] {code} 缺口:{g[0]}~{g[-1]} "
|
||||
f"网络:{t_fetch:.1f}s 写入:{t_write:.1f}s 行数:{len(rows) if rows else 0} "
|
||||
f"成功:{success} 剩余:{eta:.0f}s", flush=True)
|
||||
time.sleep(delay)
|
||||
|
||||
# 批量本地标记停牌
|
||||
if fetched_codes:
|
||||
print(f" 正在本地批量标记停牌天...", flush=True)
|
||||
_mark_suspensions_local(fetched_codes, start_date, end_date, trading_days)
|
||||
|
||||
total_time = time.time() - t_start
|
||||
print(f" 日线行情抓取完成,成功:{success} 失败:{fail} 停牌:{nodata_count} 总耗时:{total_time:.1f}s", flush=True)
|
||||
|
||||
@@ -2,20 +2,30 @@
|
||||
|
||||
BaoStock query_dividend_data() 按年度查询分红记录。
|
||||
迭代最近10年获取完整分红历史。
|
||||
跳过策略:已有当前年份分红记录的股票不再重复抓取。
|
||||
"""
|
||||
|
||||
from datetime import datetime
|
||||
import baostock as bs
|
||||
from src.baostock_conn import bs_query, code_to_bs
|
||||
from src.db import StockInfo, StockDividend, batch_upsert, get_session
|
||||
from sqlalchemy import select
|
||||
from src.db import StockDividend, batch_upsert, get_session, get_stock_codes
|
||||
from sqlalchemy import select, func
|
||||
|
||||
|
||||
def _get_stock_codes() -> list[str]:
|
||||
def _get_dividend_done(codes: list[str]) -> set[str]:
|
||||
"""查询已有完整10年分红记录的股票代码"""
|
||||
current_year = datetime.now().year
|
||||
years = [str(y) for y in range(current_year - 10, current_year + 1)]
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(select(StockInfo.code))
|
||||
return [row[0] for row in result]
|
||||
result = session.execute(
|
||||
select(StockDividend.code, func.count(StockDividend.id))
|
||||
.where(StockDividend.code.in_(codes))
|
||||
.where(StockDividend.report_date.in_(years))
|
||||
.group_by(StockDividend.code)
|
||||
)
|
||||
# 有记录即可跳过(分红非每年都有,只要有任意年份数据就说明已抓过)
|
||||
return {row[0] for row in result if row[1] > 0}
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
@@ -74,14 +84,26 @@ def fetch_dividend(symbol: str | None = None):
|
||||
"""抓取分红送转数据
|
||||
|
||||
用法:python -m src.main --dividend [--symbol 000001]
|
||||
已有分红记录的股票自动跳过。
|
||||
"""
|
||||
if symbol:
|
||||
codes = [symbol]
|
||||
else:
|
||||
codes = _get_stock_codes()
|
||||
codes = get_stock_codes()
|
||||
|
||||
# 跳过已有分红数据的股票
|
||||
if not symbol:
|
||||
done = _get_dividend_done(codes)
|
||||
if done:
|
||||
print(f" {len(done)} 只股票已有分红数据,跳过", flush=True)
|
||||
codes = [c for c in codes if c not in done]
|
||||
|
||||
total = len(codes)
|
||||
print(f"正在抓取分红送转数据,共 {total} 只股票...", flush=True)
|
||||
if total == 0:
|
||||
print(" 所有股票分红数据已完整,无需抓取", flush=True)
|
||||
return
|
||||
|
||||
print(f"正在抓取分红送转数据,需抓取 {total} 只股票...", flush=True)
|
||||
|
||||
success = 0
|
||||
for i, code in enumerate(codes):
|
||||
|
||||
+41
-15
@@ -6,23 +6,15 @@ BaoStock 按季度查询财务数据:
|
||||
- query_cash_flow_data() 现金流
|
||||
|
||||
数据以 JSON 格式存入 data 列(与现有表结构兼容)。
|
||||
跳过策略:已有 (code, report_date) 记录的季度不再重复抓取。
|
||||
"""
|
||||
|
||||
import json
|
||||
from datetime import datetime
|
||||
import baostock as bs
|
||||
from src.baostock_conn import bs_query, code_to_bs
|
||||
from src.db import StockInfo, FinancialIncome, FinancialBalance, FinancialCashflow, batch_upsert, get_session
|
||||
from sqlalchemy import select
|
||||
|
||||
|
||||
def _get_stock_codes() -> list[str]:
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(select(StockInfo.code))
|
||||
return [row[0] for row in result]
|
||||
finally:
|
||||
session.close()
|
||||
from src.db import FinancialIncome, FinancialBalance, FinancialCashflow, batch_upsert, get_session, get_stock_codes
|
||||
from sqlalchemy import select, func
|
||||
|
||||
|
||||
def _recent_quarters(n: int) -> list[tuple[int, int]]:
|
||||
@@ -39,6 +31,25 @@ def _recent_quarters(n: int) -> list[tuple[int, int]]:
|
||||
return result
|
||||
|
||||
|
||||
def _get_existing_financial(codes: list[str], quarters: list[tuple[int, int]],
|
||||
model_cls) -> set[str]:
|
||||
"""查询已有财务数据的 (code) 集合:三张表都有完整8季度数据的股票"""
|
||||
q_labels = {f"{y}-{m:02d}-{d:02d}" for y, _q in quarters
|
||||
for m, d in [(3, 31), (6, 30), (9, 30), (12, 31)]}
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(
|
||||
select(model_cls.code, func.count(model_cls.id))
|
||||
.where(model_cls.code.in_(codes))
|
||||
.where(model_cls.report_date.in_(q_labels))
|
||||
.group_by(model_cls.code)
|
||||
)
|
||||
q_count = len(q_labels)
|
||||
return {row[0] for row in result if row[1] >= q_count}
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _parse_resultset(code: str, rs, fields: list[str], year: int, quarter: int) -> list[dict]:
|
||||
"""将 BaoStock ResultData 转为 JSON 行"""
|
||||
rows = []
|
||||
@@ -65,16 +76,31 @@ def fetch_financial(symbol: str | None = None):
|
||||
"""抓取财务数据
|
||||
|
||||
用法:python -m src.main --financial [--symbol 000001]
|
||||
默认抓取所有股票最近 8 个季度。
|
||||
默认抓取所有股票最近 8 个季度。已有完整数据的股票自动跳过。
|
||||
"""
|
||||
if symbol:
|
||||
codes = [symbol]
|
||||
else:
|
||||
codes = _get_stock_codes()
|
||||
codes = get_stock_codes()
|
||||
|
||||
quarters = _recent_quarters(8)
|
||||
|
||||
# 跳过三张表都有完整季度数据的股票
|
||||
if not symbol:
|
||||
inc_done = _get_existing_financial(codes, quarters, FinancialIncome)
|
||||
bal_done = _get_existing_financial(codes, quarters, FinancialBalance)
|
||||
cf_done = _get_existing_financial(codes, quarters, FinancialCashflow)
|
||||
done = inc_done & bal_done & cf_done
|
||||
if done:
|
||||
print(f" {len(done)} 只股票财务数据已完整,跳过", flush=True)
|
||||
codes = [c for c in codes if c not in done]
|
||||
|
||||
total = len(codes)
|
||||
quarters = _recent_quarters(8)
|
||||
print(f"正在抓取财务数据,共 {total} 只股票 × {len(quarters)} 个季度...", flush=True)
|
||||
if total == 0:
|
||||
print(" 所有股票财务数据已完整,无需抓取", flush=True)
|
||||
return
|
||||
|
||||
print(f"正在抓取财务数据,需抓取 {total} 只股票 × {len(quarters)} 个季度...", flush=True)
|
||||
|
||||
success = 0
|
||||
fail = 0
|
||||
|
||||
@@ -0,0 +1,129 @@
|
||||
"""指数日线行情抓取 — 使用 BaoStock
|
||||
|
||||
主要指数:
|
||||
sh.000001 上证指数 sh.000300 沪深300
|
||||
sh.000905 中证500 sh.000852 中证1000
|
||||
sh.000688 科创50 sz.399001 深证成指
|
||||
sz.399006 创业板指 sz.399005 中小板指
|
||||
|
||||
用法:
|
||||
python -m src.main --index
|
||||
python -m src.main --index --start-date 20260101 --end-date 20260508
|
||||
"""
|
||||
|
||||
import time
|
||||
from datetime import datetime, timedelta
|
||||
import baostock as bs
|
||||
from src.baostock_conn import bs_query, bs_login
|
||||
from src.config import get_fetch_config
|
||||
from src.db import IndexDaily, batch_upsert, get_session
|
||||
from sqlalchemy import select, func
|
||||
|
||||
|
||||
# 主要指数代码 → BaoStock 格式
|
||||
INDICES = {
|
||||
"000001": ("sh", "上证指数"),
|
||||
"000300": ("sh", "沪深300"),
|
||||
"000905": ("sh", "中证500"),
|
||||
"000852": ("sh", "中证1000"),
|
||||
"000688": ("sh", "科创50"),
|
||||
"399001": ("sz", "深证成指"),
|
||||
"399006": ("sz", "创业板指"),
|
||||
"399005": ("sz", "中小板指"),
|
||||
}
|
||||
|
||||
|
||||
def _clean(val):
|
||||
if val is None:
|
||||
return None
|
||||
if isinstance(val, str) and val.strip() == "":
|
||||
return None
|
||||
return val
|
||||
|
||||
|
||||
def _get_fetched_codes(sd: str, ed: str) -> set[str]:
|
||||
"""查询数据已覆盖起始日期的指数代码"""
|
||||
session = get_session()
|
||||
try:
|
||||
from sqlalchemy import distinct
|
||||
sd_date = datetime.strptime(sd, "%Y-%m-%d").date()
|
||||
check_start = str(sd_date - timedelta(days=5))
|
||||
check_end = str(sd_date + timedelta(days=5))
|
||||
result = session.execute(
|
||||
select(distinct(IndexDaily.code))
|
||||
.where(IndexDaily.date >= check_start)
|
||||
.where(IndexDaily.date <= check_end)
|
||||
)
|
||||
return {row[0] for row in result}
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _fetch_one_index(code: str, market: str, name: str,
|
||||
sd: str, ed: str) -> list[dict] | None:
|
||||
"""抓取单个指数的日线数据"""
|
||||
bs_code = f"{market}.{code}"
|
||||
try:
|
||||
with bs_query(
|
||||
bs.query_history_k_data_plus,
|
||||
bs_code,
|
||||
"date,open,high,low,close,volume,amount,pctChg",
|
||||
start_date=sd, end_date=ed, frequency="d",
|
||||
) as rs:
|
||||
rows = []
|
||||
while rs.next():
|
||||
r = rs.get_row_data()
|
||||
rows.append({
|
||||
"code": code,
|
||||
"date": r[0],
|
||||
"open": _clean(r[1]),
|
||||
"high": _clean(r[2]),
|
||||
"low": _clean(r[3]),
|
||||
"close": _clean(r[4]),
|
||||
"volume": _clean(r[5]),
|
||||
"amount": _clean(r[6]),
|
||||
"pct_change": _clean(r[7]),
|
||||
})
|
||||
return rows if rows else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def fetch_index(start_date: str | None = None, end_date: str | None = None):
|
||||
"""抓取主要指数日线行情"""
|
||||
cfg = get_fetch_config()
|
||||
delay = cfg.get("delay", 0.1)
|
||||
|
||||
if end_date is None:
|
||||
end_date = datetime.now().strftime("%Y%m%d")
|
||||
if start_date is None:
|
||||
start_date = "19901219"
|
||||
|
||||
sd = f"{start_date[:4]}-{start_date[4:6]}-{start_date[6:8]}"
|
||||
ed = f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"
|
||||
|
||||
# 跳过已有数据的指数
|
||||
fetched = _get_fetched_codes(sd, ed)
|
||||
to_fetch = {k: v for k, v in INDICES.items() if k not in fetched}
|
||||
|
||||
if not to_fetch:
|
||||
print(f" 指数数据 {sd} ~ {ed} 已完整,跳过", flush=True)
|
||||
return
|
||||
|
||||
skip_msg = f"(跳过 {len(fetched)} 个已有数据)" if fetched else ""
|
||||
print(f"正在抓取指数日线 {sd} ~ {ed},共 {len(to_fetch)} 个{skip_msg}...", flush=True)
|
||||
|
||||
bs_login()
|
||||
success = 0
|
||||
for code, (market, name) in to_fetch.items():
|
||||
t0 = time.time()
|
||||
rows = _fetch_one_index(code, market, name, sd, ed)
|
||||
if rows:
|
||||
batch_upsert(IndexDaily, rows, ["code", "date"])
|
||||
success += 1
|
||||
print(f" {name}({code}): {len(rows)} 天, {time.time()-t0:.1f}s", flush=True)
|
||||
else:
|
||||
print(f" {name}({code}): 无数据", flush=True)
|
||||
time.sleep(delay)
|
||||
|
||||
print(f" 指数数据抓取完成,成功:{success}/{len(to_fetch)}", flush=True)
|
||||
@@ -16,8 +16,8 @@ import baostock as bs
|
||||
from src.baostock_conn import bs_query, code_to_bs, bs_login
|
||||
from src.config import get_fetch_config
|
||||
from src.db import (
|
||||
StockInfo, StockMin5, StockMin15, StockMin30, StockMin60,
|
||||
batch_upsert, get_session,
|
||||
StockMin5, StockMin15, StockMin30, StockMin60,
|
||||
batch_upsert, get_session, get_stock_codes,
|
||||
)
|
||||
from sqlalchemy import select, func, distinct
|
||||
|
||||
@@ -32,15 +32,6 @@ FREQ_MODEL = {
|
||||
}
|
||||
|
||||
|
||||
def _get_stock_codes() -> list[str]:
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(select(StockInfo.code))
|
||||
return [row[0] for row in result]
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _clean(val):
|
||||
if val is None:
|
||||
return None
|
||||
@@ -121,7 +112,7 @@ def _fetch_one_freq(freq: str, start_date: str | None, end_date: str | None,
|
||||
if symbol:
|
||||
codes = [symbol]
|
||||
else:
|
||||
codes = _get_stock_codes()
|
||||
codes = get_stock_codes()
|
||||
if not codes:
|
||||
print(" 无股票列表,请先运行 --stock-info", flush=True)
|
||||
return
|
||||
|
||||
+2
-11
@@ -15,7 +15,7 @@ from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
import requests
|
||||
import baostock as bs
|
||||
from src.baostock_conn import bs_query, bs_login
|
||||
from src.db import StockInfo, StockSector, batch_upsert, get_session
|
||||
from src.db import StockSector, batch_upsert, get_session, get_stock_codes
|
||||
from sqlalchemy import select
|
||||
|
||||
|
||||
@@ -26,15 +26,6 @@ _HEADERS = {
|
||||
_F10_URL = "https://emweb.securities.eastmoney.com/PC_HSF10/CompanySurvey/CompanySurveyAjax"
|
||||
|
||||
|
||||
def _get_stock_codes() -> list[str]:
|
||||
session = get_session()
|
||||
try:
|
||||
result = session.execute(select(StockInfo.code))
|
||||
return [row[0] for row in result]
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def _fetch_industry() -> dict[str, str]:
|
||||
"""从 BaoStock 获取全部股票的行业分类"""
|
||||
result = {}
|
||||
@@ -88,7 +79,7 @@ def _fetch_region_batch(codes: list[str], workers: int = 10) -> dict[str, str]:
|
||||
|
||||
def fetch_sector(industry_only: bool = False, region_only: bool = False):
|
||||
"""抓取行业分类 + 地域分类"""
|
||||
codes = _get_stock_codes()
|
||||
codes = get_stock_codes()
|
||||
if not codes:
|
||||
print(" 无股票列表,请先运行 --stock-info", flush=True)
|
||||
return
|
||||
|
||||
+8
-1
@@ -12,6 +12,7 @@
|
||||
python -m src.main --intraday --symbol 000001 --freq 30
|
||||
python -m src.main --sector # 行业+地域分类
|
||||
python -m src.main --sector --industry-only # 仅行业分类
|
||||
python -m src.main --index # 指数日线(上证/沪深300/创业板等)
|
||||
"""
|
||||
|
||||
import argparse
|
||||
@@ -32,6 +33,7 @@ def main():
|
||||
choices=["5", "15", "30", "60", "all"],
|
||||
help="分钟K线频率(默认5,all=全部)")
|
||||
parser.add_argument("--sector", action="store_true", help="抓取行业+地域分类")
|
||||
parser.add_argument("--index", action="store_true", help="抓取指数日线行情")
|
||||
parser.add_argument("--industry-only", action="store_true", help="仅抓取行业分类")
|
||||
parser.add_argument("--region-only", action="store_true", help="仅抓取地域分类")
|
||||
parser.add_argument("--start-date", type=str, help="开始日期 YYYYMMDD")
|
||||
@@ -42,7 +44,8 @@ def main():
|
||||
|
||||
if not any([args.stock_info, args.trading_day, args.daily,
|
||||
args.financial, args.dividend, args.intraday,
|
||||
args.sector, args.industry_only, args.region_only]):
|
||||
args.sector, args.industry_only, args.region_only,
|
||||
args.index]):
|
||||
parser.print_help()
|
||||
return
|
||||
|
||||
@@ -80,6 +83,10 @@ def main():
|
||||
from src.fetchers.sector import fetch_sector
|
||||
fetch_sector(industry_only=args.industry_only, region_only=args.region_only)
|
||||
|
||||
if args.index:
|
||||
from src.fetchers.index import fetch_index
|
||||
fetch_index(start_date=args.start_date, end_date=args.end_date)
|
||||
|
||||
print("全部任务完成", flush=True)
|
||||
finally:
|
||||
bs_logout()
|
||||
|
||||
Reference in New Issue
Block a user