193 lines
8.4 KiB
Python
193 lines
8.4 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
from datetime import date, datetime, time, timedelta
|
|
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.repository import MarketRepository
|
|
from backend.database.connection import Database
|
|
from backend.features.market.sentiment import calculate_sentiment
|
|
from backend.features.market.snapshot import build_snapshot
|
|
|
|
SHANGHAI = ZoneInfo("Asia/Shanghai")
|
|
|
|
|
|
class SnapshotSyncError(RuntimeError):
|
|
pass
|
|
|
|
|
|
class MarketSnapshotService:
|
|
def __init__(
|
|
self, database: Database, repository: MarketRepository, gateway: DataGateway
|
|
) -> None:
|
|
self._database = database
|
|
self._repository = repository
|
|
self._gateway = gateway
|
|
|
|
def sync(
|
|
self, requested_date: str | None = None, now: datetime | None = None
|
|
) -> dict[str, Any]:
|
|
clock = now or datetime.now(SHANGHAI)
|
|
requested = _date(requested_date or clock.date().isoformat())
|
|
through = requested
|
|
if requested == clock.date().isoformat() and clock.time() < time(9, 15):
|
|
through = (clock.date() - timedelta(days=1)).isoformat()
|
|
elif requested == clock.date().isoformat() and clock.time() < time(15, 10):
|
|
raise SnapshotSyncError("盘中快照任务尚未开放,请保留最近真实收盘数据")
|
|
with self._database.read() as connection:
|
|
dates = self._repository.open_dates(connection, through, 2)
|
|
active_count = self._repository.active_stock_count(connection)
|
|
if len(dates) < 2:
|
|
raise SnapshotSyncError("请先同步完整交易日历")
|
|
if active_count <= 0:
|
|
raise SnapshotSyncError("请先同步股票目录")
|
|
trade_date, previous_date = dates[0], dates[1]
|
|
try:
|
|
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):
|
|
raise SnapshotSyncError("收盘日线缺失,已保留原有快照")
|
|
coverage = len(daily.rows) / active_count
|
|
if coverage < 0.98:
|
|
raise SnapshotSyncError(f"收盘日线覆盖率仅{coverage * 100:.1f}%,未写入不完整快照")
|
|
for key in ("limit_up", "limit_down", "broken", "previous_limit_up", "price_limits"):
|
|
value = inputs.get(key)
|
|
if not isinstance(value, ProviderResult) or value.metadata.coverage < 1:
|
|
raise SnapshotSyncError("涨跌停事件数据不完整,已保留原有快照")
|
|
|
|
snapshot = build_snapshot(trade_date, previous_date, inputs)
|
|
with self._database.read() as connection:
|
|
rows = self._repository.summaries(connection, previous_date, 250)
|
|
history = [json.loads(str(row["payload_json"])) for row in rows]
|
|
sentiment = calculate_sentiment(snapshot, history)
|
|
snapshot["sentiment"] = sentiment
|
|
snapshot.update(snapshot["overview"])
|
|
snapshot["temperature"] = sentiment["score"]
|
|
observed_at = datetime.combine(
|
|
date.fromisoformat(trade_date), time(15), tzinfo=SHANGHAI
|
|
).isoformat(timespec="seconds")
|
|
state = (
|
|
SnapshotState.FINAL if trade_date == clock.date().isoformat() else SnapshotState.ARCHIVE
|
|
)
|
|
with self._database.transaction() as connection:
|
|
self._repository.save_summary(
|
|
connection,
|
|
trade_date=trade_date,
|
|
observed_at=observed_at,
|
|
state=state.value,
|
|
source="tushare",
|
|
coverage=min(coverage, 1),
|
|
payload=snapshot,
|
|
)
|
|
return {
|
|
"trade_date": trade_date,
|
|
"observed_at": observed_at,
|
|
"coverage": round(min(coverage, 1), 4),
|
|
"stocks": len(daily.rows),
|
|
"limit_up": len(snapshot["limits"]),
|
|
"limit_down": len(snapshot["down_limits"]),
|
|
"broken": len(snapshot["broken"]),
|
|
"temperature": sentiment["score"],
|
|
}
|
|
|
|
def workspace(self, key: str, requested_date: str | None = None) -> dict[str, Any]:
|
|
requested = _date(requested_date or datetime.now(SHANGHAI).date().isoformat())
|
|
with self._database.read() as connection:
|
|
row = self._repository.latest_summary(connection, requested)
|
|
history_rows = self._repository.summaries(connection, requested, 60)
|
|
if row is None:
|
|
return {"trade_date": None, "message": "等待管理员首次同步真实收盘行情"}
|
|
payload = json.loads(str(row["payload_json"]))
|
|
response: dict[str, Any] = {
|
|
"trade_date": str(row["trade_date"]),
|
|
"observed_at": str(row["observed_at"]),
|
|
"carried_forward": str(row["trade_date"]) != requested,
|
|
"message": "沿用最近真实收盘快照" if str(row["trade_date"]) != requested else "",
|
|
"overview": payload.get("overview") or {},
|
|
}
|
|
if key == "emotion":
|
|
response["sentiment"] = payload.get("sentiment") or {}
|
|
response["history"] = [
|
|
_history_item(json.loads(str(item["payload_json"]))) for item in history_rows
|
|
]
|
|
elif key == "pool":
|
|
response["items"] = payload.get("limits") or []
|
|
elif key == "broken":
|
|
response["items"] = payload.get("broken") or []
|
|
elif key == "limit-down":
|
|
response["items"] = payload.get("down_limits") or []
|
|
elif key == "yesterday":
|
|
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 {}
|
|
stats = sentiment.get("stats") or {}
|
|
overview = payload.get("overview") or {}
|
|
return {
|
|
"trade_date": payload.get("trade_date"),
|
|
"temperature": sentiment.get("score"),
|
|
"phase": sentiment.get("phase"),
|
|
"direction": sentiment.get("direction"),
|
|
"positive_rate": stats.get("positive_rate"),
|
|
"seal_rate": overview.get("seal_rate"),
|
|
"limit_up": overview.get("limit_up"),
|
|
"broken": overview.get("broken"),
|
|
"limit_down": overview.get("limit_down"),
|
|
"max_height": stats.get("max_height"),
|
|
"amount": overview.get("amount"),
|
|
}
|
|
|
|
|
|
def _date(value: str) -> str:
|
|
try:
|
|
return date.fromisoformat(value).isoformat()
|
|
except ValueError as exc:
|
|
raise SnapshotSyncError("日期格式无效") from exc
|