stock-scraper / database /repository.py
sbasu2512's picture
intertwine the services
57384dd
Raw
History Blame Contribute Delete
5 kB
from datetime import date, datetime
from typing import Iterable
from database.connection import connect
class PredictionRepository:
def __init__(self, database_url: str):
self.database_url = database_url
def connection(self):
return connect(self.database_url)
# ─────────────────────────────────────────────
# Predictions
# ─────────────────────────────────────────────
def insert_predictions(
self,
predictions: Iterable[dict],
) -> int:
sql = """
INSERT INTO predictions (
model_version,
feature_version,
prediction_date,
prediction_timestamp,
symbol,
predicted_probability,
rank,
target_threshold,
target_horizon_days,
prediction_close
)
VALUES (
:model_version,
:feature_version,
:prediction_date,
:prediction_timestamp,
:symbol,
:predicted_probability,
:rank,
:target_threshold,
:target_horizon_days,
:prediction_close
)
ON CONFLICT (
prediction_date,
symbol,
model_version
)
DO NOTHING
"""
rows = list(predictions)
if not rows:
return 0
with self.connection() as conn:
with conn.cursor() as cur:
inserted = 0
for row in rows:
cur.execute(sql.replace("DO NOTHING", "DO NOTHING RETURNING id"), row)
if cur.fetchone() is not None:
inserted += 1
return inserted
# ─────────────────────────────────────────────
# Unresolved predictions
# ─────────────────────────────────────────────
def get_unresolved_predictions(self):
sql = """
SELECT *
FROM predictions
WHERE resolved_at IS NULL
ORDER BY prediction_date, symbol
"""
with self.connection() as conn:
with conn.cursor() as cur:
cur.execute(sql)
return cur.fetchall()
# ─────────────────────────────────────────────
# Resolve prediction
# ─────────────────────────────────────────────
def resolve_prediction(
self,
prediction_id: int,
entry_date: date,
entry_price: float,
evaluation_end_date: date,
max_future_close: float,
actual_return: float,
actual_label: int,
):
sql = """
UPDATE predictions
SET
entry_date = :entry_date,
entry_open = :entry_open,
evaluation_end_date = :evaluation_end_date,
max_close_5d = :max_close_5d,
actual_return = :actual_return,
actual_label = :actual_label,
resolved_at = STRFTIME('%Y-%m-%dT%H:%M:%fZ', 'now')
WHERE id = :prediction_id
"""
params = {
"prediction_id": prediction_id,
"entry_date": entry_date,
"entry_open": entry_price,
"evaluation_end_date": evaluation_end_date,
"max_close_5d": max_future_close,
"actual_return": actual_return,
"actual_label": actual_label,
}
with self.connection() as conn:
with conn.cursor() as cur:
cur.execute(sql, params)
# ─────────────────────────────────────────────
# Recent resolved predictions
# ─────────────────────────────────────────────
def get_recent_resolved(
self,
days: int = 20,
):
sql = """
SELECT *
FROM predictions
WHERE
resolved_at IS NOT NULL
AND prediction_date >= date('now', ?)
ORDER BY prediction_date DESC
"""
with self.connection() as conn:
with conn.cursor() as cur:
cur.execute(sql, (f"-{days} days",))
return cur.fetchall()