feat(selection): 补充策略执行结果查询链路
This commit is contained in:
@@ -0,0 +1,112 @@
|
||||
"""Create persisted selection run and signal result tables."""
|
||||
|
||||
from collections.abc import Sequence
|
||||
|
||||
import sqlalchemy as sa
|
||||
from alembic import op
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
|
||||
revision: str = "0002_selection_results"
|
||||
down_revision: str | None = "0001_market_data"
|
||||
branch_labels: str | Sequence[str] | None = None
|
||||
depends_on: str | Sequence[str] | None = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
"""Create the current selection-run result model."""
|
||||
|
||||
op.create_table(
|
||||
"selection_run",
|
||||
sa.Column("id", sa.String(36), primary_key=True),
|
||||
sa.Column("strategy", sa.String(64), nullable=False),
|
||||
sa.Column("target_trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("market_sync_batch_id", sa.String(36), nullable=False),
|
||||
sa.Column("status", sa.String(24), nullable=False),
|
||||
sa.Column("target_count", sa.Integer(), nullable=False),
|
||||
sa.Column("eligible_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("evaluated_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("selected_stock_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("signal_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("failed_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("coverage", sa.Numeric(8, 6), nullable=False, server_default="0"),
|
||||
sa.Column("error_type", sa.String(64)),
|
||||
sa.Column("error_message", sa.Text()),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.Column("finished_at", sa.DateTime(timezone=True)),
|
||||
sa.UniqueConstraint(
|
||||
"strategy",
|
||||
"target_trade_date",
|
||||
name="uq_selection_run_strategy_date",
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_selection_run_status_date",
|
||||
"selection_run",
|
||||
["strategy", "status", "target_trade_date"],
|
||||
)
|
||||
op.create_table(
|
||||
"selection_run_item",
|
||||
sa.Column(
|
||||
"run_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("selection_run.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("ts_code", sa.String(12), nullable=False),
|
||||
sa.Column("name", sa.String(128), nullable=False),
|
||||
sa.Column("status", sa.String(32), nullable=False),
|
||||
sa.Column("signal_count", sa.Integer(), nullable=False, server_default="0"),
|
||||
sa.Column("reason", sa.Text()),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("run_id", "ts_code"),
|
||||
)
|
||||
op.create_index("ix_selection_run_item_status", "selection_run_item", ["run_id", "status"])
|
||||
op.create_table(
|
||||
"selection_signal",
|
||||
sa.Column(
|
||||
"run_id",
|
||||
sa.String(36),
|
||||
sa.ForeignKey("selection_run.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("ts_code", sa.String(12), nullable=False),
|
||||
sa.Column("name", sa.String(128), nullable=False),
|
||||
sa.Column("target_trade_date", sa.Date(), nullable=False),
|
||||
sa.Column("strategy", sa.String(64), nullable=False),
|
||||
sa.Column("category", sa.String(64), nullable=False),
|
||||
sa.Column("close", sa.Numeric(20, 6), nullable=False),
|
||||
sa.Column("details", JSONB, nullable=False),
|
||||
sa.Column(
|
||||
"created_at",
|
||||
sa.DateTime(timezone=True),
|
||||
nullable=False,
|
||||
server_default=sa.text("now()"),
|
||||
),
|
||||
sa.PrimaryKeyConstraint("run_id", "ts_code", "category"),
|
||||
)
|
||||
op.create_index(
|
||||
"ix_selection_signal_strategy_date",
|
||||
"selection_signal",
|
||||
["strategy", "target_trade_date", "ts_code"],
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
"""Drop selection result tables in dependency-safe order."""
|
||||
|
||||
op.drop_index("ix_selection_signal_strategy_date", table_name="selection_signal")
|
||||
op.drop_table("selection_signal")
|
||||
op.drop_index("ix_selection_run_item_status", table_name="selection_run_item")
|
||||
op.drop_table("selection_run_item")
|
||||
op.drop_index("ix_selection_run_status_date", table_name="selection_run")
|
||||
op.drop_table("selection_run")
|
||||
@@ -4,9 +4,11 @@ from fastapi import APIRouter
|
||||
|
||||
from zhixing_server.interfaces.http.system import operational_router, system_router
|
||||
from zhixing_server.modules.market_data.presentation.home import home_router
|
||||
from zhixing_server.modules.selection.presentation.http import selection_router
|
||||
|
||||
api_v1_router = APIRouter(prefix="/api/v1")
|
||||
api_v1_router.include_router(system_router, prefix="/system", tags=["system"])
|
||||
api_v1_router.include_router(home_router, prefix="/home", tags=["home"])
|
||||
api_v1_router.include_router(selection_router, prefix="/selection", tags=["selection"])
|
||||
|
||||
__all__ = ["api_v1_router", "operational_router"]
|
||||
|
||||
@@ -5,6 +5,8 @@ from sqlalchemy import (
|
||||
Column,
|
||||
Date,
|
||||
DateTime,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
MetaData,
|
||||
Numeric,
|
||||
@@ -12,8 +14,10 @@ from sqlalchemy import (
|
||||
String,
|
||||
Table,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
func,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import JSONB
|
||||
|
||||
metadata = MetaData()
|
||||
|
||||
@@ -114,3 +118,78 @@ market_sync_item = Table(
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
PrimaryKeyConstraint("batch_id", "item_kind", "item_key"),
|
||||
)
|
||||
|
||||
selection_run = Table(
|
||||
"selection_run",
|
||||
metadata,
|
||||
Column("id", String(36), primary_key=True),
|
||||
Column("strategy", String(64), nullable=False),
|
||||
Column("target_trade_date", Date, nullable=False),
|
||||
Column("market_sync_batch_id", String(36), nullable=False),
|
||||
Column("status", String(24), nullable=False),
|
||||
Column("target_count", Integer, nullable=False),
|
||||
Column("eligible_count", Integer, nullable=False, server_default="0"),
|
||||
Column("evaluated_count", Integer, nullable=False, server_default="0"),
|
||||
Column("selected_stock_count", Integer, nullable=False, server_default="0"),
|
||||
Column("signal_count", Integer, nullable=False, server_default="0"),
|
||||
Column("failed_count", Integer, nullable=False, server_default="0"),
|
||||
Column("coverage", Numeric(8, 6), nullable=False, server_default="0"),
|
||||
Column("error_type", String(64)),
|
||||
Column("error_message", Text),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
Column("finished_at", DateTime(timezone=True)),
|
||||
UniqueConstraint("strategy", "target_trade_date", name="uq_selection_run_strategy_date"),
|
||||
)
|
||||
|
||||
selection_run_item = Table(
|
||||
"selection_run_item",
|
||||
metadata,
|
||||
Column(
|
||||
"run_id", String(36), ForeignKey("selection_run.id", ondelete="CASCADE"), nullable=False
|
||||
),
|
||||
Column("ts_code", String(12), nullable=False),
|
||||
Column("name", String(128), nullable=False),
|
||||
Column("status", String(32), nullable=False),
|
||||
Column("signal_count", Integer, nullable=False, server_default="0"),
|
||||
Column("reason", Text),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
PrimaryKeyConstraint("run_id", "ts_code"),
|
||||
)
|
||||
|
||||
selection_signal = Table(
|
||||
"selection_signal",
|
||||
metadata,
|
||||
Column(
|
||||
"run_id", String(36), ForeignKey("selection_run.id", ondelete="CASCADE"), nullable=False
|
||||
),
|
||||
Column("ts_code", String(12), nullable=False),
|
||||
Column("name", String(128), nullable=False),
|
||||
Column("target_trade_date", Date, nullable=False),
|
||||
Column("strategy", String(64), nullable=False),
|
||||
Column("category", String(64), nullable=False),
|
||||
Column("close", Numeric(20, 6), nullable=False),
|
||||
Column("details", JSONB, nullable=False),
|
||||
Column("created_at", DateTime(timezone=True), nullable=False, server_default=func.now()),
|
||||
PrimaryKeyConstraint("run_id", "ts_code", "category"),
|
||||
)
|
||||
|
||||
# Keep the declarative metadata aligned with the indexes created by the
|
||||
# Alembic revisions. Alembic uses this object for both offline inspection
|
||||
# and future autogeneration, so omitting these indexes would make the schema
|
||||
# appear drifted even though the migration creates them.
|
||||
Index("ix_market_daily_bar_trade_date", market_daily_bar.c.trade_date)
|
||||
Index("ix_market_daily_basic_trade_date", market_daily_basic.c.trade_date)
|
||||
Index("ix_market_sync_item_status", market_sync_item.c.batch_id, market_sync_item.c.status)
|
||||
Index(
|
||||
"ix_selection_run_status_date",
|
||||
selection_run.c.strategy,
|
||||
selection_run.c.status,
|
||||
selection_run.c.target_trade_date,
|
||||
)
|
||||
Index("ix_selection_run_item_status", selection_run_item.c.run_id, selection_run_item.c.status)
|
||||
Index(
|
||||
"ix_selection_signal_strategy_date",
|
||||
selection_signal.c.strategy,
|
||||
selection_signal.c.target_trade_date,
|
||||
selection_signal.c.ts_code,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,187 @@
|
||||
"""Application orchestration for persisted whole-universe B1 runs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
from datetime import date
|
||||
from typing import Literal, Protocol
|
||||
|
||||
from ..domain.models import SelectionEvaluation
|
||||
from ..domain.runs import (
|
||||
SelectionExecutionSource,
|
||||
SelectionRerunRequired,
|
||||
SelectionRun,
|
||||
SelectionRunInProgress,
|
||||
SelectionRunItem,
|
||||
SelectionRunStatus,
|
||||
SelectionRunStore,
|
||||
SelectionUniverseReader,
|
||||
)
|
||||
from .evaluate import EvaluateZhixingB1
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
StrategyName = Literal["zhixing_b1"]
|
||||
_FAILURE_STATUSES = {"insufficient_history", "missing_target_bar", "data_error"}
|
||||
|
||||
|
||||
class SelectionEvaluator(Protocol):
|
||||
"""Minimal single-stock evaluator required by the batch orchestrator."""
|
||||
|
||||
def execute(self, ts_code: str, target_trade_date: date) -> SelectionEvaluation: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PreparedSelectionRun:
|
||||
"""A claimed run and its immutable market-data source snapshot."""
|
||||
|
||||
run: SelectionRun
|
||||
source: SelectionExecutionSource
|
||||
|
||||
|
||||
class RunZhixingB1:
|
||||
"""Prepare, execute, and query persisted Zhixing B1 result batches."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
reader: SelectionUniverseReader,
|
||||
store: SelectionRunStore,
|
||||
evaluator: SelectionEvaluator | None = None,
|
||||
) -> None:
|
||||
"""Inject storage ports and optionally a test evaluator."""
|
||||
|
||||
self.reader = reader
|
||||
self.store = store
|
||||
self.evaluator = evaluator or EvaluateZhixingB1(reader)
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
strategy: StrategyName,
|
||||
target_trade_date: date,
|
||||
*,
|
||||
rerun: bool,
|
||||
) -> PreparedSelectionRun:
|
||||
"""Validate source eligibility before claiming the rerunnable key."""
|
||||
|
||||
source = self.reader.load_execution_source(strategy, target_trade_date)
|
||||
run = self.store.prepare_run(
|
||||
strategy,
|
||||
target_trade_date,
|
||||
source,
|
||||
rerun=rerun,
|
||||
)
|
||||
return PreparedSelectionRun(run=run, source=source)
|
||||
|
||||
def execute(self, prepared: PreparedSelectionRun) -> None:
|
||||
"""Evaluate every eligible stock and converge the persisted run status.
|
||||
|
||||
This method is the boundary used by FastAPI's in-process background
|
||||
task. An unexpected batch-level error is recorded before the worker
|
||||
returns so the UI never mistakes a lost worker exception for success.
|
||||
"""
|
||||
|
||||
evaluated_count = 0
|
||||
selected_stock_count = 0
|
||||
signal_count = 0
|
||||
failed_count = 0
|
||||
try:
|
||||
for stock in prepared.source.stocks:
|
||||
try:
|
||||
evaluation = self.evaluator.execute(
|
||||
stock.ts_code,
|
||||
prepared.source.target_trade_date,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - isolate one stock from the batch
|
||||
logger.exception(
|
||||
"selection_item_failed run_id=%s ts_code=%s",
|
||||
prepared.run.id,
|
||||
stock.ts_code,
|
||||
)
|
||||
evaluation = SelectionEvaluation(
|
||||
ts_code=stock.ts_code,
|
||||
target_trade_date=prepared.source.target_trade_date,
|
||||
status="data_error",
|
||||
reason=_safe_item_error(exc),
|
||||
)
|
||||
item = _to_item(stock.ts_code, stock.name, evaluation)
|
||||
self.store.record_item(prepared.run.id, item)
|
||||
evaluated_count += 1
|
||||
selected_stock_count += evaluation.status == "selected"
|
||||
signal_count += len(evaluation.signals)
|
||||
failed_count += evaluation.status in _FAILURE_STATUSES
|
||||
|
||||
status = _run_status(evaluated_count, failed_count)
|
||||
self.store.finish_run(
|
||||
prepared.run.id,
|
||||
status,
|
||||
evaluated_count=evaluated_count,
|
||||
selected_stock_count=selected_stock_count,
|
||||
signal_count=signal_count,
|
||||
failed_count=failed_count,
|
||||
)
|
||||
except Exception as exc: # noqa: BLE001 - worker boundary must persist failure state
|
||||
logger.exception("selection_run_failed run_id=%s", prepared.run.id)
|
||||
try:
|
||||
self.store.finish_run(
|
||||
prepared.run.id,
|
||||
"failed",
|
||||
evaluated_count=evaluated_count,
|
||||
selected_stock_count=selected_stock_count,
|
||||
signal_count=signal_count,
|
||||
failed_count=max(failed_count, 1),
|
||||
error_type="batch_error",
|
||||
error_message=str(exc),
|
||||
)
|
||||
except Exception: # noqa: BLE001 - preserve the original worker failure
|
||||
logger.exception("selection_run_failure_persist_failed run_id=%s", prepared.run.id)
|
||||
|
||||
def get_run(self, run_id: str) -> SelectionRun | None:
|
||||
"""Read one persisted run for polling."""
|
||||
|
||||
return self.store.get_run(run_id)
|
||||
|
||||
def get_latest(
|
||||
self,
|
||||
strategy: StrategyName,
|
||||
target_trade_date: date | None = None,
|
||||
) -> SelectionRun | None:
|
||||
"""Read the current result by date or the latest result for a strategy."""
|
||||
|
||||
return self.store.get_latest_run(strategy, target_trade_date)
|
||||
|
||||
|
||||
def _to_item(ts_code: str, name: str, evaluation: SelectionEvaluation) -> SelectionRunItem:
|
||||
"""Translate a single-stock domain result into a stored item."""
|
||||
|
||||
return SelectionRunItem(
|
||||
ts_code=ts_code,
|
||||
name=name or (evaluation.signals[0].name if evaluation.signals else ""),
|
||||
status=evaluation.status,
|
||||
signal_count=len(evaluation.signals),
|
||||
reason=evaluation.reason,
|
||||
signals=evaluation.signals,
|
||||
)
|
||||
|
||||
|
||||
def _run_status(evaluated_count: int, failed_count: int) -> SelectionRunStatus:
|
||||
"""Map per-stock outcomes into a visible batch status."""
|
||||
|
||||
if failed_count == 0:
|
||||
return "success"
|
||||
if evaluated_count == 0 or failed_count >= evaluated_count:
|
||||
return "failed"
|
||||
return "partial_success"
|
||||
|
||||
|
||||
def _safe_item_error(error: Exception) -> str:
|
||||
"""Keep per-stock failure context readable without persisting tracebacks."""
|
||||
|
||||
return " ".join(str(error).split())[:500] or error.__class__.__name__
|
||||
|
||||
|
||||
__all__ = [
|
||||
"PreparedSelectionRun",
|
||||
"RunZhixingB1",
|
||||
"SelectionRerunRequired",
|
||||
"SelectionRunInProgress",
|
||||
]
|
||||
@@ -6,6 +6,7 @@ from datetime import date
|
||||
from typing import Protocol
|
||||
|
||||
from .models import StockHistory
|
||||
from .runs import SelectionExecutionSource, SelectionUniverseReader
|
||||
|
||||
|
||||
class MarketDataReaderError(RuntimeError):
|
||||
@@ -16,3 +17,11 @@ class MarketDataReader(Protocol):
|
||||
"""Read qfq history sufficient for one historical strategy evaluation."""
|
||||
|
||||
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory: ...
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MarketDataReader",
|
||||
"MarketDataReaderError",
|
||||
"SelectionExecutionSource",
|
||||
"SelectionUniverseReader",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,133 @@
|
||||
"""Domain contracts for persisted historical selection runs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from typing import Literal, Protocol
|
||||
|
||||
from .models import SelectionEvaluationStatus, SelectionSignal, StockHistory
|
||||
|
||||
SelectionRunStatus = Literal["running", "success", "partial_success", "failed"]
|
||||
SelectionRunItemStatus = SelectionEvaluationStatus
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SelectionStock:
|
||||
"""One eligible current stock that will be evaluated for a run."""
|
||||
|
||||
ts_code: str
|
||||
name: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SelectionExecutionSource:
|
||||
"""Market-data batch and eligible stock snapshot used by one run."""
|
||||
|
||||
market_sync_batch_id: str
|
||||
target_trade_date: date
|
||||
target_count: int
|
||||
valid_count: int
|
||||
coverage: Decimal
|
||||
stocks: tuple[SelectionStock, ...] = field(default_factory=tuple)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SelectionRunItem:
|
||||
"""Persistable per-stock evaluation state and its independent signals."""
|
||||
|
||||
ts_code: str
|
||||
name: str
|
||||
status: SelectionRunItemStatus
|
||||
signal_count: int = 0
|
||||
reason: str | None = None
|
||||
signals: tuple[SelectionSignal, ...] = field(default_factory=tuple)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SelectionRun:
|
||||
"""A current execution attempt and its materialized result rows."""
|
||||
|
||||
id: str
|
||||
strategy: Literal["zhixing_b1"]
|
||||
target_trade_date: date
|
||||
market_sync_batch_id: str | None
|
||||
status: SelectionRunStatus
|
||||
target_count: int
|
||||
eligible_count: int
|
||||
evaluated_count: int
|
||||
selected_stock_count: int
|
||||
signal_count: int
|
||||
failed_count: int
|
||||
coverage: Decimal
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
created_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
items: tuple[SelectionRunItem, ...] = field(default_factory=tuple)
|
||||
signals: tuple[SelectionSignal, ...] = field(default_factory=tuple)
|
||||
|
||||
|
||||
class SelectionRunError(RuntimeError):
|
||||
"""Base class for expected selection-run persistence failures."""
|
||||
|
||||
|
||||
class SelectionRunInProgress(SelectionRunError):
|
||||
"""The requested strategy and date already have a running attempt."""
|
||||
|
||||
|
||||
class SelectionRerunRequired(SelectionRunError):
|
||||
"""A terminal result exists and an explicit rerun confirmation is missing."""
|
||||
|
||||
|
||||
class SelectionRunStoreError(SelectionRunError):
|
||||
"""The selection-run repository could not complete a database operation."""
|
||||
|
||||
|
||||
class SelectionRunStore(Protocol):
|
||||
"""Persistence port for current selection runs and their materialized rows."""
|
||||
|
||||
def prepare_run(
|
||||
self,
|
||||
strategy: Literal["zhixing_b1"],
|
||||
target_trade_date: date,
|
||||
source: SelectionExecutionSource,
|
||||
*,
|
||||
rerun: bool,
|
||||
) -> SelectionRun: ...
|
||||
|
||||
def record_item(self, run_id: str, item: SelectionRunItem) -> None: ...
|
||||
|
||||
def finish_run(
|
||||
self,
|
||||
run_id: str,
|
||||
status: SelectionRunStatus,
|
||||
*,
|
||||
evaluated_count: int,
|
||||
selected_stock_count: int,
|
||||
signal_count: int,
|
||||
failed_count: int,
|
||||
error_type: str | None = None,
|
||||
error_message: str | None = None,
|
||||
) -> None: ...
|
||||
|
||||
def get_run(self, run_id: str) -> SelectionRun | None: ...
|
||||
|
||||
def get_latest_run(
|
||||
self,
|
||||
strategy: Literal["zhixing_b1"],
|
||||
target_trade_date: date | None = None,
|
||||
) -> SelectionRun | None: ...
|
||||
|
||||
|
||||
class SelectionUniverseReader(Protocol):
|
||||
"""Read a qualified market-data source snapshot for one strategy run."""
|
||||
|
||||
def load_execution_source(
|
||||
self,
|
||||
strategy: str,
|
||||
target_trade_date: date,
|
||||
) -> SelectionExecutionSource: ...
|
||||
|
||||
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory: ...
|
||||
@@ -11,12 +11,17 @@ import psycopg
|
||||
from ....bootstrap.config import Settings
|
||||
from ..domain.models import SelectionBar, SelectionDailyBasic, StockHistory
|
||||
from ..domain.ports import MarketDataReaderError
|
||||
from ..domain.runs import SelectionExecutionSource, SelectionStock
|
||||
|
||||
|
||||
class SelectionReaderError(MarketDataReaderError):
|
||||
"""Database read failure with stock and target-date context."""
|
||||
|
||||
|
||||
class SelectionMarketDataNotReady(MarketDataReaderError):
|
||||
"""The requested date has no market-data batch eligible for selection."""
|
||||
|
||||
|
||||
_HISTORY_QUERY = """
|
||||
SELECT
|
||||
bar.ts_code,
|
||||
@@ -41,6 +46,41 @@ WHERE bar.ts_code = %s
|
||||
ORDER BY bar.trade_date ASC
|
||||
"""
|
||||
|
||||
_SOURCE_QUERY = """
|
||||
SELECT id, target_count, valid_count, coverage
|
||||
FROM market_sync_batch
|
||||
WHERE target_trade_date = %s
|
||||
AND strategy_eligible = true
|
||||
AND status IN ('success', 'partial_success')
|
||||
ORDER BY finished_at DESC NULLS LAST, created_at DESC, id DESC
|
||||
LIMIT 1
|
||||
"""
|
||||
|
||||
_ELIGIBLE_STOCKS_QUERY = """
|
||||
SELECT stock.ts_code, stock.name
|
||||
FROM market_stock AS stock
|
||||
WHERE stock.is_active = true
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM market_daily_bar AS bar
|
||||
WHERE bar.ts_code = stock.ts_code
|
||||
AND bar.trade_date = %s
|
||||
AND bar.source_adj = 'qfq'
|
||||
AND bar.open IS NOT NULL
|
||||
AND bar.high IS NOT NULL
|
||||
AND bar.low IS NOT NULL
|
||||
AND bar.close IS NOT NULL
|
||||
AND bar.vol IS NOT NULL
|
||||
)
|
||||
AND EXISTS (
|
||||
SELECT 1
|
||||
FROM market_daily_basic AS basic
|
||||
WHERE basic.ts_code = stock.ts_code
|
||||
AND basic.trade_date = %s
|
||||
)
|
||||
ORDER BY stock.ts_code
|
||||
"""
|
||||
|
||||
|
||||
def _as_date(value: object) -> date:
|
||||
"""Convert a PostgreSQL date-like scalar to a date."""
|
||||
@@ -122,6 +162,64 @@ class PostgresMarketDataReader:
|
||||
daily_basic={trade_date: daily_basic[trade_date] for trade_date in sorted(daily_basic)},
|
||||
)
|
||||
|
||||
def load_execution_source(
|
||||
self,
|
||||
strategy: str,
|
||||
target_trade_date: date,
|
||||
) -> SelectionExecutionSource:
|
||||
"""Load the qualified market-data snapshot for a strategy run.
|
||||
|
||||
Args:
|
||||
strategy: Supported strategy identity. The current reader accepts
|
||||
``zhixing_b1`` and keeps the parameter explicit for future
|
||||
strategy-specific eligibility rules.
|
||||
target_trade_date: Historical trading date to evaluate.
|
||||
|
||||
Returns:
|
||||
The eligible stock snapshot and its source synchronization facts.
|
||||
|
||||
Raises:
|
||||
SelectionMarketDataNotReady: If no eligible synchronization batch
|
||||
or complete active stock exists for the requested date.
|
||||
SelectionReaderError: If PostgreSQL cannot complete the read.
|
||||
"""
|
||||
|
||||
if strategy != "zhixing_b1":
|
||||
raise SelectionMarketDataNotReady(f"unsupported selection strategy: {strategy}")
|
||||
try:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
source_row = connection.execute(_SOURCE_QUERY, (target_trade_date,)).fetchone()
|
||||
if source_row is None:
|
||||
raise SelectionMarketDataNotReady(
|
||||
f"market data is not strategy-eligible for {target_trade_date.isoformat()}"
|
||||
)
|
||||
stock_rows = connection.execute(
|
||||
_ELIGIBLE_STOCKS_QUERY,
|
||||
(target_trade_date, target_trade_date),
|
||||
).fetchall()
|
||||
except SelectionMarketDataNotReady:
|
||||
raise
|
||||
except psycopg.Error as exc:
|
||||
raise SelectionReaderError(
|
||||
f"failed to load selection source at {target_trade_date.isoformat()}"
|
||||
) from exc
|
||||
|
||||
stocks = tuple(
|
||||
SelectionStock(ts_code=str(row[0]), name=str(row[1] or "")) for row in stock_rows
|
||||
)
|
||||
if not stocks:
|
||||
raise SelectionMarketDataNotReady(
|
||||
f"no eligible stocks have complete market data for {target_trade_date.isoformat()}"
|
||||
)
|
||||
return SelectionExecutionSource(
|
||||
market_sync_batch_id=str(source_row[0]),
|
||||
target_trade_date=target_trade_date,
|
||||
target_count=int(source_row[1]),
|
||||
valid_count=int(source_row[2]),
|
||||
coverage=Decimal(str(source_row[3])),
|
||||
stocks=stocks,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _map_row(
|
||||
row: tuple[object, ...],
|
||||
|
||||
@@ -0,0 +1,414 @@
|
||||
"""PostgreSQL persistence adapter for selection execution runs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from collections import defaultdict
|
||||
from collections.abc import Generator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from datetime import date, datetime
|
||||
from decimal import Decimal
|
||||
from typing import Any, Literal, cast
|
||||
from uuid import uuid4
|
||||
|
||||
import psycopg
|
||||
from psycopg.types.json import Jsonb
|
||||
|
||||
from ..domain.models import SelectionSignal, ZhixingB1Category
|
||||
from ..domain.runs import (
|
||||
SelectionExecutionSource,
|
||||
SelectionRerunRequired,
|
||||
SelectionRun,
|
||||
SelectionRunError,
|
||||
SelectionRunInProgress,
|
||||
SelectionRunItem,
|
||||
SelectionRunStatus,
|
||||
SelectionRunStore,
|
||||
SelectionRunStoreError,
|
||||
)
|
||||
from ..domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
|
||||
|
||||
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(ZHIXING_B1_SIGNAL_ORDER)}
|
||||
|
||||
|
||||
class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
"""Persist one current result attempt per strategy and target date."""
|
||||
|
||||
def __init__(self, database_url: str) -> None:
|
||||
"""Create the adapter with an injected PostgreSQL URL."""
|
||||
|
||||
self.database_url = database_url
|
||||
|
||||
def prepare_run(
|
||||
self,
|
||||
strategy: Literal["zhixing_b1"],
|
||||
target_trade_date: date,
|
||||
source: SelectionExecutionSource,
|
||||
*,
|
||||
rerun: bool,
|
||||
) -> SelectionRun:
|
||||
"""Atomically claim the business key and create a running attempt.
|
||||
|
||||
The advisory transaction lock protects the small delete-and-create
|
||||
window from duplicate HTTP requests. The long-running calculation is
|
||||
intentionally performed after this transaction is released.
|
||||
"""
|
||||
|
||||
run_id = str(uuid4())
|
||||
key = f"selection:{strategy}:{target_trade_date.isoformat()}"
|
||||
try:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute("SELECT pg_advisory_xact_lock(hashtext(%s))", (key,))
|
||||
existing = connection.execute(
|
||||
"""
|
||||
SELECT id, status
|
||||
FROM selection_run
|
||||
WHERE strategy = %s AND target_trade_date = %s
|
||||
FOR UPDATE
|
||||
""",
|
||||
(strategy, target_trade_date),
|
||||
).fetchone()
|
||||
if existing is not None:
|
||||
existing_status = str(existing[1])
|
||||
if existing_status == "running":
|
||||
raise SelectionRunInProgress(
|
||||
f"selection run is already running for {strategy} at "
|
||||
f"{target_trade_date.isoformat()}"
|
||||
)
|
||||
if not rerun:
|
||||
raise SelectionRerunRequired(
|
||||
f"rerun confirmation is required for {strategy} at "
|
||||
f"{target_trade_date.isoformat()}"
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM selection_run WHERE strategy = %s AND target_trade_date = %s",
|
||||
(strategy, target_trade_date),
|
||||
)
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO selection_run
|
||||
(
|
||||
id, strategy, target_trade_date, market_sync_batch_id,
|
||||
status, target_count, eligible_count, coverage
|
||||
)
|
||||
VALUES (%s, %s, %s, %s, 'running', %s, %s, %s)
|
||||
""",
|
||||
(
|
||||
run_id,
|
||||
strategy,
|
||||
target_trade_date,
|
||||
source.market_sync_batch_id,
|
||||
source.target_count,
|
||||
len(source.stocks),
|
||||
source.coverage,
|
||||
),
|
||||
)
|
||||
except SelectionRunError:
|
||||
raise
|
||||
except psycopg.Error as exc:
|
||||
raise SelectionRunStoreError("failed to prepare selection run") from exc
|
||||
return SelectionRun(
|
||||
id=run_id,
|
||||
strategy=strategy,
|
||||
target_trade_date=target_trade_date,
|
||||
market_sync_batch_id=source.market_sync_batch_id,
|
||||
status="running",
|
||||
target_count=source.target_count,
|
||||
eligible_count=len(source.stocks),
|
||||
evaluated_count=0,
|
||||
selected_stock_count=0,
|
||||
signal_count=0,
|
||||
failed_count=0,
|
||||
coverage=source.coverage,
|
||||
)
|
||||
|
||||
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
|
||||
"""Upsert one stock outcome and all of its independent signal rows."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO selection_run_item
|
||||
(run_id, ts_code, name, status, signal_count, reason)
|
||||
VALUES (%s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (run_id, ts_code) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
status = EXCLUDED.status,
|
||||
signal_count = EXCLUDED.signal_count,
|
||||
reason = EXCLUDED.reason
|
||||
""",
|
||||
(
|
||||
run_id,
|
||||
item.ts_code,
|
||||
item.name,
|
||||
item.status,
|
||||
item.signal_count,
|
||||
item.reason,
|
||||
),
|
||||
)
|
||||
connection.execute(
|
||||
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = %s",
|
||||
(run_id, item.ts_code),
|
||||
)
|
||||
for signal in item.signals:
|
||||
connection.execute(
|
||||
"""
|
||||
INSERT INTO selection_signal
|
||||
(
|
||||
run_id, ts_code, name, target_trade_date, strategy,
|
||||
category, close, details
|
||||
)
|
||||
VALUES (%s, %s, %s, %s, %s, %s, %s, %s)
|
||||
ON CONFLICT (run_id, ts_code, category) DO UPDATE SET
|
||||
name = EXCLUDED.name,
|
||||
close = EXCLUDED.close,
|
||||
details = EXCLUDED.details
|
||||
""",
|
||||
(
|
||||
run_id,
|
||||
signal.ts_code,
|
||||
signal.name,
|
||||
signal.target_trade_date,
|
||||
signal.strategy,
|
||||
signal.category.value,
|
||||
signal.close,
|
||||
Jsonb(dict(signal.details)),
|
||||
),
|
||||
)
|
||||
except psycopg.Error as exc:
|
||||
raise SelectionRunStoreError(
|
||||
f"failed to persist selection item {item.ts_code}"
|
||||
) from exc
|
||||
|
||||
def finish_run(
|
||||
self,
|
||||
run_id: str,
|
||||
status: SelectionRunStatus,
|
||||
*,
|
||||
evaluated_count: int,
|
||||
selected_stock_count: int,
|
||||
signal_count: int,
|
||||
failed_count: int,
|
||||
error_type: str | None = None,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
"""Persist terminal counters and an optional safe batch error."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection, connection.transaction():
|
||||
connection.execute(
|
||||
"""
|
||||
UPDATE selection_run
|
||||
SET status = %s,
|
||||
evaluated_count = %s,
|
||||
selected_stock_count = %s,
|
||||
signal_count = %s,
|
||||
failed_count = %s,
|
||||
error_type = %s,
|
||||
error_message = %s,
|
||||
finished_at = now()
|
||||
WHERE id = %s
|
||||
""",
|
||||
(
|
||||
status,
|
||||
evaluated_count,
|
||||
selected_stock_count,
|
||||
signal_count,
|
||||
failed_count,
|
||||
error_type,
|
||||
_safe_error(error_message),
|
||||
run_id,
|
||||
),
|
||||
)
|
||||
except psycopg.Error as exc:
|
||||
raise SelectionRunStoreError(f"failed to finish selection run {run_id}") from exc
|
||||
|
||||
def get_run(self, run_id: str) -> SelectionRun | None:
|
||||
"""Read one run with its item failures and signal details."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
return self._load_run(connection, run_id)
|
||||
except psycopg.Error as exc:
|
||||
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
|
||||
|
||||
def get_latest_run(
|
||||
self,
|
||||
strategy: Literal["zhixing_b1"],
|
||||
target_trade_date: date | None = None,
|
||||
) -> SelectionRun | None:
|
||||
"""Read the current run for a date or the latest date for a strategy."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
if target_trade_date is None:
|
||||
row = connection.execute(
|
||||
"""
|
||||
SELECT id
|
||||
FROM selection_run
|
||||
WHERE strategy = %s
|
||||
ORDER BY target_trade_date DESC, created_at DESC, id DESC
|
||||
LIMIT 1
|
||||
""",
|
||||
(strategy,),
|
||||
).fetchone()
|
||||
else:
|
||||
row = connection.execute(
|
||||
"""
|
||||
SELECT id
|
||||
FROM selection_run
|
||||
WHERE strategy = %s AND target_trade_date = %s
|
||||
LIMIT 1
|
||||
""",
|
||||
(strategy, target_trade_date),
|
||||
).fetchone()
|
||||
return self._load_run(connection, str(row[0])) if row else None
|
||||
except psycopg.Error as exc:
|
||||
raise SelectionRunStoreError("failed to load latest selection run") from exc
|
||||
|
||||
@staticmethod
|
||||
def _load_run(connection: Any, run_id: str) -> SelectionRun | None:
|
||||
row = connection.execute(
|
||||
"""
|
||||
SELECT
|
||||
id, strategy, target_trade_date, market_sync_batch_id, status,
|
||||
target_count, eligible_count, evaluated_count, selected_stock_count,
|
||||
signal_count, failed_count, coverage, error_type, error_message,
|
||||
created_at, finished_at
|
||||
FROM selection_run
|
||||
WHERE id = %s
|
||||
""",
|
||||
(run_id,),
|
||||
).fetchone()
|
||||
if row is None:
|
||||
return None
|
||||
item_rows = connection.execute(
|
||||
"""
|
||||
SELECT ts_code, name, status, signal_count, reason
|
||||
FROM selection_run_item
|
||||
WHERE run_id = %s
|
||||
ORDER BY ts_code
|
||||
""",
|
||||
(run_id,),
|
||||
).fetchall()
|
||||
signal_rows = connection.execute(
|
||||
"""
|
||||
SELECT
|
||||
ts_code, name, target_trade_date, strategy, category, close, details
|
||||
FROM selection_signal
|
||||
WHERE run_id = %s
|
||||
ORDER BY ts_code, category
|
||||
""",
|
||||
(run_id,),
|
||||
).fetchall()
|
||||
signals = tuple(
|
||||
sorted(
|
||||
(_signal_from_row(cast(tuple[object, ...], value)) for value in signal_rows),
|
||||
key=lambda signal: (
|
||||
signal.ts_code,
|
||||
_SIGNAL_PRIORITY.get(signal.category, len(_SIGNAL_PRIORITY)),
|
||||
),
|
||||
)
|
||||
)
|
||||
signals_by_stock: dict[str, list[SelectionSignal]] = defaultdict(list)
|
||||
for signal in signals:
|
||||
signals_by_stock[signal.ts_code].append(signal)
|
||||
items = tuple(
|
||||
SelectionRunItem(
|
||||
ts_code=str(value[0]),
|
||||
name=str(value[1] or ""),
|
||||
status=cast(
|
||||
Literal[
|
||||
"selected",
|
||||
"no_signal",
|
||||
"insufficient_history",
|
||||
"missing_target_bar",
|
||||
"data_error",
|
||||
],
|
||||
str(value[2]),
|
||||
),
|
||||
signal_count=int(value[3] or 0),
|
||||
reason=str(value[4]) if value[4] is not None else None,
|
||||
signals=tuple(signals_by_stock.get(str(value[0]), ())),
|
||||
)
|
||||
for value in item_rows
|
||||
)
|
||||
return SelectionRun(
|
||||
id=str(row[0]),
|
||||
strategy=cast(Literal["zhixing_b1"], str(row[1])),
|
||||
target_trade_date=_as_date(row[2]),
|
||||
market_sync_batch_id=str(row[3]) if row[3] is not None else None,
|
||||
status=cast(SelectionRunStatus, str(row[4])),
|
||||
target_count=int(row[5]),
|
||||
eligible_count=int(row[6]),
|
||||
evaluated_count=int(row[7]),
|
||||
selected_stock_count=int(row[8]),
|
||||
signal_count=int(row[9]),
|
||||
failed_count=int(row[10]),
|
||||
coverage=Decimal(str(row[11])),
|
||||
error_type=str(row[12]) if row[12] is not None else None,
|
||||
error_message=str(row[13]) if row[13] is not None else None,
|
||||
created_at=cast(datetime | None, row[14]),
|
||||
finished_at=cast(datetime | None, row[15]),
|
||||
items=items,
|
||||
signals=signals,
|
||||
)
|
||||
|
||||
@contextmanager
|
||||
def _connection(self) -> Generator[Any, None, None]:
|
||||
"""Translate psycopg failures without exposing driver details."""
|
||||
|
||||
try:
|
||||
with psycopg.connect(self.database_url) as connection:
|
||||
yield connection
|
||||
except psycopg.Error as exc:
|
||||
raise SelectionRunStoreError("selection database operation failed") from exc
|
||||
|
||||
|
||||
def _signal_from_row(row: tuple[object, ...]) -> SelectionSignal:
|
||||
"""Map a persisted signal row back to the domain signal model."""
|
||||
|
||||
return SelectionSignal(
|
||||
ts_code=str(row[0]),
|
||||
name=str(row[1] or ""),
|
||||
target_trade_date=_as_date(row[2]),
|
||||
strategy=cast(Literal["zhixing_b1"], str(row[3])),
|
||||
category=ZhixingB1Category(str(row[4])),
|
||||
close=float(str(row[5])),
|
||||
details=_details(row[6]),
|
||||
)
|
||||
|
||||
|
||||
def _details(value: object) -> dict[str, float | str | None]:
|
||||
"""Normalize JSONB details into the domain's scalar-only mapping."""
|
||||
|
||||
if isinstance(value, str):
|
||||
try:
|
||||
value = json.loads(value)
|
||||
except json.JSONDecodeError:
|
||||
return {}
|
||||
if not isinstance(value, Mapping):
|
||||
return {}
|
||||
values = cast(Mapping[object, object], value)
|
||||
details: dict[str, float | str | None] = {}
|
||||
for key, item in values.items():
|
||||
if item is None or isinstance(item, str):
|
||||
details[str(key)] = item
|
||||
elif isinstance(item, (int, float)) and not isinstance(item, bool):
|
||||
details[str(key)] = float(item)
|
||||
return details
|
||||
|
||||
|
||||
def _as_date(value: object) -> date:
|
||||
if isinstance(value, datetime):
|
||||
return value.date()
|
||||
if isinstance(value, date):
|
||||
return value
|
||||
return date.fromisoformat(str(value)[:10])
|
||||
|
||||
|
||||
def _safe_error(message: str | None) -> str | None:
|
||||
if message is None:
|
||||
return None
|
||||
return " ".join(message.split())[:500]
|
||||
@@ -0,0 +1,259 @@
|
||||
"""HTTP presentation for persisted strategy execution results."""
|
||||
|
||||
from datetime import date, datetime
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, status
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhixing_server.bootstrap.config import Settings, get_settings
|
||||
from zhixing_server.modules.selection.application.run import (
|
||||
RunZhixingB1,
|
||||
)
|
||||
from zhixing_server.modules.selection.domain.runs import (
|
||||
SelectionRerunRequired,
|
||||
SelectionRun,
|
||||
SelectionRunInProgress,
|
||||
SelectionRunStoreError,
|
||||
)
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
PostgresMarketDataReader,
|
||||
SelectionMarketDataNotReady,
|
||||
SelectionReaderError,
|
||||
)
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_runs import (
|
||||
PostgresSelectionRunRepository,
|
||||
)
|
||||
|
||||
selection_router = APIRouter()
|
||||
|
||||
StrategyValue = Literal["zhixing_b1"]
|
||||
SelectionStatusValue = Literal[
|
||||
"no_data",
|
||||
"running",
|
||||
"success",
|
||||
"partial_success",
|
||||
"failed",
|
||||
]
|
||||
|
||||
|
||||
class SelectionRunRequest(BaseModel):
|
||||
"""Input contract for one initial run or explicit rerun."""
|
||||
|
||||
strategy: StrategyValue
|
||||
target_trade_date: date
|
||||
rerun: bool = False
|
||||
|
||||
|
||||
class SelectionRunAcceptedResponse(BaseModel):
|
||||
"""Small response returned before the background evaluation completes."""
|
||||
|
||||
run_id: str
|
||||
strategy: StrategyValue
|
||||
target_trade_date: date
|
||||
status: Literal["running"]
|
||||
|
||||
|
||||
class SelectionSignalResponse(BaseModel):
|
||||
"""One persisted independent sub-signal in the public result contract."""
|
||||
|
||||
ts_code: str
|
||||
name: str
|
||||
target_trade_date: date
|
||||
strategy: StrategyValue
|
||||
category: str
|
||||
close: float
|
||||
details: dict[str, float | str | None]
|
||||
|
||||
|
||||
class SelectionFailureResponse(BaseModel):
|
||||
"""One stock that could not produce a complete evaluation."""
|
||||
|
||||
ts_code: str
|
||||
name: str
|
||||
status: str
|
||||
reason: str | None
|
||||
|
||||
|
||||
def _empty_failures() -> list[SelectionFailureResponse]:
|
||||
"""Create a typed default list for Pydantic's strict checker."""
|
||||
|
||||
return []
|
||||
|
||||
|
||||
def _empty_signals() -> list[SelectionSignalResponse]:
|
||||
"""Create a typed default list for Pydantic's strict checker."""
|
||||
|
||||
return []
|
||||
|
||||
|
||||
class SelectionResultsResponse(BaseModel):
|
||||
"""Batch summary and materialized signals consumed by the Web feature."""
|
||||
|
||||
strategy: StrategyValue
|
||||
target_trade_date: date | None
|
||||
run_id: str | None
|
||||
market_sync_batch_id: str | None
|
||||
status: SelectionStatusValue
|
||||
target_count: int = Field(default=0, ge=0)
|
||||
eligible_count: int = Field(default=0, ge=0)
|
||||
evaluated_count: int = Field(default=0, ge=0)
|
||||
selected_stock_count: int = Field(default=0, ge=0)
|
||||
signal_count: int = Field(default=0, ge=0)
|
||||
failed_count: int = Field(default=0, ge=0)
|
||||
coverage: float = Field(default=0, ge=0, le=1)
|
||||
error_type: str | None = None
|
||||
error_message: str | None = None
|
||||
created_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
failures: list[SelectionFailureResponse] = Field(default_factory=_empty_failures)
|
||||
signals: list[SelectionSignalResponse] = Field(default_factory=_empty_signals)
|
||||
|
||||
|
||||
def get_selection_service(
|
||||
settings: Annotated[Settings, Depends(get_settings)],
|
||||
) -> RunZhixingB1:
|
||||
"""Build one request-scoped selection application service."""
|
||||
|
||||
reader = PostgresMarketDataReader(settings)
|
||||
store = PostgresSelectionRunRepository(settings.database_url)
|
||||
return RunZhixingB1(reader, store)
|
||||
|
||||
|
||||
@selection_router.post(
|
||||
"/runs",
|
||||
response_model=SelectionRunAcceptedResponse,
|
||||
status_code=status.HTTP_202_ACCEPTED,
|
||||
)
|
||||
def trigger_selection_run(
|
||||
request: SelectionRunRequest,
|
||||
background_tasks: BackgroundTasks,
|
||||
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
||||
) -> SelectionRunAcceptedResponse:
|
||||
"""Claim a run and schedule its whole-universe evaluation."""
|
||||
|
||||
try:
|
||||
prepared = service.prepare(
|
||||
request.strategy,
|
||||
request.target_trade_date,
|
||||
rerun=request.rerun,
|
||||
)
|
||||
except SelectionRunInProgress as exc:
|
||||
raise _http_error(409, "run_in_progress", str(exc)) from exc
|
||||
except SelectionRerunRequired as exc:
|
||||
raise _http_error(409, "rerun_confirmation_required", str(exc)) from exc
|
||||
except SelectionMarketDataNotReady as exc:
|
||||
raise _http_error(422, "market_data_not_ready", str(exc)) from exc
|
||||
except (SelectionReaderError, SelectionRunStoreError) as exc:
|
||||
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
||||
|
||||
background_tasks.add_task(service.execute, prepared)
|
||||
return SelectionRunAcceptedResponse(
|
||||
run_id=prepared.run.id,
|
||||
strategy=prepared.run.strategy,
|
||||
target_trade_date=prepared.run.target_trade_date,
|
||||
status="running",
|
||||
)
|
||||
|
||||
|
||||
@selection_router.get("/runs/{run_id}", response_model=SelectionResultsResponse)
|
||||
def get_selection_run(
|
||||
run_id: str,
|
||||
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
||||
) -> SelectionResultsResponse:
|
||||
"""Return one run for asynchronous polling."""
|
||||
|
||||
try:
|
||||
run = service.get_run(run_id)
|
||||
except SelectionRunStoreError as exc:
|
||||
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
||||
if run is None:
|
||||
raise _http_error(404, "run_not_found", f"selection run not found: {run_id}")
|
||||
return _run_response(run)
|
||||
|
||||
|
||||
@selection_router.get("/results", response_model=SelectionResultsResponse)
|
||||
def get_selection_results(
|
||||
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
||||
strategy: StrategyValue = "zhixing_b1",
|
||||
target_trade_date: date | None = None,
|
||||
) -> SelectionResultsResponse:
|
||||
"""Return the current persisted result for a strategy and optional date."""
|
||||
|
||||
try:
|
||||
run = service.get_latest(strategy, target_trade_date)
|
||||
except SelectionRunStoreError as exc:
|
||||
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
||||
if run is None:
|
||||
return SelectionResultsResponse(
|
||||
strategy=strategy,
|
||||
target_trade_date=target_trade_date,
|
||||
run_id=None,
|
||||
market_sync_batch_id=None,
|
||||
status="no_data",
|
||||
coverage=0,
|
||||
)
|
||||
return _run_response(run)
|
||||
|
||||
|
||||
def _run_response(run: SelectionRun) -> SelectionResultsResponse:
|
||||
"""Translate a domain run without exposing storage-specific fields."""
|
||||
|
||||
return SelectionResultsResponse(
|
||||
strategy=run.strategy,
|
||||
target_trade_date=run.target_trade_date,
|
||||
run_id=run.id,
|
||||
market_sync_batch_id=run.market_sync_batch_id,
|
||||
status=run.status,
|
||||
target_count=run.target_count,
|
||||
eligible_count=run.eligible_count,
|
||||
evaluated_count=run.evaluated_count,
|
||||
selected_stock_count=run.selected_stock_count,
|
||||
signal_count=run.signal_count,
|
||||
failed_count=run.failed_count,
|
||||
coverage=float(run.coverage),
|
||||
error_type=run.error_type,
|
||||
error_message=run.error_message,
|
||||
created_at=run.created_at,
|
||||
finished_at=run.finished_at,
|
||||
failures=[
|
||||
SelectionFailureResponse(
|
||||
ts_code=item.ts_code,
|
||||
name=item.name,
|
||||
status=item.status,
|
||||
reason=item.reason,
|
||||
)
|
||||
for item in run.items
|
||||
if item.status in {"insufficient_history", "missing_target_bar", "data_error"}
|
||||
],
|
||||
signals=[
|
||||
SelectionSignalResponse(
|
||||
ts_code=signal.ts_code,
|
||||
name=signal.name,
|
||||
target_trade_date=signal.target_trade_date,
|
||||
strategy=signal.strategy,
|
||||
category=signal.category.value,
|
||||
close=signal.close,
|
||||
details=dict(signal.details),
|
||||
)
|
||||
for signal in run.signals
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
def _http_error(code: int, error_type: str, message: str) -> HTTPException:
|
||||
"""Create the project's explicit, safe error envelope."""
|
||||
|
||||
return HTTPException(
|
||||
status_code=code,
|
||||
detail={"code": error_type, "message": message},
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"SelectionResultsResponse",
|
||||
"SelectionRunAcceptedResponse",
|
||||
"SelectionRunRequest",
|
||||
"get_selection_service",
|
||||
"selection_router",
|
||||
]
|
||||
@@ -33,6 +33,9 @@ def test_postgres_migration_creates_market_data_contract(
|
||||
"market_daily_basic",
|
||||
"market_sync_batch",
|
||||
"market_sync_item",
|
||||
"selection_run",
|
||||
"selection_run_item",
|
||||
"selection_signal",
|
||||
} <= tables
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
@@ -0,0 +1,207 @@
|
||||
"""HTTP contracts for triggering and querying persisted selection runs."""
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from zhixing_server.bootstrap.app import create_app
|
||||
from zhixing_server.modules.selection.application.run import PreparedSelectionRun
|
||||
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
|
||||
from zhixing_server.modules.selection.domain.runs import (
|
||||
SelectionExecutionSource,
|
||||
SelectionRerunRequired,
|
||||
SelectionRun,
|
||||
SelectionRunInProgress,
|
||||
SelectionStock,
|
||||
)
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
SelectionMarketDataNotReady,
|
||||
)
|
||||
from zhixing_server.modules.selection.presentation.http import get_selection_service
|
||||
|
||||
TARGET = date(2026, 8, 8)
|
||||
|
||||
|
||||
class FakeSelectionService:
|
||||
def __init__(self, run: SelectionRun | None = None) -> None:
|
||||
self.run = run
|
||||
self.executed = False
|
||||
self.mode = "ok"
|
||||
|
||||
def prepare(
|
||||
self,
|
||||
strategy: str,
|
||||
target_trade_date: date,
|
||||
*,
|
||||
rerun: bool,
|
||||
) -> PreparedSelectionRun:
|
||||
if self.mode == "in_progress":
|
||||
raise SelectionRunInProgress("already running")
|
||||
if self.mode == "rerun_required":
|
||||
raise SelectionRerunRequired("confirm rerun")
|
||||
if self.mode == "market_data_not_ready":
|
||||
raise SelectionMarketDataNotReady("market data is not ready")
|
||||
run = self.run or _run("run-http", "running")
|
||||
return PreparedSelectionRun(
|
||||
run=run,
|
||||
source=SelectionExecutionSource(
|
||||
market_sync_batch_id="market-run-1",
|
||||
target_trade_date=target_trade_date,
|
||||
target_count=1,
|
||||
valid_count=1,
|
||||
coverage=Decimal("1"),
|
||||
stocks=(SelectionStock("000001.SZ", "平安银行"),),
|
||||
),
|
||||
)
|
||||
|
||||
def execute(self, prepared: PreparedSelectionRun) -> None:
|
||||
self.executed = True
|
||||
|
||||
def get_run(self, run_id: str) -> SelectionRun | None:
|
||||
return self.run if self.run and self.run.id == run_id else None
|
||||
|
||||
def get_latest(
|
||||
self,
|
||||
strategy: str,
|
||||
target_trade_date: date | None = None,
|
||||
) -> SelectionRun | None:
|
||||
if self.run is None:
|
||||
return None
|
||||
if target_trade_date is not None and self.run.target_trade_date != target_trade_date:
|
||||
return None
|
||||
return self.run
|
||||
|
||||
|
||||
def _run(run_id: str, status: str) -> SelectionRun:
|
||||
signal = SelectionSignal(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
target_trade_date=TARGET,
|
||||
strategy="zhixing_b1",
|
||||
category=ZhixingB1Category.ORIGINAL_B1,
|
||||
close=10.5,
|
||||
details={"j": 12.0},
|
||||
)
|
||||
from zhixing_server.modules.selection.domain.runs import SelectionRunItem
|
||||
|
||||
return SelectionRun(
|
||||
id=run_id,
|
||||
strategy="zhixing_b1",
|
||||
target_trade_date=TARGET,
|
||||
market_sync_batch_id="market-run-1",
|
||||
status=status, # type: ignore[arg-type]
|
||||
target_count=1,
|
||||
eligible_count=1,
|
||||
evaluated_count=1,
|
||||
selected_stock_count=1,
|
||||
signal_count=1,
|
||||
failed_count=0,
|
||||
coverage=Decimal("1"),
|
||||
items=(
|
||||
SelectionRunItem(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
status="selected",
|
||||
signal_count=1,
|
||||
signals=(signal,),
|
||||
),
|
||||
),
|
||||
signals=(signal,),
|
||||
)
|
||||
|
||||
|
||||
def _client(service: FakeSelectionService) -> TestClient:
|
||||
app = create_app()
|
||||
app.dependency_overrides[get_selection_service] = lambda: service
|
||||
return TestClient(app)
|
||||
|
||||
|
||||
def test_trigger_returns_accepted_run_and_schedules_execution() -> None:
|
||||
service = FakeSelectionService()
|
||||
|
||||
response = _client(service).post(
|
||||
"/api/v1/selection/runs",
|
||||
json={
|
||||
"strategy": "zhixing_b1",
|
||||
"target_trade_date": "2026-08-08",
|
||||
"rerun": False,
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 202
|
||||
assert response.json()["status"] == "running"
|
||||
assert response.json()["target_trade_date"] == "2026-08-08"
|
||||
assert service.executed is True
|
||||
|
||||
|
||||
def test_trigger_requires_explicit_rerun_confirmation() -> None:
|
||||
service = FakeSelectionService()
|
||||
service.mode = "rerun_required"
|
||||
|
||||
response = _client(service).post(
|
||||
"/api/v1/selection/runs",
|
||||
json={"strategy": "zhixing_b1", "target_trade_date": "2026-08-08"},
|
||||
)
|
||||
|
||||
assert response.status_code == 409
|
||||
assert response.json()["detail"]["code"] == "rerun_confirmation_required"
|
||||
|
||||
|
||||
def test_trigger_rejects_a_duplicate_running_request() -> None:
|
||||
service = FakeSelectionService()
|
||||
service.mode = "in_progress"
|
||||
|
||||
response = _client(service).post(
|
||||
"/api/v1/selection/runs",
|
||||
json={"strategy": "zhixing_b1", "target_trade_date": "2026-08-08"},
|
||||
)
|
||||
|
||||
assert response.status_code == 409
|
||||
assert response.json()["detail"]["code"] == "run_in_progress"
|
||||
|
||||
|
||||
def test_trigger_rejects_unqualified_market_data() -> None:
|
||||
service = FakeSelectionService()
|
||||
service.mode = "market_data_not_ready"
|
||||
|
||||
response = _client(service).post(
|
||||
"/api/v1/selection/runs",
|
||||
json={"strategy": "zhixing_b1", "target_trade_date": "2026-08-08"},
|
||||
)
|
||||
|
||||
assert response.status_code == 422
|
||||
assert response.json()["detail"]["code"] == "market_data_not_ready"
|
||||
|
||||
|
||||
def test_query_returns_no_data_without_fabricating_a_result() -> None:
|
||||
response = _client(FakeSelectionService()).get(
|
||||
"/api/v1/selection/results?strategy=zhixing_b1&target_trade_date=2026-08-08"
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["status"] == "no_data"
|
||||
assert response.json()["signals"] == []
|
||||
|
||||
|
||||
def test_query_returns_persisted_signal_details() -> None:
|
||||
response = _client(FakeSelectionService(_run("run-http", "success"))).get(
|
||||
"/api/v1/selection/results?strategy=zhixing_b1&target_trade_date=2026-08-08"
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
body = response.json()
|
||||
assert body["run_id"] == "run-http"
|
||||
assert body["signal_count"] == 1
|
||||
assert body["signals"][0]["category"] == "zhixing_b1_original_b1"
|
||||
assert body["signals"][0]["details"] == {"j": 12.0}
|
||||
|
||||
|
||||
def test_run_polling_returns_the_persisted_terminal_result() -> None:
|
||||
response = _client(FakeSelectionService(_run("run-http", "success"))).get(
|
||||
"/api/v1/selection/runs/run-http"
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert response.json()["status"] == "success"
|
||||
assert response.json()["signals"][0]["category"] == "zhixing_b1_original_b1"
|
||||
@@ -1,6 +1,7 @@
|
||||
"""PostgreSQL reader contract tests using a fake connection."""
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import cast
|
||||
|
||||
import psycopg
|
||||
@@ -8,6 +9,7 @@ import pytest
|
||||
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_reader import (
|
||||
PostgresMarketDataReader,
|
||||
SelectionMarketDataNotReady,
|
||||
)
|
||||
|
||||
|
||||
@@ -82,3 +84,77 @@ def test_reader_parameterizes_target_and_maps_left_join(monkeypatch: pytest.Monk
|
||||
assert connection.parameters == ("000001.SZ", date(2024, 1, 3))
|
||||
assert "source_adj = 'qfq'" in cast(str, connection.query)
|
||||
assert "trade_date <= %s" in cast(str, connection.query)
|
||||
|
||||
|
||||
class SourceConnection:
|
||||
def __init__(self, source_row: tuple[object, ...] | None) -> None:
|
||||
self.source_row = source_row
|
||||
self.queries: list[tuple[str, tuple[object, ...]]] = []
|
||||
|
||||
def __enter__(self) -> "SourceConnection":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
def execute(self, query: str, parameters: tuple[object, ...]) -> "SourceResult":
|
||||
self.queries.append((query, parameters))
|
||||
if "FROM market_sync_batch" in query:
|
||||
return SourceResult(row=self.source_row)
|
||||
return SourceResult(rows=[("000001.SZ", "平安银行")])
|
||||
|
||||
|
||||
class SourceResult:
|
||||
def __init__(
|
||||
self,
|
||||
row: tuple[object, ...] | None = None,
|
||||
rows: list[tuple[object, ...]] | None = None,
|
||||
) -> None:
|
||||
self.row = row
|
||||
self.rows = rows or []
|
||||
|
||||
def fetchone(self) -> tuple[object, ...] | None:
|
||||
return self.row
|
||||
|
||||
def fetchall(self) -> list[tuple[object, ...]]:
|
||||
return self.rows
|
||||
|
||||
|
||||
def test_reader_loads_only_eligible_stocks_from_finished_market_batch(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
) -> None:
|
||||
connection = SourceConnection(("market-run-1", 2, 2, Decimal("1")))
|
||||
|
||||
def connect(database_url: str) -> SourceConnection:
|
||||
assert database_url == "postgresql://test"
|
||||
return connection
|
||||
|
||||
monkeypatch.setattr(psycopg, "connect", connect)
|
||||
|
||||
source = PostgresMarketDataReader("postgresql://test").load_execution_source(
|
||||
"zhixing_b1",
|
||||
date(2026, 8, 8),
|
||||
)
|
||||
|
||||
assert source.market_sync_batch_id == "market-run-1"
|
||||
assert source.target_count == 2
|
||||
assert source.coverage == Decimal("1")
|
||||
assert source.stocks[0].ts_code == "000001.SZ"
|
||||
assert "strategy_eligible = true" in connection.queries[0][0]
|
||||
assert "source_adj = 'qfq'" in connection.queries[1][0]
|
||||
|
||||
|
||||
def test_reader_rejects_date_without_eligible_market_batch(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = SourceConnection(None)
|
||||
|
||||
def connect(database_url: str) -> SourceConnection:
|
||||
assert database_url == "postgresql://test"
|
||||
return connection
|
||||
|
||||
monkeypatch.setattr(psycopg, "connect", connect)
|
||||
|
||||
with pytest.raises(SelectionMarketDataNotReady):
|
||||
PostgresMarketDataReader("postgresql://test").load_execution_source(
|
||||
"zhixing_b1",
|
||||
date(2026, 8, 8),
|
||||
)
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
"""Persistence transaction tests for selection runs."""
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
|
||||
import psycopg
|
||||
import pytest
|
||||
from psycopg.types.json import Jsonb
|
||||
|
||||
from zhixing_server.modules.selection.domain.models import SelectionSignal, ZhixingB1Category
|
||||
from zhixing_server.modules.selection.domain.runs import (
|
||||
SelectionExecutionSource,
|
||||
SelectionRerunRequired,
|
||||
SelectionRunInProgress,
|
||||
SelectionRunItem,
|
||||
SelectionStock,
|
||||
)
|
||||
from zhixing_server.modules.selection.domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
|
||||
from zhixing_server.modules.selection.infrastructure.postgres_runs import (
|
||||
PostgresSelectionRunRepository,
|
||||
)
|
||||
|
||||
TARGET = date(2026, 8, 8)
|
||||
|
||||
|
||||
class FakeResult:
|
||||
def __init__(self, row: tuple[object, ...] | None = None) -> None:
|
||||
self.row = row
|
||||
|
||||
def fetchone(self) -> tuple[object, ...] | None:
|
||||
return self.row
|
||||
|
||||
|
||||
class FakeTransaction:
|
||||
def __enter__(self) -> "FakeTransaction":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(self, existing: tuple[object, ...] | None) -> None:
|
||||
self.existing = existing
|
||||
self.statements: list[tuple[str, tuple[object, ...]]] = []
|
||||
|
||||
def __enter__(self) -> "FakeConnection":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
def transaction(self) -> FakeTransaction:
|
||||
return FakeTransaction()
|
||||
|
||||
def execute(self, query: str, parameters: tuple[object, ...]) -> FakeResult:
|
||||
self.statements.append((query, parameters))
|
||||
if "SELECT id, status" in query:
|
||||
return FakeResult(self.existing)
|
||||
return FakeResult()
|
||||
|
||||
|
||||
def _source() -> SelectionExecutionSource:
|
||||
return SelectionExecutionSource(
|
||||
market_sync_batch_id="market-run-1",
|
||||
target_trade_date=TARGET,
|
||||
target_count=1,
|
||||
valid_count=1,
|
||||
coverage=Decimal("1"),
|
||||
stocks=(SelectionStock("000001.SZ", "平安银行"),),
|
||||
)
|
||||
|
||||
|
||||
def _repository(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
connection: FakeConnection,
|
||||
) -> PostgresSelectionRunRepository:
|
||||
def connect(database_url: str) -> FakeConnection:
|
||||
assert database_url == "postgresql://test"
|
||||
return connection
|
||||
|
||||
monkeypatch.setattr(psycopg, "connect", connect)
|
||||
return PostgresSelectionRunRepository("postgresql://test")
|
||||
|
||||
|
||||
def test_prepare_claims_new_business_key(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(None)
|
||||
run = _repository(monkeypatch, connection).prepare_run(
|
||||
"zhixing_b1",
|
||||
TARGET,
|
||||
_source(),
|
||||
rerun=False,
|
||||
)
|
||||
|
||||
assert run.status == "running"
|
||||
assert run.target_trade_date == TARGET
|
||||
assert any("INSERT INTO selection_run" in query for query, _ in connection.statements)
|
||||
assert not any("DELETE FROM selection_run" in query for query, _ in connection.statements)
|
||||
|
||||
|
||||
def test_prepare_requires_confirmation_for_terminal_run(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(("old-run", "success"))
|
||||
repository = _repository(monkeypatch, connection)
|
||||
|
||||
with pytest.raises(SelectionRerunRequired):
|
||||
repository.prepare_run("zhixing_b1", TARGET, _source(), rerun=False)
|
||||
|
||||
assert not any("DELETE FROM selection_run" in query for query, _ in connection.statements)
|
||||
|
||||
|
||||
def test_prepare_rejects_duplicate_running_run(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(("old-run", "running"))
|
||||
repository = _repository(monkeypatch, connection)
|
||||
|
||||
with pytest.raises(SelectionRunInProgress):
|
||||
repository.prepare_run("zhixing_b1", TARGET, _source(), rerun=True)
|
||||
|
||||
|
||||
def test_prepare_rerun_deletes_old_result_before_insert(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(("old-run", "failed"))
|
||||
run = _repository(monkeypatch, connection).prepare_run(
|
||||
"zhixing_b1",
|
||||
TARGET,
|
||||
_source(),
|
||||
rerun=True,
|
||||
)
|
||||
|
||||
statements = [query for query, _ in connection.statements]
|
||||
assert "DELETE FROM selection_run WHERE strategy = %s AND target_trade_date = %s" in statements
|
||||
assert run.id != "old-run"
|
||||
|
||||
|
||||
def test_record_item_persists_independent_signals_as_jsonb(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = FakeConnection(None)
|
||||
repository = _repository(monkeypatch, connection)
|
||||
signal = SelectionSignal(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
target_trade_date=TARGET,
|
||||
strategy="zhixing_b1",
|
||||
category=ZhixingB1Category.ORIGINAL_B1,
|
||||
close=10.5,
|
||||
details={"j": 12.0},
|
||||
)
|
||||
|
||||
repository.record_item(
|
||||
"run-1",
|
||||
SelectionRunItem(
|
||||
ts_code="000001.SZ",
|
||||
name="平安银行",
|
||||
status="selected",
|
||||
signal_count=1,
|
||||
signals=(signal,),
|
||||
),
|
||||
)
|
||||
|
||||
signal_insert = next(
|
||||
parameters
|
||||
for query, parameters in connection.statements
|
||||
if "INSERT INTO selection_signal" in query
|
||||
)
|
||||
assert isinstance(signal_insert[-1], Jsonb)
|
||||
|
||||
|
||||
class LoadConnection:
|
||||
def __enter__(self) -> "LoadConnection":
|
||||
return self
|
||||
|
||||
def __exit__(self, *args: object) -> None:
|
||||
return None
|
||||
|
||||
def execute(self, query: str, parameters: tuple[object, ...]) -> "LoadResult":
|
||||
if "FROM selection_run\n" in query:
|
||||
return LoadResult(
|
||||
row=(
|
||||
"run-1",
|
||||
"zhixing_b1",
|
||||
TARGET,
|
||||
"market-run-1",
|
||||
"success",
|
||||
1,
|
||||
1,
|
||||
1,
|
||||
1,
|
||||
2,
|
||||
0,
|
||||
Decimal("1"),
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
None,
|
||||
)
|
||||
)
|
||||
if "FROM selection_run_item" in query:
|
||||
return LoadResult(rows=[("000001.SZ", "平安银行", "selected", 2, None)])
|
||||
return LoadResult(
|
||||
rows=[
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
TARGET,
|
||||
"zhixing_b1",
|
||||
ZHIXING_B1_SIGNAL_ORDER[-1].value,
|
||||
Decimal("10.5"),
|
||||
{},
|
||||
),
|
||||
(
|
||||
"000001.SZ",
|
||||
"平安银行",
|
||||
TARGET,
|
||||
"zhixing_b1",
|
||||
ZHIXING_B1_SIGNAL_ORDER[0].value,
|
||||
Decimal("10.5"),
|
||||
{},
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
|
||||
class LoadResult:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
row: tuple[object, ...] | None = None,
|
||||
rows: list[tuple[object, ...]] | None = None,
|
||||
) -> None:
|
||||
self.row = row
|
||||
self.rows = rows or []
|
||||
|
||||
def fetchone(self) -> tuple[object, ...] | None:
|
||||
return self.row
|
||||
|
||||
def fetchall(self) -> list[tuple[object, ...]]:
|
||||
return self.rows
|
||||
|
||||
|
||||
def test_get_run_orders_signals_by_formula_priority(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
connection = LoadConnection()
|
||||
|
||||
def connect(database_url: str) -> LoadConnection:
|
||||
assert database_url == "postgresql://test"
|
||||
return connection
|
||||
|
||||
monkeypatch.setattr(psycopg, "connect", connect)
|
||||
run = PostgresSelectionRunRepository("postgresql://test").get_run("run-1")
|
||||
|
||||
assert run is not None
|
||||
assert [signal.category for signal in run.signals] == [
|
||||
ZHIXING_B1_SIGNAL_ORDER[0],
|
||||
ZHIXING_B1_SIGNAL_ORDER[-1],
|
||||
]
|
||||
@@ -0,0 +1,228 @@
|
||||
"""Application tests for persisted whole-universe selection runs."""
|
||||
|
||||
from datetime import date
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
|
||||
from zhixing_server.modules.selection.application.run import (
|
||||
PreparedSelectionRun,
|
||||
RunZhixingB1,
|
||||
)
|
||||
from zhixing_server.modules.selection.domain.models import (
|
||||
SelectionEvaluation,
|
||||
SelectionSignal,
|
||||
StockHistory,
|
||||
)
|
||||
from zhixing_server.modules.selection.domain.runs import (
|
||||
SelectionExecutionSource,
|
||||
SelectionRun,
|
||||
SelectionRunItem,
|
||||
SelectionRunStatus,
|
||||
SelectionStock,
|
||||
)
|
||||
|
||||
TARGET = date(2026, 8, 8)
|
||||
|
||||
|
||||
class FakeReader:
|
||||
def __init__(self, source: SelectionExecutionSource) -> None:
|
||||
self.source = source
|
||||
|
||||
def load_execution_source(
|
||||
self,
|
||||
strategy: str,
|
||||
target_trade_date: date,
|
||||
) -> SelectionExecutionSource:
|
||||
assert strategy == "zhixing_b1"
|
||||
assert target_trade_date == TARGET
|
||||
return self.source
|
||||
|
||||
def load_history(self, ts_code: str, target_trade_date: date) -> StockHistory:
|
||||
raise AssertionError("the fake evaluator should be used")
|
||||
|
||||
|
||||
class FakeStore:
|
||||
def __init__(self) -> None:
|
||||
self.items: list[SelectionRunItem] = []
|
||||
self.finished: tuple[str, SelectionRunStatus, dict[str, object]] | None = None
|
||||
|
||||
def prepare_run(
|
||||
self,
|
||||
strategy: Literal["zhixing_b1"],
|
||||
target_trade_date: date,
|
||||
source: SelectionExecutionSource,
|
||||
*,
|
||||
rerun: bool,
|
||||
) -> SelectionRun:
|
||||
assert strategy == "zhixing_b1"
|
||||
assert target_trade_date == TARGET
|
||||
assert rerun is False
|
||||
return SelectionRun(
|
||||
id="run-1",
|
||||
strategy="zhixing_b1",
|
||||
target_trade_date=TARGET,
|
||||
market_sync_batch_id=source.market_sync_batch_id,
|
||||
status="running",
|
||||
target_count=source.target_count,
|
||||
eligible_count=len(source.stocks),
|
||||
evaluated_count=0,
|
||||
selected_stock_count=0,
|
||||
signal_count=0,
|
||||
failed_count=0,
|
||||
coverage=source.coverage,
|
||||
)
|
||||
|
||||
def record_item(self, run_id: str, item: SelectionRunItem) -> None:
|
||||
assert run_id == "run-1"
|
||||
self.items.append(item)
|
||||
|
||||
def finish_run(
|
||||
self,
|
||||
run_id: str,
|
||||
status: SelectionRunStatus,
|
||||
*,
|
||||
evaluated_count: int,
|
||||
selected_stock_count: int,
|
||||
signal_count: int,
|
||||
failed_count: int,
|
||||
error_type: str | None = None,
|
||||
error_message: str | None = None,
|
||||
) -> None:
|
||||
kwargs: dict[str, object] = {
|
||||
"evaluated_count": evaluated_count,
|
||||
"selected_stock_count": selected_stock_count,
|
||||
"signal_count": signal_count,
|
||||
"failed_count": failed_count,
|
||||
}
|
||||
if error_type is not None:
|
||||
kwargs["error_type"] = error_type
|
||||
if error_message is not None:
|
||||
kwargs["error_message"] = error_message
|
||||
self.finished = (run_id, status, kwargs)
|
||||
|
||||
def get_run(self, run_id: str):
|
||||
return None
|
||||
|
||||
def get_latest_run(self, strategy: str, target_trade_date: date | None = None):
|
||||
return None
|
||||
|
||||
|
||||
class FakeEvaluator:
|
||||
def __init__(self, results: dict[str, SelectionEvaluation]) -> None:
|
||||
self.results = results
|
||||
|
||||
def execute(self, ts_code: str, target_trade_date: date) -> SelectionEvaluation:
|
||||
return self.results[ts_code]
|
||||
|
||||
|
||||
class RaisingEvaluator:
|
||||
def execute(self, ts_code: str, target_trade_date: date) -> SelectionEvaluation:
|
||||
if ts_code == "600000.SH":
|
||||
raise RuntimeError("temporary evaluator failure")
|
||||
return SelectionEvaluation(ts_code, target_trade_date, "no_signal")
|
||||
|
||||
|
||||
def _source() -> SelectionExecutionSource:
|
||||
return SelectionExecutionSource(
|
||||
market_sync_batch_id="market-run-1",
|
||||
target_trade_date=TARGET,
|
||||
target_count=2,
|
||||
valid_count=2,
|
||||
coverage=Decimal("1"),
|
||||
stocks=(
|
||||
SelectionStock("000001.SZ", "平安银行"),
|
||||
SelectionStock("600000.SH", "浦发银行"),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _signal(ts_code: str, category: str) -> SelectionSignal:
|
||||
from zhixing_server.modules.selection.domain.models import ZhixingB1Category
|
||||
|
||||
return SelectionSignal(
|
||||
ts_code=ts_code,
|
||||
name="平安银行",
|
||||
target_trade_date=TARGET,
|
||||
strategy="zhixing_b1",
|
||||
category=ZhixingB1Category(category),
|
||||
close=10.5,
|
||||
details={"j": 12.0},
|
||||
)
|
||||
|
||||
|
||||
def test_prepare_captures_market_source_and_execute_persists_all_categories() -> None:
|
||||
source = _source()
|
||||
store = FakeStore()
|
||||
evaluator = FakeEvaluator(
|
||||
{
|
||||
"000001.SZ": SelectionEvaluation(
|
||||
"000001.SZ",
|
||||
TARGET,
|
||||
"selected",
|
||||
signals=(
|
||||
_signal("000001.SZ", "zhixing_b1_original_b1"),
|
||||
_signal("000001.SZ", "zhixing_b1_pullback_white"),
|
||||
),
|
||||
),
|
||||
"600000.SH": SelectionEvaluation(
|
||||
"600000.SH",
|
||||
TARGET,
|
||||
"no_signal",
|
||||
reason="no category matched",
|
||||
),
|
||||
}
|
||||
)
|
||||
service = RunZhixingB1(FakeReader(source), store, evaluator)
|
||||
|
||||
prepared = service.prepare("zhixing_b1", TARGET, rerun=False)
|
||||
assert isinstance(prepared, PreparedSelectionRun)
|
||||
service.execute(prepared)
|
||||
|
||||
assert [item.status for item in store.items] == ["selected", "no_signal"]
|
||||
assert store.items[0].signal_count == 2
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "success")
|
||||
assert store.finished[2] == {
|
||||
"evaluated_count": 2,
|
||||
"selected_stock_count": 1,
|
||||
"signal_count": 2,
|
||||
"failed_count": 0,
|
||||
}
|
||||
|
||||
|
||||
def test_execute_marks_partial_success_when_one_stock_lacks_history() -> None:
|
||||
source = _source()
|
||||
store = FakeStore()
|
||||
evaluator = FakeEvaluator(
|
||||
{
|
||||
"000001.SZ": SelectionEvaluation("000001.SZ", TARGET, "no_signal"),
|
||||
"600000.SH": SelectionEvaluation(
|
||||
"600000.SH",
|
||||
TARGET,
|
||||
"insufficient_history",
|
||||
reason="warm-up data is incomplete",
|
||||
),
|
||||
}
|
||||
)
|
||||
service = RunZhixingB1(FakeReader(source), store, evaluator)
|
||||
|
||||
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
|
||||
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "partial_success")
|
||||
assert store.finished[2]["failed_count"] == 1
|
||||
|
||||
|
||||
def test_execute_isolates_unexpected_single_stock_failure() -> None:
|
||||
source = _source()
|
||||
store = FakeStore()
|
||||
service = RunZhixingB1(FakeReader(source), store, RaisingEvaluator())
|
||||
|
||||
service.execute(service.prepare("zhixing_b1", TARGET, rerun=False))
|
||||
|
||||
assert [item.status for item in store.items] == ["no_signal", "data_error"]
|
||||
assert store.items[1].reason == "temporary evaluator failure"
|
||||
assert store.finished is not None
|
||||
assert store.finished[0:2] == ("run-1", "partial_success")
|
||||
assert store.finished[2]["evaluated_count"] == 2
|
||||
assert store.finished[2]["failed_count"] == 1
|
||||
Reference in New Issue
Block a user