rebuild(stage-7): deliver ladder and sector rotation

This commit is contained in:
leefer
2026-07-30 03:28:24 +08:00
parent 31a53de890
commit a76d344a98
28 changed files with 1161 additions and 13 deletions
+62
View File
@@ -107,6 +107,68 @@ class DataGateway:
row = self._repository.latest_summary(connection, context.actual_date)
return {"context": context, "values": json.loads(str(row["payload_json"])) if row else None}
def snapshot_inputs(
self, trade_date: str, previous_trade_date: str
) -> dict[str, Any]:
provider = self._provider(DataSource.TUSHARE)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
return provider.snapshot_inputs(trade_date, previous_trade_date)
def sector_members(
self, trade_date: str, sector_name: str, representative: str
) -> dict[str, Any]:
with self._database.read() as connection:
cached = self._repository.sector_members(connection, trade_date, sector_name)
if cached:
return json.loads(str(cached["payload_json"]))
provider = self._provider(DataSource.TUSHARE)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
result = provider.sector_members(representative, trade_date)
if not result.rows:
raise MarketDataUnavailable("该板块暂无可核验的申万成分股")
rows = sorted(
(dict(row) for row in result.rows),
key=lambda row: (
bool(row.get("quoted")),
_number(row.get("change")),
_number(row.get("amount")),
),
reverse=True,
)
payload = {
"trade_date": trade_date,
"sector_name": str(rows[0].get("sector_name") or sector_name),
"sector_code": str(rows[0].get("sector_code") or ""),
"member_count": len(rows),
"quoted_count": sum(bool(row.get("quoted")) for row in rows),
"coverage": round(result.metadata.coverage, 4),
"items": [
{
"identifier": str(row.get("ts_code") or ""),
"code": str(row.get("ts_code") or "").split(".")[0],
"name": str(row.get("name") or ""),
"change": row.get("change"),
"open": row.get("open"),
"close": row.get("close"),
"amount": row.get("amount"),
"quoted": bool(row.get("quoted")),
}
for row in rows
],
}
with self._database.transaction() as connection:
self._repository.save_sector_members(
connection,
trade_date=trade_date,
sector_name=sector_name,
sector_code=payload["sector_code"],
observed_at=result.metadata.observed_at.isoformat(timespec="seconds"),
source=result.metadata.source.value,
coverage=result.metadata.coverage,
payload=payload,
)
return payload
def search(self, query: str) -> tuple[MarketEntity, ...]:
with self._database.read() as connection:
return self._repository.search(connection, query)
+2
View File
@@ -26,3 +26,5 @@ class MarketDataProvider(Protocol):
def snapshot_inputs(
self, trade_date: str, previous_trade_date: str
) -> dict[str, ProviderResult | dict[str, Any]]: ...
def sector_members(self, representative: str, trade_date: str) -> ProviderResult: ...
+3
View File
@@ -99,6 +99,9 @@ class EastmoneyProvider:
) -> dict[str, ProviderResult | dict[str, object]]:
raise ProviderError("The display provider cannot build market snapshots")
def sector_members(self, representative: str, trade_date: str) -> ProviderResult:
raise ProviderError("The display provider is not the constituent authority")
@staticmethod
def _secid(entity_type: str, identifier: str) -> str:
if entity_type == "index" and identifier in INDEX_CODES:
+3
View File
@@ -89,6 +89,9 @@ class IfindProvider:
) -> dict[str, ProviderResult | dict[str, Any]]:
raise ProviderError("iFinD is not the post-close snapshot authority")
def sector_members(self, representative: str, trade_date: str) -> ProviderResult:
raise ProviderError("iFinD is not the Shenwan constituent authority")
def _request(self, endpoint: str, body: dict[str, Any]) -> dict[str, Any]:
if not self.configured:
raise ProviderError("实时行情服务尚未配置")
+90
View File
@@ -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
+46
View File
@@ -188,6 +188,52 @@ class MarketRepository:
(through,),
).fetchone()
def sector_members(
self, connection: sqlite3.Connection, trade_date: str, sector_name: str
) -> sqlite3.Row | None:
return connection.execute(
"""
SELECT * FROM sector_member_snapshots
WHERE trade_date = ? AND sector_name = ?
""",
(trade_date, sector_name),
).fetchone()
def save_sector_members(
self,
connection: sqlite3.Connection,
*,
trade_date: str,
sector_name: str,
sector_code: str,
observed_at: str,
source: str,
coverage: float,
payload: dict[str, Any],
) -> None:
connection.execute(
"""
INSERT INTO sector_member_snapshots (
trade_date, sector_name, sector_code, observed_at, source, coverage, payload_json
) VALUES (?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(trade_date, sector_name) DO UPDATE SET
sector_code = excluded.sector_code,
observed_at = excluded.observed_at,
source = excluded.source,
coverage = excluded.coverage,
payload_json = excluded.payload_json
""",
(
trade_date,
sector_name,
sector_code,
observed_at,
source,
coverage,
json.dumps(payload, ensure_ascii=False, separators=(",", ":")),
),
)
def save_chart(
self,
connection: sqlite3.Connection,