rebuild(stage-11): deliver deterministic heaven workflows

This commit is contained in:
leefer
2026-07-30 07:08:13 +08:00
parent aa3f02bd59
commit 35ae079de7
49 changed files with 7208 additions and 39 deletions
+145 -29
View File
@@ -132,36 +132,9 @@ class TushareProvider:
def sector_members(self, representative: str, trade_date: str) -> ProviderResult:
target = _compact(trade_date)
memberships = self._membership_rows({"ts_code": representative})
active = [row for row in memberships if _active_on(row, target)]
if not active:
raise ProviderError("未找到该股票在目标日期的申万行业")
industry = max(
active,
key=lambda row: (
str(row.get("in_date") or ""),
str(row.get("l2_code") or ""),
),
)
industry, members = self._sector_memberships(representative, target)
sector_code = str(industry.get("l2_code") or "")
sector_name = str(industry.get("l2_name") or "").strip()
if not sector_code:
raise ProviderError("该股票缺少申万二级行业")
members = [
row
for row in self._membership_rows({"l2_code": sector_code})
if _active_on(row, target)
]
deduplicated: dict[str, dict[str, Any]] = {}
for row in members:
code = str(row.get("ts_code") or "")
current = deduplicated.get(code)
if code and (
current is None or str(row.get("in_date") or "") > str(current.get("in_date") or "")
):
deduplicated[code] = row
if not deduplicated:
raise ProviderError("该申万行业没有有效成分股")
daily = self._query(
"daily",
{"trade_date": target},
@@ -170,7 +143,8 @@ class TushareProvider:
)
quote_map = {str(row.get("ts_code") or ""): row for row in daily.rows}
rows = []
for code, member in deduplicated.items():
for member in members:
code = str(member.get("ts_code") or "")
quote = quote_map.get(code) or {}
rows.append(
{
@@ -188,6 +162,114 @@ class TushareProvider:
coverage = sum(bool(row["quoted"]) for row in rows) / len(rows)
return ProviderResult(tuple(rows), _metadata(self.source, "mixed", coverage))
def heaven_inputs(
self, representative: str, trade_date: str
) -> dict[str, ProviderResult | None]:
target = _compact(trade_date)
membership = self.sector_members(representative, trade_date)
sector_code = str(membership.rows[0].get("sector_code") or "") if membership.rows else ""
return {
"members": membership,
"daily": self._optional_query(
"daily",
{"trade_date": target},
"ts_code,trade_date,open,high,low,close,pre_close,pct_chg,vol,amount",
),
"daily_basic": self._optional_query(
"daily_basic",
{"trade_date": target},
"ts_code,trade_date,turnover_rate,volume_ratio,total_mv,circ_mv",
),
"sector_daily": self._optional_query(
"sw_daily",
{"ts_code": sector_code, "trade_date": target},
"ts_code,trade_date,name,open,high,low,close,pct_change,vol,amount,pe,pb,float_mv,total_mv",
),
"indices": self._index_rows(target),
}
def heaven_realtime_inputs(
self, representative: str, trade_date: str, previous_trade_date: str
) -> dict[str, ProviderResult | None]:
target = _compact(trade_date)
previous = _compact(previous_trade_date)
industry, members = self._sector_memberships(representative, target)
sector_code = str(industry.get("l2_code") or "")
directory = self._optional_query(
"stock_basic",
{"exchange": "", "list_status": "L"},
"ts_code,name,industry,market,list_date",
)
active_codes = tuple(
str(row.get("ts_code") or "")
for row in (directory.rows if directory else ())
if row.get("ts_code")
)
realtime_codes = (*active_codes, "000001.SH", "399001.SZ", "399006.SZ")
member_rows = tuple(
{
"sector_code": sector_code,
"sector_name": str(industry.get("l2_name") or "").strip(),
"ts_code": str(row.get("ts_code") or ""),
"name": str(row.get("name") or "").strip(),
}
for row in members
)
start = (datetime.strptime(target, "%Y%m%d") - timedelta(days=35)).strftime("%Y%m%d")
return {
"directory": directory,
"members": ProviderResult(
member_rows,
_metadata(self.source, "membership", 1 if member_rows else 0),
),
"realtime": self._optional_query(
"rt_k",
{"ts_code": ",".join(realtime_codes)},
"ts_code,name,trade_time,open,high,low,close,pre_close,vol,amount,num,pct_chg",
),
"capital": self._optional_query(
"daily_basic",
{"trade_date": previous},
"ts_code,trade_date,total_share,float_share,free_share,total_mv,circ_mv",
),
"stock_history": self._optional_query(
"daily",
{"ts_code": representative, "start_date": start, "end_date": previous},
"ts_code,trade_date,vol,amount",
),
"price_limits": self._optional_query(
"stk_limit",
{"trade_date": target},
"ts_code,trade_date,up_limit,down_limit",
),
"suspensions": self._optional_query(
"suspend_d",
{"suspend_date": target},
"ts_code,suspend_date,resume_date,suspend_timing,suspend_type",
),
"sector_realtime": self._optional_query(
"rt_sw_k",
{"ts_code": sector_code},
"ts_code,name,trade_time,close,pre_close,high,open,low,vol,amount,pct_change",
),
}
def _index_rows(self, trade_date: str) -> ProviderResult | None:
rows: list[dict[str, Any]] = []
completed = 0
for identifier in ("000001.SH", "399001.SZ", "399006.SZ"):
result = self._optional_query(
"index_daily",
{"ts_code": identifier, "trade_date": trade_date},
"ts_code,trade_date,close,pre_close,pct_chg",
)
if result is not None and result.rows:
rows.extend(result.rows)
completed += 1
if not rows:
return None
return ProviderResult(tuple(rows), _metadata(self.source, "percent", completed / 3))
def market_insight(
self,
kind: str,
@@ -443,6 +525,40 @@ class TushareProvider:
rows.extend(result.rows)
return rows
def _sector_memberships(
self, representative: str, trade_date: str
) -> tuple[dict[str, Any], list[dict[str, Any]]]:
active = [
row
for row in self._membership_rows({"ts_code": representative})
if _active_on(row, trade_date)
]
if not active:
raise ProviderError("未找到该股票在目标日期的申万行业")
industry = max(
active,
key=lambda row: (str(row.get("in_date") or ""), str(row.get("l2_code") or "")),
)
sector_code = str(industry.get("l2_code") or "")
if not sector_code:
raise ProviderError("该股票缺少申万二级行业")
rows = [
row
for row in self._membership_rows({"l2_code": sector_code})
if _active_on(row, trade_date)
]
deduplicated: dict[str, dict[str, Any]] = {}
for row in rows:
code = str(row.get("ts_code") or "")
current = deduplicated.get(code)
if code and (
current is None or str(row.get("in_date") or "") > str(current.get("in_date") or "")
):
deduplicated[code] = row
if not deduplicated:
raise ProviderError("该申万行业没有有效成分股")
return industry, list(deduplicated.values())
def _query(
self,
api_name: str,