"""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, SelectionResultQuery, 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)} _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): """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, *, 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, 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: 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.""" 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 FROM selection_run_item WHERE run_id = %s ORDER BY ts_code """, (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 {signal_filter} ORDER BY ts_code, {_SIGNAL_ORDER_SQL} LIMIT %s OFFSET %s """, tuple((*signal_parameters, query.page_size, offset)), ).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, signals_total=signal_total, ) @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 _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.""" 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]