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

415 lines
15 KiB
Python
Raw Normal View History

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