Files
xiaobaifupan/next/backend/data/gateway.py
T

573 lines
23 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 (
ChartPoint,
ChartSeries,
DataSource,
DataUsage,
MarketEntity,
ObservationMetadata,
ProviderResult,
SnapshotState,
TradeContext,
)
from backend.data.heaven import historical_payload, realtime_payload, should_use_realtime
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.data.screener_gateway import assemble_screener_inputs
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 snapshot_inputs(self, trade_date: str, previous_trade_date: str) -> dict[str, Any]:
provider = self._provider(DataSource.TUSHARE)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
return provider.snapshot_inputs(trade_date, previous_trade_date)
def realtime_snapshot_inputs(
self, trade_date: str, previous_trade_date: str
) -> dict[str, Any]:
provider = self._provider(DataSource.TUSHARE)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
directory = self.stock_directory()
identifiers = tuple(sorted(directory))
if not identifiers:
raise MarketDataUnavailable("股票目录为空,无法读取盘中行情")
return provider.realtime_market_inputs(
trade_date, previous_trade_date, identifiers
)
def trading_dates(self, through: str, limit: int = 2) -> tuple[str, ...]:
requested = _date(through)
with self._database.read() as connection:
return self._repository.open_dates(connection, requested, limit)
def stock_directory(self) -> dict[str, dict[str, Any]]:
with self._database.read() as connection:
rows = self._repository.stock_directory(connection)
return {str(row["identifier"]): dict(row) for row in rows}
def insight_inputs(
self,
kind: str,
trade_date: str,
previous_trade_date: str = "",
identifier: str = "",
) -> dict[str, Any]:
provider = self._provider(DataSource.TUSHARE)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
return provider.market_insight(kind, trade_date, previous_trade_date, identifier)
def screener_inputs(
self, trade_date: str, history_days: int = 260
) -> tuple[dict[str, Any], dict[str, float], list[str]]:
dates = self.trading_dates(trade_date, history_days)
if len(dates) < 21:
raise MarketDataUnavailable("历史交易日不足21日,无法生成选股因子")
chronological = tuple(reversed(dates))
provider = self._provider(DataSource.TUSHARE)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
raw = provider.screener_inputs(chronological)
inputs, coverage = assemble_screener_inputs(
self._database, self._repository, trade_date, chronological, raw
)
return inputs, coverage, [provider.source.value, DataSource.LOCAL.value]
def dynamic_auction(
self, identifiers: tuple[str, ...], start_time: str, end_time: str
) -> ProviderResult:
if not identifiers:
raise MarketDataUnavailable("动态竞价候选范围为空")
provider = self._provider(DataSource.IFIND)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
result = provider.realtime_snapshots(identifiers, start_time, end_time)
if not result.rows:
raise MarketDataUnavailable("当前动态竞价快照暂不可用")
return result
def event_reasons(self, trade_date: str) -> ProviderResult:
provider = self._provider(DataSource.IFIND)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
result = provider.event_reasons(trade_date)
if result.metadata.usage != DataUsage.CALCULATION:
raise MarketDataUnavailable("事件补充来源未获准参与正式数据治理")
if result.metadata.coverage <= 0:
raise MarketDataUnavailable("事件原因补充服务暂不可用")
return result
def sector_members(
self, trade_date: str, sector_name: str, representative: str
) -> dict[str, Any]:
with self._database.read() as connection:
cached = self._repository.sector_members(connection, trade_date, sector_name)
if cached:
return json.loads(str(cached["payload_json"]))
provider = self._provider(DataSource.TUSHARE)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
result = provider.sector_members(representative, trade_date)
if not result.rows:
raise MarketDataUnavailable("该板块暂无可核验的申万成分股")
rows = sorted(
(dict(row) for row in result.rows),
key=lambda row: (
bool(row.get("quoted")),
_number(row.get("change")),
_number(row.get("amount")),
),
reverse=True,
)
payload = {
"trade_date": trade_date,
"sector_name": str(rows[0].get("sector_name") or sector_name),
"sector_code": str(rows[0].get("sector_code") or ""),
"member_count": len(rows),
"quoted_count": sum(bool(row.get("quoted")) for row in rows),
"coverage": round(result.metadata.coverage, 4),
"items": [
{
"identifier": str(row.get("ts_code") or ""),
"code": str(row.get("ts_code") or "").split(".")[0],
"name": str(row.get("name") or ""),
"change": row.get("change"),
"open": row.get("open"),
"close": row.get("close"),
"amount": row.get("amount"),
"quoted": bool(row.get("quoted")),
}
for row in rows
],
}
with self._database.transaction() as connection:
self._repository.save_sector_members(
connection,
trade_date=trade_date,
sector_name=sector_name,
sector_code=payload["sector_code"],
observed_at=result.metadata.observed_at.isoformat(timespec="seconds"),
source=result.metadata.source.value,
coverage=result.metadata.coverage,
payload=payload,
)
return payload
def heaven_trend_inputs(
self, query: str, requested_date: str, now: datetime | None = None
) -> dict[str, Any]:
clock = now or datetime.now(SHANGHAI)
context = self.trade_context(requested_date, clock)
if context.actual_date is None:
raise MarketDataUnavailable("等待管理员首次同步真实收盘行情")
stock = self._resolve_stock_query(query)
provider = self._provider(DataSource.TUSHARE)
self._policy.assert_allowed(provider.source, DataUsage.CALCULATION)
if should_use_realtime(requested_date, context, clock):
dates = self.trading_dates(requested_date, 2)
if len(dates) < 2 or dates[0] != requested_date:
raise MarketDataUnavailable("目标日期不是有效交易日")
raw = provider.heaven_realtime_inputs(stock.identifier, dates[0], dates[1])
return realtime_payload(
self._database,
self._repository,
stock,
dates[0],
dates[1],
raw,
clock,
)
return historical_payload(
self._database, self._repository, stock, context.actual_date, provider
)
def search(self, query: str) -> tuple[MarketEntity, ...]:
with self._database.read() as connection:
return self._repository.search(connection, query)
def resolve_entity(self, entity_type: str, identifier: str) -> MarketEntity:
return self._resolve_entity(entity_type, identifier)
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 _resolve_stock_query(self, query: str) -> MarketEntity:
normalized = query.strip()
if not normalized:
raise MarketDataUnavailable("请输入股票代码或股票名称")
with self._database.read() as connection:
matches = tuple(
item
for item in self._repository.search(connection, normalized, 16)
if item.entity_type == "stock"
)
exact = [
item
for item in matches
if item.code.casefold() == normalized.casefold()
or item.identifier.casefold() == normalized.casefold()
or item.name.casefold() == normalized.casefold()
]
if len(exact) == 1:
return exact[0]
if len(exact) > 1:
raise MarketDataUnavailable("股票名称存在重名,请输入六位代码")
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:
if value is None or value == "":
return None
try:
number = float(value)
except (TypeError, ValueError):
return None
return number if number == number else None