rebuild(stage-11): deliver deterministic heaven workflows

This commit is contained in:
leefer
2026-07-30 07:08:13 +08:00
parent aa3f02bd59
commit 35ae079de7
49 changed files with 7208 additions and 39 deletions
+59 -2
View File
@@ -16,6 +16,7 @@ from backend.data.contracts import (
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
@@ -217,6 +218,34 @@ class DataGateway:
)
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)
@@ -310,6 +339,29 @@ class DataGateway:
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,
@@ -485,5 +537,10 @@ def _number(value: Any) -> float:
def _optional_number(value: Any) -> float | None:
number = _number(value)
return number if number > 0 else 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