Files
zhixing-system/zhixing-server/src/zhixing_server/modules/selection/infrastructure/postgres_runs.py
T

687 lines
25 KiB
Python
Raw Normal View History

"""PostgreSQL persistence adapter for selection execution runs."""
from __future__ import annotations
import json
from collections import defaultdict
2026-08-12 09:45:16 +08:00
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,
SelectionRunInProgress,
SelectionRunItem,
SelectionRunStatus,
SelectionRunStore,
SelectionRunStoreError,
)
from ..domain.zhixing_b1 import ZHIXING_B1_SIGNAL_ORDER
2026-08-12 09:45:16 +08:00
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",
}
2026-08-12 09:45:16 +08:00
_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)
2026-08-12 09:45:16 +08:00
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
2026-08-12 09:45:16 +08:00
"""
_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."""
2026-08-12 09:45:16 +08:00
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
2026-08-12 09:45:16 +08:00
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:
2026-08-12 09:45:16 +08:00
"""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.
2026-08-12 09:45:16 +08:00
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,
2026-08-12 09:45:16 +08:00
)
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(
2026-08-12 09:45:16 +08:00
"DELETE FROM selection_signal WHERE run_id = %s AND ts_code = ANY(%s)",
(run_id, codes),
)
2026-08-12 09:45:16 +08:00
_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(
2026-08-12 09:45:16 +08:00
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,
) -> 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())
except psycopg.Error as exc:
raise SelectionRunStoreError(f"failed to load selection run {run_id}") from exc
def get_latest_run(
self,
strategy: SelectionStrategyName,
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."""
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())
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,
query: SelectionResultQuery,
) -> 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)
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:
2026-08-12 09:45:16 +08:00
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
2026-08-12 09:45:16 +08:00
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) -> 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.
"""
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]}%")
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]