feat(selection): 支持选股结果服务端分页
This commit is contained in:
@@ -11,6 +11,7 @@ from ..domain.models import SelectionEvaluation
|
||||
from ..domain.runs import (
|
||||
SelectionExecutionSource,
|
||||
SelectionRerunRequired,
|
||||
SelectionResultQuery,
|
||||
SelectionRun,
|
||||
SelectionRunInProgress,
|
||||
SelectionRunItem,
|
||||
@@ -135,19 +136,26 @@ class RunZhixingB1:
|
||||
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:
|
||||
def get_run(
|
||||
self,
|
||||
run_id: str,
|
||||
*,
|
||||
query: SelectionResultQuery | None = None,
|
||||
) -> SelectionRun | None:
|
||||
"""Read one persisted run for polling."""
|
||||
|
||||
return self.store.get_run(run_id)
|
||||
return self.store.get_run(run_id, query=query)
|
||||
|
||||
def get_latest(
|
||||
self,
|
||||
strategy: StrategyName,
|
||||
target_trade_date: date | None = None,
|
||||
*,
|
||||
query: SelectionResultQuery | 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)
|
||||
return self.store.get_latest_run(strategy, target_trade_date, query=query)
|
||||
|
||||
|
||||
def _to_item(ts_code: str, name: str, evaluation: SelectionEvaluation) -> SelectionRunItem:
|
||||
|
||||
@@ -11,6 +11,17 @@ from .models import SelectionEvaluationStatus, SelectionSignal, StockHistory
|
||||
|
||||
SelectionRunStatus = Literal["running", "success", "partial_success", "failed"]
|
||||
SelectionRunItemStatus = SelectionEvaluationStatus
|
||||
SelectionSignalCategoryFilter = Literal["pullback", "oversold", "original"]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SelectionResultQuery:
|
||||
"""Validated query options for a paged selection-result read."""
|
||||
|
||||
page: int = 1
|
||||
page_size: int = 10
|
||||
search: str | None = None
|
||||
category: SelectionSignalCategoryFilter | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
@@ -67,6 +78,7 @@ class SelectionRun:
|
||||
finished_at: datetime | None = None
|
||||
items: tuple[SelectionRunItem, ...] = field(default_factory=tuple)
|
||||
signals: tuple[SelectionSignal, ...] = field(default_factory=tuple)
|
||||
signals_total: int | None = None
|
||||
|
||||
|
||||
class SelectionRunError(RuntimeError):
|
||||
@@ -112,12 +124,19 @@ class SelectionRunStore(Protocol):
|
||||
error_message: str | None = None,
|
||||
) -> None: ...
|
||||
|
||||
def get_run(self, run_id: str) -> SelectionRun | None: ...
|
||||
def get_run(
|
||||
self,
|
||||
run_id: str,
|
||||
*,
|
||||
query: SelectionResultQuery | None = None,
|
||||
) -> SelectionRun | None: ...
|
||||
|
||||
def get_latest_run(
|
||||
self,
|
||||
strategy: Literal["zhixing_b1"],
|
||||
target_trade_date: date | None = None,
|
||||
*,
|
||||
query: SelectionResultQuery | None = None,
|
||||
) -> SelectionRun | None: ...
|
||||
|
||||
|
||||
|
||||
+68
-9
@@ -18,6 +18,7 @@ from ..domain.models import SelectionSignal, ZhixingB1Category
|
||||
from ..domain.runs import (
|
||||
SelectionExecutionSource,
|
||||
SelectionRerunRequired,
|
||||
SelectionResultQuery,
|
||||
SelectionRun,
|
||||
SelectionRunError,
|
||||
SelectionRunInProgress,
|
||||
@@ -29,6 +30,19 @@ from ..domain.runs import (
|
||||
from ..domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
|
||||
|
||||
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(ZHIXING_B1_SIGNAL_ORDER)}
|
||||
_CATEGORY_PREFIXES = {
|
||||
"pullback": "zhixing_b1_pullback_",
|
||||
"oversold": "zhixing_b1_oversold_",
|
||||
"original": "zhixing_b1_original_b1",
|
||||
}
|
||||
_SIGNAL_ORDER_SQL = (
|
||||
"CASE category "
|
||||
+ " ".join(
|
||||
f"WHEN '{category.value}' THEN {index}"
|
||||
for index, category in enumerate(ZHIXING_B1_SIGNAL_ORDER)
|
||||
)
|
||||
+ f" ELSE {len(ZHIXING_B1_SIGNAL_ORDER)} END"
|
||||
)
|
||||
|
||||
|
||||
class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
@@ -224,12 +238,17 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
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."""
|
||||
def get_run(
|
||||
self,
|
||||
run_id: str,
|
||||
*,
|
||||
query: SelectionResultQuery | None = None,
|
||||
) -> SelectionRun | None:
|
||||
"""Read one run with filtered, paged signals and item failures."""
|
||||
|
||||
try:
|
||||
with self._connection() as connection:
|
||||
return self._load_run(connection, run_id)
|
||||
return self._load_run(connection, run_id, query or SelectionResultQuery())
|
||||
except psycopg.Error as exc:
|
||||
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
|
||||
|
||||
@@ -237,6 +256,8 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
self,
|
||||
strategy: Literal["zhixing_b1"],
|
||||
target_trade_date: date | None = None,
|
||||
*,
|
||||
query: SelectionResultQuery | None = None,
|
||||
) -> SelectionRun | None:
|
||||
"""Read the current run for a date or the latest date for a strategy."""
|
||||
|
||||
@@ -263,12 +284,20 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
""",
|
||||
(strategy, target_trade_date),
|
||||
).fetchone()
|
||||
return self._load_run(connection, str(row[0])) if row else None
|
||||
return (
|
||||
self._load_run(connection, str(row[0]), query or SelectionResultQuery())
|
||||
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:
|
||||
def _load_run(
|
||||
connection: Any,
|
||||
run_id: str,
|
||||
query: SelectionResultQuery,
|
||||
) -> SelectionRun | None:
|
||||
row = connection.execute(
|
||||
"""
|
||||
SELECT
|
||||
@@ -292,15 +321,23 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
""",
|
||||
(run_id,),
|
||||
).fetchall()
|
||||
signal_filter, signal_parameters = _signal_filter(query, run_id)
|
||||
signal_total_row = connection.execute(
|
||||
f"SELECT COUNT(*) FROM selection_signal WHERE {signal_filter}",
|
||||
tuple(signal_parameters),
|
||||
).fetchone()
|
||||
signal_total = int(signal_total_row[0] or 0) if signal_total_row else 0
|
||||
offset = (query.page - 1) * query.page_size
|
||||
signal_rows = connection.execute(
|
||||
"""
|
||||
f"""
|
||||
SELECT
|
||||
ts_code, name, target_trade_date, strategy, category, close, details
|
||||
FROM selection_signal
|
||||
WHERE run_id = %s
|
||||
ORDER BY ts_code, category
|
||||
WHERE {signal_filter}
|
||||
ORDER BY ts_code, {_SIGNAL_ORDER_SQL}
|
||||
LIMIT %s OFFSET %s
|
||||
""",
|
||||
(run_id,),
|
||||
tuple((*signal_parameters, query.page_size, offset)),
|
||||
).fetchall()
|
||||
signals = tuple(
|
||||
sorted(
|
||||
@@ -353,6 +390,7 @@ class PostgresSelectionRunRepository(SelectionRunStore):
|
||||
finished_at=cast(datetime | None, row[15]),
|
||||
items=items,
|
||||
signals=signals,
|
||||
signals_total=signal_total,
|
||||
)
|
||||
|
||||
@contextmanager
|
||||
@@ -380,6 +418,27 @@ def _signal_from_row(row: tuple[object, ...]) -> SelectionSignal:
|
||||
)
|
||||
|
||||
|
||||
def _signal_filter(query: SelectionResultQuery, run_id: str) -> tuple[str, list[object]]:
|
||||
"""Build the parameterized WHERE clause shared by count and page reads."""
|
||||
|
||||
clauses = ["run_id = %s"]
|
||||
parameters: list[object] = [run_id]
|
||||
if query.search:
|
||||
pattern = f"%{_escape_like(query.search)}%"
|
||||
clauses.append("(name ILIKE %s ESCAPE '\\' OR ts_code ILIKE %s ESCAPE '\\')")
|
||||
parameters.extend((pattern, pattern))
|
||||
if query.category:
|
||||
clauses.append("category LIKE %s")
|
||||
parameters.append(f"{_CATEGORY_PREFIXES[query.category]}%")
|
||||
return " AND ".join(clauses), parameters
|
||||
|
||||
|
||||
def _escape_like(value: str) -> str:
|
||||
"""Escape user wildcards before placing text inside a SQL LIKE pattern."""
|
||||
|
||||
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
|
||||
|
||||
|
||||
def _details(value: object) -> dict[str, float | str | None]:
|
||||
"""Normalize JSONB details into the domain's scalar-only mapping."""
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
from datetime import date, datetime
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, status
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, status
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhixing_server.bootstrap.config import Settings, get_settings
|
||||
@@ -12,6 +12,7 @@ from zhixing_server.modules.selection.application.run import (
|
||||
)
|
||||
from zhixing_server.modules.selection.domain.runs import (
|
||||
SelectionRerunRequired,
|
||||
SelectionResultQuery,
|
||||
SelectionRun,
|
||||
SelectionRunInProgress,
|
||||
SelectionRunStoreError,
|
||||
@@ -88,7 +89,7 @@ def _empty_signals() -> list[SelectionSignalResponse]:
|
||||
|
||||
|
||||
class SelectionResultsResponse(BaseModel):
|
||||
"""Batch summary and materialized signals consumed by the Web feature."""
|
||||
"""Batch summary and one filtered page of signals consumed by the Web feature."""
|
||||
|
||||
strategy: StrategyValue
|
||||
target_trade_date: date | None
|
||||
@@ -106,6 +107,9 @@ class SelectionResultsResponse(BaseModel):
|
||||
error_message: str | None = None
|
||||
created_at: datetime | None = None
|
||||
finished_at: datetime | None = None
|
||||
page: int = Field(default=1, ge=1)
|
||||
page_size: int = Field(default=10, ge=1, le=100)
|
||||
signals_total: int = Field(default=0, ge=0)
|
||||
failures: list[SelectionFailureResponse] = Field(default_factory=_empty_failures)
|
||||
signals: list[SelectionSignalResponse] = Field(default_factory=_empty_signals)
|
||||
|
||||
@@ -160,16 +164,21 @@ def trigger_selection_run(
|
||||
def get_selection_run(
|
||||
run_id: str,
|
||||
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
||||
page: Annotated[int, Query(ge=1)] = 1,
|
||||
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
||||
search: Annotated[str | None, Query(max_length=100)] = None,
|
||||
category: Literal["pullback", "oversold", "original"] | None = None,
|
||||
) -> SelectionResultsResponse:
|
||||
"""Return one run for asynchronous polling."""
|
||||
|
||||
try:
|
||||
run = service.get_run(run_id)
|
||||
query = _result_query(page, page_size, search, category)
|
||||
run = service.get_run(run_id, query=query)
|
||||
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)
|
||||
return _run_response(run, query=query)
|
||||
|
||||
|
||||
@selection_router.get("/results", response_model=SelectionResultsResponse)
|
||||
@@ -177,11 +186,16 @@ def get_selection_results(
|
||||
service: Annotated[RunZhixingB1, Depends(get_selection_service)],
|
||||
strategy: StrategyValue = "zhixing_b1",
|
||||
target_trade_date: date | None = None,
|
||||
page: Annotated[int, Query(ge=1)] = 1,
|
||||
page_size: Annotated[int, Query(ge=1, le=100)] = 10,
|
||||
search: Annotated[str | None, Query(max_length=100)] = None,
|
||||
category: Literal["pullback", "oversold", "original"] | None = None,
|
||||
) -> SelectionResultsResponse:
|
||||
"""Return the current persisted result for a strategy and optional date."""
|
||||
|
||||
try:
|
||||
run = service.get_latest(strategy, target_trade_date)
|
||||
query = _result_query(page, page_size, search, category)
|
||||
run = service.get_latest(strategy, target_trade_date, query=query)
|
||||
except SelectionRunStoreError as exc:
|
||||
raise _http_error(503, "selection_storage_unavailable", str(exc)) from exc
|
||||
if run is None:
|
||||
@@ -192,11 +206,14 @@ def get_selection_results(
|
||||
market_sync_batch_id=None,
|
||||
status="no_data",
|
||||
coverage=0,
|
||||
page=query.page,
|
||||
page_size=query.page_size,
|
||||
signals_total=0,
|
||||
)
|
||||
return _run_response(run)
|
||||
return _run_response(run, query=query)
|
||||
|
||||
|
||||
def _run_response(run: SelectionRun) -> SelectionResultsResponse:
|
||||
def _run_response(run: SelectionRun, *, query: SelectionResultQuery) -> SelectionResultsResponse:
|
||||
"""Translate a domain run without exposing storage-specific fields."""
|
||||
|
||||
return SelectionResultsResponse(
|
||||
@@ -216,6 +233,9 @@ def _run_response(run: SelectionRun) -> SelectionResultsResponse:
|
||||
error_message=run.error_message,
|
||||
created_at=run.created_at,
|
||||
finished_at=run.finished_at,
|
||||
page=query.page,
|
||||
page_size=query.page_size,
|
||||
signals_total=run.signals_total if run.signals_total is not None else run.signal_count,
|
||||
failures=[
|
||||
SelectionFailureResponse(
|
||||
ts_code=item.ts_code,
|
||||
@@ -241,6 +261,23 @@ def _run_response(run: SelectionRun) -> SelectionResultsResponse:
|
||||
)
|
||||
|
||||
|
||||
def _result_query(
|
||||
page: int,
|
||||
page_size: int,
|
||||
search: str | None,
|
||||
category: Literal["pullback", "oversold", "original"] | None,
|
||||
) -> SelectionResultQuery:
|
||||
"""Normalize HTTP query values before handing them to the selection port."""
|
||||
|
||||
normalized_search = search.strip() if search else None
|
||||
return SelectionResultQuery(
|
||||
page=page,
|
||||
page_size=page_size,
|
||||
search=normalized_search or None,
|
||||
category=category,
|
||||
)
|
||||
|
||||
|
||||
def _http_error(code: int, error_type: str, message: str) -> HTTPException:
|
||||
"""Create the project's explicit, safe error envelope."""
|
||||
|
||||
|
||||
Reference in New Issue
Block a user