Files
zhixing-system/zhixing-server/src/zhixing_server/modules/selection/presentation/http.py
T

375 lines
12 KiB
Python

"""HTTP presentation for persisted strategy execution results."""
import atexit
import threading
from datetime import date, datetime
from typing import Annotated, Literal
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, 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.models import SelectionSignal
from zhixing_server.modules.selection.domain.runs import (
SelectionRerunRequired,
SelectionResultQuery,
SelectionRun,
SelectionRunInProgress,
SelectionRunStoreError,
)
from zhixing_server.modules.selection.infrastructure.postgres_pool import SelectionPostgresPool
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()
_SELECTION_POOL_CACHE_LOCK = threading.Lock()
_SELECTION_POOL_CACHE: dict[tuple[str, int], SelectionPostgresPool] = {}
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 SelectionStockResponse(BaseModel):
"""One selected stock with all independently persisted signals."""
ts_code: str
name: str
target_trade_date: date
strategy: StrategyValue
close: float
signals: list[SelectionSignalResponse] = Field(default_factory=_empty_signals)
def _empty_stocks() -> list[SelectionStockResponse]:
"""Create a typed default stock-result list."""
return []
class SelectionResultsResponse(BaseModel):
"""Batch summary and one filtered page of selected stocks."""
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
page: int = Field(default=1, ge=1)
page_size: int = Field(default=10, ge=1, le=100)
stocks_total: int = Field(default=0, ge=0)
failures: list[SelectionFailureResponse] = Field(default_factory=_empty_failures)
stocks: list[SelectionStockResponse] = Field(default_factory=_empty_stocks)
def get_selection_service(
settings: Annotated[Settings, Depends(get_settings)],
) -> RunZhixingB1:
"""Build the selection service on top of process-scoped shared resources."""
pool = get_selection_postgres_pool(settings)
reader = PostgresMarketDataReader(settings, pool=pool)
store = PostgresSelectionRunRepository(settings.database_url, pool=pool)
return RunZhixingB1(
reader,
store,
max_workers=settings.selection_max_workers,
batch_size=settings.selection_batch_size,
)
def get_selection_postgres_pool(settings: Settings) -> SelectionPostgresPool:
"""Return the cached bounded pool shared by selection adapters."""
key = (settings.database_url, settings.selection_max_workers + 2)
with _SELECTION_POOL_CACHE_LOCK:
pool = _SELECTION_POOL_CACHE.get(key)
if pool is None:
pool = SelectionPostgresPool(
settings.database_url,
max_connections=key[1],
)
_SELECTION_POOL_CACHE[key] = pool
return pool
def _close_cached_selection_pools() -> None:
"""Close all process-cached selection pools during interpreter shutdown."""
with _SELECTION_POOL_CACHE_LOCK:
pools = tuple(_SELECTION_POOL_CACHE.values())
_SELECTION_POOL_CACHE.clear()
for pool in pools:
pool.close()
atexit.register(_close_cached_selection_pools)
@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)],
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:
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, query=query)
@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,
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:
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:
return SelectionResultsResponse(
strategy=strategy,
target_trade_date=target_trade_date,
run_id=None,
market_sync_batch_id=None,
status="no_data",
coverage=0,
page=query.page,
page_size=query.page_size,
stocks_total=0,
)
return _run_response(run, query=query)
def _run_response(run: SelectionRun, *, query: SelectionResultQuery) -> SelectionResultsResponse:
"""Translate a domain run without exposing storage-specific fields."""
signals_by_stock: dict[str, list[SelectionSignalResponse]] = {}
for signal in run.signals:
signals_by_stock.setdefault(signal.ts_code, []).append(_signal_response(signal))
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,
page=query.page,
page_size=query.page_size,
stocks_total=(
run.stocks_total if run.stocks_total is not None else run.selected_stock_count
),
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"}
],
stocks=[
SelectionStockResponse(
ts_code=signals[0].ts_code,
name=signals[0].name,
target_trade_date=signals[0].target_trade_date,
strategy=signals[0].strategy,
close=signals[0].close,
signals=signals,
)
for signals in signals_by_stock.values()
],
)
def _signal_response(signal: SelectionSignal) -> SelectionSignalResponse:
"""Map a domain signal while preserving its category-specific details."""
return 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),
)
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."""
return HTTPException(
status_code=code,
detail={"code": error_type, "message": message},
)
__all__ = [
"SelectionResultsResponse",
"SelectionRunAcceptedResponse",
"SelectionRunRequest",
"SelectionStockResponse",
"get_selection_service",
"get_selection_postgres_pool",
"selection_router",
]