rebuild(stage-11): deliver deterministic heaven workflows
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user