374 lines
17 KiB
Python
374 lines
17 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.data.sentiment import calculate_sentiment
|
|
from backend.database.connection import Database
|
|
from backend.features.market.events import apply_event_revisions
|
|
from backend.features.market.snapshot import build_realtime_inputs, 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"],
|
|
"source_set": ["tushare", "local"],
|
|
"output_version": "market-summary-v1",
|
|
}
|
|
|
|
def sync_realtime(
|
|
self, requested_date: str | None = None, now: datetime | None = None
|
|
) -> dict[str, Any]:
|
|
clock = now or datetime.now(SHANGHAI)
|
|
target = _date(requested_date or clock.date().isoformat())
|
|
if target != clock.date().isoformat():
|
|
raise SnapshotSyncError("盘中任务只允许同步当前交易日")
|
|
local_time = clock.time().replace(tzinfo=None)
|
|
in_window = time(9, 15) <= local_time < time(11, 35) or time(
|
|
12, 55
|
|
) <= local_time < time(15, 5)
|
|
if not in_window:
|
|
raise SnapshotSyncError("当前不在盘中行情刷新窗口")
|
|
with self._database.read() as connection:
|
|
dates = self._repository.open_dates(connection, target, 2)
|
|
active_count = self._repository.active_stock_count(connection)
|
|
if len(dates) < 2 or dates[0] != target:
|
|
raise SnapshotSyncError("当前日期不是有效交易日")
|
|
if active_count <= 0:
|
|
raise SnapshotSyncError("请先同步股票目录")
|
|
try:
|
|
raw = self._gateway.realtime_snapshot_inputs(target, dates[1])
|
|
except (ProviderError, MarketDataUnavailable) as exc:
|
|
raise SnapshotSyncError("盘中行情读取失败,已保留最后成功快照") from exc
|
|
daily = raw.get("daily")
|
|
price_limits = raw.get("price_limits")
|
|
if not isinstance(daily, ProviderResult):
|
|
raise SnapshotSyncError("盘中行情缺失,已保留最后成功快照")
|
|
coverage = len(daily.rows) / active_count
|
|
if coverage < 0.9:
|
|
raise SnapshotSyncError(
|
|
f"盘中行情覆盖率仅{coverage * 100:.1f}%,已保留最后成功快照"
|
|
)
|
|
if not isinstance(price_limits, ProviderResult):
|
|
raise SnapshotSyncError("盘中涨跌停价格缺失,已保留最后成功快照")
|
|
limit_coverage = len(price_limits.rows) / max(len(daily.rows), 1)
|
|
if limit_coverage < 0.95:
|
|
raise SnapshotSyncError("盘中涨跌停价格覆盖不足,已保留最后成功快照")
|
|
inputs = build_realtime_inputs(raw, self._gateway.stock_directory())
|
|
snapshot = build_snapshot(target, dates[1], inputs)
|
|
with self._database.read() as connection:
|
|
rows = self._repository.summaries(connection, dates[1], 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 = clock.isoformat(timespec="seconds")
|
|
with self._database.transaction() as connection:
|
|
self._repository.save_summary(
|
|
connection,
|
|
trade_date=target,
|
|
observed_at=observed_at,
|
|
state=SnapshotState.REALTIME.value,
|
|
source="tushare",
|
|
coverage=min(coverage, 1),
|
|
payload=snapshot,
|
|
)
|
|
return {
|
|
"trade_date": target,
|
|
"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"],
|
|
"source_set": ["tushare", "local"],
|
|
"output_version": "market-summary-realtime-v1",
|
|
}
|
|
|
|
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"]))
|
|
with self._database.read() as connection:
|
|
revisions = self._repository.event_revisions(connection, str(row["trade_date"]))
|
|
apply_event_revisions(payload, revisions)
|
|
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 entity_detail(
|
|
self, entity_type: str, identifier: str, requested_date: str | None = None
|
|
) -> dict[str, Any]:
|
|
context = self._gateway.trade_context(requested_date)
|
|
if context.actual_date is None:
|
|
raise SnapshotSyncError("等待管理员首次同步真实收盘行情")
|
|
entity = self._gateway.resolve_entity(entity_type, identifier)
|
|
series = self._gateway.chart(entity_type, entity.identifier, "day")
|
|
eligible = [point for point in series.points if point.time <= context.actual_date]
|
|
if not eligible:
|
|
raise SnapshotSyncError("所选日期之前没有可核验的真实行情")
|
|
current = eligible[-1]
|
|
previous = eligible[-2].close if len(eligible) > 1 else series.previous_close
|
|
change = (current.close / previous - 1) * 100 if previous else None
|
|
|
|
with self._database.read() as connection:
|
|
summary = self._repository.latest_summary(connection, context.actual_date)
|
|
factor_row = connection.execute(
|
|
"""
|
|
SELECT factor_values.payload_json
|
|
FROM screener_factor_values AS factor_values
|
|
JOIN screener_factor_snapshots AS snapshots
|
|
ON snapshots.id = factor_values.snapshot_id
|
|
WHERE factor_values.identifier = ? AND snapshots.trade_date <= ?
|
|
ORDER BY snapshots.trade_date DESC, snapshots.id DESC LIMIT 1
|
|
""",
|
|
(entity.identifier, context.actual_date),
|
|
).fetchone()
|
|
factors = json.loads(str(factor_row["payload_json"])) if factor_row else {}
|
|
snapshot = json.loads(str(summary["payload_json"])) if summary else {}
|
|
if summary:
|
|
with self._database.read() as connection:
|
|
revisions = self._repository.event_revisions(
|
|
connection, str(summary["trade_date"])
|
|
)
|
|
apply_event_revisions(snapshot, revisions)
|
|
event = (
|
|
_entity_event(snapshot, entity.identifier, entity.code)
|
|
if entity_type == "stock"
|
|
else None
|
|
)
|
|
metrics = [
|
|
{"key": "open", "label": "开盘", "value": current.open, "unit": "元"},
|
|
{"key": "high", "label": "最高", "value": current.high, "unit": "元"},
|
|
{"key": "low", "label": "最低", "value": current.low, "unit": "元"},
|
|
{"key": "amount", "label": "成交额", "value": current.amount, "unit": "元"},
|
|
{"key": "volume", "label": "成交量", "value": current.volume, "unit": "股"},
|
|
]
|
|
for key, label, unit in (
|
|
("turnover_rate", "换手率", "%"),
|
|
("return_5d", "近5日", "%"),
|
|
("return_20d", "近20日", "%"),
|
|
("total_mv_billion", "总市值", "亿"),
|
|
("circ_mv_billion", "流通市值", "亿"),
|
|
):
|
|
metrics.append({"key": key, "label": label, "value": factors.get(key), "unit": unit})
|
|
money_flow = None
|
|
if entity_type == "stock":
|
|
money_flow = {
|
|
"available": any(
|
|
factors.get(key) is not None
|
|
for key in ("net_flow_million", "large_flow_million", "net_flow_5d_million")
|
|
),
|
|
"net_million": factors.get("net_flow_million"),
|
|
"large_million": factors.get("large_flow_million"),
|
|
"net_5d_million": factors.get("net_flow_5d_million"),
|
|
"flow_to_circ_mv_5d": factors.get("flow_to_circ_mv_5d"),
|
|
}
|
|
return {
|
|
"entity": {
|
|
"entity_type": entity.entity_type,
|
|
"identifier": entity.identifier,
|
|
"code": entity.code,
|
|
"name": factors.get("name") or entity.name,
|
|
"sector": factors.get("sector") or entity.sector,
|
|
},
|
|
"trade_date": current.time,
|
|
"observed_at": series.metadata.observed_at,
|
|
"price": current.close,
|
|
"previous_close": previous,
|
|
"change": round(change, 4) if change is not None else None,
|
|
"metrics": metrics,
|
|
"money_flow": money_flow,
|
|
"event": event,
|
|
}
|
|
|
|
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 _entity_event(
|
|
snapshot: dict[str, Any], identifier: str, code: str
|
|
) -> dict[str, Any] | None:
|
|
for key, status in (("limits", "涨停"), ("broken", "炸板"), ("down_limits", "跌停")):
|
|
for row in snapshot.get(key) or []:
|
|
if str(row.get("identifier") or "") == identifier or str(row.get("code") or "") == code:
|
|
return {
|
|
"status": status,
|
|
"reason": str(row.get("reason") or ""),
|
|
"streak": row.get("streak"),
|
|
"first_time": row.get("first_time"),
|
|
"last_time": row.get("last_time"),
|
|
"open_times": row.get("open_times"),
|
|
"seal_amount": row.get("seal_amount"),
|
|
}
|
|
return None
|
|
|
|
|
|
def _date(value: str) -> str:
|
|
try:
|
|
return date.fromisoformat(value).isoformat()
|
|
except ValueError as exc:
|
|
raise SnapshotSyncError("日期格式无效") from exc
|