Initial commit: Trading Portal - FastAPI + React + PostgreSQL
This commit is contained in:
Executable
Executable
+367
@@ -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
|
||||
Executable
+82
@@ -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
|
||||
Executable
+332
@@ -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()
|
||||
Executable
+544
@@ -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)
|
||||
Executable
+1479
File diff suppressed because it is too large
Load Diff
Executable
+304
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
Executable
+300
@@ -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)
|
||||
Executable
+2328
File diff suppressed because it is too large
Load Diff
@@ -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),
|
||||
)
|
||||
Executable
+76
@@ -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")
|
||||
Reference in New Issue
Block a user