573 lines
23 KiB
Python
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
|