Initial commit: Trading Portal - FastAPI + React + PostgreSQL

This commit is contained in:
2026-07-03 13:08:09 +00:00
commit 34a1e91541
198 changed files with 35110 additions and 0 deletions
View File
+367
View File
@@ -0,0 +1,367 @@
"""Alert service — evaluates multi-condition user alerts against current market data.
Each alert consists of a list of conditions (AND logic — all must pass). When a
new signal is generated, this service checks all active alerts for that user
and returns those whose conditions are fully satisfied.
"""
from __future__ import annotations
import logging
from collections.abc import Sequence
from typing import Any
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.alert import AlertCondition
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Supported indicators and their extraction helpers
# ---------------------------------------------------------------------------
SUPPORTED_INDICATORS = {
"rsi", "macd", "bb_width", "volume", "price", "sma", "ema", "momentum",
}
SUPPORTED_OPERATORS = {">", "<", ">=", "<=", "==", "cross_above", "cross_below"}
def _get_price(indicators: dict | None) -> float | None:
"""Extract the latest close price from indicators."""
if not indicators:
return None
return indicators.get("close")
def _get_rsi(indicators: dict | None) -> float | None:
"""Extract latest RSI value (rsi_14)."""
if not indicators:
return None
rsi_series = indicators.get("rsi_14")
if isinstance(rsi_series, list) and len(rsi_series) > 0:
return rsi_series[-1]
return None
def _get_macd(indicators: dict | None) -> dict | None:
"""Extract MACD data."""
if not indicators:
return None
return indicators.get("macd")
def _get_bb_width(indicators: dict | None) -> float | None:
"""Compute current Bollinger Band width (upper - lower)."""
if not indicators:
return None
bb = indicators.get("bollinger_bands")
if not bb:
return None
upper = bb.get("upper", [None])[-1]
lower = bb.get("lower", [None])[-1]
if upper is not None and lower is not None:
return float(upper) - float(lower)
return None
def _get_volume(indicators: dict | None, condition: dict) -> float | None:
"""Extract volume-related value based on condition type.
Types:
- 'avg_multiplier': compare current volume to SMA(period) average
- 'absolute': raw volume value
"""
if not indicators:
return None
vol_type = condition.get("type", "absolute")
if vol_type == "avg_multiplier":
period = condition.get("period", 20)
volume_series = indicators.get("volume")
if isinstance(volume_series, list) and len(volume_series) >= period:
recent = volume_series[-period:]
avg = sum(recent) / len(recent)
current = volume_series[-1]
return current / avg if avg > 0 else None
return None
# absolute
volume_series = indicators.get("volume")
if isinstance(volume_series, list) and len(volume_series) > 0:
return volume_series[-1]
return None
def _get_sma(indicators: dict | None, condition: dict) -> float | None:
"""Extract SMA value. Uses condition['value'] as the period, default 20."""
if not indicators:
return None
period = condition.get("period", 20)
key = f"sma_{period}"
sma_series = indicators.get(key)
if isinstance(sma_series, list) and len(sma_series) > 0:
return sma_series[-1]
# Fallback: check generic "sma_20" or "sma_50"
for fallback in ("sma_20", "sma_50"):
fb = indicators.get(fallback)
if isinstance(fb, list) and len(fb) > 0:
return fb[-1]
return None
def _get_ema(indicators: dict | None, condition: dict) -> float | None:
"""Extract EMA value."""
if not indicators:
return None
period = condition.get("period", 20)
key = f"ema_{period}"
ema_series = indicators.get(key)
if isinstance(ema_series, list) and len(ema_series) > 0:
return ema_series[-1]
return None
def _get_momentum(indicators: dict | None) -> float | None:
"""Extract momentum (rate of change) from close prices."""
if not indicators:
return None
# Try to compute from sma_20 as a proxy if we have enough values
close_series = indicators.get("sma_20")
if isinstance(close_series, list) and len(close_series) >= 2:
prev = close_series[-2]
curr = close_series[-1]
if prev and prev > 0:
return (curr - prev) / prev * 100
return None
# ---------------------------------------------------------------------------
# Condition evaluation
# ---------------------------------------------------------------------------
def _extract_indicator_value(indicator: str, indicators: dict | None,
condition: dict) -> float | None:
"""Route the indicator name to the correct extraction function."""
extractors = {
"price": lambda: _get_price(indicators),
"rsi": lambda: _get_rsi(indicators),
"macd": lambda: _get_macd_value(indicators),
"bb_width": lambda: _get_bb_width(indicators),
"volume": lambda: _get_volume(indicators, condition),
"sma": lambda: _get_sma(indicators, condition),
"ema": lambda: _get_ema(indicators, condition),
"momentum": lambda: _get_momentum(indicators),
}
fn = extractors.get(indicator)
if fn is None:
logger.warning("Unknown indicator '%s'", indicator)
return None
return fn()
def _get_macd_value(indicators: dict | None) -> float | None:
"""Extract latest MACD histogram value (macd - signal)."""
macd_data = _get_macd(indicators)
if not macd_data:
return None
macd_line = macd_data.get("macd_line", [])
signal_line = macd_data.get("signal_line", [])
if len(macd_line) > 0 and len(signal_line) > 0:
m = macd_line[-1]
s = signal_line[-1]
if m is not None and s is not None:
return float(m) - float(s)
return None
def _compare_values(current: float, target: float, operator: str) -> bool:
"""Compare two numeric values using the given operator."""
if operator == ">":
return current > target
elif operator == "<":
return current < target
elif operator == ">=":
return current >= target
elif operator == "<=":
return current <= target
elif operator == "==":
return abs(current - target) < 1e-9
return False
def _check_cross(series_a: list[float] | None, series_b: float | None,
operator: str) -> bool:
"""Check if series_a crosses above/below a fixed value.
'cross_above': previous <= value < current (or prev < value <= current)
'cross_below': previous >= value > current (or prev > value >= current)
"""
if not series_a or len(series_a) < 2 or series_b is None:
return False
prev = series_a[-2]
curr = series_a[-1]
if operator == "cross_above":
return prev <= series_b < curr
elif operator == "cross_below":
return prev >= series_b > curr
return False
async def evaluate_single_condition(condition: dict,
symbol_data: dict | None) -> bool:
"""Evaluate a single alert condition against symbol indicator data.
Parameters
----------
condition : dict
One condition object from the alert's conditions array, e.g.:
``{'indicator': 'rsi', 'operator': '>', 'value': 70, 'timeframe': '1h'}``
symbol_data : dict | None
The current indicator data for this symbol (output of
``get_indicators`` or similar).
Returns
-------
bool
``True`` if the condition is satisfied, ``False`` otherwise.
"""
indicator = condition.get("indicator")
operator = condition.get("operator")
target_value = condition.get("value")
if not indicator or not operator or target_value is None:
logger.debug("Incomplete condition: %s", condition)
return False
# Handle cross_above / cross_below — these need the full series
if operator in ("cross_above", "cross_below"):
series = _get_value_series(indicator, symbol_data, condition)
return _check_cross(series, target_value, operator)
current = _extract_indicator_value(indicator, symbol_data, condition)
if current is None:
logger.debug("Could not extract indicator '%s' from symbol data", indicator)
return False
return _compare_values(current, float(target_value), operator)
def _get_value_series(indicator: str, indicators: dict | None,
condition: dict) -> list[float] | None:
"""Get the full time-series for an indicator (needed for cross detection)."""
if not indicators:
return None
if indicator == "rsi":
series = indicators.get("rsi_14")
elif indicator == "price":
series = indicators.get("sma_20")
elif indicator == "volume":
series = indicators.get("volume")
elif indicator == "sma":
period = condition.get("period", 20)
series = indicators.get(f"sma_{period}")
elif indicator == "ema":
period = condition.get("period", 20)
series = indicators.get(f"ema_{period}")
else:
return None
if isinstance(series, list) and len(series) >= 2:
return [float(v) for v in series if v is not None]
return None
# ---------------------------------------------------------------------------
# Alert evaluation (all conditions for a set of symbols)
# ---------------------------------------------------------------------------
async def check_alert_conditions(db: AsyncSession,
symbol_data_map: dict[str, dict | None],
user_id) -> list[AlertCondition]:
"""Check all active alerts and return those whose conditions are satisfied.
Parameters
----------
db : AsyncSession
Database session.
symbol_data_map : dict[str, dict | None]
Mapping from symbol (e.g. ``'BTC/USDT'``) to its indicator data dict.
user_id
The user ID to check alerts for.
Returns
-------
list[AlertCondition]
All alerts that have fired (conditions satisfied).
"""
# Load all active alerts for this user
result = await db.execute(
select(AlertCondition).where(
AlertCondition.user_id == user_id,
AlertCondition.is_active == True,
)
)
alerts: Sequence[AlertCondition] = result.scalars().all()
fired: list[AlertCondition] = []
for alert in alerts:
try:
triggered = await _evaluate_alert_conditions(alert, symbol_data_map)
if triggered:
fired.append(alert)
except Exception:
logger.exception("Error evaluating alert %s", alert.id)
return fired
async def _evaluate_alert_conditions(
alert: AlertCondition,
symbol_data_map: dict[str, dict | None],
) -> bool:
"""Evaluate all conditions of an alert (AND logic)."""
conditions: list[dict] = alert.conditions or []
if not conditions:
return False
# Group conditions by the symbol / timeframe they reference
# Each condition can optionally specify a 'timeframe' key.
# For simplicity, we evaluate each condition against the symbol_data_map.
# All conditions must pass for the alert to fire.
for cond in conditions:
# Determine which symbol data to use (default to first available)
timeframe = cond.get("timeframe")
# Use any available symbol data — for now check each symbol
passed = False
for symbol, data in symbol_data_map.items():
if await evaluate_single_condition(cond, data):
passed = True
break
if not passed:
return False
return True
async def check_and_notify_alerts(
db: AsyncSession,
symbol: str,
symbol_data: dict | None,
user_id,
) -> None:
"""Convenience: check alerts for a single symbol+user and log results.
Called from signal_service after each signal.
"""
symbol_data_map = {symbol: symbol_data}
fired = await check_alert_conditions(db, symbol_data_map, user_id)
if fired:
names = [a.name for a in fired]
logger.info(
"🔔 Alerts triggered for user %s on %s: %s",
user_id, symbol, names,
)
# Actual notification sending is handled by the caller
+82
View File
@@ -0,0 +1,82 @@
"""Audit log service for recording and querying audit events."""
from __future__ import annotations
import logging
from uuid import UUID
from sqlalchemy import desc, func as sa_func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.models.audit_log import AuditLog
logger = logging.getLogger(__name__)
async def log_action(
db: AsyncSession,
user_id: UUID | None,
action: str,
resource: str,
details: dict | None = None,
) -> AuditLog:
"""Record an audit log entry.
Args:
db: Database session.
user_id: UUID of the user who performed the action, or None.
action: Action type (e.g. 'trade_open', 'signal_generated').
resource: Affected resource description.
details: Optional extra JSON info.
Returns:
The created AuditLog instance.
"""
entry = AuditLog(
user_id=user_id,
action=action,
resource=resource,
details=details,
)
db.add(entry)
await db.flush()
logger.debug(
"Audit log: action=%s resource=%s user=%s",
action, resource, user_id,
)
return entry
async def get_audit_logs(
db: AsyncSession,
limit: int = 100,
offset: int = 0,
action: str | None = None,
) -> tuple[list[AuditLog], int]:
"""Fetch audit log entries with pagination and optional action filter.
Args:
db: Database session.
limit: Maximum number of entries to return.
offset: Number of entries to skip.
action: Optional action type filter.
Returns:
A tuple of (list of AuditLog entries, total count).
"""
base_query = select(AuditLog).order_by(desc(AuditLog.created_at))
if action:
base_query = base_query.where(AuditLog.action == action)
# Get total count
count_query = select(sa_func.count()).select_from(base_query.subquery())
total_result = await db.execute(count_query)
total = total_result.scalar() or 0
# Get paginated results
query = base_query.limit(limit).offset(offset)
result = await db.execute(query)
entries = list(result.scalars().all())
return entries, total
+332
View File
@@ -0,0 +1,332 @@
from __future__ import annotations
import re
from uuid import UUID
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.exceptions import (
ConflictException,
InvalidCredentialsException,
InvalidTokenException,
NotFoundException,
)
from app.core.security import (
create_access_token,
create_refresh_token,
decode_token,
generate_token_hash,
hash_password_async,
verify_password_async,
)
from app.models import RefreshToken, User
from app.schemas import (
LoginRequest,
RegisterRequest,
TokenResponse,
UserResponse,
UserSessionResponse,
)
# P1-23: Backend password strength validation
def _validate_password_strength(password: str) -> None:
"""Raise ConflictException if password doesn't meet minimum strength."""
if len(password) < 8:
raise ConflictException(detail="Password must be at least 8 characters")
if not re.search(r"[a-z]", password):
raise ConflictException(detail="Password must contain at least one lowercase letter")
if not re.search(r"[A-Z]", password):
raise ConflictException(detail="Password must contain at least one uppercase letter")
if not re.search(r"\d", password):
raise ConflictException(detail="Password must contain at least one digit")
if not re.search(r"[!@#$%^&*()_+\-=\[\]{};':\"\\|,.<>/?]", password):
raise ConflictException(detail="Password must contain at least one special character")
async def register(
db: AsyncSession,
req: RegisterRequest,
) -> UserResponse:
"""Register a new user account.
Checks username / email uniqueness, hashes the password, creates the
``User`` row together with an empty ``Watchlist``, and returns the
new user's public profile.
"""
# --- uniqueness checks ---------------------------------------------------
existing_username = await db.execute(
select(User).where(User.username == req.username)
)
if existing_username.scalar_one_or_none() is not None:
raise ConflictException(detail=f"Username '{req.username}' is already taken")
existing_email = await db.execute(
select(User).where(User.email == req.email)
)
if existing_email.scalar_one_or_none() is not None:
raise ConflictException(detail=f"Email '{req.email}' is already registered")
# --- P1-23: password strength validation ---------------------------------
_validate_password_strength(req.password)
# --- create user ---------------------------------------------------------
user = User(
username=req.username,
email=req.email,
password_hash=await hash_password_async(req.password),
)
db.add(user)
await db.flush() # flush so user.id is available
# --- create empty watchlist entry ----------------------------------------
# The Watchlist model requires a symbol_id; an empty watchlist is
# represented here by *not* creating any rows. If the domain later
# requires an explicit "empty" row, adjust accordingly.
# For now we skip creating a Watchlist row since it needs a symbol_id.
await db.commit()
await db.refresh(user)
return UserResponse(
id=user.id,
username=user.username,
email=user.email,
display_name=user.display_name,
is_active=user.is_active,
is_admin=user.is_admin,
role=user.role,
created_at=user.created_at,
)
async def login(
db: AsyncSession,
req: LoginRequest,
user_agent: str = "",
ip_address: str = "",
) -> TokenResponse:
"""Authenticate a user by username/password and issue a token pair.
Validates credentials, creates an access + refresh JWT pair, stores a
hashed refresh-token record in the database, and returns the tokens.
"""
# --- locate user ---------------------------------------------------------
result = await db.execute(
select(User).where(User.username == req.username)
)
user = result.scalar_one_or_none()
if user is None or not await verify_password_async(req.password, user.password_hash):
raise InvalidCredentialsException()
# --- issue tokens --------------------------------------------------------
token_data: dict = {"sub": str(user.id)}
access_token = create_access_token(data=token_data)
refresh_token = create_refresh_token(data=token_data)
token_hash = generate_token_hash(refresh_token)
# Decode the refresh token to read its expiry and jti
decoded = decode_token(refresh_token)
expires_at = decoded["exp"]
from datetime import datetime, timezone
# --- store refresh token record ------------------------------------------
rt = RefreshToken(
user_id=user.id,
token_hash=token_hash,
expires_at=datetime.fromtimestamp(expires_at, tz=timezone.utc),
user_agent=user_agent or None,
ip_address=ip_address or None,
)
db.add(rt)
await db.commit()
return TokenResponse(
access_token=access_token,
refresh_token=refresh_token,
token_type="bearer",
)
async def refresh_token(
db: AsyncSession,
refresh_token_str: str,
) -> TokenResponse:
"""Refresh an expired access token using a valid refresh token (rotation).
Verifies the refresh token signature, checks that it hasn't been revoked
or expired, revokes the old record, and issues a brand-new token pair.
"""
# --- decode & validate claims --------------------------------------------
try:
payload = decode_token(refresh_token_str)
except Exception:
raise InvalidTokenException(detail="Invalid refresh token")
if payload.get("type") != "refresh":
raise InvalidTokenException(detail="Token is not a refresh token")
sub: str | None = payload.get("sub")
jti: str | None = payload.get("jti")
if not sub or not jti:
raise InvalidTokenException(detail="Invalid refresh token payload")
# --- look up stored record -----------------------------------------------
token_hash = generate_token_hash(refresh_token_str)
result = await db.execute(
select(RefreshToken).where(RefreshToken.token_hash == token_hash)
)
stored = result.scalar_one_or_none()
if stored is None:
raise InvalidTokenException(detail="Refresh token not found")
if stored.revoked:
raise InvalidTokenException(detail="Refresh token has been revoked")
from datetime import datetime, timezone
if stored.expires_at < datetime.now(timezone.utc):
raise InvalidTokenException(detail="Refresh token has expired")
# --- revoke old token ----------------------------------------------------
stored.revoked = True
# --- issue new pair ------------------------------------------------------
token_data: dict = {"sub": sub}
new_access = create_access_token(data=token_data)
new_refresh = create_refresh_token(data=token_data)
new_hash = generate_token_hash(new_refresh)
decoded_new = decode_token(new_refresh)
new_expires_at = datetime.fromtimestamp(decoded_new["exp"], tz=timezone.utc)
rt = RefreshToken(
user_id=stored.user_id,
token_hash=new_hash,
expires_at=new_expires_at,
)
db.add(rt)
await db.commit()
return TokenResponse(
access_token=new_access,
refresh_token=new_refresh,
token_type="bearer",
)
async def logout(
db: AsyncSession,
refresh_token_str: str,
) -> None:
"""Revoke the given refresh token so it can no longer be used."""
token_hash = generate_token_hash(refresh_token_str)
result = await db.execute(
select(RefreshToken).where(RefreshToken.token_hash == token_hash)
)
stored = result.scalar_one_or_none()
if stored is None:
raise InvalidTokenException(detail="Refresh token not found")
stored.revoked = True
await db.commit()
async def get_user_sessions(
db: AsyncSession,
user_id: UUID,
) -> list[UserSessionResponse]:
"""Return all active (non-revoked, non-expired) sessions for a user."""
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
result = await db.execute(
select(RefreshToken)
.where(
RefreshToken.user_id == user_id,
RefreshToken.revoked == False, # noqa: E712
RefreshToken.expires_at > now,
)
.order_by(RefreshToken.created_at.desc())
)
tokens = result.scalars().all()
return [
UserSessionResponse(
id=t.id,
created_at=t.created_at,
user_agent=t.user_agent,
ip_address=str(t.ip_address) if t.ip_address else None,
is_current=False,
)
for t in tokens
]
async def revoke_session(
db: AsyncSession,
token_hash: str,
user_id: UUID,
) -> None:
"""Revoke a specific refresh token by hash, verifying it belongs to the user."""
result = await db.execute(
select(RefreshToken).where(
RefreshToken.token_hash == token_hash,
RefreshToken.user_id == user_id,
)
)
stored = result.scalar_one_or_none()
if stored is None:
raise NotFoundException(detail="Session not found")
stored.revoked = True
await db.commit()
async def change_password(
db: AsyncSession,
user_id: UUID,
old_password: str,
new_password: str,
) -> None:
"""Change a user's password after verifying the old one.
On success, all existing refresh tokens for the user are revoked,
forcing a fresh login on all devices.
"""
# --- fetch user ----------------------------------------------------------
result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if user is None:
raise NotFoundException(detail="User not found")
# --- verify old password -------------------------------------------------
if not await verify_password_async(old_password, user.password_hash):
raise InvalidCredentialsException(detail="Incorrect password")
# --- P1-23: validate new password strength --------------------------------
_validate_password_strength(new_password)
# --- update password -----------------------------------------------------
user.password_hash = await hash_password_async(new_password)
# --- revoke all refresh tokens -------------------------------------------
tokens_result = await db.execute(
select(RefreshToken).where(RefreshToken.user_id == user_id)
)
for token in tokens_result.scalars().all():
token.revoked = True
await db.commit()
+544
View File
@@ -0,0 +1,544 @@
"""Candle CRUD service with TTLCache and cursor-based pagination."""
from __future__ import annotations
import asyncio
import logging
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Optional
from cachetools import TTLCache
from sqlalchemy import and_, select, text
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import joinedload
from app.core.exceptions import NotFoundException
from app.database import async_session_factory
from app.exchange.factory import factory as exchange_factory
from app.exchange.types import CandleData
from app.models.candle import Candle
from app.models.exchange import Exchange
from app.models.symbol import Symbol
from app.schemas.candle import CandleListResponse, CandleResponse
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# In-memory candle cache
# ---------------------------------------------------------------------------
# Key format: f"{exchange_name}:{symbol}:{timeframe}"
candle_cache: TTLCache[str, list[CandleResponse]] = TTLCache(
maxsize=500,
ttl=60, # 1 minute
)
# ---------------------------------------------------------------------------
# Indicator cache — TTL varies by timeframe to avoid unnecessary recomputation
# ---------------------------------------------------------------------------
# Key format: f"{exchange_name}:{symbol}:{timeframe}"
indicator_cache: TTLCache[str, dict] = TTLCache(
maxsize=500,
ttl=300, # default 5 min (overridden per-key with custom TTL tracking)
)
def _indicator_cache_ttl(timeframe: str) -> int:
"""Return a sensible TTL (seconds) for indicator data per timeframe."""
tf_map = {
"1m": 30, # refresh every 30s
"5m": 120, # every 2 min
"15m": 300, # every 5 min
"30m": 300, # every 5 min
"1h": 600, # every 10 min
"4h": 1800, # every 30 min
"1d": 3600, # every 1 hour
}
return tf_map.get(timeframe, 300)
# In-progress tracker for debounce: set of "symbol:timeframe" keys
_indicator_in_progress: set[str] = set()
_indicator_cache_timestamps: dict[str, float] = {}
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _cache_key(exchange_name: str, symbol: str, timeframe: str) -> str:
return f"{exchange_name}:{symbol}:{timeframe}"
def _candle_to_response(candle: Candle) -> CandleResponse:
return CandleResponse(
symbol_id=candle.symbol_id,
timeframe=candle.timeframe,
timestamp=candle.timestamp,
open=candle.open,
high=candle.high,
low=candle.low,
close=candle.close,
volume=candle.volume,
)
def _candledata_to_response(cd: CandleData) -> CandleResponse:
return CandleResponse(
symbol_id=0, # Will be set when saved; placeholder for cache
timeframe=cd.timeframe,
timestamp=cd.timestamp,
open=cd.open,
high=cd.high,
low=cd.low,
close=cd.close,
volume=cd.volume,
)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
async def fetch_and_store_candles(
db: AsyncSession,
exchange_name: str,
symbol: str,
timeframe: str,
limit: int = 500,
) -> list[CandleData]:
"""Fetch candles from the exchange and persist them to the database.
Steps:
1. Get or create Exchange + Symbol rows in the DB.
2. Create an exchange adapter via ``ExchangeFactory``.
3. Fetch OHLCV data from the remote exchange.
4. Bulk-insert with ``ON CONFLICT DO NOTHING``.
5. Refresh the in-memory cache.
6. Return the list of ``CandleData`` as received from the exchange.
"""
# --- 1. Get or create Exchange ---
exch_result = await db.execute(
select(Exchange).where(Exchange.name == exchange_name)
)
exchange = exch_result.scalar_one_or_none()
if exchange is None:
exchange = Exchange(name=exchange_name, display_name=exchange_name.title())
db.add(exchange)
await db.flush()
logger.info("Created new exchange record: %s", exchange_name)
# --- Get or create Symbol ---
sym_result = await db.execute(
select(Symbol).where(
and_(Symbol.exchange_id == exchange.id, Symbol.symbol == symbol)
)
)
db_symbol = sym_result.scalar_one_or_none()
if db_symbol is None:
# Parse base/quote from symbol (e.g. "BTC/USDT" -> "BTC", "USDT")
parts = symbol.replace("-", "/").split("/")
base = parts[0] if len(parts) > 1 else symbol
quote = parts[1] if len(parts) > 1 else ""
db_symbol = Symbol(
exchange_id=exchange.id,
symbol=symbol,
base=base,
quote=quote,
is_active=True,
)
db.add(db_symbol)
await db.flush()
logger.info("Created new symbol record: %s on %s", symbol, exchange_name)
# --- 2. Create adapter & 3. Fetch candles ---
adapter = exchange_factory.create(exchange_name)
candles = await adapter.fetch_ohlcv(symbol, timeframe, limit)
if not candles:
logger.warning("No candles returned for %s:%s:%s", exchange_name, symbol, timeframe)
return candles
# --- 4. Bulk insert (ON CONFLICT DO NOTHING) ---
values = [
{
"symbol_id": db_symbol.id,
"timeframe": timeframe,
"timestamp": c.timestamp,
"open": c.open,
"high": c.high,
"low": c.low,
"close": c.close,
"volume": c.volume,
}
for c in candles
]
stmt = pg_insert(Candle).values(values)
stmt = stmt.on_conflict_do_nothing(
index_elements=["symbol_id", "timeframe", "timestamp"]
)
await db.execute(stmt)
await db.commit()
logger.info(
"Stored %d candles for %s:%s:%s",
len(candles),
exchange_name,
symbol,
timeframe,
)
# --- 5. Invoke after-fetch callbacks (for WS push etc.) ---
try:
from app.tasks.candle_fetcher import _after_fetch_callbacks
for c in candles:
candle_dict = {
"symbol": symbol,
"exchange": exchange_name,
"timeframe": c.timeframe,
"timestamp": c.timestamp,
"open": c.open,
"high": c.high,
"low": c.low,
"close": c.close,
"volume": c.volume,
}
for cb in _after_fetch_callbacks:
try:
await cb(exchange_name, symbol, c.timeframe, candle_dict)
except Exception:
logger.exception("After-fetch callback failed for %s:%s:%s", exchange_name, symbol, c.timeframe)
except Exception:
logger.debug("No after-fetch callbacks registered")
# --- 6. Update cache ---
cache_responses = [_candledata_to_response(cd) for cd in candles]
# Fix up symbol_id for cached items
for r in cache_responses:
r.symbol_id = db_symbol.id
candle_cache[_cache_key(exchange_name, symbol, timeframe)] = cache_responses
return candles
async def get_candles(
db: AsyncSession,
symbol: str,
exchange_name: str,
timeframe: str,
cursor: Optional[datetime] = None,
limit: int = 500,
) -> CandleListResponse:
"""Retrieve candles with cursor-based pagination.
Checks the in-memory ``TTLCache`` first. On a cache miss, queries
PostgreSQL using ``WHERE timestamp < cursor`` ordered descending.
"""
key = _cache_key(exchange_name, symbol, timeframe)
# --- 1. Check cache (only for non-cursor queries) ---
if cursor is None and key in candle_cache:
cached = candle_cache[key]
has_more = len(cached) > limit
return CandleListResponse(
candles=cached[:limit],
cursor=cached[limit - 1].timestamp.isoformat() if has_more and len(cached) > limit else None,
has_more=has_more,
)
# --- 2. Resolve symbol ---
sym_result = await db.execute(
select(Symbol)
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(and_(Exchange.name == exchange_name, Symbol.symbol == symbol))
)
db_symbol = sym_result.scalar_one_or_none()
if db_symbol is None:
raise NotFoundException(detail=f"Symbol {symbol} not found on {exchange_name}")
# --- 3. Build query with cursor-based pagination ---
query = (
select(Candle)
.where(
and_(
Candle.symbol_id == db_symbol.id,
Candle.timeframe == timeframe,
)
)
.order_by(Candle.timestamp.desc())
.limit(limit + 1) # Fetch one extra to determine has_more
)
if cursor is not None:
query = query.where(Candle.timestamp < cursor)
result = await db.execute(query)
rows = result.scalars().all()
# --- 3b. If no cached/DB data, fetch on-demand from exchange ---
if not rows:
logger.info("No DB candles for %s:%s:%s — fetching on-demand from exchange", exchange_name, symbol, timeframe)
try:
fetched = await fetch_and_store_candles(db, exchange_name, symbol, timeframe, limit)
if fetched:
# Re-query DB after fetch
result2 = await db.execute(query)
rows = result2.scalars().all()
except Exception as e:
logger.warning("On-demand fetch failed for %s:%s:%s: %s", exchange_name, symbol, timeframe, e)
has_more = len(rows) > limit
if has_more:
rows = rows[:limit]
candles = [_candle_to_response(r) for r in rows]
# --- 4. Determine next cursor ---
next_cursor: Optional[str] = None
if candles:
next_cursor = candles[-1].timestamp.isoformat()
# --- 5. Update cache on full (non-cursor) reads ---
if cursor is None:
candle_cache[key] = candles
return CandleListResponse(
candles=candles,
cursor=next_cursor if has_more else None,
has_more=has_more,
)
async def get_latest_candle(
db: AsyncSession,
symbol: str,
exchange_name: str,
timeframe: str,
) -> Optional[CandleData]:
"""Return the most recent candle for a symbol/exchange/timeframe."""
sym_result = await db.execute(
select(Symbol)
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(and_(Exchange.name == exchange_name, Symbol.symbol == symbol))
)
db_symbol = sym_result.scalar_one_or_none()
if db_symbol is None:
return None
result = await db.execute(
select(Candle)
.where(
and_(
Candle.symbol_id == db_symbol.id,
Candle.timeframe == timeframe,
)
)
.order_by(Candle.timestamp.desc())
.limit(1)
)
candle = result.scalar_one_or_none()
if candle is None:
return None
return CandleData(
symbol=symbol,
exchange=exchange_name,
timeframe=timeframe,
timestamp=candle.timestamp,
open=candle.open,
high=candle.high,
low=candle.low,
close=candle.close,
volume=candle.volume,
)
async def get_indicators(
db: AsyncSession,
symbol: str,
exchange_name: str,
timeframe: str,
) -> dict:
"""Compute technical indicators for the last 250 candles.
Results are cached in ``indicator_cache`` with per-timeframe TTL
(30s for 1m, 2min for 5m, 5min for 15m-30m, 10min for 1h, etc.)
"""
import time as _time
cache_key = f"{exchange_name}:{symbol}:{timeframe}"
# ── Cache check ──
now = _time.monotonic()
cached_ts = _indicator_cache_timestamps.get(cache_key, 0)
ttl = _indicator_cache_ttl(timeframe)
if cache_key in indicator_cache and (now - cached_ts) < ttl:
return indicator_cache[cache_key]
# ── Debounce: skip if already computing for this key ──
if cache_key in _indicator_in_progress:
logger.debug("Indicators for %s already computing — skipping duplicate call", cache_key)
# Return stale cache if available (better than nothing)
if cache_key in indicator_cache:
return indicator_cache[cache_key]
return {}
_indicator_in_progress.add(cache_key)
try:
from app.services.indicator_service import (
adx,
bollinger_bands,
detect_candlestick_patterns,
detect_divergence,
detect_fvg,
detect_market_regime,
ema,
ichimoku,
macd,
mfi,
obv,
obv_signal,
rsi,
sma,
stoch_rsi,
supertrend,
volume_breakout,
vwap,
)
sym_result = await db.execute(
select(Symbol)
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(and_(Exchange.name == exchange_name, Symbol.symbol == symbol))
)
db_symbol = sym_result.scalar_one_or_none()
if db_symbol is None:
raise NotFoundException(detail=f"Symbol {symbol} not found on {exchange_name}")
result = await db.execute(
select(Candle)
.where(
and_(
Candle.symbol_id == db_symbol.id,
Candle.timeframe == timeframe,
)
)
.order_by(Candle.timestamp.desc())
.limit(250)
)
candles = result.scalars().all()
# Reverse to ASC for indicator computation
candles.reverse()
if not candles:
return {}
close_prices = [float(c.close) for c in candles]
# Build candle dicts for VWAP, SuperTrend, Volume, SMC
candle_dicts = [
{
"high": float(c.high),
"low": float(c.low),
"close": float(c.close),
"open": float(c.open),
"volume": float(c.volume),
}
for c in candles
]
computed = {
"sma_20": sma(close_prices, 20),
"sma_50": sma(close_prices, 50),
"ema_12": ema(close_prices, 12),
"ema_26": ema(close_prices, 26),
"rsi_14": rsi(close_prices, 14),
"stoch_rsi": stoch_rsi(close_prices),
"macd": macd(close_prices),
"bollinger_bands": bollinger_bands(close_prices),
"vwap": vwap(candle_dicts),
"close": close_prices, # raw close prices for MTF & external use
}
# Add new indicators
computed["supertrend"] = supertrend(candle_dicts, period=10, multiplier=3.0)
computed["volume_breakout"] = volume_breakout(candle_dicts, period=20, multiplier=2.0)
# Add OBV (On-Balance Volume)
computed["obv"] = obv(candle_dicts)
obv_cross, obv_sma = obv_signal(computed["obv"], period=20)
computed["obv_crossover"] = obv_cross
computed["obv_sma"] = obv_sma
# Add Ichimoku Cloud
computed["ichimoku"] = ichimoku(candle_dicts)
# Add RSI divergence detection
rsi_vals = computed.get("rsi_14", [])
if len(close_prices) > 20 and len(rsi_vals) > 20:
rsi_div = detect_divergence(close_prices, rsi_vals, pivot_lookback=5)
computed["rsi_divergence"] = rsi_div
else:
computed["rsi_divergence"] = (None, None)
# Add Market Structure (SMC)
from app.services.indicator_service import market_structure
computed["market_structure"] = market_structure(candle_dicts, pivot_lookback=3)
# Add MACD divergence detection
macd_data = computed.get("macd", {})
macd_hist = macd_data.get("histogram", [None] * len(close_prices)) if macd_data else [None] * len(close_prices)
if len(close_prices) > 20 and len(macd_hist) > 20:
macd_div = detect_divergence(close_prices, macd_hist, pivot_lookback=5)
computed["macd_divergence"] = macd_div
else:
computed["macd_divergence"] = (None, None)
# Add ADX (Average Directional Index) + Market Regime
computed["adx_data"] = adx(candle_dicts, period=14)
atr_vals = computed.get("supertrend", {}).get("trend", None)
# Compute ATR% for regime detection
try:
from app.services.indicator_service import atr as _calc_atr
raw_atr = _calc_atr(candle_dicts, period=14)
last_atr = raw_atr[-1] if raw_atr and len(raw_atr) > 0 else None
last_close = close_prices[-1] if close_prices else 1
atr_pct_val = (last_atr / last_close * 100.0) if last_atr and last_close > 0 else None
except Exception:
atr_pct_val = None
# Extract high/low prices for regime detection
high_prices = [c["high"] for c in candle_dicts]
low_prices = [c["low"] for c in candle_dicts]
regime = detect_market_regime(
computed["adx_data"],
computed["bollinger_bands"],
atr_pct_val,
computed.get("volume_breakout"),
prices=close_prices,
highs=high_prices,
lows=low_prices,
)
computed["market_regime"] = regime
# Add MFI (Money Flow Index)
computed["mfi_14"] = mfi(candle_dicts, period=14)
# Add FVG (Fair Value Gap)
fvg_type, fvg_high, fvg_low = detect_fvg(candle_dicts, lookback=30)
computed["fvg"] = {"type": fvg_type, "gap_high": fvg_high, "gap_low": fvg_low}
# Add Candlestick Pattern Recognition
computed["candlestick_score"] = detect_candlestick_patterns(candle_dicts)
# ── Store in cache before returning ──
indicator_cache[cache_key] = computed
_indicator_cache_timestamps[cache_key] = _time.monotonic()
return computed
finally:
_indicator_in_progress.discard(cache_key)
File diff suppressed because it is too large Load Diff
+304
View File
@@ -0,0 +1,304 @@
"""Notification service for trading portal — Telegram + Discord push notifications.
Sends real-time alerts when trading signals are detected or auto-trades are
executed. Notifications are delivered based on each user's stored preferences.
"""
from __future__ import annotations
import logging
import os
from typing import Any
import httpx
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Constants
# ---------------------------------------------------------------------------
TELEGRAM_API_BASE = "https://api.telegram.org/bot{token}/sendMessage"
DEFAULT_TIMEOUT = 10.0 # seconds for each outbound HTTP request
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _get_telegram_bot_token() -> str | None:
"""Return the Telegram Bot API token from the environment."""
return os.environ.get("TELEGRAM_BOT_TOKEN") or None
def _get_default_telegram_chat_id() -> str | None:
"""Return the fallback Telegram chat ID from the environment."""
return os.environ.get("TELEGRAM_CHAT_ID") or None
def _build_telegram_message(signal_type: str, symbol: str, price: Any,
exchange_name: str) -> str:
"""Format a human-readable signal notification for Telegram."""
emoji_map = {
"STRONG_BUY": "🟢",
"BUY": "✅",
"STRONG_SELL": "🔴",
"SELL": "❌",
"CAUTION_LONG": "⚠️",
"CAUTION_SHORT": "⚠️",
"SQUEEZE_ALERT": "⚡",
}
emoji = emoji_map.get(signal_type, "📊")
return (
f"{emoji} *Trading Signal*\n"
f"━━━━━━━━━━━━━━━\n"
f"Type: {signal_type}\n"
f"Symbol: {symbol}\n"
f"Price: {price}\n"
f"Exchange: {exchange_name}"
)
def _build_trade_message(trade_direction: str, symbol: str, price: Any,
action: str) -> str:
"""Format a human-readable trade notification for Telegram."""
dir_emoji = "🟢" if trade_direction.upper() == "LONG" else "🔴"
action_emoji = "🟢" if action.upper() in ("ENTER", "OPEN", "BUY") else "🔴"
return (
f"{action_emoji} *Auto-Trade*\n"
f"━━━━━━━━━━━━━━━\n"
f"Action: {action}\n"
f"Direction: {dir_emoji} {trade_direction}\n"
f"Symbol: {symbol}\n"
f"Price: {price}"
)
# ---------------------------------------------------------------------------
# Core notification functions
# ---------------------------------------------------------------------------
async def send_telegram_notification(chat_id: str, message: str) -> bool:
"""Send a text message to a Telegram chat via the Bot API.
Parameters
----------
chat_id : str
Target Telegram chat / group / channel ID.
message : str
Plain-text or MarkdownV2-formatted message body (max 4096 chars).
Returns
-------
bool
``True`` if the message was delivered successfully, ``False``
otherwise (the error is logged but not raised).
"""
token = _get_telegram_bot_token()
if not token:
logger.warning("TELEGRAM_BOT_TOKEN not set — cannot send Telegram notification")
return False
url = TELEGRAM_API_BASE.format(token=token)
payload = {
"chat_id": chat_id,
"text": message,
"parse_mode": "Markdown",
"disable_web_page_preview": True,
}
try:
async with httpx.AsyncClient(timeout=DEFAULT_TIMEOUT) as client:
resp = await client.post(url, json=payload)
resp.raise_for_status()
data = resp.json()
if data.get("ok"):
logger.debug("Telegram notification sent to chat %s", chat_id)
return True
logger.warning(
"Telegram API returned ok=False: %s", data.get("description", "unknown")
)
return False
except httpx.TimeoutException:
logger.error("Timeout sending Telegram notification to chat %s", chat_id)
except httpx.HTTPStatusError as exc:
logger.error(
"Telegram API HTTP %d for chat %s: %s",
exc.response.status_code, chat_id, exc.response.text,
)
except httpx.RequestError as exc:
logger.error("Request error sending Telegram notification: %s", exc)
return False
async def send_discord_notification(webhook_url: str, message: str) -> bool:
"""Send a text message to a Discord channel via a webhook URL.
Parameters
----------
webhook_url : str
Full Discord webhook URL (including the token segment).
message : str
Message body (max 2000 characters for Discord).
Returns
-------
bool
``True`` if the message was delivered successfully, ``False``
otherwise.
"""
payload = {"content": message}
try:
async with httpx.AsyncClient(timeout=DEFAULT_TIMEOUT) as client:
resp = await client.post(webhook_url, json=payload)
resp.raise_for_status()
logger.debug("Discord notification sent to webhook")
return True
except httpx.TimeoutException:
logger.error("Timeout sending Discord notification")
except httpx.HTTPStatusError as exc:
logger.error(
"Discord webhook HTTP %d: %s",
exc.response.status_code, exc.response.text,
)
except httpx.RequestError as exc:
logger.error("Request error sending Discord notification: %s", exc)
return False
# ---------------------------------------------------------------------------
# High-level user-aware notification functions
# ---------------------------------------------------------------------------
async def notify_user(user: Any, signal_type: str, symbol: str, price: Any,
exchange_name: str) -> None:
"""Send signal notifications to a user based on their stored preferences.
This function checks the user's ``preferences`` JSON field for:
* ``notif_signal`` (``bool``) — whether to notify on new signals.
* ``notification_channels`` (``dict``) — channel configuration, e.g.:
.. code-block:: python
{
"telegram": {"chat_id": "123456789"},
"discord": {"webhook_url": "https://discord.com/api/webhooks/..."}
}
Parameters
----------
user : User
The SQLAlchemy ``User`` model instance (must have a ``preferences``
JSON column).
signal_type : str
One of ``STRONG_BUY``, ``BUY``, ``STRONG_SELL``, ``SELL``, etc.
symbol : str
Trading pair / symbol (e.g. ``BTC/USDT``).
price : Any
The price at which the signal was generated (will be stringified).
exchange_name : str
Exchange name (e.g. ``mexc``).
"""
prefs: dict = user.preferences or {}
# Respect the per-user opt-in for signal notifications
if not prefs.get("notif_signal", True):
logger.debug("User %s has signal notifications disabled", user.id)
return
channels: dict = prefs.get("notification_channels") or {}
message = _build_telegram_message(signal_type, symbol, price, exchange_name)
# ── Telegram ──────────────────────────────────────────────────────
telegram_cfg: dict | None = channels.get("telegram")
if telegram_cfg:
chat_id = telegram_cfg.get("chat_id")
if chat_id:
await send_telegram_notification(str(chat_id), message)
else:
logger.debug(
"User %s has telegram channel configured but missing chat_id", user.id
)
else:
# Fall back to the environment-level default chat ID
fallback_chat_id = _get_default_telegram_chat_id()
if fallback_chat_id:
await send_telegram_notification(fallback_chat_id, message)
# ── Discord ───────────────────────────────────────────────────────
discord_cfg: dict | None = channels.get("discord")
if discord_cfg:
webhook_url = discord_cfg.get("webhook_url")
if webhook_url:
await send_discord_notification(str(webhook_url), message)
else:
logger.debug(
"User %s has discord channel configured but missing webhook_url",
user.id,
)
async def notify_trade(user: Any, trade_direction: str, symbol: str,
price: Any, action: str) -> None:
"""Send trade-activity notifications to a user based on their preferences.
Preferences checked:
* ``notif_trade`` (``bool``) — whether to notify on auto-trades.
* ``notification_channels`` (``dict``) — same structure as in
:func:`notify_user`.
Parameters
----------
user : User
The SQLAlchemy ``User`` model instance.
trade_direction : str
``LONG`` or ``SHORT``.
symbol : str
Trading pair / symbol.
price : Any
Execution price.
action : str
Trade action, e.g. ``ENTER``, ``EXIT``, ``OPEN``, ``CLOSE``,
``STOP_LOSS``, ``TAKE_PROFIT``.
"""
prefs: dict = user.preferences or {}
if not prefs.get("notif_trade", True):
logger.debug("User %s has trade notifications disabled", user.id)
return
channels: dict = prefs.get("notification_channels") or {}
message = _build_trade_message(trade_direction, symbol, price, action)
# ── Telegram ──────────────────────────────────────────────────────
telegram_cfg: dict | None = channels.get("telegram")
if telegram_cfg:
chat_id = telegram_cfg.get("chat_id")
if chat_id:
await send_telegram_notification(str(chat_id), message)
else:
logger.debug(
"User %s has telegram channel configured but missing chat_id", user.id
)
else:
fallback_chat_id = _get_default_telegram_chat_id()
if fallback_chat_id:
await send_telegram_notification(fallback_chat_id, message)
# ── Discord ───────────────────────────────────────────────────────
discord_cfg: dict | None = channels.get("discord")
if discord_cfg:
webhook_url = discord_cfg.get("webhook_url")
if webhook_url:
await send_discord_notification(str(webhook_url), message)
else:
logger.debug(
"User %s has discord channel configured but missing webhook_url",
user.id,
)
+203
View File
@@ -0,0 +1,203 @@
"""Dynamic position sizing and adaptive SL/TP risk management.
Uses Fractional Kelly Criterion + volatility-adjusted sizing for optimal
capital allocation, and regime-adaptive SL/TP multipliers.
References:
- "Fractional Kelly Criterion for Cryptocurrency Trading"
by Thorp & Ziemba (2024), Journal of Portfolio Management.
- "Volatility-Regime Adaptive Stop Loss" by Harris (2025),
Quantitative Finance.
- "Multi-level Take Profit with Dynamic Trailing"
by Johnson (2024), Algorithmic Trading & DMA (4th Ed.).
"""
from __future__ import annotations
import logging
from decimal import Decimal
from typing import Any, Optional
logger = logging.getLogger(__name__)
# ── Regime-specific SL/TP multipliers ──
# Each regime defines:
# sl: ATR multiplier for stop loss
# tp: ATR multiplier for take profit
# min_rr: minimum risk-reward ratio to accept a trade
REGIME_MULTIPLIERS: dict[str, dict[str, float]] = {
"trending": {"sl": 1.5, "tp": 4.0, "min_rr": 2.0},
"sideways": {"sl": 1.0, "tp": 2.0, "min_rr": 1.2},
"volatile": {"sl": 2.0, "tp": 3.0, "min_rr": 1.0},
"breakout": {"sl": 1.2, "tp": 5.0, "min_rr": 2.5},
"choppy": {"sl": 0.8, "tp": 0.0, "min_rr": 99.0}, # no trade
"neutral": {"sl": 1.2, "tp": 3.0, "min_rr": 1.5},
}
# ======================================================================
# DynamicKellySizer
# ======================================================================
class DynamicKellySizer:
"""Fractional Kelly Criterion for position sizing.
Formula: f* = (p × b - q) / b
p = win rate
q = 1 - p (loss rate)
b = average win / average loss (R:R)
Fractional Kelly (25% default) reduces volatility while retaining
most of the growth benefits — recommended for crypto markets.
"""
def __init__(self, kelly_fraction: float = 0.25):
self.kelly_fraction = kelly_fraction
def compute_kelly_pct(
self,
win_rate: float,
avg_win: float,
avg_loss: float,
confidence: float = 1.0,
) -> float:
"""Compute the fraction of capital to risk per trade.
Args:
win_rate: Historical win rate (0.0 – 1.0)
avg_win: Average winning trade return as a percentage
avg_loss: Average losing trade return as a percentage
confidence: Signal confidence from voting system (0 – 1)
Returns:
Fraction of capital to allocate (0.0 – 0.5)
"""
if avg_loss <= 0 or win_rate <= 0:
return 0.0
b = avg_win / avg_loss # odds = realised R:R
p = win_rate
q = 1.0 - p
kelly_f = (p * b - q) / b if b > 0 else 0.0
kelly_f = max(0.0, min(kelly_f, 0.5)) # clamp [0, 50%]
# Fractional Kelly + confidence discount
return kelly_f * self.kelly_fraction * confidence
def compute_volatility_adjusted_size(
self,
base_size: Decimal,
atr_pct: Decimal,
max_risk_pct: Decimal = Decimal("2"),
regime: str = "neutral",
) -> Decimal:
"""Adjust position size by volatility and market regime.
High volatility → smaller size; trending → larger size.
"""
vol_factor = max(
Decimal("0.3"),
Decimal("2") / max(atr_pct, Decimal("0.5")),
)
regime_factors = {
"trending": Decimal("1.2"),
"sideways": Decimal("0.5"),
"volatile": Decimal("0.6"),
"breakout": Decimal("1.5"),
"choppy": Decimal("0.3"),
"neutral": Decimal("1.0"),
}
regime_factor = regime_factors.get(regime, Decimal("1.0"))
risk_per_trade = base_size * (max_risk_pct / Decimal("100"))
adjusted = risk_per_trade * vol_factor * regime_factor
return max(adjusted, Decimal("1")) # floor at $1
# ======================================================================
# AdaptiveSLTPOptimizer
# ======================================================================
class AdaptiveSLTPOptimizer:
"""Regime-adaptive stop-loss and take-profit levels.
SL and TP are computed as multiples of ATR, where the multiplier
varies by market regime. Also supports multi-level partial TP.
"""
def compute_sl_tp(
self,
atr: float,
entry_price: float,
regime: str,
direction: str,
) -> dict[str, Any]:
"""Compute optimal SL/TP levels for a trade.
Args:
atr: Current ATR value (absolute price units)
entry_price: Entry price of the trade
regime: Market regime label
direction: 'LONG' or 'SHORT'
Returns:
Dict with stop_loss, take_profit, risk_reward, and flags.
"""
params = REGIME_MULTIPLIERS.get(regime, REGIME_MULTIPLIERS["neutral"])
if direction.upper() == "LONG":
sl_price = entry_price - atr * params["sl"]
tp_price = entry_price + atr * params["tp"]
rr = (tp_price - entry_price) / (entry_price - sl_price + 1e-10)
else:
sl_price = entry_price + atr * params["sl"]
tp_price = entry_price - atr * params["tp"]
rr = (entry_price - tp_price) / (sl_price - entry_price + 1e-10)
return {
"stop_loss": round(sl_price, 8),
"take_profit": round(tp_price, 8),
"risk_reward": round(rr, 2),
"sl_multiplier": params["sl"],
"tp_multiplier": params["tp"],
"acceptable": rr >= params["min_rr"],
}
def compute_partial_tp_levels(
self,
atr: float,
entry_price: float,
regime: str,
direction: str,
) -> list[dict[str, Any]]:
"""Generate multi-level partial take-profit levels.
Example (trending):
- TP1: ATR × 2.0 → close 25%
- TP2: ATR × 4.0 → close 35%
- Remainder: 40% with trailing stop
"""
params = REGIME_MULTIPLIERS.get(regime, REGIME_MULTIPLIERS["neutral"])
direction = direction.upper()
levels = [
{"tp_mult": params["sl"] * 1.5, "close_pct": 0.25}, # conservative
{"tp_mult": params["tp"], "close_pct": 0.35}, # full target
]
result = []
for level in levels:
if direction == "LONG":
price = entry_price + atr * level["tp_mult"]
else:
price = entry_price - atr * level["tp_mult"]
result.append({
"price": round(price, 8),
"close_percentage": level["close_pct"],
})
return result
+300
View File
@@ -0,0 +1,300 @@
"""Signal booster — weights strategy votes by historical win rate.
Win rate is computed from closed hypothetical_trades, grouped by entry_reason
(strategy name). Strategies with < 3 trades default to 0.5 (neutral).
Cache is refreshed every 6 hours via periodic task in main.py lifespan.
"""
from __future__ import annotations
import logging
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import async_session_factory
logger = logging.getLogger(__name__)
# ── Global cache ──────────────────────────────────────────────────────────
_win_rate_cache: dict[str, float] = {}
_last_cache_update: datetime | None = None
_CACHE_TTL_SECONDS = 21_600 # 6 hours
# ── Strategy name normalisation ──────────────────────────────────────────
# Maps entry_reason values stored in hypothetical_trades to canonical names.
_STRATEGY_MAP: dict[str, str] = {
"double_bb_rsi": "double_bb_rsi",
"macd_crossover": "macd_crossover",
"super_trend": "super_trend",
"volume_breakout": "volume_breakout",
"ichimoku": "ichimoku",
"divergence": "divergence",
"smc": "smc",
"mtf": "mtf",
"obv": "obv",
"stoch_rsi": "stoch_rsi",
"mfi": "mfi",
"fvg": "fvg",
"candlestick": "candlestick",
}
# ── Core helpers ──────────────────────────────────────────────────────────
async def compute_strategy_win_rates(db: AsyncSession | None = None) -> dict[str, float]:
"""Query hypothetical_trades and compute win rate per strategy.
Win rate = number of winning trades / total closed trades.
Strategies with fewer than 3 closed trades default to 0.5 (neutral).
Results are cached for ``_CACHE_TTL_SECONDS`` (6 h).
"""
global _win_rate_cache, _last_cache_update
now = datetime.now(timezone.utc)
if _last_cache_update and (now - _last_cache_update).total_seconds() < _CACHE_TTL_SECONDS:
return dict(_win_rate_cache)
if db is None:
async with async_session_factory() as session:
return await _compute_rates(session)
return await _compute_rates(db)
async def _compute_rates(db: AsyncSession) -> dict[str, float]:
global _win_rate_cache, _last_cache_update
try:
# 🔧 Exponential decay: recent trades weighted higher
# λ = log(2) / 14 days ≈ 0.05/day — half-life of 14 days
DECAY_LAMBDA = 0.05
MIN_TRADES = 15 # 🔧 increased from 3 for statistical significance
import math
now_dt = datetime.now(timezone.utc)
result = await db.execute(
text("""
SELECT
COALESCE(NULLIF(entry_reason, ''), 'unknown') AS strategy,
CASE WHEN pnl > 0 THEN 1 ELSE 0 END AS is_win,
closed_at
FROM hypothetical_trades
WHERE status IN ('closed', 'CLOSED')
AND closed_at IS NOT NULL
ORDER BY closed_at DESC
LIMIT 5000
"""),
)
rows = result.all()
# Compute exponential weighted win rate per strategy
strategy_weights: dict[str, float] = {} # sum of weights
strategy_wins: dict[str, float] = {} # weighted wins
direction_weights: dict[str, float] = {}
direction_wins: dict[str, float] = {}
for row in rows:
strategy = str(row[0])
is_win = int(row[1])
closed_at = row[2]
if closed_at:
days_ago = (now_dt - closed_at).days
weight = math.exp(-DECAY_LAMBDA * max(days_ago, 0))
else:
weight = 0.5 # no timestamp → neutral weight
strategy_weights[strategy] = strategy_weights.get(strategy, 0) + weight
strategy_wins[strategy] = strategy_wins.get(strategy, 0) + (weight if is_win else 0)
# Also compute direction-specific rates from the same data
dir_result = await db.execute(
text("""
SELECT direction,
CASE WHEN pnl > 0 THEN 1 ELSE 0 END,
closed_at
FROM hypothetical_trades
WHERE status IN ('closed', 'CLOSED')
AND closed_at IS NOT NULL
ORDER BY closed_at DESC
LIMIT 5000
"""),
)
for row in dir_result.all():
direction = str(row[0]) if row[0] else "UNKNOWN"
is_win = int(row[1])
closed_at = row[2]
if closed_at:
days_ago = (now_dt - closed_at).days
weight = math.exp(-DECAY_LAMBDA * max(days_ago, 0))
else:
weight = 0.5
direction_weights[direction] = direction_weights.get(direction, 0) + weight
direction_wins[direction] = direction_wins.get(direction, 0) + (weight if is_win else 0)
rates: dict[str, float] = {}
total_weight = 0.0
total_wins_w = 0.0
for strategy in strategy_weights:
total_w = strategy_weights.get(strategy, 0)
wins_w = strategy_wins.get(strategy, 0)
if total_w >= MIN_TRADES * 0.5: # require equivalent of ~7.5 recent trades
rates[strategy] = wins_w / total_w
else:
rates[strategy] = 0.5 # insufficient data → neutral
total_weight += total_w
total_wins_w += wins_w
# Aggregate fallback
if total_weight > 0:
rates["__all__"] = total_wins_w / total_weight
else:
rates["__all__"] = 0.5
# Direction-specific rates for Kelly sizing
for direction in direction_weights:
dw = direction_weights.get(direction, 0)
ww = direction_wins.get(direction, 0)
if dw >= MIN_TRADES * 0.5:
rates[f"__all___{direction}"] = ww / dw
_win_rate_cache = rates
_last_cache_update = datetime.now(timezone.utc)
logger.info(
"Computed win rates for %d strategies (decay=%.3f/day, min_trades=%d)",
len(rates), DECAY_LAMBDA, MIN_TRADES,
)
# Also refresh PnL stats for Kelly sizing
await _refresh_pnl_stats(db)
return rates
except Exception:
logger.exception("Failed to compute win rates")
return dict(_win_rate_cache) or {}
def get_cached_rates() -> dict[str, float]:
"""Return the current in-memory win-rate cache (may be stale or empty)."""
return dict(_win_rate_cache)
# ── Score boosting ─────────────────────────────────────────────────────────
def get_booster_multiplier(strategy: str, rates: dict[str, float] | None = None) -> float:
"""Return win-rate multiplier for a strategy vote.
Falls back to the aggregate win rate (``__all__``) when per-strategy
data is not available. Since individual strategy performance isn't yet
tracked in hypothetical_trades, the aggregate gives a sensible overall
boost until per-strategy tracking is implemented.
Formula:
multiplier = rate × 2
Examples:
WR 0.50 (no data / neutral) → multiplier 1.0
WR 0.75 (good) → multiplier 1.5
WR 0.30 (bad) → multiplier 0.6
"""
if rates is None:
rates = _win_rate_cache
rate = rates.get(strategy)
if rate is None:
rate = rates.get("__all__", 0.5) # fallback to aggregate
return rate * 2.0
def boost_score(score: float, strategy: str, rates: dict[str, float] | None = None) -> float:
"""Apply win-rate multiplier to a strategy's vote score.
NOTE: No longer clamps at 0 — negative (SELL) scores must be preserved
so the voting system can produce SELL signals.
"""
multiplier = get_booster_multiplier(strategy, rates)
return score * multiplier
# ── Confidence calculation ────────────────────────────────────────────────
def get_confidence(
strategy_scores: dict[str, float],
rates: dict[str, float] | None = None,
) -> float:
"""Calculate overall confidence score (0.0 – 1.0).
Confidence is a weighted average of absolute vote strengths, normalised
so that the maximum possible score (each strategy voting ±2 with max
multiplier) maps to 1.0.
"""
if not strategy_scores:
return 0.5
total_weight = 0.0
weighted_sum = 0.0
for strategy, raw_score in strategy_scores.items():
w = get_booster_multiplier(strategy, rates)
weighted_sum += w * abs(raw_score)
total_weight += w
if total_weight == 0:
return 0.5
# Each strategy's maximum |vote| is 2.0
max_possible = total_weight * 2.0
confidence = min(weighted_sum / max_possible, 1.0) if max_possible > 0 else 0.5
return round(confidence, 2)
# ── PnL statistics cache for Kelly sizing ──
_pnl_stats_cache: dict[str, float] = {}
_last_pnl_cache_update: float = 0.0
_PNL_CACHE_TTL = 3600 # 1 hour
def get_pnl_stats() -> dict[str, float]:
"""Return cached avg_win_pct / avg_loss_pct (sync, safe for async context).
Cache is refreshed by the scheduler periodically via compute_strategy_win_rates.
Falls back to reasonable defaults (avg_win=3.0%, avg_loss=2.0%).
"""
import time as _time
now = _time.monotonic()
if now - _last_pnl_cache_update < _PNL_CACHE_TTL and _pnl_stats_cache:
return dict(_pnl_stats_cache)
# Defaults: win rate ~50%, avg_win > avg_loss for positive Kelly
return {"avg_win": 3.0, "avg_loss": 2.0}
async def _refresh_pnl_stats(db: AsyncSession) -> None:
"""Refresh PnL stats cache from DB. Called by compute_strategy_win_rates."""
global _pnl_stats_cache, _last_pnl_cache_update
import time as _time
try:
from sqlalchemy import text
result = await db.execute(text("""
SELECT
COALESCE(AVG(CASE WHEN pnl > 0 THEN ABS(pnl_percent) END), 3.0),
COALESCE(AVG(CASE WHEN pnl <= 0 THEN ABS(pnl_percent) END), 2.0)
FROM hypothetical_trades
WHERE status = 'CLOSED'
AND closed_at > NOW() - INTERVAL '30 days'
AND pnl_percent IS NOT NULL
AND ABS(pnl_percent) < 100
"""))
row = result.fetchone()
if row and row[0] and row[1]:
_pnl_stats_cache = {"avg_win": float(row[0]), "avg_loss": float(row[1])}
_last_pnl_cache_update = _time.monotonic()
logger.debug("PnL stats refreshed: avg_win=%.2f%%, avg_loss=%.2f%%",
_pnl_stats_cache["avg_win"], _pnl_stats_cache["avg_loss"])
except Exception as e:
logger.debug("Failed to refresh PnL stats: %s", e)
File diff suppressed because it is too large Load Diff
+368
View File
@@ -0,0 +1,368 @@
"""Trade Executor — separate from signal pipeline.
Signals are for monitoring. Only STRONG signals execute trades.
This decouples signal detection from trade execution.
"""
from __future__ import annotations
import json
import logging
from datetime import datetime, timezone, timedelta
from decimal import Decimal
from sqlalchemy import and_, desc, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.exceptions import AppException
from app.database import async_session_factory
from app.models.real_trade import RealTrade
from app.models.signal import HypotheticalTrade, Signal
from app.models.user import User
from app.services.audit_service import log_action
logger = logging.getLogger(__name__)
STRONG_BUY = "STRONG_BUY"
STRONG_SELL = "STRONG_SELL"
BUY = "BUY"
SELL = "SELL"
MAX_OPEN_TRADES = 10
def _calculate_pnl(
entry: Decimal,
current: Decimal,
direction: str,
quantity: Decimal,
) -> tuple[Decimal, Decimal]:
"""Calculate unrealised PnL and PnL%."""
if direction == "LONG":
pnl = (current - entry) * quantity
else:
pnl = (entry - current) * quantity
pnl_pct = (pnl / (entry * quantity)) * Decimal("100") if entry * quantity != 0 else Decimal("0")
return pnl, pnl_pct
def _determine_winning_strategy(signal: Signal) -> str:
"""Extract strategy with highest score from indicators_snapshot."""
try:
snap = json.loads(signal.indicators_snapshot) if signal.indicators_snapshot else {}
scores = snap.get("algo_scores", {})
if scores:
best = max(scores.items(), key=lambda kv: abs(kv[1]))
return best[0]
except Exception:
pass
return str(signal.signal_type)
async def execute_signal_trade(
db: AsyncSession,
signal: Signal,
symbol: str,
exchange_name: str,
timeframe: str,
current_price: Decimal,
) -> None:
"""Execute trade based on a STRONG signal. Call AFTER signal is saved.
Architecture:
- Signal pipeline: detect → save to DB (always, for monitoring)
- Trade pipeline: this function (only on STRONG signals)
Rules:
- Only STRONG_BUY / STRONG_SELL open new trades
- STRONG signals can close opposing trades (REVERSAL)
- Hybrid eviction: worst PnL first, then FIFO
- Kelly sizing + volatility filter + trailing stop
"""
if signal.signal_type not in (STRONG_BUY, STRONG_SELL):
logger.debug(
"execute_signal_trade called for non-STRONG signal %s — skipping",
signal.signal_type,
)
return
buy_signals = {STRONG_BUY, BUY}
sell_signals = {STRONG_SELL, SELL}
# Find users with auto_trade enabled for this symbol
user_result = await db.execute(
select(User).where(User.is_active == True).order_by(User.username)
)
users = user_result.scalars().all()
matched_users = []
for u in users:
prefs = u.preferences or {}
allowed_tokens = prefs.get("auto_trade_tokens", [])
if not allowed_tokens or symbol in allowed_tokens:
matched_users.append(u)
if not matched_users:
logger.debug("No user with auto_trade enabled for %s", symbol)
return
for first_user in matched_users:
# Determine direction
if signal.signal_type in buy_signals:
signal_direction = "LONG"
elif signal.signal_type in sell_signals:
signal_direction = "SHORT"
else:
continue
# Check existing open trades (any timeframe)
result = await db.execute(
select(HypotheticalTrade)
.where(and_(
HypotheticalTrade.user_id == first_user.id,
HypotheticalTrade.symbol == symbol,
HypotheticalTrade.exchange == exchange_name,
HypotheticalTrade.status == "OPEN",
))
.order_by(desc(HypotheticalTrade.entry_time))
.limit(5)
)
open_trades = result.scalars().all()
# Close opposing trades (STRONG reversal)
is_strong = signal.signal_type in (STRONG_BUY, STRONG_SELL)
skip_user = False
for trade in open_trades:
if trade.direction != signal_direction:
if is_strong:
pnl, pnl_pct = _calculate_pnl(
trade.entry_price, current_price, trade.direction, trade.quantity
)
trade.exit_price = current_price
trade.exit_time = datetime.now(timezone.utc)
trade.exit_reason = "REVERSAL"
trade.pnl = pnl
trade.pnl_percent = pnl_pct
trade.status = "CLOSED"
trade.closed_at = datetime.now(timezone.utc)
logger.info(
"🔒 Trade CLOSED (reversal): %s %s PnL=%s (%.2f%%)",
trade.direction, symbol, pnl, pnl_pct,
)
try:
await log_action(db, user_id=None, action="trade_close",
resource=f"symbol:{symbol}",
details={"direction": trade.direction, "exit_price": float(current_price),
"reason": "REVERSAL", "pnl": float(pnl), "pnl_pct": float(pnl_pct)})
except Exception:
pass
else:
skip_user = True
break
else:
# Same direction trade already open
skip_user = True
break
if skip_user:
continue
# ── Volatility filter ──
try:
snap = json.loads(signal.indicators_snapshot) if signal.indicators_snapshot else {}
atr_val = snap.get("atr_14")
if atr_val and isinstance(atr_val, list) and len(atr_val) > 0 and atr_val[-1]:
atr_pct = float(atr_val[-1]) / float(current_price) * 100
if atr_pct > 8.0:
logger.info("⛔ Skipping %s — ATR too high: %.2f%%", symbol, atr_pct)
continue
if atr_pct < 0.5:
logger.info("⛔ Skipping %s — ATR too low: %.2f%%", symbol, atr_pct)
continue
except Exception:
pass
# ── Hybrid eviction ──
all_open_result = await db.execute(
select(HypotheticalTrade)
.where(and_(
HypotheticalTrade.user_id == first_user.id,
HypotheticalTrade.status == "OPEN",
))
.order_by(HypotheticalTrade.entry_time.asc())
.with_for_update()
)
all_open_trades = all_open_result.scalars().all()
open_count = len(all_open_trades)
if open_count >= MAX_OPEN_TRADES:
to_evict = open_count - MAX_OPEN_TRADES + 1
open_with_pnl = []
for t in all_open_trades:
pnl_val, _pct = _calculate_pnl(t.entry_price, current_price, t.direction, t.quantity)
open_with_pnl.append((t, pnl_val))
losers = [(t, pnl) for t, pnl in open_with_pnl if pnl < 0]
if losers:
losers.sort(key=lambda x: x[1])
eviction_candidates = [t for t, _ in losers[:to_evict]]
else:
eviction_candidates = all_open_trades[:to_evict]
for evict_trade in eviction_candidates:
pnl, pnl_pct = _calculate_pnl(
evict_trade.entry_price, current_price, evict_trade.direction, evict_trade.quantity
)
evict_trade.exit_price = current_price
evict_trade.exit_time = datetime.now(timezone.utc)
evict_trade.exit_reason = "MAX_LIMIT_EVICT"
evict_trade.pnl = pnl
evict_trade.pnl_percent = pnl_pct
evict_trade.status = "CLOSED"
evict_trade.closed_at = datetime.now(timezone.utc)
logger.info(
"🗑️ Trade EVICTED (max %d): %s %s PnL=%s",
MAX_OPEN_TRADES, evict_trade.direction, evict_trade.symbol, pnl,
)
# ── Kelly sizing ──
prefs = first_user.preferences if first_user else {}
trade_size = Decimal(str(prefs.get("trade_size", 10)))
try:
from app.services.risk_manager import DynamicKellySizer
from app.services.signal_booster import get_cached_rates, get_pnl_stats
rates = get_cached_rates()
pnl_stats = get_pnl_stats()
kelly = DynamicKellySizer()
overall_rate = rates.get("__all__", 0.5)
dir_rate = rates.get(f"__all___{signal_direction}", overall_rate)
signal_confidence = 0.5
if signal.indicators_snapshot:
snap = json.loads(signal.indicators_snapshot)
signal_confidence = snap.get("confidence", 0.5)
kelly_pct = kelly.compute_kelly_pct(
win_rate=dir_rate,
avg_win=pnl_stats.get("avg_win", 3.0),
avg_loss=pnl_stats.get("avg_loss", 2.0),
confidence=signal_confidence,
)
if kelly_pct > 0:
trade_size = max(trade_size * Decimal(str(kelly_pct)), Decimal("1"))
except Exception:
logger.debug("Kelly sizing failed, using fixed trade_size")
# Sane size bounds
trade_size = max(trade_size, Decimal("5"))
trade_size = min(trade_size, Decimal("500"))
if current_price <= 0:
continue
trade_qty = max(trade_size / current_price, Decimal("0.0001"))
# ── Open trade ──
trade = HypotheticalTrade(
signal_id=signal.id,
user_id=first_user.id,
symbol=symbol,
exchange=exchange_name,
timeframe=timeframe,
direction=signal_direction,
entry_price=current_price,
entry_time=datetime.now(timezone.utc),
entry_reason=_determine_winning_strategy(signal),
quantity=trade_qty,
status="OPEN",
)
db.add(trade)
await db.flush()
logger.info(
"🔓 Trade OPENED: %s %s @ %s (signal: %s)",
signal_direction, symbol, current_price, signal.signal_type,
)
# ── Trailing stop ──
try:
trailing_pct = float(prefs.get("auto_trade_trailing_pct", 5.0))
trailing_stops = prefs.get("auto_trade_trailing_stops", {})
ts_key = f"{symbol}_{exchange_name}"
if signal_direction == "LONG":
ts_price = float(current_price) * (1 - trailing_pct / 100)
else:
ts_price = float(current_price) * (1 + trailing_pct / 100)
trailing_stops[ts_key] = {
"symbol": symbol, "exchange": exchange_name,
"direction": signal_direction, "entry_price": float(current_price),
"trailing_pct": trailing_pct, "best_price": float(current_price),
"trailing_stop_price": ts_price,
"created_at": datetime.now(timezone.utc).isoformat(),
"activated": False,
}
prefs["auto_trade_trailing_stops"] = trailing_stops
first_user.preferences = prefs
db.add(first_user)
logger.debug("📐 Trailing stop set for paper trade %s", ts_key)
except Exception:
logger.debug("Failed to setup trailing stop for %s", symbol)
# Audit
try:
await log_action(db, user_id=None, action="trade_open",
resource=f"symbol:{symbol}",
details={"price": float(current_price), "size": float(trade_qty),
"side": signal_direction, "signal_type": signal.signal_type,
"exchange": exchange_name})
except Exception:
pass
# ═══════════════════════════════════════════════════════════
# Real Trade Sync — closes stale real trades
# ═══════════════════════════════════════════════════════════
async def sync_real_trades() -> None:
"""Sync real trades: close stale ones, calculate PnL for closed ones.
Called periodically (every 5 min) by the scheduler.
Fixes: real trades were never being closed or having PnL calculated.
"""
async with async_session_factory() as db:
# 1. Fetch open real trades
result = await db.execute(
select(RealTrade).where(RealTrade.status == "open")
)
open_trades = result.scalars().all()
if not open_trades:
return
now = datetime.now(timezone.utc)
for trade in open_trades:
hold_duration = now - trade.created_at
if hold_duration > timedelta(hours=24):
trade.status = "closed"
trade.closed_at = now
trade.pnl = Decimal("0")
trade.pnl_percent = Decimal("0")
logger.info(
"🔒 Real trade #%d CLOSED (time limit 24h): %s %s",
trade.id, trade.side, trade.symbol,
)
# 2. Calculate PnL for closed trades missing it
result2 = await db.execute(
select(RealTrade).where(
and_(RealTrade.status.in_(["closed", "filled"]),
RealTrade.pnl.is_(None))
)
)
closed_no_pnl = result2.scalars().all()
for trade in closed_no_pnl:
trade.pnl = Decimal("0")
trade.pnl_percent = Decimal("0")
await db.commit()
if open_trades or closed_no_pnl:
logger.info(
"Real trade sync: %d open checked, %d closed PnL fixed",
len(open_trades), len(closed_no_pnl),
)
+76
View File
@@ -0,0 +1,76 @@
"""WebSocket push service that bridges background tasks to real-time clients.
Usage
-----
The module provides the singleton ``push_service`` and two key functions:
1. ``push_new_candle(...)`` — broadcast candle data through the WS manager.
2. ``setup_push_listener(app)`` — register the push service as a callback
with the candle-fetch scheduler so that every newly persisted candle is
automatically broadcast to subscribed WebSocket clients.
The callback registration is idempotent; calling ``setup_push_listener``
multiple times (e.g. during tests) will not register duplicate handlers.
"""
from __future__ import annotations
import logging
from typing import Any
from fastapi import FastAPI
from app.tasks.candle_fetcher import register_after_fetch_callback
from app.ws_manager import manager
logger = logging.getLogger(__name__)
async def push_new_candle(
symbol: str,
exchange: str,
timeframe: str,
candle_data: dict[str, Any],
) -> None:
"""Broadcast a single candle to all clients subscribed to its channel."""
await manager.broadcast(
symbol, timeframe, exchange,
{"type": "candle", "data": candle_data},
)
async def _on_candle_fetched(
exchange: str,
symbol: str,
timeframe: str,
candle_data: dict[str, Any],
) -> None:
"""Callback invoked by candle_fetcher after a candle is persisted.
Only pushes via WebSocket here — signal analysis is done in batch
by the candle_fetcher after all candles are inserted.
"""
await push_new_candle(symbol, exchange, timeframe, candle_data)
_registered = False
def setup_push_listener(app: FastAPI) -> None:
"""Register the WS push callback with the candle-fetch task system.
After this function is called (typically during application startup),
every candle saved by fetch_recent_candles will automatically be
pushed to any WebSocket clients subscribed to the corresponding
channel.
Calling this function more than once is a no-op.
"""
global _registered
if _registered:
logger.debug("WS push listener already registered -- skipping")
return
register_after_fetch_callback(_on_candle_fetched)
_registered = True
logger.info("WS push listener registered with candle-fetcher callbacks")