415 lines
15 KiB
Python
415 lines
15 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
|
||
|
|
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,
|
||
|
|
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)}
|
||
|
|
|
||
|
|
|
||
|
|
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) -> SelectionRun | None:
|
||
|
|
"""Read one run with its item failures and signal details."""
|
||
|
|
|
||
|
|
try:
|
||
|
|
with self._connection() as connection:
|
||
|
|
return self._load_run(connection, run_id)
|
||
|
|
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,
|
||
|
|
) -> 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])) 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) -> 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_rows = connection.execute(
|
||
|
|
"""
|
||
|
|
SELECT
|
||
|
|
ts_code, name, target_trade_date, strategy, category, close, details
|
||
|
|
FROM selection_signal
|
||
|
|
WHERE run_id = %s
|
||
|
|
ORDER BY ts_code, category
|
||
|
|
""",
|
||
|
|
(run_id,),
|
||
|
|
).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,
|
||
|
|
)
|
||
|
|
|
||
|
|
@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 _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]
|