rebuild(stage-7): deliver ladder and sector rotation
This commit is contained in:
@@ -130,6 +130,82 @@ class TushareProvider:
|
||||
)
|
||||
return datasets
|
||||
|
||||
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 ""),
|
||||
),
|
||||
)
|
||||
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},
|
||||
"ts_code,trade_date,open,close,pct_chg,amount",
|
||||
unit="mixed",
|
||||
)
|
||||
quote_map = {str(row.get("ts_code") or ""): row for row in daily.rows}
|
||||
rows = []
|
||||
for code, member in deduplicated.items():
|
||||
quote = quote_map.get(code) or {}
|
||||
rows.append(
|
||||
{
|
||||
"sector_code": sector_code,
|
||||
"sector_name": sector_name,
|
||||
"ts_code": code,
|
||||
"name": str(member.get("name") or "").strip(),
|
||||
"change": _number(quote.get("pct_chg")) if quote else None,
|
||||
"open": _number(quote.get("open")) if quote else None,
|
||||
"close": _number(quote.get("close")) if quote else None,
|
||||
"amount": _number(quote.get("amount")) * 1000 if quote else None,
|
||||
"quoted": bool(quote),
|
||||
}
|
||||
)
|
||||
coverage = sum(bool(row["quoted"]) for row in rows) / len(rows)
|
||||
return ProviderResult(tuple(rows), _metadata(self.source, "mixed", coverage))
|
||||
|
||||
def _membership_rows(self, params: dict[str, str]) -> list[dict[str, Any]]:
|
||||
rows: list[dict[str, Any]] = []
|
||||
fields = (
|
||||
"l1_code,l1_name,l2_code,l2_name,l3_code,l3_name,"
|
||||
"ts_code,name,in_date,out_date,is_new"
|
||||
)
|
||||
for is_new in ("Y", "N"):
|
||||
result = self._query(
|
||||
"index_member_all",
|
||||
{**params, "is_new": is_new},
|
||||
fields,
|
||||
unit="membership",
|
||||
empty_is_complete=True,
|
||||
)
|
||||
rows.extend(result.rows)
|
||||
return rows
|
||||
|
||||
def _query(
|
||||
self,
|
||||
api_name: str,
|
||||
@@ -197,3 +273,17 @@ def _compact(value: str) -> str:
|
||||
def _display(value: str) -> str:
|
||||
compact = _compact(value)
|
||||
return f"{compact[:4]}-{compact[4:6]}-{compact[6:]}"
|
||||
|
||||
|
||||
def _active_on(row: dict[str, Any], trade_date: str) -> bool:
|
||||
start = str(row.get("in_date") or "")
|
||||
end = str(row.get("out_date") or "")
|
||||
return (not start or start <= trade_date) and (not end or end > trade_date)
|
||||
|
||||
|
||||
def _number(value: Any) -> float:
|
||||
try:
|
||||
number = float(value)
|
||||
return number if number == number else 0.0
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
|
||||
Reference in New Issue
Block a user