rebuild(stage-5): establish market data gateway and charts
This commit is contained in:
@@ -0,0 +1,379 @@
|
||||
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 (
|
||||
ChartPoint,
|
||||
ChartSeries,
|
||||
DataSource,
|
||||
DataUsage,
|
||||
MarketEntity,
|
||||
ObservationMetadata,
|
||||
SnapshotState,
|
||||
TradeContext,
|
||||
)
|
||||
from backend.data.policy import DataSourcePolicy
|
||||
from backend.data.providers.base import MarketDataProvider, ProviderError
|
||||
from backend.data.quality import DataQualityError, require_quality
|
||||
from backend.data.repository import MarketRepository
|
||||
from backend.database.connection import Database
|
||||
|
||||
SHANGHAI = ZoneInfo("Asia/Shanghai")
|
||||
|
||||
|
||||
class MarketDataUnavailable(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class DataGateway:
|
||||
def __init__(
|
||||
self,
|
||||
database: Database,
|
||||
repository: MarketRepository,
|
||||
providers: tuple[MarketDataProvider, ...],
|
||||
policy: DataSourcePolicy,
|
||||
) -> None:
|
||||
self._database = database
|
||||
self._repository = repository
|
||||
self._providers = {provider.source: provider for provider in providers}
|
||||
self._policy = policy
|
||||
|
||||
def refresh_reference(self, now: datetime | None = None) -> dict[str, int | str]:
|
||||
clock = now or datetime.now(SHANGHAI)
|
||||
provider = self._provider(DataSource.TUSHARE)
|
||||
start = (clock.date() - timedelta(days=370)).isoformat()
|
||||
end = (clock.date() + timedelta(days=40)).isoformat()
|
||||
calendar = require_quality("calendar", provider.calendar(start, end))
|
||||
entities = require_quality("entities", provider.entities())
|
||||
observed_at = clock.isoformat(timespec="seconds")
|
||||
with self._database.transaction() as connection:
|
||||
self._repository.replace_calendar(
|
||||
connection, calendar.rows, calendar.metadata.source.value, observed_at
|
||||
)
|
||||
self._repository.replace_stocks(
|
||||
connection, entities.rows, entities.metadata.source.value, observed_at
|
||||
)
|
||||
return {
|
||||
"calendar_days": len(calendar.rows),
|
||||
"entities": len(entities.rows),
|
||||
"observed_at": observed_at,
|
||||
}
|
||||
|
||||
def trade_context(
|
||||
self, requested_date: str | None = None, now: datetime | None = None
|
||||
) -> TradeContext:
|
||||
clock = now or datetime.now(SHANGHAI)
|
||||
requested = _date(requested_date or clock.date().isoformat())
|
||||
with self._database.read() as connection:
|
||||
summary = self._repository.latest_summary(connection, requested)
|
||||
if summary is None:
|
||||
return TradeContext(
|
||||
requested_date=requested,
|
||||
actual_date=None,
|
||||
previous_date=None,
|
||||
observed_at=None,
|
||||
state=None,
|
||||
carried_forward=False,
|
||||
message="等待管理员首次同步真实行情",
|
||||
)
|
||||
actual = str(summary["trade_date"])
|
||||
observed_at = datetime.fromisoformat(str(summary["observed_at"]))
|
||||
state = SnapshotState(str(summary["state"]))
|
||||
carried = actual != requested
|
||||
with self._database.read() as connection:
|
||||
dates = self._repository.open_dates(connection, actual, 2)
|
||||
return TradeContext(
|
||||
requested_date=requested,
|
||||
actual_date=actual,
|
||||
previous_date=dates[1] if len(dates) > 1 else None,
|
||||
observed_at=observed_at,
|
||||
state=state,
|
||||
carried_forward=carried,
|
||||
message="沿用最近真实收盘快照" if carried else "",
|
||||
)
|
||||
|
||||
def latest_available_date(self, now: datetime | None = None) -> str:
|
||||
context = self.trade_context(now=now)
|
||||
return context.actual_date or (now or datetime.now(SHANGHAI)).date().isoformat()
|
||||
|
||||
def summary(self, requested_date: str | None = None) -> dict[str, Any]:
|
||||
context = self.trade_context(requested_date)
|
||||
if context.actual_date is None:
|
||||
return {"context": context, "values": None}
|
||||
with self._database.read() as connection:
|
||||
row = self._repository.latest_summary(connection, context.actual_date)
|
||||
return {"context": context, "values": json.loads(str(row["payload_json"])) if row else None}
|
||||
|
||||
def search(self, query: str) -> tuple[MarketEntity, ...]:
|
||||
with self._database.read() as connection:
|
||||
return self._repository.search(connection, query)
|
||||
|
||||
def chart(
|
||||
self,
|
||||
entity_type: str,
|
||||
identifier: str,
|
||||
interval: str,
|
||||
now: datetime | None = None,
|
||||
) -> ChartSeries:
|
||||
if entity_type not in {"stock", "sector", "theme", "index"}:
|
||||
raise MarketDataUnavailable("不支持的行情标的类型")
|
||||
if interval not in {"day", "minute"}:
|
||||
raise MarketDataUnavailable("不支持的行情周期")
|
||||
entity = self._resolve_entity(entity_type, identifier)
|
||||
clock = now or datetime.now(SHANGHAI)
|
||||
with self._database.read() as connection:
|
||||
cached = self._repository.chart(connection, entity_type, entity.identifier, interval)
|
||||
if cached and not _chart_stale(cached, interval, clock):
|
||||
return _stored_chart(entity, cached)
|
||||
try:
|
||||
series = self._load_chart(entity, interval, clock)
|
||||
except (ProviderError, DataQualityError) as exc:
|
||||
if cached:
|
||||
return _stored_chart(entity, cached)
|
||||
raise MarketDataUnavailable("当前没有可用的真实行情数据") from exc
|
||||
self._save_chart(series)
|
||||
return series
|
||||
|
||||
def _load_chart(self, entity: MarketEntity, interval: str, clock: datetime) -> ChartSeries:
|
||||
dataset = "daily_chart" if interval == "day" else "minute_chart"
|
||||
target_dates = self._chart_dates(clock)
|
||||
errors: list[Exception] = []
|
||||
for source in self._policy.candidates(dataset, DataUsage.DISPLAY):
|
||||
provider = self._providers.get(source)
|
||||
if provider is None or not provider.configured:
|
||||
continue
|
||||
self._policy.assert_allowed(source, DataUsage.DISPLAY)
|
||||
for target in target_dates:
|
||||
try:
|
||||
raw = (
|
||||
provider.daily(entity.entity_type, entity.identifier, target)
|
||||
if interval == "day"
|
||||
else provider.minute(entity.entity_type, entity.identifier, target)
|
||||
)
|
||||
require_quality(dataset, raw)
|
||||
normalized = _normalize_chart(entity, interval, raw.rows, raw.metadata, clock)
|
||||
if normalized.points:
|
||||
return normalized
|
||||
except (ProviderError, DataQualityError, ValueError) as exc:
|
||||
errors.append(exc)
|
||||
if interval == "day":
|
||||
break
|
||||
raise MarketDataUnavailable("当前没有可用的真实行情数据") from (
|
||||
errors[-1] if errors else None
|
||||
)
|
||||
|
||||
def _chart_dates(self, clock: datetime) -> tuple[str, ...]:
|
||||
today = clock.date().isoformat()
|
||||
with self._database.read() as connection:
|
||||
dates = self._repository.open_dates(connection, today, 8)
|
||||
if dates:
|
||||
return dates
|
||||
return tuple((clock.date() - timedelta(days=offset)).isoformat() for offset in range(8))
|
||||
|
||||
def _resolve_entity(self, entity_type: str, identifier: str) -> MarketEntity:
|
||||
normalized = identifier.strip().upper()
|
||||
with self._database.read() as connection:
|
||||
entity = self._repository.entity(connection, entity_type, normalized)
|
||||
if entity is None and entity_type == "stock" and normalized.isdigit():
|
||||
matches = self._repository.search(connection, normalized, 4)
|
||||
entity = next(
|
||||
(
|
||||
item
|
||||
for item in matches
|
||||
if item.entity_type == "stock" and item.code == normalized
|
||||
),
|
||||
None,
|
||||
)
|
||||
if entity:
|
||||
return entity
|
||||
if entity_type == "stock" and len(normalized) == 6 and normalized.isdigit():
|
||||
suffix = (
|
||||
"BJ"
|
||||
if normalized.startswith(("4", "8", "9"))
|
||||
else "SH"
|
||||
if normalized.startswith("6")
|
||||
else "SZ"
|
||||
)
|
||||
return MarketEntity("stock", f"{normalized}.{suffix}", normalized, normalized)
|
||||
raise MarketDataUnavailable("未找到该行情标的")
|
||||
|
||||
def _save_chart(self, series: ChartSeries) -> None:
|
||||
payload = {
|
||||
"previous_close": series.previous_close,
|
||||
"points": [
|
||||
{
|
||||
"time": point.time,
|
||||
"open": point.open,
|
||||
"high": point.high,
|
||||
"low": point.low,
|
||||
"close": point.close,
|
||||
"volume": point.volume,
|
||||
"amount": point.amount,
|
||||
"average": point.average,
|
||||
}
|
||||
for point in series.points
|
||||
],
|
||||
}
|
||||
with self._database.transaction() as connection:
|
||||
self._repository.save_chart(
|
||||
connection,
|
||||
entity_type=series.entity.entity_type,
|
||||
identifier=series.entity.identifier,
|
||||
interval=series.interval,
|
||||
trade_date=series.trade_date,
|
||||
observed_at=series.metadata.observed_at.isoformat(timespec="seconds"),
|
||||
source=series.metadata.source.value,
|
||||
usage=series.metadata.usage.value,
|
||||
adjustment=series.metadata.adjustment,
|
||||
coverage=series.metadata.coverage,
|
||||
payload=payload,
|
||||
)
|
||||
|
||||
def _provider(self, source: DataSource) -> MarketDataProvider:
|
||||
provider = self._providers.get(source)
|
||||
if provider is None or not provider.configured:
|
||||
raise MarketDataUnavailable("所需行情服务尚未配置")
|
||||
return provider
|
||||
|
||||
|
||||
def _normalize_chart(
|
||||
entity: MarketEntity,
|
||||
interval: str,
|
||||
rows: tuple[dict[str, Any], ...],
|
||||
metadata: ObservationMetadata,
|
||||
clock: datetime,
|
||||
) -> ChartSeries:
|
||||
parsed: list[tuple[str, ChartPoint, float | None]] = []
|
||||
for row in rows:
|
||||
stamp = str(row.get("trade_date") or row.get("time") or row.get("trade_time") or "")
|
||||
trade_date = _row_date(stamp)
|
||||
point_time = trade_date if interval == "day" else _row_time(stamp)
|
||||
close = _number(row.get("close"))
|
||||
open_price = _number(row.get("open"))
|
||||
high = _number(row.get("high"))
|
||||
low = _number(row.get("low"))
|
||||
if not trade_date or close <= 0 or open_price <= 0 or high <= 0 or low <= 0:
|
||||
continue
|
||||
if interval == "minute" and not "09:30" <= point_time <= "15:00":
|
||||
continue
|
||||
volume = _number(row.get("volume", row.get("vol")))
|
||||
amount = _number(row.get("amount"))
|
||||
if metadata.source is DataSource.TUSHARE:
|
||||
volume *= 100
|
||||
amount *= 1000
|
||||
parsed.append(
|
||||
(
|
||||
trade_date,
|
||||
ChartPoint(
|
||||
time=point_time,
|
||||
open=open_price,
|
||||
high=high,
|
||||
low=low,
|
||||
close=close,
|
||||
volume=volume,
|
||||
amount=amount,
|
||||
average=_optional_number(row.get("avgPrice", row.get("average"))),
|
||||
),
|
||||
_optional_number(row.get("preClose", row.get("pre_close"))),
|
||||
)
|
||||
)
|
||||
parsed.sort(key=lambda item: (item[0], item[1].time))
|
||||
if not parsed:
|
||||
raise DataQualityError("chart contains no valid points")
|
||||
if interval == "minute":
|
||||
latest = parsed[-1][0]
|
||||
parsed = [item for item in parsed if item[0] == latest]
|
||||
else:
|
||||
today = clock.date().isoformat()
|
||||
if parsed[-1][0] == today and not _valid_today_bar(parsed[-1][1], clock):
|
||||
parsed.pop()
|
||||
if not parsed:
|
||||
raise DataQualityError("chart contains no completed bar")
|
||||
parsed = parsed[-90:]
|
||||
previous = parsed[0][2]
|
||||
if previous is None and interval == "day" and len(parsed) > 1:
|
||||
previous = parsed[-2][1].close
|
||||
return ChartSeries(
|
||||
entity=entity,
|
||||
interval=interval,
|
||||
trade_date=parsed[-1][0],
|
||||
previous_close=previous,
|
||||
points=tuple(item[1] for item in parsed),
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
|
||||
def _valid_today_bar(point: ChartPoint, clock: datetime) -> bool:
|
||||
if clock.time() < time(9, 30):
|
||||
return False
|
||||
return (
|
||||
point.volume > 0
|
||||
and point.amount > 0
|
||||
and point.high >= max(point.open, point.close)
|
||||
and point.low <= min(point.open, point.close)
|
||||
)
|
||||
|
||||
|
||||
def _chart_stale(row: Any, interval: str, clock: datetime) -> bool:
|
||||
observed = datetime.fromisoformat(str(row["observed_at"]))
|
||||
if interval == "day":
|
||||
return row["trade_date"] < clock.date().isoformat() and clock.time() >= time(15, 5)
|
||||
return (clock - observed.astimezone(SHANGHAI)).total_seconds() > 30
|
||||
|
||||
|
||||
def _stored_chart(entity: MarketEntity, row: Any) -> ChartSeries:
|
||||
payload = json.loads(str(row["payload_json"]))
|
||||
points = tuple(ChartPoint(**point) for point in payload.get("points") or [])
|
||||
metadata = ObservationMetadata(
|
||||
source=DataSource(str(row["source"])),
|
||||
observed_at=datetime.fromisoformat(str(row["observed_at"])),
|
||||
unit="yuan/share",
|
||||
adjustment=str(row["adjustment"]),
|
||||
freshness_seconds=0,
|
||||
coverage=float(row["coverage"]),
|
||||
state=SnapshotState.ARCHIVE,
|
||||
usage=DataUsage(str(row["usage"])),
|
||||
)
|
||||
return ChartSeries(
|
||||
entity,
|
||||
str(row["interval"]),
|
||||
str(row["trade_date"]),
|
||||
payload.get("previous_close"),
|
||||
points,
|
||||
metadata,
|
||||
)
|
||||
|
||||
|
||||
def _date(value: str) -> str:
|
||||
try:
|
||||
return date.fromisoformat(value).isoformat()
|
||||
except ValueError as exc:
|
||||
raise MarketDataUnavailable("日期格式无效") from exc
|
||||
|
||||
|
||||
def _row_date(value: str) -> str:
|
||||
compact = value[:10].replace("-", "")
|
||||
if len(compact) != 8 or not compact.isdigit():
|
||||
return ""
|
||||
return f"{compact[:4]}-{compact[4:6]}-{compact[6:]}"
|
||||
|
||||
|
||||
def _row_time(value: str) -> str:
|
||||
if " " in value:
|
||||
return value.split(" ", 1)[1][:5]
|
||||
return value[-8:-3] if len(value) >= 8 else value[:5]
|
||||
|
||||
|
||||
def _number(value: Any) -> float:
|
||||
try:
|
||||
return float(value or 0)
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
|
||||
|
||||
def _optional_number(value: Any) -> float | None:
|
||||
number = _number(value)
|
||||
return number if number > 0 else None
|
||||
Reference in New Issue
Block a user