rebuild(stage-7): deliver ladder and sector rotation
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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: ...
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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("实时行情服务尚未配置")
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user