495 lines
16 KiB
Python
495 lines
16 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import sqlite3
|
|
from datetime import datetime
|
|
from typing import Any
|
|
|
|
|
|
class ScreenerRepository:
|
|
def save_factor_snapshot(
|
|
self,
|
|
connection: sqlite3.Connection,
|
|
*,
|
|
trade_date: str,
|
|
version: str,
|
|
observed_at: str,
|
|
state: str,
|
|
sources: list[str],
|
|
coverage: dict[str, float],
|
|
rows: list[dict[str, Any]],
|
|
) -> int:
|
|
connection.execute(
|
|
"""
|
|
INSERT OR IGNORE INTO screener_factor_snapshots (
|
|
trade_date, version, observed_at, state, source_set_json,
|
|
coverage_json, created_at
|
|
) VALUES (?, ?, ?, ?, ?, ?, ?)
|
|
""",
|
|
(
|
|
trade_date,
|
|
version,
|
|
observed_at,
|
|
state,
|
|
_json(sources),
|
|
_json(coverage),
|
|
_now(),
|
|
),
|
|
)
|
|
snapshot = connection.execute(
|
|
"""
|
|
SELECT id FROM screener_factor_snapshots
|
|
WHERE trade_date = ? AND version = ?
|
|
""",
|
|
(trade_date, version),
|
|
).fetchone()
|
|
if snapshot is None:
|
|
raise RuntimeError("因子快照写入失败")
|
|
snapshot_id = int(snapshot["id"])
|
|
connection.executemany(
|
|
"""
|
|
INSERT OR REPLACE INTO screener_factor_values (
|
|
snapshot_id, identifier, code, name, sector,
|
|
listed_days, is_st, payload_json
|
|
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
|
""",
|
|
(
|
|
(
|
|
snapshot_id,
|
|
row["identifier"],
|
|
row["code"],
|
|
row["name"],
|
|
row.get("sector"),
|
|
int(row.get("listed_days") or 0),
|
|
int(bool(row.get("is_st"))),
|
|
_json(row),
|
|
)
|
|
for row in rows
|
|
),
|
|
)
|
|
return snapshot_id
|
|
|
|
def latest_factor_snapshot(
|
|
self, connection: sqlite3.Connection, through: str
|
|
) -> sqlite3.Row | None:
|
|
return connection.execute(
|
|
"""
|
|
SELECT * FROM screener_factor_snapshots
|
|
WHERE trade_date <= ? ORDER BY trade_date DESC, id DESC LIMIT 1
|
|
""",
|
|
(through,),
|
|
).fetchone()
|
|
|
|
def factor_snapshot(
|
|
self, connection: sqlite3.Connection, snapshot_id: int
|
|
) -> sqlite3.Row | None:
|
|
return connection.execute(
|
|
"SELECT * FROM screener_factor_snapshots WHERE id = ?",
|
|
(snapshot_id,),
|
|
).fetchone()
|
|
|
|
def factor_snapshots_through(
|
|
self, connection: sqlite3.Connection, through: str, limit: int
|
|
) -> tuple[sqlite3.Row, ...]:
|
|
return tuple(
|
|
connection.execute(
|
|
"""
|
|
SELECT * FROM (
|
|
SELECT snapshots.*,
|
|
ROW_NUMBER() OVER (
|
|
PARTITION BY trade_date ORDER BY id DESC
|
|
) AS revision_rank
|
|
FROM screener_factor_snapshots AS snapshots
|
|
WHERE trade_date <= ?
|
|
)
|
|
WHERE revision_rank = 1
|
|
ORDER BY trade_date DESC, id DESC
|
|
LIMIT ?
|
|
""",
|
|
(through, limit),
|
|
).fetchall()
|
|
)
|
|
|
|
def factor_rows(self, connection: sqlite3.Connection, snapshot_id: int) -> list[dict[str, Any]]:
|
|
return [
|
|
json.loads(str(row["payload_json"]))
|
|
for row in connection.execute(
|
|
"""
|
|
SELECT payload_json FROM screener_factor_values
|
|
WHERE snapshot_id = ? ORDER BY identifier
|
|
""",
|
|
(snapshot_id,),
|
|
).fetchall()
|
|
]
|
|
|
|
def begin_run(
|
|
self,
|
|
connection: sqlite3.Connection,
|
|
*,
|
|
owner_user_id: int | None,
|
|
mode: str,
|
|
strategy_id: str,
|
|
strategy_name: str,
|
|
strategy_version: int,
|
|
selection_date: str,
|
|
factor_snapshot_id: int,
|
|
) -> sqlite3.Row:
|
|
connection.execute(
|
|
"""
|
|
INSERT OR IGNORE INTO screener_runs (
|
|
owner_user_id, mode, strategy_id, strategy_name,
|
|
strategy_version, selection_date, factor_snapshot_id,
|
|
status, started_at
|
|
) VALUES (?, ?, ?, ?, ?, ?, ?, 'running', ?)
|
|
""",
|
|
(
|
|
owner_user_id,
|
|
mode,
|
|
strategy_id,
|
|
strategy_name,
|
|
strategy_version,
|
|
selection_date,
|
|
factor_snapshot_id,
|
|
_now(),
|
|
),
|
|
)
|
|
row = connection.execute(
|
|
"""
|
|
SELECT * FROM screener_runs
|
|
WHERE mode = ? AND strategy_id = ? AND selection_date = ?
|
|
AND strategy_version = ? AND factor_snapshot_id = ?
|
|
AND COALESCE(owner_user_id, 0) = COALESCE(?, 0)
|
|
""",
|
|
(
|
|
mode,
|
|
strategy_id,
|
|
selection_date,
|
|
strategy_version,
|
|
factor_snapshot_id,
|
|
owner_user_id,
|
|
),
|
|
).fetchone()
|
|
if row is None:
|
|
raise RuntimeError("选股任务写入失败")
|
|
return row
|
|
|
|
def finish_run(
|
|
self,
|
|
connection: sqlite3.Connection,
|
|
run_id: int,
|
|
*,
|
|
status: str,
|
|
coverage: float,
|
|
missing_fields: list[str],
|
|
result: list[dict[str, Any]],
|
|
error_message: str = "",
|
|
) -> None:
|
|
connection.execute(
|
|
"""
|
|
UPDATE screener_runs SET
|
|
status = ?, completed_at = ?, coverage = ?,
|
|
missing_fields_json = ?, result_json = ?, error_message = ?
|
|
WHERE id = ?
|
|
""",
|
|
(
|
|
status,
|
|
_now(),
|
|
max(0, min(coverage, 1)),
|
|
_json(missing_fields),
|
|
_json(result),
|
|
error_message,
|
|
run_id,
|
|
),
|
|
)
|
|
|
|
def save_backtest(
|
|
self,
|
|
connection: sqlite3.Connection,
|
|
run_id: int,
|
|
payload: dict[str, Any],
|
|
) -> None:
|
|
connection.execute(
|
|
"""
|
|
INSERT OR IGNORE INTO screener_run_backtests (
|
|
run_id, payload_json, created_at
|
|
) VALUES (?, ?, ?)
|
|
""",
|
|
(run_id, _json(payload), _now()),
|
|
)
|
|
|
|
def backtest(
|
|
self, connection: sqlite3.Connection, run_id: int
|
|
) -> dict[str, Any] | None:
|
|
row = connection.execute(
|
|
"SELECT payload_json FROM screener_run_backtests WHERE run_id = ?",
|
|
(run_id,),
|
|
).fetchone()
|
|
return json.loads(str(row["payload_json"])) if row else None
|
|
|
|
def latest_runs(
|
|
self,
|
|
connection: sqlite3.Connection,
|
|
mode: str,
|
|
through: str,
|
|
owner_user_id: int | None = None,
|
|
) -> tuple[sqlite3.Row, ...]:
|
|
return tuple(
|
|
connection.execute(
|
|
"""
|
|
SELECT run.* FROM screener_runs run
|
|
JOIN (
|
|
SELECT strategy_id, MAX(id) AS latest_id
|
|
FROM screener_runs
|
|
WHERE mode = ? AND selection_date = ?
|
|
AND COALESCE(owner_user_id, 0) = COALESCE(?, 0)
|
|
GROUP BY strategy_id
|
|
) latest ON latest.latest_id = run.id
|
|
ORDER BY run.strategy_id
|
|
""",
|
|
(mode, through, owner_user_id),
|
|
).fetchall()
|
|
)
|
|
|
|
def run_for_user(
|
|
self, connection: sqlite3.Connection, run_id: int, user_id: int
|
|
) -> sqlite3.Row | None:
|
|
return connection.execute(
|
|
"""
|
|
SELECT * FROM screener_runs
|
|
WHERE id = ? AND (owner_user_id IS NULL OR owner_user_id = ?)
|
|
""",
|
|
(run_id, user_id),
|
|
).fetchone()
|
|
|
|
def save_custom_strategy(
|
|
self,
|
|
connection: sqlite3.Connection,
|
|
user_id: int,
|
|
name: str,
|
|
formula: dict[str, Any],
|
|
) -> sqlite3.Row:
|
|
now = _now()
|
|
connection.execute(
|
|
"""
|
|
INSERT INTO custom_screener_strategies (
|
|
user_id, name, version, formula_json, created_at, updated_at
|
|
) VALUES (?, ?, 1, ?, ?, ?)
|
|
ON CONFLICT(user_id, name) DO UPDATE SET
|
|
version = version + 1,
|
|
formula_json = excluded.formula_json,
|
|
updated_at = excluded.updated_at
|
|
""",
|
|
(user_id, name, _json(formula), now, now),
|
|
)
|
|
row = connection.execute(
|
|
"""
|
|
SELECT * FROM custom_screener_strategies
|
|
WHERE user_id = ? AND name = ?
|
|
""",
|
|
(user_id, name),
|
|
).fetchone()
|
|
if row is None:
|
|
raise RuntimeError("自定义策略写入失败")
|
|
return row
|
|
|
|
def custom_strategies(
|
|
self, connection: sqlite3.Connection, user_id: int
|
|
) -> tuple[sqlite3.Row, ...]:
|
|
return tuple(
|
|
connection.execute(
|
|
"""
|
|
SELECT * FROM custom_screener_strategies
|
|
WHERE user_id = ? ORDER BY updated_at DESC, id DESC
|
|
""",
|
|
(user_id,),
|
|
).fetchall()
|
|
)
|
|
|
|
def custom_strategy(
|
|
self, connection: sqlite3.Connection, user_id: int, strategy_id: int
|
|
) -> sqlite3.Row | None:
|
|
return connection.execute(
|
|
"""
|
|
SELECT * FROM custom_screener_strategies
|
|
WHERE id = ? AND user_id = ?
|
|
""",
|
|
(strategy_id, user_id),
|
|
).fetchone()
|
|
|
|
def delete_custom_strategy(
|
|
self, connection: sqlite3.Connection, user_id: int, strategy_id: int
|
|
) -> bool:
|
|
cursor = connection.execute(
|
|
"DELETE FROM custom_screener_strategies WHERE id = ? AND user_id = ?",
|
|
(strategy_id, user_id),
|
|
)
|
|
return cursor.rowcount > 0
|
|
|
|
def add_track(
|
|
self,
|
|
connection: sqlite3.Connection,
|
|
*,
|
|
user_id: int,
|
|
run: sqlite3.Row,
|
|
candidate: dict[str, Any],
|
|
) -> int:
|
|
connection.execute(
|
|
"""
|
|
INSERT OR IGNORE INTO strategy_tracks (
|
|
user_id, run_id, identifier, code, name, sector,
|
|
selection_date, strategy_name, entry_price, added_at
|
|
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
|
""",
|
|
(
|
|
user_id,
|
|
int(run["id"]),
|
|
candidate["identifier"],
|
|
candidate["code"],
|
|
candidate["name"],
|
|
candidate.get("sector"),
|
|
str(run["selection_date"]),
|
|
str(run["strategy_name"]),
|
|
float(candidate["close"]),
|
|
_now(),
|
|
),
|
|
)
|
|
row = connection.execute(
|
|
"""
|
|
SELECT id FROM strategy_tracks
|
|
WHERE user_id = ? AND run_id = ? AND identifier = ?
|
|
""",
|
|
(user_id, int(run["id"]), candidate["identifier"]),
|
|
).fetchone()
|
|
if row is None:
|
|
raise RuntimeError("策略跟踪写入失败")
|
|
return int(row["id"])
|
|
|
|
def tracks(self, connection: sqlite3.Connection, user_id: int) -> tuple[sqlite3.Row, ...]:
|
|
return tuple(
|
|
connection.execute(
|
|
"""
|
|
SELECT track.*, run.mode, run.strategy_id
|
|
FROM strategy_tracks track
|
|
JOIN screener_runs run ON run.id = track.run_id
|
|
WHERE track.user_id = ? ORDER BY track.added_at DESC, track.id DESC
|
|
""",
|
|
(user_id,),
|
|
).fetchall()
|
|
)
|
|
|
|
def track_bars(self, connection: sqlite3.Connection, track_id: int) -> tuple[sqlite3.Row, ...]:
|
|
return tuple(
|
|
connection.execute(
|
|
"""
|
|
SELECT * FROM strategy_track_bars
|
|
WHERE track_id = ? ORDER BY trade_date
|
|
""",
|
|
(track_id,),
|
|
).fetchall()
|
|
)
|
|
|
|
def tracked_before(
|
|
self, connection: sqlite3.Connection, trade_date: str
|
|
) -> tuple[sqlite3.Row, ...]:
|
|
return tuple(
|
|
connection.execute(
|
|
"""
|
|
SELECT * FROM strategy_tracks
|
|
WHERE selection_date < ? ORDER BY id
|
|
""",
|
|
(trade_date,),
|
|
).fetchall()
|
|
)
|
|
|
|
def save_track_bar(
|
|
self,
|
|
connection: sqlite3.Connection,
|
|
track_id: int,
|
|
trade_date: str,
|
|
row: dict[str, Any],
|
|
) -> None:
|
|
connection.execute(
|
|
"""
|
|
INSERT INTO strategy_track_bars (
|
|
track_id, trade_date, open, high, low, close
|
|
) VALUES (?, ?, ?, ?, ?, ?)
|
|
ON CONFLICT(track_id, trade_date) DO UPDATE SET
|
|
open = excluded.open,
|
|
high = excluded.high,
|
|
low = excluded.low,
|
|
close = excluded.close
|
|
""",
|
|
(
|
|
track_id,
|
|
trade_date,
|
|
float(row["open"]),
|
|
float(row["high"]),
|
|
float(row["low"]),
|
|
float(row["close"]),
|
|
),
|
|
)
|
|
|
|
def record_track_event(
|
|
self, connection: sqlite3.Connection, track_id: int, milestone: str
|
|
) -> bool:
|
|
cursor = connection.execute(
|
|
"""
|
|
INSERT OR IGNORE INTO strategy_track_events (track_id, milestone, created_at)
|
|
VALUES (?, ?, ?)
|
|
""",
|
|
(track_id, milestone, _now()),
|
|
)
|
|
return cursor.rowcount > 0
|
|
|
|
def remove_track(self, connection: sqlite3.Connection, user_id: int, track_id: int) -> bool:
|
|
cursor = connection.execute(
|
|
"DELETE FROM strategy_tracks WHERE id = ? AND user_id = ?",
|
|
(track_id, user_id),
|
|
)
|
|
return cursor.rowcount > 0
|
|
|
|
|
|
def decode_run(row: sqlite3.Row | None) -> dict[str, Any] | None:
|
|
if row is None:
|
|
return None
|
|
result = dict(row)
|
|
result["missing_fields"] = json.loads(str(row["missing_fields_json"]))
|
|
result["items"] = json.loads(str(row["result_json"]))
|
|
result.pop("missing_fields_json", None)
|
|
result.pop("result_json", None)
|
|
return result
|
|
|
|
|
|
def decode_custom(row: sqlite3.Row) -> dict[str, Any]:
|
|
result = dict(row)
|
|
result["formula"] = json.loads(str(row["formula_json"]))
|
|
result.pop("formula_json", None)
|
|
return result
|
|
|
|
|
|
def decode_track(row: sqlite3.Row, bars: tuple[sqlite3.Row, ...]) -> dict[str, Any]:
|
|
result = dict(row)
|
|
entry = float(row["entry_price"])
|
|
closes = [float(item["close"]) for item in bars]
|
|
highs = [float(item["high"]) for item in bars]
|
|
lows = [float(item["low"]) for item in bars]
|
|
result["t1_open_return"] = (
|
|
round((float(bars[0]["open"]) / entry - 1) * 100, 2) if bars else None
|
|
)
|
|
for index in (1, 3, 5):
|
|
result[f"t{index}_return"] = (
|
|
round((closes[index - 1] / entry - 1) * 100, 2) if len(closes) >= index else None
|
|
)
|
|
result["max_gain"] = round((max(highs) / entry - 1) * 100, 2) if highs else None
|
|
result["max_drawdown"] = round((min(lows) / entry - 1) * 100, 2) if lows else None
|
|
result["observed_days"] = len(bars)
|
|
return result
|
|
|
|
|
|
def _json(value: Any) -> str:
|
|
return json.dumps(value, ensure_ascii=False, separators=(",", ":"), sort_keys=True)
|
|
|
|
|
|
def _now() -> str:
|
|
return datetime.now().astimezone().isoformat(timespec="seconds")
|