rebuild(stage-7): deliver ladder and sector rotation
This commit is contained in:
@@ -6,8 +6,8 @@ from typing import Any
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
from backend.data.contracts import ProviderResult, SnapshotState
|
||||
from backend.data.gateway import DataGateway, MarketDataUnavailable
|
||||
from backend.data.providers.base import ProviderError
|
||||
from backend.data.providers.tushare import TushareProvider
|
||||
from backend.data.repository import MarketRepository
|
||||
from backend.database.connection import Database
|
||||
from backend.features.market.sentiment import calculate_sentiment
|
||||
@@ -22,11 +22,11 @@ class SnapshotSyncError(RuntimeError):
|
||||
|
||||
class MarketSnapshotService:
|
||||
def __init__(
|
||||
self, database: Database, repository: MarketRepository, provider: TushareProvider
|
||||
self, database: Database, repository: MarketRepository, gateway: DataGateway
|
||||
) -> None:
|
||||
self._database = database
|
||||
self._repository = repository
|
||||
self._provider = provider
|
||||
self._gateway = gateway
|
||||
|
||||
def sync(
|
||||
self, requested_date: str | None = None, now: datetime | None = None
|
||||
@@ -47,8 +47,8 @@ class MarketSnapshotService:
|
||||
raise SnapshotSyncError("请先同步股票目录")
|
||||
trade_date, previous_date = dates[0], dates[1]
|
||||
try:
|
||||
inputs = self._provider.snapshot_inputs(trade_date, previous_date)
|
||||
except ProviderError as exc:
|
||||
inputs = self._gateway.snapshot_inputs(trade_date, previous_date)
|
||||
except (ProviderError, MarketDataUnavailable) as exc:
|
||||
raise SnapshotSyncError("收盘行情读取失败,已保留原有快照") from exc
|
||||
daily = inputs.get("daily")
|
||||
if not isinstance(daily, ProviderResult):
|
||||
@@ -126,10 +126,45 @@ class MarketSnapshotService:
|
||||
response["items"] = payload.get("yesterday_limits") or []
|
||||
elif key == "performance":
|
||||
response["items"] = payload.get("limit_performance") or []
|
||||
elif key == "ladder":
|
||||
response["items"] = payload.get("ladders") or []
|
||||
response["history"] = payload.get("limit_performance") or []
|
||||
elif key == "rotation":
|
||||
response["history"] = [
|
||||
{
|
||||
"trade_date": item_payload.get("trade_date"),
|
||||
"sectors": item_payload.get("sector_rotation") or [],
|
||||
}
|
||||
for item_payload in (
|
||||
json.loads(str(item["payload_json"])) for item in history_rows[-9:]
|
||||
)
|
||||
]
|
||||
else:
|
||||
raise SnapshotSyncError("不支持的市场工作区")
|
||||
return response
|
||||
|
||||
def rotation_member_target(
|
||||
self, requested_date: str | None, sector_name: str
|
||||
) -> tuple[str, str]:
|
||||
requested = _date(requested_date or datetime.now(SHANGHAI).date().isoformat())
|
||||
with self._database.read() as connection:
|
||||
row = self._repository.latest_summary(connection, requested)
|
||||
if row is None:
|
||||
raise SnapshotSyncError("等待管理员首次同步真实收盘行情")
|
||||
payload = json.loads(str(row["payload_json"]))
|
||||
sector = next(
|
||||
(
|
||||
item
|
||||
for item in payload.get("sectors") or []
|
||||
if str(item.get("name") or "") == sector_name
|
||||
),
|
||||
None,
|
||||
)
|
||||
representative = str((sector or {}).get("representative") or "")
|
||||
if not representative:
|
||||
raise SnapshotSyncError("该板块缺少可核验的代表股票")
|
||||
return str(row["trade_date"]), representative
|
||||
|
||||
|
||||
def _history_item(payload: dict[str, Any]) -> dict[str, Any]:
|
||||
sentiment = payload.get("sentiment") or {}
|
||||
|
||||
Reference in New Issue
Block a user