Files
xiaobaifupan/next/backend/features/market/events.py
T

200 lines
7.2 KiB
Python

from __future__ import annotations
import json
from datetime import date, datetime
from typing import Any
from zoneinfo import ZoneInfo
from backend.data.gateway import DataGateway
from backend.data.repository import MarketRepository
from backend.database.connection import Database
SHANGHAI = ZoneInfo("Asia/Shanghai")
EVENT_KEYS = {"limit_up": "limits", "broken": "broken", "limit_down": "down_limits"}
class MarketEventError(RuntimeError):
pass
class MarketEventService:
def __init__(
self, database: Database, repository: MarketRepository, gateway: DataGateway
) -> None:
self._database = database
self._repository = repository
self._gateway = gateway
def supplement(self, trade_date: str) -> dict[str, Any]:
target = _date(trade_date)
snapshot = self._snapshot(target)
expected = {
(str(row.get("identifier") or ""), event_type)
for event_type, key in EVENT_KEYS.items()
for row in snapshot.get(key) or []
if row.get("identifier")
}
if not expected:
return {
"trade_date": target,
"events": 0,
"updated": 0,
"coverage": 1.0,
"source_set": ["local"],
"output_version": "event-revisions-v1",
}
result = self._gateway.event_reasons(target)
now = datetime.now(SHANGHAI).isoformat(timespec="seconds")
updated = 0
matched: set[tuple[str, str]] = set()
with self._database.transaction() as connection:
current = {
(str(row["identifier"]), str(row["event_type"])): dict(row)
for row in self._repository.event_revisions(connection, target)
}
for row in result.rows:
identity = (str(row.get("identifier") or ""), str(row.get("event_type") or ""))
if identity not in expected:
continue
useful = any(
row.get(field) not in (None, "")
for field in ("reason", "first_time", "last_time", "open_times")
)
if not useful:
continue
matched.add(identity)
values = _revision_values(row)
existing = current.get(identity)
if existing and int(existing["priority"]) >= 100:
continue
if existing and all(existing.get(key) == value for key, value in values.items()):
continue
self._repository.save_event_revision(
connection,
trade_date=target,
identifier=identity[0],
event_type=identity[1],
source="ifind",
priority=20,
created_by=None,
created_at=now,
**values,
)
updated += 1
return {
"trade_date": target,
"events": len(expected),
"matched": len(matched),
"updated": updated,
"coverage": round(len(matched) / len(expected), 4),
"source_set": ["ifind", "local"],
"output_version": "event-revisions-v1",
}
def revise(
self,
*,
trade_date: str,
identifier: str,
event_type: str,
reason: str,
first_time: str,
last_time: str,
open_times: int | None,
user_id: int,
) -> dict[str, Any]:
target = _date(trade_date)
normalized = identifier.strip().upper()
if event_type not in EVENT_KEYS:
raise MarketEventError("事件类型无效")
snapshot = self._snapshot(target)
exists = any(
str(row.get("identifier") or "") == normalized
for row in snapshot.get(EVENT_KEYS[event_type]) or []
)
if not exists:
raise MarketEventError("该股票不在所选日期的对应事件池中")
normalized_reason = " ".join(reason.split())
if not 1 <= len(normalized_reason) <= 200:
raise MarketEventError("原因应为1至200个字符")
values = {
"reason": normalized_reason,
"first_time": _time(first_time),
"last_time": _time(last_time),
"open_times": open_times,
}
with self._database.transaction() as connection:
revision_id = self._repository.save_event_revision(
connection,
trade_date=target,
identifier=normalized,
event_type=event_type,
source="admin",
priority=100,
created_by=user_id,
created_at=datetime.now(SHANGHAI).isoformat(timespec="seconds"),
**values,
)
return {"id": revision_id, "trade_date": target, "identifier": normalized, **values}
def history(self, trade_date: str, identifier: str) -> list[dict[str, Any]]:
target = _date(trade_date)
with self._database.read() as connection:
rows = self._repository.event_revision_history(
connection, target, identifier.strip().upper()
)
return [dict(row) for row in rows]
def _snapshot(self, trade_date: str) -> dict[str, Any]:
with self._database.read() as connection:
row = self._repository.latest_summary(connection, trade_date)
if row is None or str(row["trade_date"]) != trade_date:
raise MarketEventError("所选日期没有正式行情快照")
return json.loads(str(row["payload_json"]))
def apply_event_revisions(
payload: dict[str, Any], revisions: tuple[Any, ...]
) -> dict[str, Any]:
index = {
(str(row["identifier"]), str(row["event_type"])): row for row in revisions
}
for event_type, key in EVENT_KEYS.items():
for item in payload.get(key) or []:
revision = index.get((str(item.get("identifier") or ""), event_type))
if revision is None:
continue
for field in ("reason", "first_time", "last_time", "open_times"):
value = revision[field]
if value not in (None, ""):
item[field] = value
item["reason_source"] = str(revision["source"])
item["revision_id"] = int(revision["id"])
return payload
def _revision_values(row: dict[str, Any]) -> dict[str, Any]:
return {
"reason": str(row.get("reason") or "").strip(),
"first_time": _time(str(row.get("first_time") or "")),
"last_time": _time(str(row.get("last_time") or "")),
"open_times": row.get("open_times") if isinstance(row.get("open_times"), int) else None,
}
def _date(value: str) -> str:
try:
return date.fromisoformat(value).isoformat()
except ValueError as exc:
raise MarketEventError("日期格式无效") from exc
def _time(value: str) -> str:
normalized = value.strip()
if not normalized:
return ""
try:
return datetime.strptime(normalized, "%H:%M").strftime("%H:%M")
except ValueError as exc:
raise MarketEventError("事件时间格式应为HH:MM") from exc