Files
signal-platform/app/services/watchlist_service.py
T

271 lines
8.4 KiB
Python

"""Watchlist service.
A purely user-curated watchlist: the user adds and removes tickers, capped at
WATCHLIST_MAX entries. Each entry is enriched on read with its composite score,
best trade setup, active S/R levels, and latest price + day-over-day move.
"""
from __future__ import annotations
import logging
from collections import defaultdict
from datetime import datetime, timezone
from sqlalchemy import func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.exceptions import DuplicateError, NotFoundError, ValidationError
from app.models.ohlcv import OHLCVRecord
from app.models.score import CompositeScore, DimensionScore
from app.models.sr_level import SRLevel
from app.models.ticker import Ticker
from app.models.trade_setup import TradeSetup
from app.models.watchlist import WatchlistEntry
logger = logging.getLogger(__name__)
WATCHLIST_MAX = 20
async def _get_ticker(db: AsyncSession, symbol: str) -> Ticker:
normalised = symbol.strip().upper()
result = await db.execute(select(Ticker).where(Ticker.symbol == normalised))
ticker = result.scalar_one_or_none()
if ticker is None:
raise NotFoundError(f"Ticker not found: {normalised}")
return ticker
async def add_manual_entry(
db: AsyncSession,
user_id: int,
symbol: str,
) -> WatchlistEntry:
"""Add a ticker to the user's watchlist.
Raises DuplicateError if already on the watchlist.
Raises ValidationError if the watchlist cap is reached.
"""
ticker = await _get_ticker(db, symbol)
existing = await db.execute(
select(WatchlistEntry).where(
WatchlistEntry.user_id == user_id,
WatchlistEntry.ticker_id == ticker.id,
)
)
if existing.scalar_one_or_none() is not None:
raise DuplicateError(f"Ticker already on watchlist: {ticker.symbol}")
count_result = await db.execute(
select(func.count()).select_from(WatchlistEntry).where(
WatchlistEntry.user_id == user_id,
)
)
total = count_result.scalar() or 0
if total >= WATCHLIST_MAX:
raise ValidationError(
f"Watchlist cap reached ({WATCHLIST_MAX}). "
"Remove an entry before adding a new one."
)
entry = WatchlistEntry(
user_id=user_id,
ticker_id=ticker.id,
entry_type="manual",
added_at=datetime.now(timezone.utc),
)
db.add(entry)
await db.commit()
await db.refresh(entry)
return entry
async def remove_entry(
db: AsyncSession,
user_id: int,
symbol: str,
) -> None:
"""Remove a ticker from the user's watchlist."""
ticker = await _get_ticker(db, symbol)
result = await db.execute(
select(WatchlistEntry).where(
WatchlistEntry.user_id == user_id,
WatchlistEntry.ticker_id == ticker.id,
)
)
entry = result.scalar_one_or_none()
if entry is None:
raise NotFoundError(f"Ticker not on watchlist: {ticker.symbol}")
await db.delete(entry)
await db.commit()
async def _enrich_entries(
db: AsyncSession,
rows: list[tuple[WatchlistEntry, str]],
) -> list[dict]:
"""Build watchlist rows from a fixed set of bulk lookups."""
if not rows:
return []
ticker_ids = [entry.ticker_id for entry, _ in rows]
comps_result = await db.execute(
select(CompositeScore).where(CompositeScore.ticker_id.in_(ticker_ids))
)
comps = {score.ticker_id: score for score in comps_result.scalars()}
dims_result = await db.execute(
select(DimensionScore).where(DimensionScore.ticker_id.in_(ticker_ids))
)
dims_by_ticker: dict[int, list[dict]] = defaultdict(list)
for score in dims_result.scalars():
dims_by_ticker[score.ticker_id].append(
{"dimension": score.dimension, "score": score.score}
)
ranked_setups = (
select(
TradeSetup.id,
func.row_number()
.over(
partition_by=TradeSetup.ticker_id,
order_by=TradeSetup.rr_ratio.desc(),
)
.label("rank"),
)
.where(TradeSetup.ticker_id.in_(ticker_ids))
.subquery()
)
setup_result = await db.execute(
select(TradeSetup)
.join(ranked_setups, TradeSetup.id == ranked_setups.c.id)
.where(ranked_setups.c.rank == 1)
)
best_setups = {setup.ticker_id: setup for setup in setup_result.scalars()}
levels_result = await db.execute(
select(SRLevel)
.where(SRLevel.ticker_id.in_(ticker_ids))
.order_by(SRLevel.ticker_id, SRLevel.strength.desc())
)
levels_by_ticker: dict[int, list[dict]] = defaultdict(list)
for level in levels_result.scalars():
levels_by_ticker[level.ticker_id].append(
{
"price_level": level.price_level,
"type": level.type,
"strength": level.strength,
}
)
ranked_prices = (
select(
OHLCVRecord.ticker_id,
OHLCVRecord.close,
OHLCVRecord.date,
func.row_number()
.over(
partition_by=OHLCVRecord.ticker_id,
order_by=OHLCVRecord.date.desc(),
)
.label("rank"),
)
.where(OHLCVRecord.ticker_id.in_(ticker_ids))
.subquery()
)
prices_result = await db.execute(
select(
ranked_prices.c.ticker_id,
ranked_prices.c.close,
ranked_prices.c.date,
)
.where(ranked_prices.c.rank <= 2)
.order_by(ranked_prices.c.ticker_id, ranked_prices.c.rank)
)
prices_by_ticker: dict[int, list[tuple[float, datetime]]] = defaultdict(list)
for ticker_id, close, price_date in prices_result.all():
prices_by_ticker[ticker_id].append((close, price_date))
entries: list[dict] = []
for entry, symbol in rows:
ticker_id = entry.ticker_id
comp = comps.get(ticker_id)
setup = best_setups.get(ticker_id)
bars = prices_by_ticker[ticker_id]
last_close = bars[0][0] if bars else None
prev_close = bars[1][0] if len(bars) > 1 else None
entries.append(
{
"symbol": symbol,
"entry_type": entry.entry_type,
"composite_score": comp.score if comp else None,
"dimensions": dims_by_ticker[ticker_id],
"rr_ratio": setup.rr_ratio if setup else None,
"rr_direction": setup.direction if setup else None,
"momentum_percentile": setup.momentum_percentile if setup else None,
"strategy_rank": setup.strategy_rank if setup else None,
"sr_levels": levels_by_ticker[ticker_id],
"last_close": last_close,
"change_pct": (
(last_close - prev_close) / prev_close * 100
if last_close is not None and prev_close
else None
),
"price_date": bars[0][1] if bars else None,
"added_at": entry.added_at,
}
)
return entries
async def get_watchlist(
db: AsyncSession,
user_id: int,
sort_by: str = "composite",
) -> list[dict]:
"""Get the user's watchlist with enriched data.
sort_by: "composite", "rr", "change", or a dimension name
(e.g. "technical", "sr_quality", "sentiment", "fundamental", "momentum").
"""
stmt = (
select(WatchlistEntry, Ticker.symbol)
.join(Ticker, WatchlistEntry.ticker_id == Ticker.id)
.where(WatchlistEntry.user_id == user_id)
)
result = await db.execute(stmt)
rows = result.all()
entries = await _enrich_entries(db, rows)
# Sort
if sort_by == "composite":
entries.sort(
key=lambda e: e["composite_score"] if e["composite_score"] is not None else -1,
reverse=True,
)
elif sort_by == "rr":
entries.sort(
key=lambda e: e["rr_ratio"] if e["rr_ratio"] is not None else -1,
reverse=True,
)
elif sort_by == "change":
entries.sort(
key=lambda e: e["change_pct"] if e["change_pct"] is not None else float("-inf"),
reverse=True,
)
else:
# Sort by a specific dimension score
def _dim_sort_key(e: dict) -> float:
for d in e["dimensions"]:
if d["dimension"] == sort_by:
return d["score"]
return -1.0
entries.sort(key=_dim_sort_key, reverse=True)
return entries