Files
zhixing-system/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py
T
yuxuanhui 7e0f13d678 feat(selection): expose run sector aggregates and sector filter on results API
- sector_radar: add batch sector-count aggregation and sector member lookup
  over the strict last-good membership snapshot (postgres + in-memory fakes)
- selection: add SelectionSectorReader port, list_sector_counts use case,
  and sector_stock_codes filtering via run identity resolution; queries stay
  inside the selection context per ADR 0001
- http: add GET /api/v1/selection/sectors and forward sector param on
  /results and /runs/{run_id}
- fix stale positional args in pattern-scoring run tests; cover new behavior
  with read-service, application, and HTTP contract tests
2026-09-05 19:48:30 +08:00

762 lines
28 KiB
Python

"""PostgreSQL persistence adapter for selection execution runs."""
from __future__ import annotations
import json
from collections import defaultdict
from collections.abc import Generator, Mapping, Sequence
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.gold_brick import GOLD_BRICK_SIGNAL_ORDER
from ..domain.models import (
GoldBrickCategory,
SelectionSignal,
SelectionSignalCategory,
SelectionStrategyName,
ZhixingB1Category,
)
from ..domain.pattern_scoring import (
ZHIXING_B1_PATTERN_CASES,
PatternScore,
PatternScoreBreakdown,
)
from ..domain.runs import (
SelectionExecutionSource,
SelectionRerunRequired,
SelectionResultQuery,
SelectionRun,
SelectionRunError,
SelectionRunIdentity,
SelectionRunInProgress,
SelectionRunItem,
SelectionRunStatus,
SelectionRunStore,
SelectionRunStoreError,
)
from ..domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
from .postgres_pool import SelectionConnectionPool, SelectionPostgresPool
_SELECTION_SIGNAL_ORDER: tuple[SelectionSignalCategory, ...] = (
*ZHIXING_B1_SIGNAL_ORDER,
*GOLD_BRICK_SIGNAL_ORDER,
)
_SIGNAL_PRIORITY = {category: index for index, category in enumerate(_SELECTION_SIGNAL_ORDER)}
_CATEGORY_PREFIXES = {
"pullback": "zhixing_b1_pullback_",
"oversold": "zhixing_b1_oversold_",
"original": "zhixing_b1_original_b1",
"resonance": "gold_brick_resonance",
}
_SIGNAL_ORDER_SQL = (
"CASE category "
+ " ".join(
f"WHEN '{category.value}' THEN {index}"
for index, category in enumerate(_SELECTION_SIGNAL_ORDER)
)
+ f" ELSE {len(_SELECTION_SIGNAL_ORDER)} END"
)
_PATTERN_CASES_BY_ID = {definition.id: definition for definition in ZHIXING_B1_PATTERN_CASES}
_STOCK_ORDER_SQL = {
"code": "item.ts_code ASC",
"score_desc": "item.score_value DESC NULLS LAST, item.ts_code ASC",
"score_asc": "item.score_value ASC NULLS LAST, item.ts_code ASC",
}
_ITEM_UPSERT = """
INSERT INTO selection_run_item
(
run_id, ts_code, name, status, signal_count, reason,
score_status, score_value, score_threshold, score_version,
match_case_id, match_case_name, match_case_breakout_date,
match_breakdown, score_reason
)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %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,
score_status = EXCLUDED.score_status,
score_value = EXCLUDED.score_value,
score_threshold = EXCLUDED.score_threshold,
score_version = EXCLUDED.score_version,
match_case_id = EXCLUDED.match_case_id,
match_case_name = EXCLUDED.match_case_name,
match_case_breakout_date = EXCLUDED.match_case_breakout_date,
match_breakdown = EXCLUDED.match_breakdown,
score_reason = EXCLUDED.score_reason
"""
_SIGNAL_UPSERT = """
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
"""
class PostgresSelectionRunRepository(SelectionRunStore):
"""Persist one current result attempt per strategy and target date."""
def __init__(
self,
database_url: str,
*,
pool: SelectionPostgresPool | SelectionConnectionPool | None = None,
) -> None:
"""Create the adapter with a URL and optional shared PostgreSQL pool."""
self.database_url = database_url
if isinstance(pool, SelectionPostgresPool):
self.pool: SelectionPostgresPool | None = pool
elif pool is not None:
self.pool = SelectionPostgresPool(
database_url,
max_connections=1,
pool=pool,
)
else:
self.pool = None
def prepare_run(
self,
strategy: SelectionStrategyName,
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:
"""Persist one item through the batch path for compatibility."""
self.record_items(run_id, (item,))
def record_items(self, run_id: str, items: Sequence[SelectionRunItem]) -> None:
"""Persist one chunk in one transaction with set-based driver calls.
Existing signal rows are removed before the upserts so retrying a
chunk cannot retain a category that disappeared from a recalculation.
``executemany`` is used for both materialized tables; the small
fallback keeps the direct fake connections used by older tests usable.
"""
if not items:
return
item_values = tuple(
(
run_id,
item.ts_code,
item.name,
item.status,
item.signal_count,
item.reason,
item.pattern_score.status,
item.pattern_score.value,
item.pattern_score.threshold,
item.pattern_score.version,
item.pattern_score.case.id if item.pattern_score.case else None,
item.pattern_score.case.name if item.pattern_score.case else None,
item.pattern_score.case.breakout_date if item.pattern_score.case else None,
(
Jsonb(item.pattern_score.breakdown.as_dict())
if item.pattern_score.breakdown
else None
),
item.pattern_score.reason,
)
for item in items
)
signal_values = tuple(
(
run_id,
signal.ts_code,
signal.name,
signal.target_trade_date,
signal.strategy,
signal.category.value,
signal.close,
Jsonb(dict(signal.details)),
)
for item in items
for signal in item.signals
)
codes = [item.ts_code for item in items]
try:
with self._connection() as connection, connection.transaction():
connection.execute(
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = ANY(%s)",
(run_id, codes),
)
_executemany(connection, _ITEM_UPSERT, item_values)
if signal_values:
_executemany(connection, _SIGNAL_UPSERT, signal_values)
except SelectionRunError:
raise
except Exception as exc: # noqa: BLE001 - redact driver/pool details
code_context = items[0].ts_code if len(items) == 1 else f"{len(items)} items"
raise SelectionRunStoreError(
f"failed to persist selection item {code_context}"
) 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,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | None = None,
) -> SelectionRun | None:
"""Read one run with filtered, stock-paged signals and item failures."""
try:
with self._connection() as connection:
return self._load_run(
connection,
run_id,
query or SelectionResultQuery(),
sector_stock_codes=sector_stock_codes,
)
except psycopg.Error as exc:
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
def get_run_identity(self, run_id: str) -> SelectionRunIdentity | None:
"""Read only a run's locator so date-dependent filters resolve first."""
try:
with self._connection() as connection:
row = connection.execute(
"SELECT id, target_trade_date FROM selection_run WHERE id = %s",
(run_id,),
).fetchone()
except psycopg.Error as exc:
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
return SelectionRunIdentity(run_id=str(row[0]), target_trade_date=row[1]) if row else None
def get_latest_run(
self,
strategy: SelectionStrategyName,
target_trade_date: date | None = None,
*,
query: SelectionResultQuery | None = None,
sector_stock_codes: Sequence[str] | 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]),
query or SelectionResultQuery(),
sector_stock_codes=sector_stock_codes,
)
if row
else None
)
except psycopg.Error as exc:
raise SelectionRunStoreError("failed to load latest selection run") from exc
def get_latest_run_identity(
self,
strategy: SelectionStrategyName,
target_trade_date: date | None = None,
) -> SelectionRunIdentity | None:
"""Read only the current run's locator so sector filters resolve first."""
try:
with self._connection() as connection:
if target_trade_date is None:
row = connection.execute(
"""
SELECT id, target_trade_date
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, target_trade_date
FROM selection_run
WHERE strategy = %s AND target_trade_date = %s
LIMIT 1
""",
(strategy, target_trade_date),
).fetchone()
except psycopg.Error as exc:
raise SelectionRunStoreError("failed to load latest selection run") from exc
return SelectionRunIdentity(run_id=str(row[0]), target_trade_date=row[1]) if row else None
@staticmethod
def _load_run(
connection: Any,
run_id: str,
query: SelectionResultQuery,
*,
sector_stock_codes: Sequence[str] | None = None,
) -> 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,
score_status, score_value, score_threshold, score_version,
match_case_id, match_case_name, match_case_breakout_date,
match_breakdown, score_reason
FROM selection_run_item
WHERE run_id = %s
ORDER BY ts_code
""",
(run_id,),
).fetchall()
stock_filter, stock_parameters = _stock_filter(
query, run_id, sector_stock_codes=sector_stock_codes
)
stock_total_row = connection.execute(
f"SELECT COUNT(*) FROM selection_run_item AS item WHERE {stock_filter}",
tuple(stock_parameters),
).fetchone()
stock_total = int(stock_total_row[0] or 0) if stock_total_row else 0
offset = (query.page - 1) * query.page_size
stock_rows = cast(
list[tuple[object, ...]],
connection.execute(
f"""
SELECT item.ts_code
FROM selection_run_item AS item
WHERE {stock_filter}
ORDER BY {_STOCK_ORDER_SQL[query.sort]}
LIMIT %s OFFSET %s
""",
tuple((*stock_parameters, query.page_size, offset)),
).fetchall(),
)
stock_codes = [str(value[0]) for value in stock_rows]
signal_rows = (
cast(
list[tuple[object, ...]],
connection.execute(
f"""
SELECT
ts_code, name, target_trade_date, strategy, category, close, details
FROM selection_signal
WHERE run_id = %s AND ts_code = ANY(%s)
ORDER BY ts_code, {_SIGNAL_ORDER_SQL}
""",
(run_id, stock_codes),
).fetchall(),
)
if stock_codes
else []
)
signals = tuple(
sorted(
(_signal_from_row(value) for value in signal_rows),
key=lambda signal: (
stock_codes.index(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",
"missing_turnover_rate",
"data_error",
],
str(value[2]),
),
signal_count=int(value[3] or 0),
reason=str(value[4]) if value[4] is not None else None,
pattern_score=_pattern_score_from_row(value[5:14]),
signals=tuple(signals_by_stock.get(str(value[0]), ())),
)
for value in item_rows
)
return SelectionRun(
id=str(row[0]),
strategy=cast(SelectionStrategyName, 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,
stocks_total=stock_total,
)
@contextmanager
def _connection(self) -> Generator[Any, None, None]:
"""Translate psycopg failures without exposing driver details."""
try:
if self.pool is None:
with psycopg.connect(self.database_url) as connection:
yield connection
else:
with self.pool.connection() as connection:
yield connection
except SelectionRunError:
raise
except Exception as exc: # noqa: BLE001 - normalize driver/pool errors
raise SelectionRunStoreError("selection database operation failed") from exc
def _executemany(connection: Any, query: str, parameters: Sequence[tuple[object, ...]]) -> None:
"""Use psycopg's batch API while retaining a minimal fake connection seam."""
executemany = getattr(connection, "executemany", None)
if callable(executemany):
executemany(query, parameters)
return
cursor_factory = getattr(connection, "cursor", None)
if callable(cursor_factory):
cursor_context = cast(Any, cursor_factory())
with cursor_context as cursor:
cursor.executemany(query, parameters)
return
for values in parameters:
connection.execute(query, values)
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(SelectionStrategyName, str(row[3])),
category=_signal_category(str(row[4])),
close=float(str(row[5])),
details=_details(row[6]),
)
def _signal_category(value: str) -> SelectionSignalCategory:
"""Map a persisted category for either supported selection strategy."""
try:
return ZhixingB1Category(value)
except ValueError:
return GoldBrickCategory(value)
def _stock_filter(
query: SelectionResultQuery,
run_id: str,
*,
sector_stock_codes: Sequence[str] | None = None,
) -> tuple[str, list[object]]:
"""Build the signal predicate used to select distinct matching stocks.
A category narrows which stocks qualify for the page. Once a stock
qualifies, the repository loads every signal for that stock so callers
can present all independently persisted categories together. Resolved
sector membership codes arrive from the sector-radar port, so the SQL
stays inside the selection context.
"""
clauses = ["item.run_id = %s", "item.status = 'selected'", "item.signal_count > 0"]
parameters: list[object] = [run_id]
if query.search:
pattern = f"%{_escape_like(query.search)}%"
clauses.append("(item.name ILIKE %s ESCAPE '\\' OR item.ts_code ILIKE %s ESCAPE '\\')")
parameters.extend((pattern, pattern))
if query.category:
clauses.append(
"EXISTS ("
"SELECT 1 FROM selection_signal AS signal "
"WHERE signal.run_id = item.run_id "
"AND signal.ts_code = item.ts_code "
"AND signal.category LIKE %s"
")"
)
parameters.append(f"{_CATEGORY_PREFIXES[query.category]}%")
if sector_stock_codes is not None:
if not sector_stock_codes:
clauses.append("FALSE")
else:
clauses.append("item.ts_code = ANY(%s)")
parameters.append(list(sector_stock_codes))
return " AND ".join(clauses), parameters
def _pattern_score_from_row(row: Sequence[object]) -> PatternScore:
"""Reconstruct a validated stock-level score from nullable item columns."""
if len(row) < 9:
return PatternScore()
status = str(row[0] or "not_executed")
if status == "not_executed":
return PatternScore()
if status == "failed":
return PatternScore.failed(str(row[8] or "pattern scoring failed"))
if status not in {"matched", "below_threshold"}:
return PatternScore.failed("persisted pattern score status is invalid")
definition = _PATTERN_CASES_BY_ID.get(str(row[4]))
breakdown = _pattern_breakdown(row[7])
if definition is None or breakdown is None:
return PatternScore.failed("persisted pattern score is incomplete")
try:
return PatternScore(
status=cast(Literal["matched", "below_threshold"], status),
value=float(str(row[1])),
threshold=float(str(row[2])),
version=str(row[3]),
case=definition,
breakdown=breakdown,
)
except (TypeError, ValueError):
return PatternScore.failed("persisted pattern score is invalid")
def _pattern_breakdown(value: object) -> PatternScoreBreakdown | None:
"""Parse the four finite JSONB score dimensions."""
if isinstance(value, str):
try:
value = json.loads(value)
except json.JSONDecodeError:
return None
if not isinstance(value, Mapping):
return None
values = cast(Mapping[object, object], value)
try:
return PatternScoreBreakdown(
trend_structure=float(str(values["trend_structure"])),
kdj_state=float(str(values["kdj_state"])),
volume_pattern=float(str(values["volume_pattern"])),
price_shape=float(str(values["price_shape"])),
)
except (KeyError, TypeError, ValueError):
return None
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."""
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]