"""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]