feat(selection): 补充策略执行结果查询链路

This commit is contained in:
yuxuanhui
2026-08-09 09:34:46 +08:00
parent e9d06df5de
commit 9c1a1eac23
31 changed files with 3566 additions and 14 deletions
@@ -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()
+207
View File
@@ -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