200 lines
7.2 KiB
Python
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
|