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
View File
+21
View File
@@ -0,0 +1,21 @@
"""V1 API router aggregation.
Add new v1 sub-routers here by importing and including them on the
top-level ``router`` so that ``main.py`` only needs a single import.
"""
from fastapi import APIRouter
router = APIRouter()
# Example — uncomment and implement when the module exists:
# from app.api.v1 import auth, users, trades
# router.include_router(auth.router, prefix="/auth", tags=["auth"])
# router.include_router(users.router, prefix="/users", tags=["users"])
# router.include_router(trades.router, prefix="/trades", tags=["trades"])
@router.get("/ping")
async def ping():
"""Minimal liveness check scoped to the v1 namespace."""
return {"message": "pong"}
+360
View File
@@ -0,0 +1,360 @@
from __future__ import annotations
import logging
import time
from fastapi import APIRouter, Depends
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_current_admin_user, get_db_session
from app.core.exceptions import ConflictException, NotFoundException
from app.core.security import hash_password
from app.database import engine
from app.exchange.factory import factory as exchange_factory
from app.models import Exchange, RefreshToken, User, Watchlist
from app.schemas import (
AdminCreateUserRequest,
AdminResetPasswordRequest,
AdminUserUpdateRequest,
DbHealth,
DetailedHealthResponse,
ExchangeCreateRequest,
ExchangeHealth,
ExchangeResponse,
ExchangeUpdateRequest,
UserResponse,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/admin", tags=["admin"])
# ---------------------------------------------------------------------------
# Users
# ---------------------------------------------------------------------------
@router.get("/users", response_model=list[UserResponse])
async def list_users(
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> list[UserResponse]:
"""List all users (admin only)."""
result = await db.execute(select(User).order_by(User.created_at))
users = result.scalars().all()
return [
UserResponse(
id=u.id,
username=u.username,
email=u.email,
display_name=u.display_name,
is_active=u.is_active,
is_admin=u.is_admin,
role=u.role,
created_at=u.created_at,
)
for u in users
]
@router.put("/users/{user_id}", response_model=UserResponse)
async def update_user(
user_id: str,
body: AdminUserUpdateRequest,
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> UserResponse:
"""Update a user (activate/deactivate, promote/demote admin, email, display_name)."""
from uuid import UUID
result = await db.execute(select(User).where(User.id == UUID(user_id)))
user = result.scalar_one_or_none()
if user is None:
raise NotFoundException(detail=f"User with id {user_id} not found")
if body.is_active is not None:
user.is_active = body.is_active
if body.is_admin is not None:
user.is_admin = body.is_admin
if body.role is not None:
user.role = body.role
if body.email is not None:
# Check uniqueness
existing = await db.execute(
select(User).where(User.email == body.email, User.id != UUID(user_id))
)
if existing.scalar_one_or_none() is not None:
raise ConflictException(detail=f"Email '{body.email}' is already in use")
user.email = body.email
if body.display_name is not None:
user.display_name = body.display_name
await db.flush()
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,
preferences=user.preferences,
created_at=user.created_at,
)
@router.post("/users", response_model=UserResponse, status_code=201)
async def create_user(
body: AdminCreateUserRequest,
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> UserResponse:
"""Create a new user (admin bypass)."""
# Uniqueness checks
existing_username = await db.execute(
select(User).where(User.username == body.username)
)
if existing_username.scalar_one_or_none() is not None:
raise ConflictException(detail=f"Username '{body.username}' is already taken")
existing_email = await db.execute(
select(User).where(User.email == body.email)
)
if existing_email.scalar_one_or_none() is not None:
raise ConflictException(detail=f"Email '{body.email}' is already registered")
user = User(
username=body.username,
email=body.email,
display_name=body.display_name,
password_hash=hash_password(body.password),
is_admin=body.is_admin,
role=body.role,
)
db.add(user)
await db.flush()
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,
)
@router.delete("/users/{user_id}", status_code=200)
async def delete_user(
user_id: str,
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> dict:
"""Delete a user (cascades to watchlists, credentials, refresh tokens)."""
from uuid import UUID
result = await db.execute(select(User).where(User.id == UUID(user_id)))
user = result.scalar_one_or_none()
if user is None:
raise NotFoundException(detail=f"User with id {user_id} not found")
await db.delete(user)
await db.commit()
return {"message": f"User '{user.username}' deleted successfully"}
@router.post("/users/{user_id}/reset-password", status_code=200)
async def reset_user_password(
user_id: str,
body: AdminResetPasswordRequest,
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> dict:
"""Admin-reset a user's password (no old password required). Revokes all sessions."""
from uuid import UUID
result = await db.execute(select(User).where(User.id == UUID(user_id)))
user = result.scalar_one_or_none()
if user is None:
raise NotFoundException(detail=f"User with id {user_id} not found")
user.password_hash = hash_password(body.new_password)
# Revoke all refresh tokens (force re-login)
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()
return {"message": "Password reset successfully. All sessions revoked."}
# ---------------------------------------------------------------------------
# Exchanges
# ---------------------------------------------------------------------------
@router.get("/exchanges", response_model=list[ExchangeResponse])
async def list_exchanges_admin(
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> list[ExchangeResponse]:
"""List all exchanges (including inactive)."""
result = await db.execute(select(Exchange).order_by(Exchange.id))
exchanges = result.scalars().all()
return [
ExchangeResponse(
id=ex.id,
name=ex.name,
display_name=ex.display_name or ex.name,
is_active=ex.is_active,
)
for ex in exchanges
]
@router.post("/exchanges", response_model=ExchangeResponse, status_code=201)
async def create_exchange(
body: ExchangeCreateRequest,
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> ExchangeResponse:
"""Add a new exchange."""
exchange = Exchange(
name=body.name,
display_name=body.display_name,
base_url=body.base_url,
ws_url=body.ws_url,
is_active=True,
)
db.add(exchange)
await db.flush()
await db.refresh(exchange)
return ExchangeResponse(
id=exchange.id,
name=exchange.name,
display_name=exchange.display_name or exchange.name,
is_active=exchange.is_active,
)
@router.put("/exchanges/{exchange_id}", response_model=ExchangeResponse)
async def update_exchange(
exchange_id: int,
body: ExchangeUpdateRequest,
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> ExchangeResponse:
"""Update exchange config."""
result = await db.execute(select(Exchange).where(Exchange.id == exchange_id))
exchange = result.scalar_one_or_none()
if exchange is None:
raise NotFoundException(detail=f"Exchange with id {exchange_id} not found")
if body.display_name is not None:
exchange.display_name = body.display_name
if body.base_url is not None:
exchange.base_url = body.base_url
if body.ws_url is not None:
exchange.ws_url = body.ws_url
if body.is_active is not None:
exchange.is_active = body.is_active
await db.flush()
await db.refresh(exchange)
return ExchangeResponse(
id=exchange.id,
name=exchange.name,
display_name=exchange.display_name or exchange.name,
is_active=exchange.is_active,
)
# ---------------------------------------------------------------------------
# Health
# ---------------------------------------------------------------------------
async def _check_exchange_connection(exchange_name: str) -> bool:
"""Try to connect to an exchange and report whether it succeeded."""
try:
import asyncio
adapter = exchange_factory.create(exchange_name)
loop = asyncio.get_running_loop()
await loop.run_in_executor(None, lambda: adapter.client.load_markets())
return True
except Exception:
return False
@router.get("/health/detailed", response_model=DetailedHealthResponse)
async def detailed_health(
db: AsyncSession = Depends(get_db_session),
_admin: User = Depends(get_current_admin_user),
) -> DetailedHealthResponse:
"""Enhanced health with DB pool status and exchange connection statuses."""
# --- DB health ---
db_connected = False
db_latency = 0.0
try:
start = time.time()
from sqlalchemy import text
await db.execute(text("SELECT 1"))
db_latency = (time.time() - start) * 1000
db_connected = True
except Exception:
db_connected = False
pool = engine.pool
pool_size = pool.size() if hasattr(pool, "size") else 0
# --- Exchange health ---
result = await db.execute(
select(Exchange).where(Exchange.is_active == True) # noqa: E712
)
active_exchanges = result.scalars().all()
exchange_connections: list[ExchangeHealth] = []
for ex in active_exchanges:
try:
connected = await asyncio.wait_for(
_check_exchange_connection(ex.name), timeout=3.0
)
except (asyncio.TimeoutError, Exception):
connected = False
exchange_connections.append(
ExchangeHealth(exchange=ex.name, connected=connected)
)
# --- Overall status ---
all_ok = db_connected and all(ec.connected for ec in exchange_connections)
status = "healthy" if all_ok else "degraded"
return DetailedHealthResponse(
status=status,
version="1.0.0",
uptime=round(time.time() - _start_time, 2),
db=DbHealth(
connected=db_connected,
latency_ms=round(db_latency, 2),
pool_size=pool_size,
),
exchange_connections=exchange_connections,
)
# Track application start time for uptime reporting (reuse same pattern)
_start_time: float = time.time()
+216
View File
@@ -0,0 +1,216 @@
"""API routes for user-defined multi-condition alerts."""
from __future__ import annotations
import logging
from typing import Any
from fastapi import APIRouter, Depends, HTTPException
from fastapi.responses import Response
from pydantic import BaseModel, Field
from sqlalchemy import select, delete
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.core.deps import get_current_active_user
from app.models.alert import AlertCondition
from app.models.user import User
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/alerts", tags=["alerts"])
# ---------------------------------------------------------------------------
# Pydantic schemas
# ---------------------------------------------------------------------------
class ConditionObject(BaseModel):
indicator: str = Field(
...,
description="One of: rsi, macd, bb_width, volume, price, sma, ema, momentum",
)
operator: str = Field(
...,
description="One of: >, <, >=, <=, ==, cross_above, cross_below",
)
value: float = Field(..., description="Threshold value")
timeframe: str | None = Field(None, description="e.g. 1h, 15m")
type: str | None = Field(None, description="e.g. avg_multiplier, absolute")
period: int | None = Field(None, description="Lookback period, e.g. 20")
class AlertCreateRequest(BaseModel):
name: str = Field(..., max_length=100, description="Human-readable alert name")
conditions: list[ConditionObject] = Field(
..., min_length=1,
description="Array of condition objects (AND logic)",
)
notify_platform: str = Field(
"telegram",
description="'telegram', 'discord', or 'both'",
)
class AlertUpdateRequest(BaseModel):
name: str | None = Field(None, max_length=100)
conditions: list[ConditionObject] | None = None
notify_platform: str | None = None
is_active: bool | None = None
class AlertResponse(BaseModel):
id: int
user_id: str
name: str
conditions: list[dict[str, Any]]
notify_platform: str
is_active: bool
created_at: str
updated_at: str
class AlertListResponse(BaseModel):
alerts: list[AlertResponse]
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _alert_to_response(a: AlertCondition) -> AlertResponse:
return AlertResponse(
id=a.id,
user_id=str(a.user_id),
name=a.name,
conditions=a.conditions if isinstance(a.conditions, list) else [],
notify_platform=a.notify_platform,
is_active=a.is_active,
created_at=a.created_at.isoformat() if a.created_at else "",
updated_at=a.updated_at.isoformat() if a.updated_at else "",
)
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@router.get("", response_model=AlertListResponse)
async def list_alerts(
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_active_user),
) -> AlertListResponse:
"""List all alerts for the current user."""
result = await db.execute(
select(AlertCondition)
.where(AlertCondition.user_id == current_user.id)
.order_by(AlertCondition.created_at.desc())
)
alerts = result.scalars().all()
return AlertListResponse(
alerts=[_alert_to_response(a) for a in alerts]
)
@router.post("", response_model=AlertResponse, status_code=201)
async def create_alert(
body: AlertCreateRequest,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_active_user),
) -> AlertResponse:
"""Create a new multi-condition alert."""
if len(body.name.strip()) == 0:
raise HTTPException(status_code=422, detail="Alert name cannot be empty")
if not body.conditions:
raise HTTPException(status_code=422, detail="At least one condition required")
# Validate supported indicators and operators
supported_indicators = {
"rsi", "macd", "bb_width", "volume", "price", "sma", "ema", "momentum",
}
supported_operators = {">", "<", ">=", "<=", "==", "cross_above", "cross_below"}
for cond in body.conditions:
if cond.indicator not in supported_indicators:
raise HTTPException(
status_code=422,
detail=f"Unsupported indicator '{cond.indicator}'. Supported: {supported_indicators}",
)
if cond.operator not in supported_operators:
raise HTTPException(
status_code=422,
detail=f"Unsupported operator '{cond.operator}'. Supported: {supported_operators}",
)
alert = AlertCondition(
user_id=current_user.id,
name=body.name.strip(),
conditions=[c.model_dump() for c in body.conditions],
notify_platform=body.notify_platform,
is_active=True,
)
db.add(alert)
await db.flush()
await db.refresh(alert)
logger.info("Alert created: id=%d name=%s user=%s", alert.id, alert.name, current_user.id)
return _alert_to_response(alert)
@router.put("/{alert_id}", response_model=AlertResponse)
async def update_alert(
alert_id: int,
body: AlertUpdateRequest,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_active_user),
) -> AlertResponse:
"""Update an existing alert."""
result = await db.execute(
select(AlertCondition).where(
AlertCondition.id == alert_id,
AlertCondition.user_id == current_user.id,
)
)
alert = result.scalar_one_or_none()
if not alert:
raise HTTPException(status_code=404, detail="Alert not found")
if body.name is not None:
if len(body.name.strip()) == 0:
raise HTTPException(status_code=422, detail="Alert name cannot be empty")
alert.name = body.name.strip()
if body.conditions is not None:
alert.conditions = [c.model_dump() for c in body.conditions]
if body.notify_platform is not None:
alert.notify_platform = body.notify_platform
if body.is_active is not None:
alert.is_active = body.is_active
await db.flush()
await db.refresh(alert)
logger.info("Alert updated: id=%d user=%s", alert.id, current_user.id)
return _alert_to_response(alert)
@router.delete("/{alert_id}", status_code=204)
async def delete_alert(
alert_id: int,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_active_user),
) -> Response:
"""Delete an alert."""
result = await db.execute(
select(AlertCondition).where(
AlertCondition.id == alert_id,
AlertCondition.user_id == current_user.id,
)
)
alert = result.scalar_one_or_none()
if not alert:
raise HTTPException(status_code=404, detail="Alert not found")
await db.delete(alert)
await db.flush()
logger.info("Alert deleted: id=%d user=%s", alert_id, current_user.id)
return Response(status_code=204)
+410
View File
@@ -0,0 +1,410 @@
"""Dashboard Analytics API for trading portal.
Provides aggregated statistics and visualisation data drawn from
hypothetical (signal) trades and real trades.
"""
from __future__ import annotations
import logging
from datetime import date, datetime, timedelta, timezone
from decimal import Decimal
from typing import Any
from fastapi import APIRouter, Depends, Query
from sqlalchemy import text
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_current_active_user, get_db_session
from app.models.user import User
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/analytics", tags=["analytics"])
@router.get("")
async def analytics_root():
"""Redirect to dashboard and performance endpoints."""
return {
"message": "Analytics API",
"endpoints": {
"dashboard": "/api/v1/analytics/dashboard",
"performance": "/api/v1/analytics/performance",
}
}
@router.get("/performance")
async def get_performance(db: AsyncSession = Depends(get_db_session)):
"""Performance summary: win rate, PnL, profit factor."""
from sqlalchemy import text
result = await db.execute(text("""
SELECT
COUNT(*) FILTER (WHERE status='CLOSED') as total_trades,
COUNT(*) FILTER (WHERE status='CLOSED' AND pnl > 0) as wins,
COUNT(*) FILTER (WHERE status='CLOSED' AND pnl <= 0) as losses,
COALESCE(SUM(pnl) FILTER (WHERE status='CLOSED'), 0) as total_pnl,
COALESCE(SUM(pnl) FILTER (WHERE status='CLOSED' AND pnl > 0), 0) as total_profit,
COALESCE(ABS(SUM(pnl) FILTER (WHERE status='CLOSED' AND pnl <= 0)), 0) as total_loss
FROM hypothetical_trades
WHERE ABS(COALESCE(pnl_percent,0)) < 100
"""))
row = result.fetchone()
if not row:
return {"total_trades": 0, "win_rate": 0, "profit_factor": 0}
total = row[0] or 0
wins = row[1] or 0
total_pnl = float(row[3] or 0)
profit = float(row[4] or 0)
loss = float(row[5] or 1)
return {
"total_trades": total,
"wins": wins,
"losses": row[2] or 0,
"win_rate": round(wins / total * 100, 1) if total > 0 else 0,
"total_pnl": round(total_pnl, 2),
"profit_factor": round(profit / loss, 2) if loss > 0 else 0,
}
# ---------------------------------------------------------------------------
# Helper – run a raw SQL query and return a list of dicts
# ---------------------------------------------------------------------------
async def _fetchall(db: AsyncSession, sql: str, **params: Any) -> list[dict[str, Any]]:
result = await db.execute(text(sql), params)
columns = result.keys()
return [dict(zip(columns, row)) for row in result.fetchall()]
async def _fetchone(db: AsyncSession, sql: str, **params: Any) -> dict[str, Any] | None:
result = await db.execute(text(sql), params)
row = result.fetchone()
if row is None:
return None
columns = result.keys()
return dict(zip(columns, row))
# ---------------------------------------------------------------------------
# Dashboard endpoint
# ---------------------------------------------------------------------------
@router.get("/dashboard")
async def get_dashboard(
days: int = Query(90, ge=1, le=365, description="Look-back period in days"),
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user),
) -> dict[str, Any]:
"""Return aggregated dashboard statistics for the authenticated user."""
uid = current_user.id
cutoff = datetime.now(timezone.utc) - timedelta(days=days)
# ---- 1. PnL history (daily) ------------------------------------------------
pnl_history_sql = """
WITH combined AS (
SELECT
DATE(exit_time AT TIME ZONE 'UTC') AS trade_date,
COALESCE(pnl, 0) AS pnl
FROM hypothetical_trades
WHERE user_id = :uid AND exit_time >= :cutoff AND status = 'CLOSED'
UNION ALL
SELECT
DATE(closed_at AT TIME ZONE 'UTC') AS trade_date,
COALESCE(pnl, 0) AS pnl
FROM real_trades
WHERE user_id = :uid AND closed_at >= :cutoff AND status IN ('filled', 'cancelled')
),
daily AS (
SELECT trade_date, SUM(pnl::numeric) AS pnl
FROM combined
WHERE trade_date IS NOT NULL
GROUP BY trade_date
ORDER BY trade_date
)
SELECT
trade_date::text AS "date",
ROUND(pnl::numeric, 2) AS pnl,
ROUND(SUM(pnl::numeric) OVER (ORDER BY trade_date), 2) AS cumulative
FROM daily
ORDER BY trade_date
"""
pnl_history = await _fetchall(db, pnl_history_sql, uid=uid, cutoff=cutoff)
# ---- 2. Equity curve (daily balance change) --------------------------------
equity_sql = """
WITH combined AS (
SELECT
DATE(exit_time AT TIME ZONE 'UTC') AS trade_date,
COALESCE(pnl, 0) AS pnl
FROM hypothetical_trades
WHERE user_id = :uid AND exit_time >= :cutoff AND status = 'CLOSED'
UNION ALL
SELECT
DATE(closed_at AT TIME ZONE 'UTC') AS trade_date,
COALESCE(pnl, 0) AS pnl
FROM real_trades
WHERE user_id = :uid AND closed_at >= :cutoff AND status IN ('filled', 'cancelled')
),
daily AS (
SELECT trade_date, SUM(pnl::numeric) AS pnl
FROM combined
WHERE trade_date IS NOT NULL
GROUP BY trade_date
ORDER BY trade_date
)
SELECT
trade_date::text AS "date",
ROUND(
1000.0 + COALESCE(SUM(pnl::numeric) OVER (ORDER BY trade_date), 0),
2
) AS equity
FROM daily
ORDER BY trade_date
"""
equity_curve = await _fetchall(db, equity_sql, uid=uid, cutoff=cutoff)
# ---- 3. Drawdown -----------------------------------------------------------
drawdown_sql = """
WITH combined AS (
SELECT
DATE(exit_time AT TIME ZONE 'UTC') AS trade_date,
COALESCE(pnl, 0) AS pnl
FROM hypothetical_trades
WHERE user_id = :uid AND exit_time >= :cutoff AND status = 'CLOSED'
UNION ALL
SELECT
DATE(closed_at AT TIME ZONE 'UTC') AS trade_date,
COALESCE(pnl, 0) AS pnl
FROM real_trades
WHERE user_id = :uid AND closed_at >= :cutoff AND status IN ('filled', 'cancelled')
),
daily AS (
SELECT trade_date, SUM(pnl::numeric) AS pnl
FROM combined
WHERE trade_date IS NOT NULL
GROUP BY trade_date
ORDER BY trade_date
),
equity AS (
SELECT
trade_date,
1000.0 + COALESCE(SUM(pnl::numeric) OVER (ORDER BY trade_date), 0) AS equity
FROM daily
),
peaks AS (
SELECT
trade_date,
equity,
MAX(equity) OVER (ORDER BY trade_date) AS peak
FROM equity
)
SELECT
MIN(
CASE WHEN peak > 0
THEN ROUND(((peak - equity) / peak * 100)::numeric, 2)
ELSE 0
END
) AS max_drawdown_pct,
(
SELECT trade_date::text FROM peaks
WHERE (CASE WHEN peak > 0
THEN ((peak - equity) / peak * 100)
ELSE 0
END)
= (SELECT MIN(CASE WHEN peak > 0
THEN ((peak - equity) / peak * 100)
ELSE 0
END) FROM peaks)
LIMIT 1
) AS max_drawdown_date
FROM peaks
"""
dd_row = await _fetchone(db, drawdown_sql, uid=uid, cutoff=cutoff)
max_dd_pct: float = 0.0
max_dd_date: str | None = None
if dd_row and dd_row.get("max_drawdown_pct") is not None:
max_dd_pct = float(dd_row["max_drawdown_pct"])
max_dd_date = dd_row.get("max_drawdown_date")
# Current drawdown: last known equity vs its peak
if equity_curve:
last_equity = float(equity_curve[-1]["equity"])
running_peak = max(float(e["equity"]) for e in equity_curve)
current_dd_pct = round((running_peak - last_equity) / running_peak * 100, 2) if running_peak > 0 else 0.0
else:
current_dd_pct = 0.0
drawdown = {
"max_drawdown_pct": abs(max_dd_pct),
"current_drawdown_pct": abs(current_dd_pct),
"max_drawdown_date": max_dd_date or "",
}
# ---- 4. Win rate by period -------------------------------------------------
win_rate_sql = """
WITH combined AS (
SELECT
'hyp' AS src,
DATE(exit_time AT TIME ZONE 'UTC') AS trade_date,
COALESCE(pnl, 0) AS pnl
FROM hypothetical_trades
WHERE user_id = :uid AND status = 'CLOSED' AND exit_time IS NOT NULL
UNION ALL
SELECT
'real' AS src,
DATE(closed_at AT TIME ZONE 'UTC') AS trade_date,
COALESCE(pnl, 0) AS pnl
FROM real_trades
WHERE user_id = :uid AND status IN ('filled', 'cancelled') AND closed_at IS NOT NULL
)
SELECT
'daily' AS period,
COUNT(*) AS trades,
COUNT(*) FILTER (WHERE pnl > 0) AS wins,
ROUND(
(COUNT(*) FILTER (WHERE pnl > 0)::numeric / NULLIF(COUNT(*), 0)) * 100, 2
) AS win_rate,
ROUND(COALESCE(SUM(pnl::numeric), 0), 2) AS pnl
FROM combined
WHERE trade_date >= CURRENT_DATE - INTERVAL '1 day'
UNION ALL
SELECT
'weekly' AS period,
COUNT(*) AS trades,
COUNT(*) FILTER (WHERE pnl > 0) AS wins,
ROUND(
(COUNT(*) FILTER (WHERE pnl > 0)::numeric / NULLIF(COUNT(*), 0)) * 100, 2
) AS win_rate,
ROUND(COALESCE(SUM(pnl::numeric), 0), 2) AS pnl
FROM combined
WHERE trade_date >= CURRENT_DATE - INTERVAL '7 days'
UNION ALL
SELECT
'monthly' AS period,
COUNT(*) AS trades,
COUNT(*) FILTER (WHERE pnl > 0) AS wins,
ROUND(
(COUNT(*) FILTER (WHERE pnl > 0)::numeric / NULLIF(COUNT(*), 0)) * 100, 2
) AS win_rate,
ROUND(COALESCE(SUM(pnl::numeric), 0), 2) AS pnl
FROM combined
WHERE trade_date >= CURRENT_DATE - INTERVAL '30 days'
"""
wr_rows = await _fetchall(db, win_rate_sql, uid=uid)
win_rate_by_period: dict[str, dict[str, Any]] = {}
period_map = {"daily": "daily", "weekly": "weekly", "monthly": "monthly"}
for row in wr_rows:
period = row["period"]
win_rate_by_period[period] = {
"trades": row["trades"],
"wins": row["wins"],
"win_rate": float(row["win_rate"]) if row["win_rate"] is not None else 0,
"pnl": float(row["pnl"]),
}
# ---- 5. Best & worst trade -------------------------------------------------
best_worst_sql = """
WITH combined AS (
SELECT symbol, pnl::numeric AS pnl, entry_time::text AS entry_time
FROM hypothetical_trades
WHERE user_id = :uid AND status = 'CLOSED' AND pnl IS NOT NULL
UNION ALL
SELECT symbol, pnl::numeric AS pnl, created_at::text AS entry_time
FROM real_trades
WHERE user_id = :uid AND status IN ('filled', 'cancelled') AND pnl IS NOT NULL
)
SELECT
COALESCE(
(SELECT json_build_object('symbol', symbol, 'pnl', ROUND(pnl, 2), 'date', entry_time)
FROM combined ORDER BY pnl DESC LIMIT 1),
'{}'::json
) AS best,
COALESCE(
(SELECT json_build_object('symbol', symbol, 'pnl', ROUND(pnl, 2), 'date', entry_time)
FROM combined ORDER BY pnl ASC LIMIT 1),
'{}'::json
) AS worst
"""
bw = await _fetchone(db, best_worst_sql, uid=uid)
best_trade: dict[str, Any] = {"symbol": "", "pnl": 0.0, "date": ""}
worst_trade: dict[str, Any] = {"symbol": "", "pnl": 0.0, "date": ""}
if bw:
if bw.get("best") and isinstance(bw["best"], dict) and bw["best"].get("symbol"):
best_trade = {
"symbol": bw["best"]["symbol"],
"pnl": float(bw["best"]["pnl"]),
"date": bw["best"]["date"][:10] if bw["best"].get("date") else "",
}
if bw.get("worst") and isinstance(bw["worst"], dict) and bw["worst"].get("symbol"):
worst_trade = {
"symbol": bw["worst"]["symbol"],
"pnl": float(bw["worst"]["pnl"]),
"date": bw["worst"]["date"][:10] if bw["worst"].get("date") else "",
}
# ---- 6. Aggregate stats ----------------------------------------------------
agg_sql = """
WITH combined AS (
SELECT pnl::numeric AS pnl, direction, entry_time, exit_time
FROM hypothetical_trades
WHERE user_id = :uid AND status = 'CLOSED' AND pnl IS NOT NULL
UNION ALL
SELECT pnl::numeric AS pnl, side AS direction, created_at AS entry_time, closed_at AS exit_time
FROM real_trades
WHERE user_id = :uid AND status IN ('filled', 'cancelled') AND pnl IS NOT NULL
)
SELECT
COUNT(*) AS total_trades,
ROUND(COALESCE(SUM(pnl), 0), 2) AS total_pnl,
ROUND(COALESCE(AVG(pnl), 0), 2) AS avg_trade_pnl,
ROUND(
COALESCE(
SUM(pnl) FILTER (WHERE pnl > 0) / NULLIF(ABS(SUM(pnl) FILTER (WHERE pnl < 0)), 0),
0
), 2
) AS profit_factor,
ROUND(
COALESCE(
AVG(
EXTRACT(EPOCH FROM (exit_time - entry_time)) / 3600.0
), 0
), 2
) AS avg_hold_time_hours
FROM combined
"""
agg = await _fetchone(db, agg_sql, uid=uid)
total_trades: int = 0
total_pnl: float = 0.0
avg_trade_pnl: float = 0.0
profit_factor: float = 0.0
avg_hold_time_hours: float = 0.0
if agg:
total_trades = agg.get("total_trades") or 0
total_pnl = float(agg.get("total_pnl") or 0)
avg_trade_pnl = float(agg.get("avg_trade_pnl") or 0)
profit_factor = float(agg.get("profit_factor") or 0)
avg_hold_time_hours = float(agg.get("avg_hold_time_hours") or 0)
# ---- 7. Open trades count --------------------------------------------------
open_sql = """
SELECT COUNT(*) AS cnt FROM hypothetical_trades
WHERE user_id = :uid AND status = 'OPEN'
"""
open_row = await _fetchone(db, open_sql, uid=uid)
open_trades = open_row["cnt"] if open_row else 0
return {
"pnl_history": pnl_history,
"equity_curve": equity_curve,
"drawdown": drawdown,
"win_rate_by_period": win_rate_by_period,
"best_trade": best_trade,
"worst_trade": worst_trade,
"total_trades": total_trades,
"total_pnl": total_pnl,
"open_trades": open_trades,
"avg_trade_pnl": avg_trade_pnl,
"profit_factor": profit_factor,
"avg_hold_time_hours": avg_hold_time_hours,
}
+116
View File
@@ -0,0 +1,116 @@
"""API routes for Audit Log management."""
from __future__ import annotations
import logging
from uuid import UUID
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.core.deps import get_current_active_user, get_current_admin_user
from app.models.user import User
from app.models.audit_log import AuditLog
from app.services.audit_service import (
get_audit_logs as _get_audit_logs,
log_action as _log_action,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/audit", tags=["audit"])
# ---------------------------------------------------------------------------
# Schemas
# ---------------------------------------------------------------------------
class AuditLogEntryResponse(BaseModel):
id: int
user_id: str | None = None
action: str
resource: str
details: dict | None = None
created_at: str
@classmethod
def from_orm(cls, entry: AuditLog) -> AuditLogEntryResponse:
return cls(
id=entry.id,
user_id=str(entry.user_id) if entry.user_id else None,
action=entry.action,
resource=entry.resource,
details=entry.details,
created_at=entry.created_at.isoformat(),
)
class AuditLogListResponse(BaseModel):
items: list[AuditLogEntryResponse]
total: int
limit: int
offset: int
class AuditLogCreateRequest(BaseModel):
action: str
resource: str
details: dict | None = None
# ---------------------------------------------------------------------------
# Endpoints
# ---------------------------------------------------------------------------
@router.get("", response_model=AuditLogListResponse)
async def list_audit_logs(
limit: int = Query(100, ge=1, le=500),
offset: int = Query(0, ge=0),
action: str | None = Query(None),
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_admin_user),
) -> AuditLogListResponse:
"""List audit log entries. Admin only.
Supports pagination and optional action type filtering.
"""
entries, total = await _get_audit_logs(
db, limit=limit, offset=offset, action=action
)
return AuditLogListResponse(
items=[AuditLogEntryResponse.from_orm(e) for e in entries],
total=total,
limit=limit,
offset=offset,
)
@router.post("", response_model=AuditLogEntryResponse)
async def create_audit_log(
body: AuditLogCreateRequest,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_active_user),
) -> AuditLogEntryResponse:
"""Log an action manually. Available to any authenticated user."""
entry = await _log_action(
db,
user_id=current_user.id,
action=body.action,
resource=body.resource,
details=body.details,
)
return AuditLogEntryResponse.from_orm(entry)
@router.get("/logs", response_model=AuditLogListResponse)
async def list_audit_logs_alias(
db: AsyncSession = Depends(get_db),
action: str | None = Query(None),
limit: int = Query(50, ge=1, le=200),
offset: int = Query(0, ge=0),
):
"""Alias for GET /audit — same as root endpoint."""
return await list_audit_logs(db=db, action=action, limit=limit, offset=offset)
+150
View File
@@ -0,0 +1,150 @@
from __future__ import annotations
from uuid import UUID
from fastapi import APIRouter, Depends, Request, status
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_current_user, get_db_session
from app.core.exceptions import AppException
from app.models import User
from app.schemas import (
ChangePasswordRequest,
LoginRequest,
LogoutRequest,
RefreshRequest,
RegisterRequest,
TokenResponse,
UserResponse,
UserSessionResponse,
UserUpdateRequest,
)
from app.services import auth_service
router = APIRouter(prefix="/auth", tags=["auth"])
@router.post("/register", status_code=status.HTTP_201_CREATED, response_model=UserResponse)
async def register(
req: RegisterRequest,
db: AsyncSession = Depends(get_db_session),
) -> UserResponse:
"""Register a new user account."""
return await auth_service.register(db=db, req=req)
@router.post("/login", response_model=TokenResponse)
async def login(
req: LoginRequest,
request: Request,
db: AsyncSession = Depends(get_db_session),
) -> TokenResponse:
"""Authenticate with username/password and receive a token pair."""
return await auth_service.login(
db=db,
req=req,
user_agent=request.headers.get("User-Agent", ""),
ip_address=request.client.host if request.client else "",
)
@router.post("/refresh", response_model=TokenResponse)
async def refresh(
req: RefreshRequest,
db: AsyncSession = Depends(get_db_session),
) -> TokenResponse:
"""Exchange a valid refresh token for a new token pair (rotation)."""
return await auth_service.refresh_token(db=db, refresh_token_str=req.refresh_token)
@router.post("/logout", status_code=status.HTTP_200_OK)
async def logout(
req: LogoutRequest,
db: AsyncSession = Depends(get_db_session),
) -> dict:
"""Revoke a refresh token so it can no longer be used."""
await auth_service.logout(db=db, refresh_token_str=req.refresh_token)
return {"message": "Successfully logged out"}
@router.get("/me", response_model=UserResponse)
async def get_me(
current_user: User = Depends(get_current_user),
) -> UserResponse:
"""Return the authenticated user's profile."""
return UserResponse(
id=current_user.id,
username=current_user.username,
email=current_user.email,
display_name=current_user.display_name,
is_active=current_user.is_active,
is_admin=current_user.is_admin,
role=current_user.role,
preferences=current_user.preferences,
created_at=current_user.created_at,
)
@router.put("/me", response_model=UserResponse)
async def update_me(
req: UserUpdateRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> UserResponse:
"""Update the authenticated user's profile (email, display_name, preferences)."""
if req.email is not None:
current_user.email = req.email
if req.display_name is not None:
current_user.display_name = req.display_name
if req.preferences is not None:
current_user.preferences = req.preferences
await db.commit()
await db.refresh(current_user)
return UserResponse(
id=current_user.id,
username=current_user.username,
email=current_user.email,
display_name=current_user.display_name,
is_active=current_user.is_active,
is_admin=current_user.is_admin,
role=current_user.role,
preferences=current_user.preferences,
created_at=current_user.created_at,
)
@router.get("/sessions", response_model=list[UserSessionResponse])
async def list_sessions(
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> list[UserSessionResponse]:
"""Return all active (non-revoked, non-expired) sessions for the current user."""
return await auth_service.get_user_sessions(db=db, user_id=current_user.id)
@router.delete("/sessions/{token_hash}", status_code=status.HTTP_200_OK)
async def revoke_session(
token_hash: str,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> dict:
"""Revoke a specific session by its token hash."""
await auth_service.revoke_session(db=db, token_hash=token_hash, user_id=current_user.id)
return {"message": "Session revoked"}
@router.post("/change-password", status_code=status.HTTP_200_OK)
async def change_password(
req: ChangePasswordRequest,
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db_session),
) -> dict:
"""Change the current user's password (revokes all other sessions)."""
await auth_service.change_password(
db=db,
user_id=current_user.id,
old_password=req.old_password,
new_password=req.new_password,
)
return {"message": "Password changed successfully"}
+416
View File
@@ -0,0 +1,416 @@
"""Backtest API endpoint — run backtest and return JSON results."""
import asyncio
import logging
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from collections import defaultdict
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy import select, and_, func
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.models.candle import Candle
from app.models.symbol import Symbol
from app.models.exchange import Exchange
from app.services.indicator_service import (
bollinger_bands, rsi, sma, macd, supertrend,
volume_breakout, ichimoku, detect_divergence, market_structure,
)
from app.services.signal_service import (
_classify_signal_combined,
STRONG_BUY, BUY, STRONG_SELL, SELL,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/backtest", tags=["backtest"])
TRADE_SIZE = Decimal("10")
MAX_HOLD_CANDLES = 48
@router.get("")
async def backtest_root():
"""Redirect to run endpoint."""
return {"message": "Use GET /backtest/run to run a backtest"}
async def _run_backtest(
db: AsyncSession,
symbol: str,
exchange: str,
timeframe: str = "30m",
days: int = 7,
trade_size: Decimal = Decimal("10"),
) -> dict:
"""Run backtest and return structured results."""
result = await db.execute(
select(Symbol)
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(and_(Exchange.name == exchange, Symbol.symbol == symbol))
)
db_symbol = result.scalar_one_or_none()
if not db_symbol:
return {"error": f"Symbol {symbol} not found on {exchange}"}
cutoff = datetime.now(timezone.utc) - timedelta(days=days)
result = await db.execute(
select(Candle)
.where(and_(
Candle.symbol_id == db_symbol.id,
Candle.timeframe == timeframe,
Candle.timestamp >= cutoff,
))
.order_by(Candle.timestamp.asc())
)
candles = list(result.scalars().all())
# Min candles for warmup: BB(20) + RSI(14) + some room = 30
MIN_CANDLES = 30
if len(candles) < MIN_CANDLES:
if candles:
span_hours = (candles[-1].timestamp - candles[0].timestamp).total_seconds() / 3600
if span_hours >= 24:
avail = f"{span_hours/24:.0f}d"
else:
avail = f"{span_hours:.0f}h"
return {"error": f"Only ~{avail} data available (need at least {MIN_CANDLES} candles for {timeframe}). Try more days or a higher timeframe (4h)."}
return {"error": f"No candle data found for {symbol} on {timeframe}. The exchange may not support this pair."}
# MTF config
tf_minutes = {"15m": 15, "30m": 30, "1h": 60, "4h": 240}
main_minutes = tf_minutes.get(timeframe, 30)
mtf_config = []
for mtf_tf, mtf_minutes, mtf_w in [("15m", 15, 0.5), ("1h", 60, 1.5), ("4h", 240, 2.0)]:
if mtf_tf == timeframe:
continue
mult = mtf_minutes // main_minutes
if mult >= 1 and len(candles) >= mult * MIN_CANDLES:
mtf_config.append((mtf_tf, mult, mtf_w))
# ── Pre-compute all candle data & indicators ONCE ──
candle_dicts_full = [
{"high": float(c.high), "low": float(c.low),
"close": float(c.close), "open": float(c.open),
"volume": float(c.volume)}
for c in candles
]
close_prices_full = [float(c.close) for c in candles]
# Pre-compute indicators on full dataset (O(n) instead of O(n²))
bb_full = bollinger_bands(close_prices_full) or {}
rsi_full = rsi(close_prices_full) or []
sma_full = sma(close_prices_full, 20) or []
macd_full = macd(close_prices_full) or {}
st_full = supertrend(candle_dicts_full) or {}
vb_full = volume_breakout(candle_dicts_full) or []
ichi_full = ichimoku(candle_dicts_full) or {}
smc_full = market_structure(candle_dicts_full) or {}
# Pre-compute divergence ONCE (uses full arrays, indexes match)
rsi_div_full = detect_divergence(close_prices_full, rsi_full)
macd_hist_full = macd_full.get("histogram", []) if macd_full else []
macd_div_full = detect_divergence(close_prices_full, macd_hist_full)
# Pre-build MTF candles ONCE per MTF config
mtf_precomputed = []
for mtf_name, mtf_mult, mtf_w in mtf_config:
mtf_candles_list = []
for j in range(0, len(candle_dicts_full) - mtf_mult + 1, mtf_mult):
chunk = candle_dicts_full[j:j + mtf_mult]
mtf_candles_list.append({
"open": chunk[0]["open"],
"high": max(c["high"] for c in chunk),
"low": min(c["low"] for c in chunk),
"close": chunk[-1]["close"],
"volume": sum(c["volume"] for c in chunk),
})
if len(mtf_candles_list) >= MIN_CANDLES:
mtf_p = [c["close"] for c in mtf_candles_list]
mtf_precomputed.append({
"name": mtf_name,
"weight": mtf_w,
"mult": mtf_mult,
"candles_list": mtf_candles_list,
"close_prices": mtf_p,
"bb": bollinger_bands(mtf_p) or {},
"rsi": rsi(mtf_p) or [],
"sma": sma(mtf_p, 20) or [],
"macd": macd(mtf_p) or {},
"st": supertrend(mtf_candles_list) or {},
"vb": volume_breakout(mtf_candles_list) or [],
"ichi": ichimoku(mtf_candles_list) or {},
"smc": market_structure(mtf_candles_list) or {},
})
all_signals = []
trades = []
current_position = None
for i in range(MIN_CANDLES, len(candles)):
candle = candles[i]
latest_close = close_prices_full[i]
timestamp = candle.timestamp.isoformat()
# Slice pre-computed arrays (O(i) but ~100x faster than recomputing)
clip = i + 1
def _safe_slice(v):
return v[:clip] if v is not None and hasattr(v, '__getitem__') else v
bb_data = {k: _safe_slice(v) for k, v in bb_full.items()} if bb_full else {}
rsi_data = rsi_full[:clip] if rsi_full else []
sma_data = sma_full[:clip] if sma_full else []
macd_data = {k: _safe_slice(v) for k, v in macd_full.items()} if macd_full else {}
st_data = {k: _safe_slice(v) for k, v in st_full.items()} if st_full else {}
vb_data = vb_full[:clip] if vb_full else []
ichi_data = {k: _safe_slice(v) for k, v in ichi_full.items()} if ichi_full else {}
smc_data = {k: _safe_slice(v) for k, v in smc_full.items()} if smc_full else {}
# MTF votes — use precomputed MTF indicators, sliced to current MTF candle index
mtf_votes = []
for mtf in mtf_precomputed:
# Which MTF candle corresponds to main candle i?
mtf_idx = i // mtf["mult"]
if mtf_idx < MIN_CANDLES or mtf_idx >= len(mtf["close_prices"]):
continue
clip_mtf = mtf_idx + 1
def _safe_slice_mtf(v):
return v[:clip_mtf] if v is not None and hasattr(v, '__getitem__') else v
mtf_s, *_ = _classify_signal_combined(
mtf["close_prices"][mtf_idx],
{k: _safe_slice_mtf(v) for k, v in mtf["bb"].items()},
mtf["rsi"][:clip_mtf],
mtf["sma"][:clip_mtf],
{k: _safe_slice_mtf(v) for k, v in mtf["macd"].items()} if mtf["macd"] else None,
{k: _safe_slice_mtf(v) for k, v in mtf["st"].items()} if mtf["st"] else None,
mtf["vb"][:clip_mtf] if mtf["vb"] else None,
{k: _safe_slice_mtf(v) for k, v in mtf["ichi"].items()} if mtf["ichi"] else None,
(None, None), (None, None),
{k: _safe_slice_mtf(v) for k, v in mtf["smc"].items()} if mtf["smc"] else None,
)
if mtf_s:
mtf_votes.append((mtf_s, "", mtf["weight"]))
signal_type, strength, *_ = _classify_signal_combined(
latest_close, bb_data, rsi_data, sma_data,
macd_data, st_data, vb_data, ichi_data,
rsi_div_full, macd_div_full, smc_data, mtf_votes or None,
)
if signal_type:
all_signals.append({
"time": timestamp, "signal": signal_type,
"strength": strength or "", "price": latest_close,
})
# PnL simulation
if signal_type in (STRONG_BUY, BUY):
if current_position and current_position["direction"] == "SHORT":
if signal_type == STRONG_BUY:
entry = current_position["entry_price"]
qty = current_position["quantity"]
pnl = (entry - latest_close) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
})
trades.append(current_position)
current_position = None
else:
continue
if not current_position:
qty = float(trade_size) / latest_close
current_position = {
"direction": "LONG", "entry_price": latest_close,
"entry_time": timestamp, "quantity": qty,
"entry_signal": signal_type, "entry_index": i, "status": "OPEN",
}
elif signal_type in (STRONG_SELL, SELL):
if current_position and current_position["direction"] == "LONG":
if signal_type == STRONG_SELL:
entry = current_position["entry_price"]
qty = current_position["quantity"]
pnl = (latest_close - entry) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "REVERSAL",
})
trades.append(current_position)
current_position = None
else:
continue
if not current_position:
qty = float(trade_size) / latest_close
current_position = {
"direction": "SHORT", "entry_price": latest_close,
"entry_time": timestamp, "quantity": qty,
"entry_signal": signal_type, "entry_index": i, "status": "OPEN",
}
# Time limit
if current_position and current_position["status"] == "OPEN":
hold = i - current_position["entry_index"]
if hold >= MAX_HOLD_CANDLES:
entry = current_position["entry_price"]
qty = current_position["quantity"]
if current_position["direction"] == "LONG":
pnl = (latest_close - entry) * qty
else:
pnl = (entry - latest_close) * qty
current_position.update({
"exit_price": latest_close, "exit_time": timestamp,
"pnl": pnl, "status": "CLOSED", "exit_reason": "TIME_LIMIT",
})
trades.append(current_position)
current_position = None
# Close final position
if current_position and current_position["status"] == "OPEN":
last_close = float(candles[-1].close)
entry = current_position["entry_price"]
qty = current_position["quantity"]
if current_position["direction"] == "LONG":
pnl = (last_close - entry) * qty
else:
pnl = (entry - last_close) * qty
current_position.update({
"exit_price": last_close,
"exit_time": candles[-1].timestamp.isoformat(),
"pnl": pnl, "status": "CLOSED", "exit_reason": "END_OF_DATA",
})
trades.append(current_position)
# Compute stats
counts = defaultdict(int)
for s in all_signals:
counts[s["signal"]] += 1
closed_trades = [t for t in trades if t.get("status") == "CLOSED"]
winning_trades = [t for t in closed_trades if t.get("pnl", 0) > 0]
losing_trades = [t for t in closed_trades if t.get("pnl", 0) <= 0]
total_pnl = sum(t.get("pnl", 0) for t in closed_trades)
gross_profit = sum(t.get("pnl", 0) for t in winning_trades)
gross_loss = sum(t.get("pnl", 0) for t in losing_trades)
win_rate = round(len(winning_trades) / len(closed_trades) * 100, 1) if closed_trades else 0
profit_factor = round(abs(gross_profit / gross_loss), 2) if gross_loss != 0 else None
avg_win = round(gross_profit / len(winning_trades), 2) if winning_trades else None
avg_loss = round(gross_loss / len(losing_trades), 2) if losing_trades else None
best_trade = max(closed_trades, key=lambda t: t.get("pnl", 0)) if closed_trades else None
worst_trade = min(closed_trades, key=lambda t: t.get("pnl", 0)) if closed_trades else None
return {
"symbol": symbol,
"exchange": exchange,
"timeframe": timeframe,
"days": days,
"candles_count": len(candles),
"signal_counts": dict(counts),
"total_signals": len(all_signals),
"recent_signals": all_signals[-15:],
"trades": {
"total": len(closed_trades),
"wins": len(winning_trades),
"losses": len(losing_trades),
"win_rate": win_rate,
"total_pnl": round(total_pnl, 2),
"profit_factor": profit_factor,
"avg_win": avg_win,
"avg_loss": avg_loss,
"best_trade": {
"direction": best_trade.get("direction"),
"entry_price": round(best_trade["entry_price"], 4),
"exit_price": round(best_trade["exit_price"], 4),
"pnl": round(best_trade["pnl"], 2),
"entry_signal": best_trade.get("entry_signal"),
} if best_trade else None,
"worst_trade": {
"direction": worst_trade.get("direction"),
"entry_price": round(worst_trade["entry_price"], 4),
"exit_price": round(worst_trade["exit_price"], 4),
"pnl": round(worst_trade["pnl"], 2),
"entry_signal": worst_trade.get("entry_signal"),
} if worst_trade else None,
"per_signal": {},
"recent": closed_trades[-10:],
},
}
@router.get("/symbols")
async def get_backtest_symbols(
exchange: str = Query(None, description="Exchange name filter (e.g., binance, bybit)"),
db: AsyncSession = Depends(get_db),
):
"""Return symbols with sufficient candles (>=30 in each of 30m/1h/4h/1d) for backtesting."""
MIN_CANDLES = 30
TFS = ["30m", "1h", "4h", "1d"]
# Subquery: symbol_id + timeframe that have >= MIN_CANDLES
conditions = [Candle.timeframe.in_(TFS)]
if exchange:
conditions.append(Exchange.name == exchange)
base_query = (
select(Candle.symbol_id, Candle.timeframe, func.count().label("cnt"))
.join(Symbol, Symbol.id == Candle.symbol_id)
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(and_(*conditions))
.group_by(Candle.symbol_id, Candle.timeframe)
.having(func.count() >= MIN_CANDLES)
).subquery()
# Symbols that have ALL 4 timeframes with >= MIN_CANDLES
query = (
select(Symbol.symbol, Exchange.name.label("exchange"))
.select_from(base_query)
.join(Symbol, Symbol.id == base_query.c.symbol_id)
.join(Exchange, Exchange.id == Symbol.exchange_id)
.group_by(Symbol.symbol, Exchange.name)
.having(func.count(func.distinct(base_query.c.timeframe)) == len(TFS))
.order_by(Symbol.symbol)
)
result = await db.execute(query)
rows = result.all()
return {
"exchange": exchange or "all",
"count": len(rows),
"symbols": [
{"symbol": r[0], "exchange": r[1]}
for r in rows
],
}
@router.get("/run")
async def run_backtest(
symbol: str = Query("BTC/USDT"),
exchange: str = Query("mexc"),
timeframe: str = Query("30m"),
days: int = Query(7),
trade_size: float = Query(10.0),
db: AsyncSession = Depends(get_db),
):
"""Run backtest and return JSON results."""
result = await _run_backtest(
db, symbol, exchange, timeframe, days, Decimal(str(trade_size))
)
if "error" in result:
raise HTTPException(status_code=400, detail=result["error"])
return result
@router.post("/run")
async def run_backtest_post(
symbol: str = Query("BTC/USDT"),
exchange: str = Query("mexc"),
timeframe: str = Query("30m"),
days: int = Query(7),
trade_size: float = Query(10.0),
db: AsyncSession = Depends(get_db),
):
"""Alias for GET /backtest/run — supports POST method."""
return await run_backtest(symbol, exchange, timeframe, days, trade_size, db)
+228
View File
@@ -0,0 +1,228 @@
"""API routes for backtest history (save/load per user) and comparison."""
from __future__ import annotations
import json
import logging
from datetime import datetime, timezone
from uuid import UUID, uuid4
from decimal import Decimal
from typing import List
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, desc, delete, text
from app.database import get_db
from app.core.deps import get_current_user
from app.models.user import User as UserModel
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/backtest", tags=["backtest_history"])
class CompareRequest(BaseModel):
"""Request body for comparing multiple backtest results."""
ids: list[str]
@router.post("/save")
async def save_backtest(
symbol: str = Query(...),
exchange: str = Query("mexc"),
timeframe: str = Query(...),
days: int = Query(...),
trade_size: float = Query(10.0),
total_trades: int = Query(0),
wins: int = Query(0),
losses: int = Query(0),
win_rate: float = Query(0.0),
total_pnl: float = Query(0.0),
profit_factor: float | None = Query(None),
avg_win: float | None = Query(None),
avg_loss: float | None = Query(None),
db: AsyncSession = Depends(get_db),
current_user: UserModel = Depends(get_current_user),
):
"""Save a backtest result to the user's history."""
bt_id = uuid4()
raw_result = {
"symbol": symbol, "exchange": exchange, "timeframe": timeframe,
"days": days, "trade_size": trade_size,
"total_trades": total_trades, "wins": wins, "losses": losses,
"win_rate": win_rate, "total_pnl": total_pnl,
"profit_factor": profit_factor, "avg_win": avg_win, "avg_loss": avg_loss,
}
await db.execute(
text("""
INSERT INTO user_backtests
(id, user_id, symbol, exchange, timeframe, days, trade_size,
total_trades, wins, losses, win_rate, total_pnl,
profit_factor, avg_win, avg_loss, result_json, created_at)
VALUES (:id, :uid, :symbol, :exchange, :tf, :days, :ts,
:tt, :w, :l, :wr, :tp,
:pf, :aw, :al, :rj, :ca)
"""),
{
"id": bt_id, "uid": current_user.id,
"symbol": symbol, "exchange": exchange, "tf": timeframe,
"days": days, "ts": trade_size,
"tt": total_trades, "w": wins, "l": losses,
"wr": win_rate, "tp": total_pnl,
"pf": profit_factor, "aw": avg_win, "al": avg_loss,
"rj": json.dumps(raw_result, default=str),
"ca": datetime.now(timezone.utc),
}
)
await db.commit()
return {"id": str(bt_id), "message": "Backtest saved"}
@router.get("/history")
async def list_backtests(
limit: int = Query(50, ge=1, le=200),
db: AsyncSession = Depends(get_db),
current_user: UserModel = Depends(get_current_user),
):
"""List backtest history for the current user."""
result = await db.execute(
text("""
SELECT id, symbol, exchange, timeframe, days, trade_size,
total_trades, wins, losses, win_rate, total_pnl,
profit_factor, avg_win, avg_loss, created_at
FROM user_backtests
WHERE user_id = :uid
ORDER BY created_at DESC
LIMIT :limit
"""),
{"uid": current_user.id, "limit": limit}
)
rows = result.fetchall()
return [
{
"id": str(r[0]),
"symbol": r[1], "exchange": r[2],
"timeframe": r[3], "days": r[4],
"trade_size": float(r[5]) if r[5] else None,
"total_trades": r[6], "wins": r[7], "losses": r[8],
"win_rate": float(r[9]) if r[9] else None,
"total_pnl": float(r[10]) if r[10] else None,
"profit_factor": float(r[11]) if r[11] else None,
"avg_win": float(r[12]) if r[12] else None,
"avg_loss": float(r[13]) if r[13] else None,
"created_at": r[14].isoformat() if r[14] else None,
}
for r in rows
]
@router.delete("/{bt_id}")
async def delete_backtest(
bt_id: str,
db: AsyncSession = Depends(get_db),
current_user: UserModel = Depends(get_current_user),
):
"""Delete a backtest result."""
try:
bt_uuid = UUID(bt_id)
except ValueError:
raise HTTPException(400, "Invalid ID")
result = await db.execute(
text("DELETE FROM user_backtests WHERE id = :id AND user_id = :uid"),
{"id": bt_uuid, "uid": current_user.id}
)
await db.commit()
if result.rowcount == 0:
raise HTTPException(404, "Backtest not found")
return {"message": "Deleted"}
@router.post("/compare")
async def compare_backtests(
req: CompareRequest,
db: AsyncSession = Depends(get_db),
current_user: UserModel = Depends(get_current_user),
):
"""Compare multiple backtest results by IDs, enriched with computed fields."""
if not req.ids:
raise HTTPException(status_code=400, detail="ids list is required")
# Validate and deduplicate IDs
uuids: list[UUID] = []
seen: set[str] = set()
for raw_id in req.ids:
if raw_id in seen:
continue
seen.add(raw_id)
try:
uuids.append(UUID(raw_id))
except ValueError:
raise HTTPException(400, f"Invalid UUID: {raw_id}")
if not uuids:
raise HTTPException(400, "No valid IDs provided")
result = await db.execute(
text("""
SELECT id, symbol, exchange, timeframe, days, trade_size,
total_trades, wins, losses, win_rate, total_pnl,
profit_factor, avg_win, avg_loss, created_at
FROM user_backtests
WHERE id = ANY(:ids) AND user_id = :uid
ORDER BY created_at DESC
"""),
{"ids": uuids, "uid": current_user.id},
)
rows = result.fetchall()
if not rows:
raise HTTPException(404, "No backtest records found for the given IDs")
# Build base records (same format as /history)
records = []
for r in rows:
total_pnl_val = float(r[10]) if r[10] is not None else 0.0
trade_size_val = float(r[5]) if r[5] is not None else 10.0
# Equity curve: [1.0, 1.0 + pnl_ratio]
# pnl_ratio = total_pnl / trade_size (approximate % return)
pnl_ratio = total_pnl_val / trade_size_val if trade_size_val != 0 else 0.0
equity_curve_series = [1.0, round(1.0 + pnl_ratio, 6)]
# Prediction based on PnL direction
if total_pnl_val > 0:
prediction = "BUY"
elif total_pnl_val < 0:
prediction = "SELL"
else:
prediction = "NEUTRAL"
records.append({
"id": str(r[0]),
"symbol": r[1],
"exchange": r[2],
"timeframe": r[3],
"days": r[4],
"trade_size": trade_size_val,
"total_trades": r[6],
"wins": r[7],
"losses": r[8],
"win_rate": float(r[9]) if r[9] is not None else None,
"total_pnl": total_pnl_val,
"profit_factor": float(r[11]) if r[11] is not None else None,
"avg_win": float(r[12]) if r[12] is not None else None,
"avg_loss": float(r[13]) if r[13] is not None else None,
"created_at": r[14].isoformat() if r[14] else None,
"equity_curve_series": equity_curve_series,
"prediction": prediction,
})
# Add rank: sort by total_pnl descending, 1 = best
records.sort(key=lambda rec: rec["total_pnl"], reverse=True)
for rank_idx, rec in enumerate(records, start=1):
rec["rank"] = rank_idx
return records
+476
View File
@@ -0,0 +1,476 @@
"""API endpoints for user credentials (exchange API keys) management."""
from __future__ import annotations
import logging
from decimal import Decimal
from fastapi import APIRouter, Depends, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import joinedload
from app.core.deps import get_current_user, get_db_session
from app.core.exceptions import NotFoundException
from app.core.security import (
decrypt_api_key,
encrypt_api_key,
)
from app.models import Exchange, ExchangeCredential, User
from app.schemas import (
CredentialCreateRequest,
CredentialResponse,
CredentialUpdateRequest,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/credentials", tags=["credentials"])
@router.get("", response_model=list[CredentialResponse])
async def list_credentials(
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> list[CredentialResponse]:
"""List all API keys for the current user (keys are masked)."""
result = await db.execute(
select(ExchangeCredential)
.options(joinedload(ExchangeCredential.exchange))
.where(ExchangeCredential.user_id == current_user.id)
.order_by(ExchangeCredential.created_at.desc())
)
creds = result.scalars().all()
return [
CredentialResponse(
id=c.id,
exchange_id=c.exchange_id,
exchange_name=c.exchange.name if c.exchange else "unknown",
api_key=c.api_key,
is_testnet=c.is_testnet,
is_active=c.is_active,
)
for c in creds
]
@router.post("", response_model=CredentialResponse, status_code=status.HTTP_201_CREATED)
async def create_credential(
body: CredentialCreateRequest,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> CredentialResponse:
"""Add a new API key for an exchange."""
# Verify exchange exists
result = await db.execute(select(Exchange).where(Exchange.id == body.exchange_id))
exchange = result.scalar_one_or_none()
if exchange is None:
raise NotFoundException(detail=f"Exchange with id {body.exchange_id} not found")
# Encrypt the secret
secret_enc, iv = encrypt_api_key(body.api_secret)
credential = ExchangeCredential(
user_id=current_user.id,
exchange_id=body.exchange_id,
api_key=body.api_key,
api_secret_enc=secret_enc,
api_secret_iv=iv,
passphrase=body.passphrase,
is_testnet=body.is_testnet,
)
db.add(credential)
await db.flush()
await db.refresh(credential)
return CredentialResponse(
id=credential.id,
exchange_id=credential.exchange_id,
exchange_name=exchange.display_name or exchange.name,
api_key=credential.api_key,
is_testnet=credential.is_testnet,
is_active=credential.is_active,
)
@router.put("/{credential_id}", response_model=CredentialResponse)
async def update_credential(
credential_id: str,
body: CredentialUpdateRequest,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> CredentialResponse:
"""Update API key, secret, passphrase, or deactivate."""
from uuid import UUID
result = await db.execute(
select(ExchangeCredential).where(
ExchangeCredential.id == UUID(credential_id),
ExchangeCredential.user_id == current_user.id,
)
)
cred = result.scalar_one_or_none()
if cred is None:
raise NotFoundException(detail="Credential not found")
if body.api_key is not None:
cred.api_key = body.api_key
if body.api_secret is not None:
secret_enc, iv = encrypt_api_key(body.api_secret)
cred.api_secret_enc = secret_enc
cred.api_secret_iv = iv
if body.passphrase is not None:
cred.passphrase = body.passphrase
if body.is_active is not None:
cred.is_active = body.is_active
await db.flush()
await db.refresh(cred)
return CredentialResponse(
id=cred.id,
exchange_id=cred.exchange_id,
exchange_name=cred.exchange.display_name if cred.exchange else "unknown",
api_key=cred.api_key,
is_testnet=cred.is_testnet,
is_active=cred.is_active,
)
@router.delete("/{credential_id}", status_code=200)
async def delete_credential(
credential_id: str,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> dict:
"""Delete an API key."""
from uuid import UUID
result = await db.execute(
select(ExchangeCredential).where(
ExchangeCredential.id == UUID(credential_id),
ExchangeCredential.user_id == current_user.id,
)
)
cred = result.scalar_one_or_none()
if cred is None:
raise NotFoundException(detail="Credential not found")
await db.delete(cred)
await db.commit()
return {"message": "API key deleted"}
@router.post("/{credential_id}/test")
async def test_credential(
credential_id: str,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> dict:
"""Test an API key by fetching markets from the exchange."""
from uuid import UUID
result = await db.execute(
select(ExchangeCredential)
.options(joinedload(ExchangeCredential.exchange))
.where(
ExchangeCredential.id == UUID(credential_id),
ExchangeCredential.user_id == current_user.id,
)
)
cred = result.scalar_one_or_none()
if cred is None:
raise NotFoundException(detail="Credential not found")
try:
# Decrypt secret
secret = decrypt_api_key(cred.api_secret_enc, iv_hex=cred.api_secret_iv)
# Try to create exchange adapter and test
from app.exchange.factory import factory as exchange_factory
exchange_name = cred.exchange.name if cred.exchange else "mexc"
# Quick test: try to load markets with credentials
adapter = exchange_factory.create(
exchange_name,
api_key=cred.api_key,
api_secret=secret,
)
import asyncio
loop = asyncio.get_running_loop()
markets = await loop.run_in_executor(
None, lambda: adapter.client.load_markets()
)
return {
"success": True,
"message": f"Connected to {exchange_name.upper()} successfully",
"markets": min(len(markets), 100),
}
except Exception as e:
logger.warning(f"Credential test failed: {e}")
return {
"success": False,
"message": f"Connection failed: {str(e)[:200]}",
}
# ── Account info endpoints ──
async def _get_adapter_for_credential(
credential_id: str, db: AsyncSession, current_user: User
):
"""Helper: decrypt & create exchange adapter for a credential."""
from uuid import UUID
result = await db.execute(
select(ExchangeCredential)
.options(joinedload(ExchangeCredential.exchange))
.where(
ExchangeCredential.id == UUID(credential_id),
ExchangeCredential.user_id == current_user.id,
)
)
cred = result.scalar_one_or_none()
if cred is None:
raise NotFoundException(detail="Credential not found")
secret = decrypt_api_key(cred.api_secret_enc, iv_hex=cred.api_secret_iv)
exchange_name = cred.exchange.name if cred.exchange else "mexc"
from app.exchange.factory import factory as exchange_factory
adapter = exchange_factory.create(
exchange_name,
api_key=cred.api_key,
api_secret=secret,
)
return cred, adapter, exchange_name
def _parse_balances_from_ccxt(raw_balance: dict) -> tuple[list, Decimal]:
"""Parse CCXT balance dict into BalanceData list + USDT total."""
from decimal import Decimal as D
balances: list = []
usdt_total = D("0")
for asset, info in raw_balance.get("total", {}).items():
if asset in ("info", "free", "used", "total"):
continue
free = D(str(raw_balance.get("free", {}).get(asset, 0)))
used = D(str(raw_balance.get("used", {}).get(asset, 0)))
total = D(str(info))
if total > 0 or free > 0:
balances.append({
"asset": asset,
"free": float(free),
"used": float(used),
"total": float(total),
})
if asset == "USDT":
usdt_total = total
return balances, usdt_total
def _parse_open_orders_from_ccxt(raw_orders: list) -> list:
"""Parse CCXT open orders into OpenOrderData dicts."""
from decimal import Decimal as D
orders: list = []
for raw in raw_orders or []:
orders.append({
"order_id": str(raw.get("id", "")),
"symbol": raw.get("symbol", ""),
"side": raw.get("side", ""),
"order_type": raw.get("type", ""),
"amount": float(D(str(raw.get("amount", 0)))),
"filled": float(D(str(raw.get("filled", 0)))),
"price": float(D(str(raw["price"]))) if raw.get("price") else None,
"average": float(D(str(raw["average"]))) if raw.get("average") else None,
"status": raw.get("status", "open"),
"timestamp": raw.get("timestamp"),
})
return orders
def _parse_positions_from_ccxt(raw_positions: list) -> list:
"""Parse CCXT positions into PositionData dicts."""
from decimal import Decimal as D
positions: list = []
for raw in raw_positions or []:
contracts = D(str(raw.get("contracts", 0)))
if contracts == 0:
continue
positions.append({
"symbol": raw.get("symbol", ""),
"side": raw.get("side", ""),
"contracts": float(contracts),
"entry_price": float(D(str(raw["entryPrice"]))) if raw.get("entryPrice") else None,
"mark_price": float(D(str(raw["markPrice"]))) if raw.get("markPrice") else None,
"unrealized_pnl": float(D(str(raw["unrealizedPnl"]))) if raw.get("unrealizedPnl") else None,
"leverage": float(D(str(raw["leverage"]))) if raw.get("leverage") else None,
"liquidation_price": float(D(str(raw["liquidationPrice"]))) if raw.get("liquidationPrice") else None,
})
return positions
@router.get("/{credential_id}/summary")
async def credential_summary(
credential_id: str,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> dict:
"""Get full account summary: balances, open orders, positions.
Uses CCXT client directly (bypasses rate limiter) to avoid blocking by
scheduler candle-fetch traffic that shares the GlobalRateLimiter singleton.
CCXT's built-in enableRateLimit provides per-client throttling.
"""
cred, adapter, exchange_name = await _get_adapter_for_credential(
credential_id, db, current_user
)
import asyncio as _asyncio
import ccxt as _ccxt
# Create a FRESH CCXT client — factory-cached adapters can have stale
# thread-local state that breaks when used from run_in_executor.
# Credential endpoints are called rarely; the ~8s init cost is acceptable.
secret = decrypt_api_key(cred.api_secret_enc, iv_hex=cred.api_secret_iv)
exchange_id = cred.exchange.name if cred.exchange else "binance"
_ccxt_class = getattr(_ccxt, exchange_id, None)
if _ccxt_class is None:
from app.core.exceptions import NotFoundException
raise NotFoundException(detail=f"Exchange {exchange_id} not supported by CCXT")
client = _ccxt_class({
"apiKey": cred.api_key,
"secret": secret,
"enableRateLimit": True,
"options": {"warnOnFetchOpenOrdersWithoutSymbol": False},
})
async def _fetch_and_parse(label: str, fetch_fn, parse_fn):
try:
loop = _asyncio.get_running_loop()
raw = await _asyncio.wait_for(
loop.run_in_executor(None, fetch_fn),
timeout=10.0,
)
result, extra = parse_fn(raw)
return label, "ok", result, extra, None
except _asyncio.TimeoutError:
logger.warning("Credential %s/%s: timed out after 10s", exchange_name, label)
return label, "timeout", [], Decimal("0"), f"{label} timed out after 10s"
except Exception as e:
err_msg = str(e)[:300]
logger.warning("Credential %s/%s failed: %s", exchange_name, label, err_msg)
return label, "error", [], Decimal("0"), err_msg
# Fetch sequentially (CCXT rate limiter is not thread-safe on the same client)
balance_label, balance_status, balances, usdt_balance, balance_err = await _fetch_and_parse("balance",
lambda: client.fetch_balance(),
lambda raw: _parse_balances_from_ccxt(raw))
orders_label, orders_status, open_orders, _, orders_err = await _fetch_and_parse("orders",
lambda: client.fetch_open_orders(symbol=None),
lambda raw: (_parse_open_orders_from_ccxt(raw), None))
positions_label, positions_status, positions, _, positions_err = await _fetch_and_parse("positions",
lambda: client.fetch_positions(symbols=None),
lambda raw: (_parse_positions_from_ccxt(raw), None))
# Collect results
open_orders: list = open_orders if orders_status == "ok" else []
positions: list = positions if positions_status == "ok" else []
errors: list = []
if balance_status != "ok":
errors.append({"endpoint": "balance", "status": balance_status, "detail": balance_err})
if orders_status != "ok":
errors.append({"endpoint": "orders", "status": orders_status, "detail": orders_err})
if positions_status != "ok":
errors.append({"endpoint": "positions", "status": positions_status, "detail": positions_err})
return {
"exchange_name": exchange_name.upper(),
"usdt_balance": float(usdt_balance),
"usdt_estimate": f"≈ ${float(usdt_balance):,.2f}",
"balance_count": len(balances),
"open_orders_count": len(open_orders),
"positions_count": len(positions),
"balances": balances,
"open_orders": open_orders,
"positions": positions,
"recent_trades": [], # requires symbol — not available in summary
"errors": errors if errors else None,
}
@router.get("/{credential_id}/balance")
async def credential_balance(
credential_id: str,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> dict:
"""Get token balances for this API key."""
_, adapter, exchange_name = await _get_adapter_for_credential(
credential_id, db, current_user
)
balance = await adapter.fetch_balance()
balances = [b.model_dump() for b in balance.balances]
usdt = Decimal("0")
for b in balance.balances:
if b.asset == "USDT":
usdt = b.total
return {
"exchange": exchange_name.upper(),
"usdt_balance": float(usdt),
"balances": balances,
"count": len(balances),
}
@router.get("/{credential_id}/orders")
async def credential_orders(
credential_id: str,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> dict:
"""Get open orders for this API key."""
_, adapter, exchange_name = await _get_adapter_for_credential(
credential_id, db, current_user
)
orders = await adapter.fetch_open_orders()
return {
"exchange": exchange_name.upper(),
"orders": [o.model_dump() for o in orders],
"count": len(orders),
}
@router.get("/{credential_id}/positions")
async def credential_positions(
credential_id: str,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> dict:
"""Get open positions for this API key."""
_, adapter, exchange_name = await _get_adapter_for_credential(
credential_id, db, current_user
)
positions = await adapter.fetch_positions()
return {
"exchange": exchange_name.upper(),
"positions": [p.model_dump() for p in positions],
"count": len(positions),
}
+114
View File
@@ -0,0 +1,114 @@
from __future__ import annotations
from fastapi import APIRouter, Body, Depends
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import joinedload
from app.core.deps import get_current_active_user, get_current_admin_user, get_db_session
from app.core.exceptions import NotFoundException
from app.models import Exchange, Symbol, User
from app.schemas import ExchangeResponse, SymbolResponse
from app.services import candle_service
from app.tasks.exchange_sync import sync_exchange_symbols
router = APIRouter(prefix="/exchanges", tags=["exchanges"])
@router.get("", response_model=list[ExchangeResponse])
async def list_exchanges(
active_only: bool = True,
db: AsyncSession = Depends(get_db_session),
) -> list[ExchangeResponse]:
"""Return all exchanges, optionally filtered to active only."""
query = select(Exchange)
if active_only:
query = query.where(Exchange.is_active == True) # noqa: E712
result = await db.execute(query)
exchanges = result.scalars().all()
return [
ExchangeResponse(
id=ex.id,
name=ex.name,
display_name=ex.display_name or ex.name,
is_active=ex.is_active,
)
for ex in exchanges
]
@router.get("/{exchange_id}/symbols", response_model=list[SymbolResponse])
async def list_exchange_symbols(
exchange_id: int,
db: AsyncSession = Depends(get_db_session),
) -> list[SymbolResponse]:
"""Return all symbols for a specific exchange."""
# Verify exchange exists
exch_result = await db.execute(select(Exchange).where(Exchange.id == exchange_id))
exchange = exch_result.scalar_one_or_none()
if exchange is None:
raise NotFoundException(detail=f"Exchange with id {exchange_id} not found")
result = await db.execute(
select(Symbol)
.options(joinedload(Symbol.exchange))
.where(Symbol.exchange_id == exchange_id)
)
symbols = result.unique().scalars().all()
return [
SymbolResponse(
id=s.id,
exchange_id=s.exchange_id,
symbol=s.symbol,
base=s.base,
quote=s.quote,
is_active=s.is_active,
)
for s in symbols
]
@router.post("/{exchange_id}/sync")
async def sync_exchange_symbols_endpoint(
exchange_id: int,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_admin_user),
) -> dict:
"""Sync symbols for an exchange (admin only)."""
# Look up exchange by id
result = await db.execute(select(Exchange).where(Exchange.id == exchange_id))
exchange = result.scalar_one_or_none()
if exchange is None:
raise NotFoundException(detail=f"Exchange with id {exchange_id} not found")
count = await sync_exchange_symbols(db=db, exchange_name=exchange.name)
return {"message": f"Synced {count} symbols"}
@router.post("/{exchange_id}/fetch_candles")
async def fetch_candles(
exchange_id: int,
symbol: str = Body(..., embed=True, description="Trading pair symbol, e.g. BTC/USDT"),
timeframe: str = Body("1h", embed=True, description="Candle timeframe"),
limit: int = Body(500, embed=True, description="Number of candles to fetch"),
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_active_user),
) -> dict:
"""Fetch and store candles for a symbol on an exchange."""
# Look up exchange by id
result = await db.execute(select(Exchange).where(Exchange.id == exchange_id))
exchange = result.scalar_one_or_none()
if exchange is None:
raise NotFoundException(detail=f"Exchange with id {exchange_id} not found")
candles = await candle_service.fetch_and_store_candles(
db=db,
exchange_name=exchange.name,
symbol=symbol,
timeframe=timeframe,
limit=limit,
)
return {"message": f"Fetched {len(candles)} candles"}
+155
View File
@@ -0,0 +1,155 @@
"""API routes for placing orders on connected exchanges."""
from __future__ import annotations
import logging
from decimal import Decimal
from uuid import UUID
from fastapi import APIRouter, Depends, status
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_current_user, get_db_session
from app.core.exceptions import NotFoundException, ValidationException
from app.core.security import decrypt_api_key
from app.exchange.factory import factory as exchange_factory
from app.exchange.types import BalanceResponse, OrderData, OrderRequest
from app.models import Exchange, ExchangeCredential, User
from app.models.real_trade import RealTrade
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/orders", tags=["orders"])
@router.post("/place", response_model=OrderData)
async def place_order(
req: OrderRequest,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> OrderData:
"""Place an order on a connected exchange.
Uses the user's saved API credentials for the exchange.
The exchange name is inferred from the symbol or passed explicitly.
"""
# Currently only market orders are supported via this endpoint
if req.order_type not in ("market", "limit"):
raise ValidationException(detail=f"Unsupported order type: {req.order_type}")
if req.order_type == "limit" and req.price is None:
raise ValidationException(detail="Price is required for limit orders")
if req.amount <= 0:
raise ValidationException(detail="Amount must be positive")
# Infer exchange name from symbol prefix heuristics, or default to mexc
# In a more advanced setup the user would specify exchange_id in the request
exchange_name = "mexc"
# Find the user's active credential for this exchange
result = await db.execute(
select(ExchangeCredential)
.join(Exchange, Exchange.id == ExchangeCredential.exchange_id)
.where(
ExchangeCredential.user_id == current_user.id,
Exchange.name == exchange_name,
ExchangeCredential.is_active == True,
)
)
cred = result.scalar_one_or_none()
if cred is None:
raise NotFoundException(
detail=f"No active API key found for {exchange_name}. "
f"Go to Profile → API Keys to add one."
)
# Decrypt the stored API key/secret
try:
# api_key is stored as plaintext (masked in responses), api_secret is encrypted
api_secret = decrypt_api_key(cred.api_secret_enc, iv_hex=cred.api_secret_iv)
except Exception as e:
logger.error(f"Failed to decrypt API key for user {current_user.id}: {e}")
raise ValidationException(detail="Failed to decrypt stored API key. Please re-add your credentials.")
# Create the exchange adapter with credentials
try:
adapter = exchange_factory.create(
exchange_name,
api_key=cred.api_key,
api_secret=api_secret,
testnet=cred.is_testnet,
)
except ValueError as e:
raise NotFoundException(detail=str(e))
# Place the order
try:
order = await adapter.create_order(req)
logger.info(
"Order placed: user=%s exchange=%s symbol=%s side=%s amount=%s",
current_user.username, exchange_name, req.symbol, req.side, req.amount,
)
# Persist to real_trades
real_trade = RealTrade(
user_id=current_user.id,
exchange=exchange_name,
symbol=req.symbol,
side=req.side,
order_type=req.order_type,
amount=req.amount,
price=req.price,
filled_amount=order.filled,
status=order.status,
order_id=order.order_id,
)
db.add(real_trade)
await db.flush()
logger.info("Real trade #%d persisted for order %s", real_trade.id, order.order_id)
return order
except Exception as e:
logger.error(f"Order failed for user {current_user.id}: {e}")
raise ValidationException(detail=f"Order failed: {str(e)[:300]}")
@router.get("/balance", response_model=BalanceResponse)
async def get_balance(
exchange_name: str = "mexc",
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_user),
) -> BalanceResponse:
"""Fetch account balance from a connected exchange."""
result = await db.execute(
select(ExchangeCredential)
.join(Exchange, Exchange.id == ExchangeCredential.exchange_id)
.where(
ExchangeCredential.user_id == current_user.id,
Exchange.name == exchange_name,
ExchangeCredential.is_active == True,
)
)
cred = result.scalar_one_or_none()
if cred is None:
raise NotFoundException(
detail=f"No active API key found for {exchange_name}. "
f"Go to Profile → API Keys to add one."
)
try:
api_secret = decrypt_api_key(cred.api_secret_enc, iv_hex=cred.api_secret_iv)
except Exception as e:
logger.error(f"Failed to decrypt API key: {e}")
raise ValidationException(detail="Failed to decrypt stored API key.")
adapter = exchange_factory.create(
exchange_name,
api_key=cred.api_key,
api_secret=api_secret,
testnet=cred.is_testnet,
)
try:
return await adapter.fetch_balance()
except Exception as e:
logger.error(f"Balance fetch failed for user {current_user.id}: {e}")
raise ValidationException(detail=f"Balance fetch failed: {str(e)[:300]}")
+134
View File
@@ -0,0 +1,134 @@
"""API routes for real trade tracking."""
from __future__ import annotations
import logging
from datetime import datetime, timedelta, timezone
from decimal import Decimal
from typing import Optional
from fastapi import APIRouter, Depends, Query
from sqlalchemy import and_, desc, func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.deps import get_current_trader_user, get_current_viewer_user, get_db_session
from app.models import User
from app.models.real_trade import RealTrade
from app.schemas.real_trade import (
RealTradeCreateRequest,
RealTradeListResponse,
RealTradeResponse,
WinRatePeriod,
WinRateResponse,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/real-trades", tags=["real-trades"])
@router.get("", response_model=RealTradeListResponse)
async def list_real_trades(
symbol: Optional[str] = Query(None, description="Filter by symbol (e.g. BTC/USDT)"),
status: Optional[str] = Query(None, description="open / filled / cancelled"),
limit: int = Query(100, ge=1, le=500),
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_viewer_user),
):
"""Get real trades for the current user."""
query = (
select(RealTrade)
.where(RealTrade.user_id == current_user.id)
.order_by(desc(RealTrade.created_at))
.limit(limit)
)
if symbol:
query = query.where(RealTrade.symbol == symbol)
if status:
query = query.where(RealTrade.status == status)
result = await db.execute(query)
trades = result.scalars().all()
trades_resp = [RealTradeResponse.model_validate(t) for t in trades]
# Aggregate stats for filled / closed trades with PnL
closed_with_pnl = [t for t in trades if t.status in ("filled", "cancelled") and t.pnl is not None]
total_pnl: Optional[float] = None
win_rate: Optional[float] = None
if closed_with_pnl:
total_pnl = float(sum(t.pnl for t in closed_with_pnl))
wins = sum(1 for t in closed_with_pnl if t.pnl > 0)
win_rate = round(wins / len(closed_with_pnl) * 100, 1) if closed_with_pnl else 0.0
return RealTradeListResponse(
trades=trades_resp,
total=len(trades_resp),
total_pnl=total_pnl,
win_rate=win_rate,
)
@router.post("", response_model=RealTradeResponse, status_code=201)
async def create_real_trade(
req: RealTradeCreateRequest,
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_trader_user),
):
"""Persist a real trade record (called after placing an order on an exchange)."""
trade = RealTrade(
user_id=current_user.id,
exchange=req.exchange,
symbol=req.symbol,
side=req.side,
order_type=req.order_type,
amount=Decimal(str(req.amount)),
price=Decimal(str(req.price)) if req.price is not None else None,
filled_amount=Decimal(str(req.filled_amount)),
status=req.status,
order_id=req.order_id,
)
db.add(trade)
await db.flush()
await db.commit()
await db.refresh(trade)
logger.info(
"Real trade #%d created: %s %s %s on %s",
trade.id, trade.side, trade.amount, trade.symbol, trade.exchange,
)
return RealTradeResponse.model_validate(trade)
@router.get("/win-rate", response_model=WinRateResponse)
async def get_win_rate(
db: AsyncSession = Depends(get_db_session),
current_user: User = Depends(get_current_viewer_user),
):
"""Win-rate breakdown for daily, weekly, and monthly periods."""
now = datetime.now(timezone.utc)
periods = {
"daily": now - timedelta(days=1),
"weekly": now - timedelta(days=7),
"monthly": now - timedelta(days=30),
}
result: dict[str, WinRatePeriod] = {}
for label, start in periods.items():
query = select(RealTrade).where(
and_(
RealTrade.user_id == current_user.id,
RealTrade.created_at >= start,
RealTrade.status.in_(["filled", "cancelled"]),
RealTrade.pnl.isnot(None),
)
)
rows = await db.execute(query)
trades = rows.scalars().all()
total = len(trades)
wins = sum(1 for t in trades if t.pnl > 0)
rate = round(wins / total * 100, 1) if total > 0 else 0.0
result[label] = WinRatePeriod(trades=total, wins=wins, win_rate=rate)
return WinRateResponse(**result)
+56
View File
@@ -0,0 +1,56 @@
from __future__ import annotations
from fastapi import APIRouter, Depends
from app.api.v1.auth import router as auth_router
from app.api.v1.admin import router as admin_router
from app.api.v1.credentials import router as credentials_router
from app.api.v1.exchanges import router as exchanges_router
from app.api.v1.symbols import router as symbols_router
from app.api.v1.signals import router as signals_router
from app.api.v1.backtest import router as backtest_router
from app.api.v1.backtest_history import router as backtest_history_router
from app.api.v1.watchlist import router as watchlist_router
from app.api.v1.orders import router as orders_router
from app.api.v1.real_trades import router as real_trades_router
from app.api.v1.strategies import router as strategies_router
from app.api.v1.analytics import router as analytics_router
from app.api.v1.audit import router as audit_router
from app.api.v1.alerts import router as alerts_router
api_router = APIRouter(prefix="/api/v1")
api_router.include_router(auth_router)
api_router.include_router(symbols_router)
api_router.include_router(exchanges_router)
api_router.include_router(admin_router)
api_router.include_router(credentials_router)
api_router.include_router(signals_router)
api_router.include_router(backtest_router)
api_router.include_router(backtest_history_router)
api_router.include_router(watchlist_router)
api_router.include_router(orders_router)
api_router.include_router(real_trades_router)
api_router.include_router(strategies_router)
api_router.include_router(analytics_router)
api_router.include_router(audit_router)
api_router.include_router(alerts_router)
# ── Users endpoint (no prefix, directly on api_router) ──
from app.core.deps import get_current_user as _get_current_user
from app.models.user import User as _User
@api_router.get("/users/me")
async def get_current_user_info(
current_user: _User = Depends(_get_current_user),
):
"""Return the currently authenticated user's profile."""
return {
"id": str(current_user.id),
"username": current_user.username,
"email": current_user.email,
"display_name": current_user.display_name,
"is_active": current_user.is_active,
"is_admin": current_user.is_admin,
}
+71
View File
@@ -0,0 +1,71 @@
"""API routes for trading signals, hypothetical trades, and reviews."""
from __future__ import annotations
import logging
from typing import Optional
from fastapi import APIRouter, Depends, Query
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.core.deps import get_current_user
from app.models.user import User
from app.services.signal_service import (
get_recent_signals,
get_review,
get_trade_history,
)
from app.schemas.signal import (
ReviewResponse,
SignalListResponse,
SignalResponse,
TradeListResponse,
TradeResponse,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/signals", tags=["signals"])
@router.get("", response_model=SignalListResponse)
async def list_signals(
symbol: Optional[str] = Query(None, description="Filter by symbol (e.g. BTC/USDT)"),
limit: int = Query(50, ge=1, le=200),
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""Get the most recent trading signals."""
signals = await get_recent_signals(db, symbol=symbol, limit=limit)
return SignalListResponse(signals=signals, total=len(signals))
@router.get("/trades", response_model=TradeListResponse)
async def list_trades(
symbol: Optional[str] = Query(None, description="Filter by symbol"),
status: Optional[str] = Query(None, description="OPEN or CLOSED"),
limit: int = Query(100, ge=1, le=500),
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""Get hypothetical trade history for the current user."""
trades, total_pnl, win_rate = await get_trade_history(
db, symbol=symbol, status=status, limit=limit, user_id=current_user.id
)
return TradeListResponse(
trades=trades,
total=len(trades),
total_pnl=total_pnl,
win_rate=win_rate,
)
@router.get("/review", response_model=ReviewResponse)
async def get_period_review(
period: str = Query("weekly", regex="^(weekly|monthly)$"),
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_user),
):
"""Get a weekly or monthly performance review for the current user."""
review = await get_review(db, period=period, user_id=current_user.id)
return ReviewResponse(**review)
+127
View File
@@ -0,0 +1,127 @@
"""API routes for per-user strategy configuration."""
from __future__ import annotations
import logging
from typing import Any
from fastapi import APIRouter, Depends
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import get_db
from app.core.deps import get_current_active_user
from app.models.user import User
from app.schemas.strategy import (
STRATEGY_DISPLAY,
STRATEGY_NAMES,
StrategyConfigRequest,
StrategyEntry,
StrategyListResponse,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/strategies", tags=["strategies"])
def _get_enabled_strategies(preferences: dict | None) -> list[str]:
"""Extract the enabled_strategies list from user preferences.
Returns ALL strategies by default when not configured.
"""
if not preferences:
return list(STRATEGY_NAMES)
enabled = preferences.get("enabled_strategies")
if enabled is None:
return list(STRATEGY_NAMES)
return enabled
def _get_thresholds(preferences: dict | None) -> dict[str, float]:
"""Extract the thresholds dict from user preferences."""
if not preferences:
return {}
return preferences.get("thresholds", {})
@router.get("", response_model=StrategyListResponse)
async def list_strategies(
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_active_user),
) -> StrategyListResponse:
"""Get the list of all available strategies with their enabled status
for the current user, based on user preferences."""
prefs: dict[str, Any] = current_user.preferences or {}
enabled_list = _get_enabled_strategies(prefs)
enabled_set = set(enabled_list)
entries = [
StrategyEntry(
name=name,
display_name=STRATEGY_DISPLAY.get(name, name),
enabled=name in enabled_set,
)
for name in STRATEGY_NAMES
]
thresholds = _get_thresholds(prefs)
return StrategyListResponse(strategies=entries, thresholds=thresholds)
@router.post("", response_model=StrategyListResponse)
async def update_strategies_post(
body: StrategyConfigRequest,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_active_user),
) -> StrategyListResponse:
"""Alias for PUT /strategies — supports POST method."""
return await update_strategies(body=body, db=db, current_user=current_user)
@router.put("", response_model=StrategyListResponse)
async def update_strategies(
body: StrategyConfigRequest,
db: AsyncSession = Depends(get_db),
current_user: User = Depends(get_current_active_user),
) -> StrategyListResponse:
"""Update the current user's enabled_strategies and/or thresholds
in their preferences JSON column."""
prefs: dict[str, Any] = current_user.preferences or {}
if body.enabled_strategies is not None:
# Validate that all provided strategy names are known
valid_names = set(STRATEGY_NAMES)
for name in body.enabled_strategies:
if name not in valid_names:
from fastapi import HTTPException
raise HTTPException(
status_code=422,
detail=f"Unknown strategy name: {name}. Valid: {STRATEGY_NAMES}",
)
prefs["enabled_strategies"] = body.enabled_strategies
if body.thresholds is not None:
prefs["thresholds"] = body.thresholds
current_user.preferences = prefs
await db.flush()
logger.info(
"Strategy prefs updated for user %s: enabled=%s, thresholds=%s",
current_user.username,
prefs.get("enabled_strategies"),
prefs.get("thresholds"),
)
# Build response
enabled_list = _get_enabled_strategies(prefs)
enabled_set = set(enabled_list)
entries = [
StrategyEntry(
name=name,
display_name=STRATEGY_DISPLAY.get(name, name),
enabled=name in enabled_set,
)
for name in STRATEGY_NAMES
]
thresholds = _get_thresholds(prefs)
return StrategyListResponse(strategies=entries, thresholds=thresholds)
+181
View File
@@ -0,0 +1,181 @@
from __future__ import annotations
from datetime import datetime
from typing import Optional
from fastapi import APIRouter, Depends, Query
from sqlalchemy import and_, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import joinedload
from app.core.deps import get_db_session
from app.core.exceptions import ValidationException
from app.models import Candle, Exchange, Symbol
from app.schemas import CandleListResponse, SymbolResponse, SymbolSearchResponse
from app.services import candle_service
router = APIRouter(prefix="/symbols", tags=["symbols"])
@router.get("", response_model=list[SymbolResponse])
async def list_symbols(
exchange: Optional[str] = Query(None, description="Filter by exchange name"),
active_only: bool = Query(True, description="Only return active symbols"),
limit: int = Query(50, ge=1, le=1000, description="Max symbols to return"),
offset: int = Query(0, ge=0, description="Offset for pagination"),
db: AsyncSession = Depends(get_db_session),
) -> list[SymbolResponse]:
"""Return symbols with pagination, optionally filtered by exchange and/or active status."""
query = select(Symbol).options(joinedload(Symbol.exchange))
if exchange:
query = query.join(Exchange, Exchange.id == Symbol.exchange_id).where(
Exchange.name == exchange
)
if active_only:
query = query.where(Symbol.is_active == True) # noqa: E712
query = query.limit(limit).offset(offset)
result = await db.execute(query)
symbols = result.unique().scalars().all()
return [
SymbolResponse(
id=s.id,
exchange_id=s.exchange_id,
symbol=s.symbol,
base=s.base,
quote=s.quote,
is_active=s.is_active,
)
for s in symbols
]
@router.get("/search", response_model=SymbolSearchResponse)
async def search_symbols(
q: str = Query(..., min_length=1, description="Search query"),
exchange: Optional[str] = Query(None, description="Filter by exchange name"),
limit: int = Query(50, ge=1, le=1000, description="Max results"),
db: AsyncSession = Depends(get_db_session),
) -> SymbolSearchResponse:
"""Search symbols by name (case-insensitive partial match)."""
query = select(Symbol).options(joinedload(Symbol.exchange))
query = query.where(Symbol.symbol.ilike(f"%{q}%"))
if exchange:
query = query.join(Exchange, Exchange.id == Symbol.exchange_id).where(
Exchange.name == exchange
)
query = query.limit(limit)
result = await db.execute(query)
symbols = result.unique().scalars().all()
return SymbolSearchResponse(
symbols=[
SymbolResponse(
id=s.id,
exchange_id=s.exchange_id,
symbol=s.symbol,
base=s.base,
quote=s.quote,
is_active=s.is_active,
)
for s in symbols
]
)
@router.get("/candles", response_model=CandleListResponse)
async def get_candles_endpoint(
symbol: str = Query(..., description="Trading pair symbol, e.g. BTC/USDT"),
exchange: str = Query("binance", description="Exchange name"),
timeframe: str = Query("1h", description="Candle timeframe"),
cursor: Optional[str] = Query(None, description="ISO datetime cursor for pagination"),
limit: int = Query(500, ge=1, le=1000, description="Max candles to return"),
db: AsyncSession = Depends(get_db_session),
) -> CandleListResponse:
"""Get candles for a symbol with cursor-based pagination."""
cursor_dt: Optional[datetime] = None
if cursor is not None:
try:
cursor_dt = datetime.fromisoformat(cursor)
except ValueError:
raise ValidationException(
detail=f"Invalid cursor format: {cursor!r}. Expected ISO datetime string."
)
return await candle_service.get_candles(
db=db,
symbol=symbol,
exchange_name=exchange,
timeframe=timeframe,
cursor=cursor_dt,
limit=limit,
)
@router.get("/{base}/{quote}/candles", response_model=CandleListResponse)
async def get_candles_by_path(
base: str,
quote: str,
exchange: str = Query("binance", description="Exchange name"),
timeframe: str = Query("1h", description="Candle timeframe"),
cursor: Optional[str] = Query(None, description="ISO datetime cursor for pagination"),
limit: int = Query(500, ge=1, le=1000, description="Max candles to return"),
db: AsyncSession = Depends(get_db_session),
) -> CandleListResponse:
"""Get candles — symbol via path (e.g. BTC/USDT/candles)."""
symbol = f"{base}/{quote}"
cursor_dt: Optional[datetime] = None
if cursor is not None:
try:
cursor_dt = datetime.fromisoformat(cursor)
except ValueError:
raise ValidationException(
detail=f"Invalid cursor format: {cursor!r}. Expected ISO datetime string."
)
return await candle_service.get_candles(
db=db,
symbol=symbol,
exchange_name=exchange,
timeframe=timeframe,
cursor=cursor_dt,
limit=limit,
)
@router.get("/{base}/{quote}/indicators")
async def get_indicators_by_path(
base: str,
quote: str,
exchange: str = Query("binance", description="Exchange name"),
timeframe: str = Query("1h", description="Candle timeframe"),
db: AsyncSession = Depends(get_db_session),
) -> dict:
"""Return computed technical indicators — symbol via path."""
symbol = f"{base}/{quote}"
return await candle_service.get_indicators(
db=db,
symbol=symbol,
exchange_name=exchange,
timeframe=timeframe,
)
@router.get("/indicators")
async def get_indicators_endpoint(
symbol: str = Query(..., description="Trading pair symbol, e.g. BTC/USDT"),
exchange: str = Query("binance", description="Exchange name"),
timeframe: str = Query("1h", description="Candle timeframe"),
db: AsyncSession = Depends(get_db_session),
) -> dict:
"""Return computed technical indicators for a symbol."""
return await candle_service.get_indicators(
db=db,
symbol=symbol,
exchange_name=exchange,
timeframe=timeframe,
)
+196
View File
@@ -0,0 +1,196 @@
"""API routes for user watchlist management."""
from __future__ import annotations
import logging
from uuid import UUID
from fastapi import APIRouter, Depends, Query
from sqlalchemy import and_, case, delete, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import joinedload
from app.core.deps import get_db_session, get_current_active_user
from app.core.exceptions import AppException, NotFoundException
from app.models import Exchange, Symbol, Watchlist
from app.models.user import User as UserModel
from app.schemas import (
SymbolResponse,
WatchlistCreateRequest,
WatchlistResponse,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/watchlist", tags=["watchlist"])
@router.get("", response_model=list[WatchlistResponse])
async def list_watchlist(
db: AsyncSession = Depends(get_db_session),
current_user: UserModel = Depends(get_current_active_user),
):
"""List all symbols in the user's watchlist."""
result = await db.execute(
select(Watchlist)
.options(joinedload(Watchlist.symbol).joinedload(Symbol.exchange))
.where(Watchlist.user_id == current_user.id)
.order_by(Watchlist.sort_order, Watchlist.created_at)
)
items = result.unique().scalars().all()
return [
WatchlistResponse(
id=item.id,
symbol_id=item.symbol_id,
symbol=item.symbol.symbol,
exchange=item.symbol.exchange.name if item.symbol.exchange else "",
label=item.label,
sort_order=item.sort_order,
)
for item in items
]
@router.post("", response_model=WatchlistResponse, status_code=201)
async def add_to_watchlist(
body: WatchlistCreateRequest,
db: AsyncSession = Depends(get_db_session),
current_user: UserModel = Depends(get_current_active_user),
):
"""Add a symbol to the user's watchlist."""
# Check symbol exists
result = await db.execute(
select(Symbol)
.options(joinedload(Symbol.exchange))
.where(Symbol.id == body.symbol_id)
)
symbol = result.unique().scalar_one_or_none()
if symbol is None:
raise NotFoundException(detail=f"Symbol id={body.symbol_id} not found")
# Check not already in watchlist
existing = await db.execute(
select(Watchlist).where(
and_(
Watchlist.user_id == current_user.id,
Watchlist.symbol_id == body.symbol_id,
)
)
)
if existing.scalar_one_or_none():
raise AppException(
status_code=409,
detail="Symbol already in watchlist",
code="duplicate_entry",
)
# Get next sort_order
max_order = await db.execute(
select(Watchlist.sort_order)
.where(Watchlist.user_id == current_user.id)
.order_by(Watchlist.sort_order.desc())
.limit(1)
)
next_order = (max_order.scalar_one_or_none() or 0) + 1
item = Watchlist(
user_id=current_user.id,
symbol_id=body.symbol_id,
label=body.label,
sort_order=body.sort_order if body.sort_order is not None else next_order,
)
db.add(item)
await db.flush()
await db.refresh(item, ["symbol"])
return WatchlistResponse(
id=item.id,
symbol_id=item.symbol_id,
symbol=symbol.symbol,
exchange=symbol.exchange.name if symbol.exchange else "",
label=item.label,
sort_order=item.sort_order,
)
@router.delete("/{watchlist_id}", status_code=204)
async def remove_from_watchlist(
watchlist_id: UUID,
db: AsyncSession = Depends(get_db_session),
current_user: UserModel = Depends(get_current_active_user),
):
"""Remove a symbol from the user's watchlist."""
result = await db.execute(
select(Watchlist).where(
and_(
Watchlist.id == watchlist_id,
Watchlist.user_id == current_user.id,
)
)
)
item = result.scalar_one_or_none()
if item is None:
raise NotFoundException(detail="Watchlist entry not found")
await db.delete(item)
await db.flush()
@router.get("/all-symbols", response_model=list[SymbolResponse])
async def list_all_symbols(
exchange: str = Query("mexc", description="Exchange name"),
q: str = Query("", description="Search filter"),
limit: int = Query(100, ge=1, le=1000, description="Max symbols"),
db: AsyncSession = Depends(get_db_session),
current_user: UserModel = Depends(get_current_active_user),
):
"""List all available symbols for adding to watchlist, with optional search."""
query = (
select(Symbol)
.options(joinedload(Symbol.exchange))
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(
and_(
Exchange.name == exchange,
Symbol.is_active == True, # noqa: E712
)
)
.limit(limit)
)
if q:
like_q = f"%{q}%"
exact_q = q.upper()
# Search by base AND symbol-start to avoid cross-pair matches (AR/ETH)
query = query.where(
or_(
Symbol.base.ilike(like_q),
Symbol.symbol.ilike(f"{exact_q}/%"),
)
)
# Order: exact base match first, prefix, then rest
exact_expr = case(
(Symbol.base == exact_q, 0),
(Symbol.base.ilike(f"{exact_q}%"), 1),
else_=2,
)
query = query.order_by(exact_expr, Symbol.symbol)
else:
# Default: show only USDT/USDC pairs for cleaner list
query = query.where(Symbol.quote.in_(["USDT", "USDC"]))
query = query.order_by(Symbol.symbol)
query = query.limit(500 if q else 200)
result = await db.execute(query)
symbols = result.unique().scalars().all()
return [
SymbolResponse(
id=s.id,
exchange_id=s.exchange_id,
symbol=s.symbol,
base=s.base,
quote=s.quote,
is_active=s.is_active,
)
for s in symbols
]
View File
+188
View File
@@ -0,0 +1,188 @@
"""
WebSocket endpoint for real-time candle and ticker updates.
Authentication via JWT token passed as a query parameter (``?token=...``).
Once authenticated, clients can subscribe/unsubscribe to symbol+timeframe+exchange
channels and receive live candle updates pushed by the background scheduler.
"""
from __future__ import annotations
import json
import logging
from fastapi import APIRouter, WebSocket, WebSocketDisconnect
from app.core.security import decode_token
from app.ws_manager import manager
logger = logging.getLogger(__name__)
router = APIRouter()
@router.websocket("/ws/v1/candles")
async def candle_websocket(websocket: WebSocket) -> None:
"""
WebSocket endpoint for real-time candle updates.
**Authentication:** via query parameter ``?token=JWT``
**Client → Server messages:**
.. code-block:: json
{"action": "subscribe", "symbol": "BTC/USDT", "timeframe": "1h", "exchange": "mexc"}
{"action": "unsubscribe", "symbol": "BTC/USDT", "timeframe": "1h", "exchange": "mexc"}
**Server → Client messages:**
.. code-block:: json
{"type": "connection", "status": "connected"}
{"type": "connection", "status": "authenticated"}
{"type": "candle", "data": {"symbol": "...", "timeframe": "...", ...}}
{"type": "ticker", "data": {"symbol": "...", "price": 123.45, ...}}
{"type": "error", "message": "..."}
"""
# ------------------------------------------------------------------
# 1. Accept the connection
# ------------------------------------------------------------------
await websocket.accept()
logger.info("WebSocket connection accepted: %s", id(websocket))
# ------------------------------------------------------------------
# 2. Send connection status
# ------------------------------------------------------------------
await _send_json(websocket, {"type": "connection", "status": "connected"})
# ------------------------------------------------------------------
# 3. Extract JWT token: try subprotocol header first (secure),
# then fall back to query parameter (legacy compat)
# ------------------------------------------------------------------
token: str | None = None
# Try Sec-WebSocket-Protocol header (subprotocol-based auth)
subprotocols = websocket.headers.get("sec-websocket-protocol", "")
for sp in subprotocols.split(","):
sp = sp.strip()
if sp.startswith("token,"):
token = sp.split(",", 1)[1].strip()
break
elif sp.startswith("token-"):
token = sp[6:].strip()
break
# Fallback: query parameter (backward compat, less secure)
if not token:
token = websocket.query_params.get("token")
if not token:
await _send_json(websocket, {"type": "error", "message": "Missing token query parameter"})
await websocket.close(code=4001)
return
# ------------------------------------------------------------------
# 4. Verify JWT token
# ------------------------------------------------------------------
try:
decode_token(token)
except Exception as exc:
error_msg = str(exc) if str(exc) else "Invalid or expired token"
await _send_json(websocket, {"type": "error", "message": error_msg})
await websocket.close(code=4001)
return
# ------------------------------------------------------------------
# 5. Send authenticated status
# ------------------------------------------------------------------
await _send_json(websocket, {"type": "connection", "status": "authenticated"})
logger.info("WebSocket %s authenticated", id(websocket))
# ------------------------------------------------------------------
# 6. Register this websocket with the connection manager
# (it starts with no subscriptions — the client will subscribe below)
# ------------------------------------------------------------------
# manager.subscribe(...) is called per-subscription in the message loop
# ------------------------------------------------------------------
# 7. Message loop — handle subscribe / unsubscribe
# ------------------------------------------------------------------
try:
while True:
raw = await websocket.receive_text()
try:
data = json.loads(raw)
except json.JSONDecodeError:
await _send_json(websocket, {"type": "error", "message": "Invalid JSON"})
continue
action: str | None = data.get("action")
symbol: str | None = data.get("symbol")
timeframe: str | None = data.get("timeframe")
exchange: str | None = data.get("exchange")
if not action:
await _send_json(websocket, {"type": "error", "message": "Missing 'action' field"})
continue
if action not in ("subscribe", "unsubscribe"):
await _send_json(
websocket,
{"type": "error", "message": f"Unknown action: {action}"},
)
continue
if not all([symbol, timeframe, exchange]):
await _send_json(
websocket,
{
"type": "error",
"message": "Missing one or more required fields: symbol, timeframe, exchange",
},
)
continue
if action == "subscribe":
await manager.subscribe(websocket, symbol, timeframe, exchange)
logger.info(
"WebSocket %s subscribed to %s/%s/%s",
id(websocket),
exchange,
symbol,
timeframe,
)
elif action == "unsubscribe":
await manager.unsubscribe(websocket, symbol, timeframe, exchange)
logger.info(
"WebSocket %s unsubscribed from %s/%s/%s",
id(websocket),
exchange,
symbol,
timeframe,
)
except WebSocketDisconnect:
logger.info("WebSocket %s disconnected", id(websocket))
except Exception:
logger.exception("Unexpected error in WebSocket handler %s", id(websocket))
finally:
# ------------------------------------------------------------------
# 8. Clean up — remove from all subscriptions
# ------------------------------------------------------------------
await manager.unsubscribe_all(websocket)
logger.info("WebSocket %s cleaned up (unsubscribed from all channels)", id(websocket))
# ====================================================================
# Internal helpers
# ====================================================================
async def _send_json(websocket: WebSocket, data: dict) -> None:
"""Send a JSON-serialisable dict to the websocket, ignoring errors."""
try:
await websocket.send_json(data)
except Exception:
logger.debug("Failed to send JSON to WebSocket %s (may be disconnected)", id(websocket))
+43
View File
@@ -0,0 +1,43 @@
from __future__ import annotations
from pydantic_settings import BaseSettings, SettingsConfigDict
class Settings(BaseSettings):
"""Application settings loaded from environment variables / .env file."""
# Database
DATABASE_URL: str = "postgresql+asyncpg://trading:trading_secret@db:5432/trading_portal"
# JWT
JWT_PRIVATE_KEY_PATH: str = "/run/secrets/jwt_private.pem"
JWT_PUBLIC_KEY_PATH: str = "/run/secrets/jwt_public_key.pem"
JWT_PUBLIC_KEYS_DIR: str = "/run/secrets/jwt_public_keys" # directory of valid public keys
JWT_ACCESS_TOKEN_EXPIRE_MINUTES: int = 15
JWT_REFRESH_TOKEN_EXPIRE_DAYS: int = 7
# Encryption
ENCRYPTION_KEY: str = ""
# Server
HOST: str = "0.0.0.0"
PORT: int = 8000
# Logging
LOG_LEVEL: str = "INFO"
# CORS
CORS_ORIGINS: str = ""
# Demo user (from .env)
demo_user: str = "demo"
demo_pass: str = "demo1234"
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
extra="allow",
)
settings = Settings()
View File
+104
View File
@@ -0,0 +1,104 @@
from __future__ import annotations
from collections.abc import AsyncGenerator
from uuid import UUID
from fastapi import Depends, Header, Request
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.core.exceptions import InvalidCredentialsException, InvalidTokenException
from app.core.security import decode_token
from app.database import get_db
from app.models import User
async def get_db_session() -> AsyncGenerator[AsyncSession, None]:
"""Provide an async SQLAlchemy database session via FastAPI dependency."""
async for session in get_db():
yield session
async def get_current_user(
request: Request,
db: AsyncSession = Depends(get_db_session),
) -> User:
"""
Extract the Bearer token from the Authorization header, decode it,
and fetch the corresponding user from the database.
Raises ``InvalidCredentialsException`` if the token is missing,
invalid, or the user is not found.
"""
authorization: str | None = request.headers.get("Authorization")
if not authorization or not authorization.startswith("Bearer "):
raise InvalidCredentialsException(
detail="Missing or malformed Authorization header"
)
token = authorization.removeprefix("Bearer ").strip()
try:
payload = decode_token(token)
except Exception:
raise InvalidTokenException(detail="Invalid or expired token")
sub: str | None = payload.get("sub")
if sub is None:
raise InvalidTokenException(detail="Token payload missing subject")
user_id: UUID
try:
user_id = UUID(sub)
except ValueError:
raise InvalidTokenException(detail="Invalid token subject format")
result = await db.execute(select(User).where(User.id == user_id))
user = result.scalar_one_or_none()
if user is None:
raise InvalidCredentialsException(detail="User not found")
return user
async def get_current_active_user(
current_user: User = Depends(get_current_user),
) -> User:
"""Return the current user if active, otherwise raise an error."""
if not current_user.is_active:
raise InvalidCredentialsException(detail="Inactive user account")
return current_user
async def get_current_trader_user(
current_user: User = Depends(get_current_active_user),
) -> User:
"""Return the current user if they are a trader or admin, otherwise raise an error."""
if current_user.role not in ("admin", "trader"):
from app.core.exceptions import ForbiddenException
raise ForbiddenException(detail="Trader or admin privileges required")
return current_user
async def get_current_viewer_user(
current_user: User = Depends(get_current_active_user),
) -> User:
"""Return the current user if active (any role can view)."""
return current_user
async def get_current_admin_user(
current_user: User = Depends(get_current_active_user),
) -> User:
"""Return the current user if they are an admin, otherwise raise an error.
P1-22: Checks both is_admin AND role for frontend/backend consistency.
"""
if not current_user.is_admin and current_user.role != "admin":
from app.core.exceptions import ForbiddenException
raise ForbiddenException(detail="Admin privileges required")
return current_user
+74
View File
@@ -0,0 +1,74 @@
from __future__ import annotations
class AppException(Exception):
"""Base exception for all application-level errors."""
status_code: int = 500
detail: str = "Internal server error"
code: str = "internal_error"
def __init__(
self,
status_code: int | None = None,
detail: str | None = None,
code: str | None = None,
) -> None:
if status_code is not None:
self.status_code = status_code
if detail is not None:
self.detail = detail
if code is not None:
self.code = code
super().__init__(self.detail)
class NotFoundException(AppException):
status_code: int = 404
code: str = "not_found"
detail: str = "Resource not found"
class AuthException(AppException):
status_code: int = 401
code: str = "auth_error"
detail: str = "Authentication error"
class InvalidCredentialsException(AuthException):
code: str = "invalid_credentials"
detail: str = "Invalid username or password"
class TokenExpiredException(AuthException):
code: str = "token_expired"
detail: str = "Token has expired"
class InvalidTokenException(AuthException):
code: str = "invalid_token"
detail: str = "Invalid token"
class ForbiddenException(AppException):
status_code: int = 403
code: str = "forbidden"
detail: str = "Forbidden"
class ValidationException(AppException):
status_code: int = 422
code: str = "validation_error"
detail: str = "Validation error"
class RateLimitException(AppException):
status_code: int = 429
code: str = "rate_limit"
detail: str = "Rate limit exceeded"
class ConflictException(AppException):
status_code: int = 409
code: str = "conflict"
detail: str = "Resource already exists"
+73
View File
@@ -0,0 +1,73 @@
from __future__ import annotations
import time
import structlog
from starlette.requests import Request
from starlette.responses import Response
from starlette.types import ASGIApp, Receive, Scope, Send
logger = structlog.get_logger(__name__)
class RequestLoggingMiddleware:
"""ASGI middleware that logs every request with method, path, status code,
and duration using structlog."""
def __init__(self, app: ASGIApp) -> None:
self.app = app
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
if scope["type"] != "http":
await self.app(scope, receive, send)
return
start = time.perf_counter()
request = Request(scope)
# Wrap send to capture the response status code
status_code: int | None = None
async def send_wrapper(message: dict) -> None:
nonlocal status_code
if message["type"] == "http.response.start":
status_code = message["status"]
await send(message)
try:
await self.app(scope, receive, send_wrapper)
except Exception:
duration = time.perf_counter() - start
logger.error(
"request_error",
method=request.method,
path=request.url.path,
status_code=500,
duration_ms=round(duration * 1000, 2),
)
raise
else:
duration = time.perf_counter() - start
if (status_code or 0) >= 400:
logger.warning(
"request_complete",
method=request.method,
path=request.url.path,
status_code=status_code,
duration_ms=round(duration * 1000, 2),
)
else:
logger.info(
"request_complete",
method=request.method,
path=request.url.path,
status_code=status_code,
duration_ms=round(duration * 1000, 2),
)
def register_middleware(app: ASGIApp) -> None:
"""Convenience helper — add the middleware to a FastAPI app."""
from app.core.middleware import RequestLoggingMiddleware # noqa: F811
app.add_middleware(RequestLoggingMiddleware) # type: ignore[arg-type]
+351
View File
@@ -0,0 +1,351 @@
"""
Security module for the trading portal backend.
Provides password hashing, JWT token management (RS256),
AES-256-CBC encryption for API key storage, and token utilities.
"""
from __future__ import annotations
import hashlib
import uuid
from datetime import datetime, timedelta, timezone
from typing import Optional, Tuple
from fastapi import HTTPException
from jose import JWTError, jwt
from jose.exceptions import ExpiredSignatureError
from passlib.context import CryptContext
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
from cryptography.hazmat.backends import default_backend
from app.config import settings
# ---------------------------------------------------------------------------
# Password hashing
# ---------------------------------------------------------------------------
_pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
def hash_password(password: str) -> str:
"""Hash a plaintext password using bcrypt (synchronous — use in thread pool)."""
return _pwd_context.hash(password)
def verify_password(plain: str, hashed: str) -> bool:
"""Verify a plaintext password against a bcrypt hash (synchronous — use in thread pool)."""
return _pwd_context.verify(plain, hashed)
async def hash_password_async(password: str) -> str:
"""Hash a plaintext password using bcrypt (async, runs in thread pool)."""
import asyncio
loop = asyncio.get_running_loop()
return await loop.run_in_executor(None, _pwd_context.hash, password)
async def verify_password_async(plain: str, hashed: str) -> bool:
"""Verify a plaintext password against a bcrypt hash (async, runs in thread pool)."""
import asyncio
loop = asyncio.get_running_loop()
return await loop.run_in_executor(None, _pwd_context.verify, plain, hashed)
# ---------------------------------------------------------------------------
# JWT RS256 token management — with multi-key rotation support
# ---------------------------------------------------------------------------
#
# PRINCIPLE:
# - Tokens are SIGNED with the current private key and carry a "kid"
# (Key ID = SHA-256 fingerprint of the public key) in the JWT header.
# - Tokens are VERIFIED against ALL public keys in the keys directory.
# ANY valid key can decode the token → old tokens survive key rotation.
# - When rotating: add new key pair, old tokens remain valid until they
# expire naturally. Remove old public keys only after all their tokens
# have expired.
#
# Directory layout:
# /run/secrets/jwt_private.pem ← CURRENT signing key
# /run/secrets/jwt_public_keys/*.pem ← ALL valid public keys
import hashlib
import os
from functools import lru_cache
ALGORITHM = "RS256"
def _load_private_key() -> str:
"""Read the current RSA private key PEM for signing new tokens."""
try:
with open(settings.JWT_PRIVATE_KEY_PATH, "r") as f:
return f.read()
except FileNotFoundError:
raise HTTPException(
status_code=500,
detail=f"JWT private key not found at {settings.JWT_PRIVATE_KEY_PATH}",
)
except OSError as exc:
raise HTTPException(
status_code=500,
detail=f"Failed to read JWT private key: {exc}",
)
def _compute_kid(public_key_pem: str) -> str:
"""Return a Key ID (SHA-256 fingerprint) for a public key PEM."""
return hashlib.sha256(public_key_pem.strip().encode()).hexdigest()[:16]
@lru_cache(maxsize=1)
def _cached_public_key() -> str:
"""Return the CURRENT public key (used for kid computation). Cached."""
keys = _load_all_public_keys()
if not keys:
raise HTTPException(status_code=500, detail="No valid JWT public keys found")
return keys[-1][1] # newest key
def _load_all_public_keys() -> list[tuple[str, str]]:
"""
Load ALL valid public keys from the keys directory.
Returns:
List of (kid, pem_content) tuples, sorted by filename for determinism.
The store allows multiple keys to coexist — old keys still validate
tokens signed before the last rotation.
"""
keys_dir = settings.JWT_PUBLIC_KEYS_DIR
keys: list[tuple[str, str]] = []
try:
for filename in sorted(os.listdir(keys_dir)):
if not filename.endswith(".pem"):
continue
filepath = os.path.join(keys_dir, filename)
try:
with open(filepath, "r") as f:
pem = f.read().strip()
if pem:
kid = _compute_kid(pem)
keys.append((kid, pem))
except OSError:
continue # skip unreadable files
except FileNotFoundError:
pass # no directory yet — handled gracefully
except NotADirectoryError:
raise HTTPException(
status_code=500,
detail=f"JWT_PUBLIC_KEYS_DIR ({keys_dir}) is not a directory",
)
if not keys:
raise HTTPException(
status_code=500,
detail=f"No valid JWT public keys found in {keys_dir}",
)
return keys
def create_access_token(
data: dict,
expires_delta: Optional[timedelta] = None,
) -> str:
"""Create a short-lived JWT access token (RS256).
Includes ``kid`` header so the verifier knows which key to try first.
"""
to_encode = data.copy()
now = datetime.now(timezone.utc)
if expires_delta is not None:
expire = now + expires_delta
else:
expire = now + timedelta(minutes=settings.JWT_ACCESS_TOKEN_EXPIRE_MINUTES)
to_encode.update({"iat": now, "exp": expire, "sub": str(data["sub"])})
private_key = _load_private_key()
public_key = _cached_public_key()
kid = _compute_kid(public_key)
headers = {"kid": kid}
return jwt.encode(to_encode, private_key, algorithm=ALGORITHM, headers=headers)
def create_refresh_token(data: dict) -> str:
"""Create a long-lived JWT refresh token (RS256).
Includes ``kid``, ``type: refresh``, and a unique ``jti``.
"""
to_encode = data.copy()
now = datetime.now(timezone.utc)
expire = now + timedelta(days=settings.JWT_REFRESH_TOKEN_EXPIRE_DAYS)
to_encode.update(
{
"iat": now,
"exp": expire,
"sub": str(data["sub"]),
"type": "refresh",
"jti": generate_jti(),
}
)
private_key = _load_private_key()
public_key = _cached_public_key()
kid = _compute_kid(public_key)
headers = {"kid": kid}
return jwt.encode(to_encode, private_key, algorithm=ALGORITHM, headers=headers)
def decode_token(token: str) -> dict:
"""Decode and verify a JWT token using ALL known public keys.
Tries every valid public key in the directory. If ANY key verifies the
token, it is valid — this is how key rotation works without breaking
existing sessions.
Optimisation: the token's ``kid`` header is used to try the matching key
first before falling back to a full linear scan.
"""
# 1. Extract kid from token header (without verifying signature yet)
try:
unverified_header = jwt.get_unverified_header(token)
token_kid = unverified_header.get("kid")
except JWTError:
token_kid = None
# 2. Load all valid public keys
all_keys = _load_all_public_keys() # [(kid, pem), ...]
# 3. If we have a kid, try the matching key first
if token_kid:
for kid, pem in all_keys:
if kid == token_kid:
try:
return jwt.decode(token, pem, algorithms=[ALGORITHM])
except ExpiredSignatureError:
raise HTTPException(
status_code=401,
detail="Token has expired",
headers={"WWW-Authenticate": "Bearer"},
)
except JWTError:
pass # key mismatch — fall through to full scan
# 4. Fallback: try all keys (handles tokens without kid, or kid mismatch)
for kid, pem in all_keys:
try:
return jwt.decode(token, pem, algorithms=[ALGORITHM])
except ExpiredSignatureError:
raise HTTPException(
status_code=401,
detail="Token has expired",
headers={"WWW-Authenticate": "Bearer"},
)
except JWTError:
continue # try next key
# 5. No key worked
raise HTTPException(
status_code=401,
detail="Invalid or expired token",
headers={"WWW-Authenticate": "Bearer"},
)
# ---------------------------------------------------------------------------
# AES-256-CBC encryption (for API key storage)
# ---------------------------------------------------------------------------
_BACKEND = default_backend()
def generate_encryption_key() -> str:
"""Generate a random 32-byte (256-bit) hex-encoded encryption key.
Print the key to stdout so it can be copied into the ``.env`` file.
"""
key = uuid.uuid4().hex + uuid.uuid4().hex # 64 hex chars = 32 bytes
print(f"Encryption key (save in .env as ENCRYPTION_KEY={key}): {key}")
return key
def _resolve_key(key_hex: Optional[str] = None) -> bytes:
"""Return the AES key as bytes, falling back to settings."""
raw = key_hex if key_hex is not None else settings.ENCRYPTION_KEY
if not raw:
raise HTTPException(
status_code=500,
detail="Encryption key not configured. Set ENCRYPTION_KEY in .env",
)
return bytes.fromhex(raw)
def encrypt_api_key(
api_key: str,
key_hex: Optional[str] = None,
) -> Tuple[str, str]:
"""Encrypt an API key with AES-256-CBC.
Returns ``(ciphertext_hex, iv_hex)``.
"""
key = _resolve_key(key_hex)
iv = uuid.uuid4().bytes # 16 random bytes
cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=_BACKEND)
encryptor = cipher.encryptor()
# Pad plaintext to AES block size (16 bytes) using PKCS7
plaintext_bytes = api_key.encode("utf-8")
pad_len = 16 - (len(plaintext_bytes) % 16)
padded = plaintext_bytes + bytes([pad_len] * pad_len)
ciphertext = encryptor.update(padded) + encryptor.finalize()
return ciphertext.hex(), iv.hex()
def decrypt_api_key(
ciphertext_hex: str,
iv_hex: str,
key_hex: Optional[str] = None,
) -> str:
"""Decrypt an AES-256-CBC encrypted API key.
Returns the original plaintext string.
"""
key = _resolve_key(key_hex)
ciphertext = bytes.fromhex(ciphertext_hex)
iv = bytes.fromhex(iv_hex)
cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=_BACKEND)
decryptor = cipher.decryptor()
padded = decryptor.update(ciphertext) + decryptor.finalize()
# Remove PKCS7 padding
pad_len = padded[-1]
if pad_len < 1 or pad_len > 16:
raise HTTPException(
status_code=500,
detail="Decryption failed: invalid padding",
)
plaintext_bytes = padded[:-pad_len]
return plaintext_bytes.decode("utf-8")
# ---------------------------------------------------------------------------
# Token utilities
# ---------------------------------------------------------------------------
def generate_token_hash(token: str) -> str:
"""Return the SHA-256 hex digest of a token string."""
return hashlib.sha256(token.encode("utf-8")).hexdigest()
def generate_jti() -> str:
"""Return a UUID4 hex string for use as a JWT token ID."""
return uuid.uuid4().hex
+64
View File
@@ -0,0 +1,64 @@
from __future__ import annotations
import logging
from collections.abc import AsyncGenerator
from sqlalchemy.exc import InterfaceError
from sqlalchemy.ext.asyncio import (
AsyncSession,
async_sessionmaker,
create_async_engine,
)
from sqlalchemy.orm import DeclarativeBase
from app.config import settings
logger = logging.getLogger(__name__)
engine = create_async_engine(
settings.DATABASE_URL,
echo=(settings.LOG_LEVEL == "DEBUG"),
pool_pre_ping=True,
pool_size=60,
max_overflow=20,
pool_timeout=5,
pool_recycle=600,
connect_args={
"server_settings": {
"idle_in_transaction_session_timeout": "60000",
"statement_timeout": "10000",
}
},
)
async_session_factory = async_sessionmaker(
engine,
class_=AsyncSession,
expire_on_commit=False,
)
class Base(DeclarativeBase):
"""Declarative base for all ORM models."""
async def get_db() -> AsyncGenerator[AsyncSession, None]:
"""Provide async DB session. Handles closed-connection errors gracefully
to prevent exception storms that spike CPU to 500%+."""
async with async_session_factory() as session:
try:
yield session
await session.commit()
except InterfaceError:
logger.warning("Session commit skipped: connection already closed")
except Exception:
try:
await session.rollback()
except InterfaceError:
logger.warning("Session rollback skipped: connection already closed")
raise
finally:
try:
await session.close()
except InterfaceError:
pass # already closed
View File
+344
View File
@@ -0,0 +1,344 @@
from __future__ import annotations
import logging
from abc import ABC, abstractmethod
from typing import Any, Optional
import ccxt
from app.exchange.rate_limiter import GlobalRateLimiter, RateLimiter
from app.exchange.types import (
BalanceData,
BalanceResponse,
CandleData,
CandleValidationError,
OpenOrderData,
OrderData,
OrderRequest,
PositionData,
SymbolInfo,
TickerData,
TradeData,
)
logger = logging.getLogger(__name__)
class AbstractExchange(ABC):
"""Abstract base class for exchange adapters."""
def __init__(self, api_key: str = "", api_secret: str = "", testnet: bool = False) -> None:
self.api_key = api_key
self.api_secret = api_secret
self.testnet = testnet
self._client: Optional[ccxt.Exchange] = None
self._rate_limiter: Optional[RateLimiter] = None
def _init_ccxt(self) -> ccxt.Exchange:
"""Initialize and return a CCXT exchange client.
Subclasses must override this to configure the specific exchange.
"""
raise NotImplementedError("Subclasses must implement _init_ccxt")
@property
def client(self) -> ccxt.Exchange:
if self._client is None:
self._client = self._init_ccxt()
return self._client
@property
def rate_limiter(self) -> RateLimiter:
if self._rate_limiter is None:
self._rate_limiter = GlobalRateLimiter(self.get_name())
return self._rate_limiter
@abstractmethod
def get_name(self) -> str:
...
@abstractmethod
def get_base_url(self) -> str:
"""Return the REST API base URL."""
...
@abstractmethod
def get_ws_url(self) -> str:
"""Return the WebSocket URL."""
...
async def fetch_ohlcv(self, symbol: str, timeframe: str = "1h", limit: int = 500) -> list[CandleData]:
"""Fetch OHLCV candle data from the exchange."""
await self.rate_limiter.acquire()
raw = await self._async_fetch_ohlcv(symbol, timeframe, limit)
candles: list[CandleData] = []
for entry in raw:
ts, o, h, l, c, v = entry
open_dec = Decimal(str(o))
high_dec = Decimal(str(h))
low_dec = Decimal(str(l))
close_dec = Decimal(str(c))
volume_dec = Decimal(str(v))
if not (low_dec <= open_dec <= high_dec and low_dec <= close_dec <= high_dec):
continue
candles.append(
CandleData(
symbol=symbol,
exchange=self.get_name(),
timeframe=timeframe,
timestamp=datetime.fromtimestamp(ts / 1000, tz=timezone.utc),
open=open_dec,
high=high_dec,
low=low_dec,
close=close_dec,
volume=volume_dec,
)
)
return candles
async def _async_fetch_ohlcv(self, symbol: str, timeframe: str, limit: int, since: int | None = None) -> list[list[Any]]:
"""Run the synchronous CCXT fetch_ohlcv in a thread pool."""
import asyncio
loop = asyncio.get_running_loop()
kwargs = dict(symbol=symbol, timeframe=timeframe, limit=limit)
if since is not None:
kwargs["since"] = since
return await loop.run_in_executor(
None,
lambda: self.client.fetch_ohlcv(**kwargs),
)
async def fetch_ticker(self, symbol: str) -> TickerData:
"""Fetch ticker data from the exchange."""
await self.rate_limiter.acquire()
raw = await self._async_fetch_ticker(symbol)
return TickerData(
symbol=symbol,
exchange=self.get_name(),
bid=Decimal(str(raw.get("bid", 0))),
ask=Decimal(str(raw.get("ask", 0))),
last=Decimal(str(raw.get("last", 0))),
volume_24h=Decimal(str(raw.get("baseVolume", 0))),
change_24h=Decimal(str(raw.get("percentage", 0))) if raw.get("percentage") is not None else None,
timestamp=datetime.fromtimestamp(raw["timestamp"] / 1000, tz=timezone.utc) if raw.get("timestamp") else datetime.now(tz=timezone.utc),
)
async def _async_fetch_ticker(self, symbol: str) -> dict[str, Any]:
"""Run the synchronous CCXT fetch_ticker in a thread pool."""
import asyncio
loop = asyncio.get_running_loop()
return await loop.run_in_executor(
None,
lambda: self.client.fetch_ticker(symbol),
)
async def fetch_symbols(self) -> list[SymbolInfo]:
"""Fetch all available trading pairs from the exchange."""
await self.rate_limiter.acquire()
import asyncio
loop = asyncio.get_running_loop()
markets = await loop.run_in_executor(None, lambda: self.client.load_markets())
result: list[SymbolInfo] = []
for sym, info in markets.items():
if info.get("active", True):
result.append(
SymbolInfo(
symbol=sym,
base=info.get("base", ""),
quote=info.get("quote", ""),
exchange=self.get_name(),
is_active=info.get("active", True),
)
)
return result
# ──────────────────────────────────────────────
# Order placement
# ──────────────────────────────────────────────
async def create_order(self, req: OrderRequest) -> OrderData:
"""Place an order on the exchange.
Subclasses may override to add exchange-specific logic (e.g. leverage,
position side). Default implementation uses CCXT's create_order.
"""
import asyncio
await self.rate_limiter.acquire()
params: dict[str, Any] = {}
if req.reduce_only:
params["reduceOnly"] = True
if req.position_side:
params["positionSide"] = req.position_side.upper()
def _place() -> dict[str, Any]:
return self.client.create_order(
symbol=req.symbol,
type=req.order_type,
side=req.side,
amount=float(req.amount),
price=float(req.price) if req.price else None,
params=params,
)
loop = asyncio.get_running_loop()
raw = await loop.run_in_executor(None, _place)
return OrderData(
exchange=self.get_name(),
symbol=raw.get("symbol", req.symbol),
order_id=str(raw.get("id", "")),
client_order_id=raw.get("clientOrderId"),
side=raw.get("side", req.side),
order_type=raw.get("type", req.order_type),
amount=Decimal(str(raw.get("amount", float(req.amount)))),
filled=Decimal(str(raw.get("filled", 0))),
price=Decimal(str(raw["price"])) if raw.get("price") else req.price,
average=Decimal(str(raw["average"])) if raw.get("average") else None,
status=raw.get("status", "open"),
timestamp=datetime.fromtimestamp(raw["timestamp"] / 1000, tz=timezone.utc) if raw.get("timestamp") else datetime.now(tz=timezone.utc),
raw=raw,
)
# ──────────────────────────────────────────────
# Balance
# ──────────────────────────────────────────────
async def fetch_balance(self) -> BalanceResponse:
"""Fetch the full account balance from the exchange."""
import asyncio
await self.rate_limiter.acquire()
def _fetch() -> dict[str, Any]:
return self.client.fetch_balance()
loop = asyncio.get_running_loop()
raw = await loop.run_in_executor(None, _fetch)
balances: list[BalanceData] = []
for asset, info in raw.get("total", {}).items():
if asset == "info" or asset == "free" or asset == "used" or asset == "total":
continue
free = Decimal(str(raw.get("free", {}).get(asset, 0)))
used = Decimal(str(raw.get("used", {}).get(asset, 0)))
total = Decimal(str(info))
if total > 0 or free > 0:
balances.append(BalanceData(asset=asset, free=free, used=used, total=total))
return BalanceResponse(
exchange=self.get_name(),
balances=balances,
timestamp=datetime.now(tz=timezone.utc),
)
# ──────────────────────────────────────────────
# Open Orders, Trades, Positions
# ──────────────────────────────────────────────
async def fetch_open_orders(self, symbol: str | None = None) -> list[OpenOrderData]:
"""Fetch open orders from the exchange."""
import asyncio
await self.rate_limiter.acquire()
def _fetch() -> list[dict[str, Any]]:
return self.client.fetch_open_orders(symbol=symbol)
loop = asyncio.get_running_loop()
raw_orders = await loop.run_in_executor(None, _fetch)
orders: list[OpenOrderData] = []
for raw in (raw_orders or []):
orders.append(OpenOrderData(
order_id=str(raw.get("id", "")),
symbol=raw.get("symbol", ""),
side=raw.get("side", ""),
order_type=raw.get("type", ""),
amount=Decimal(str(raw.get("amount", 0))),
filled=Decimal(str(raw.get("filled", 0))),
price=Decimal(str(raw["price"])) if raw.get("price") else None,
average=Decimal(str(raw["average"])) if raw.get("average") else None,
status=raw.get("status", "open"),
timestamp=datetime.fromtimestamp(raw["timestamp"] / 1000, tz=timezone.utc)
if raw.get("timestamp") else None,
))
return orders
async def fetch_my_trades(self, symbol: str | None = None, limit: int = 10) -> list[TradeData]:
"""Fetch recent filled trades from the exchange."""
import asyncio
await self.rate_limiter.acquire()
def _fetch() -> list[dict[str, Any]]:
return self.client.fetch_my_trades(symbol=symbol, limit=limit)
loop = asyncio.get_running_loop()
raw_trades = await loop.run_in_executor(None, _fetch)
trades: list[TradeData] = []
for raw in (raw_trades or []):
cost = None
if raw.get("cost"):
cost = Decimal(str(raw["cost"]))
fee_val = None
fee_currency = None
if raw.get("fee"):
fee_val = Decimal(str(raw["fee"].get("cost", 0)))
fee_currency = raw["fee"].get("currency")
trades.append(TradeData(
trade_id=str(raw.get("id", "")),
symbol=raw.get("symbol", ""),
side=raw.get("side", ""),
amount=Decimal(str(raw.get("amount", 0))),
price=Decimal(str(raw.get("price", 0))),
cost=cost,
fee=fee_val,
fee_currency=fee_currency,
timestamp=datetime.fromtimestamp(raw["timestamp"] / 1000, tz=timezone.utc)
if raw.get("timestamp") else None,
))
return trades
async def fetch_positions(self, symbols: list[str] | None = None) -> list[PositionData]:
"""Fetch open positions (futures/derivatives) from the exchange."""
import asyncio
await self.rate_limiter.acquire()
def _fetch() -> list[dict[str, Any]]:
return self.client.fetch_positions(symbols=symbols)
loop = asyncio.get_running_loop()
raw_positions = await loop.run_in_executor(None, _fetch)
positions: list[PositionData] = []
for raw in (raw_positions or []):
contracts = Decimal(str(raw.get("contracts", 0)))
if contracts == 0:
continue
positions.append(PositionData(
symbol=raw.get("symbol", ""),
side=raw.get("side", ""),
contracts=contracts,
entry_price=Decimal(str(raw["entryPrice"])) if raw.get("entryPrice") else None,
mark_price=Decimal(str(raw["markPrice"])) if raw.get("markPrice") else None,
unrealized_pnl=Decimal(str(raw["unrealizedPnl"])) if raw.get("unrealizedPnl") else None,
leverage=Decimal(str(raw["leverage"])) if raw.get("leverage") else None,
liquidation_price=Decimal(str(raw["liquidationPrice"])) if raw.get("liquidationPrice") else None,
percentage=Decimal(str(raw["percentage"])) if raw.get("percentage") else None,
))
return positions
from decimal import Decimal
from datetime import datetime, timezone
+37
View File
@@ -0,0 +1,37 @@
from __future__ import annotations
import ccxt
from app.exchange.base import AbstractExchange
class BinanceAdapter(AbstractExchange):
"""Exchange adapter for Binance."""
def get_name(self) -> str:
return "binance"
def get_base_url(self) -> str:
return "https://api.binance.com"
def get_ws_url(self) -> str:
return "wss://stream.binance.com:9443/ws"
def _init_ccxt(self) -> ccxt.Exchange:
options: dict = {
"apiKey": self.api_key,
"secret": self.api_secret,
"rateLimit": 1200,
"enableRateLimit": True,
"options": {
"warnOnFetchOpenOrdersWithoutSymbol": False,
},
}
if self.testnet:
options["urls"] = {
"api": {
"public": "https://testnet.binance.vision/api",
"private": "https://testnet.binance.vision/api",
}
}
return ccxt.binance(options)
+26
View File
@@ -0,0 +1,26 @@
from __future__ import annotations
import ccxt
from app.exchange.base import AbstractExchange
class BingXAdapter(AbstractExchange):
"""Exchange adapter for BingX."""
def get_name(self) -> str:
return "bingx"
def get_base_url(self) -> str:
return "https://api.bingx.com"
def get_ws_url(self) -> str:
return "wss://open-api-ws.bingx.com/market"
def _init_ccxt(self) -> ccxt.Exchange:
options: dict = {
"apiKey": self.api_key,
"secret": self.api_secret,
"enableRateLimit": True,
}
return ccxt.bingx(options)
+33
View File
@@ -0,0 +1,33 @@
from __future__ import annotations
import ccxt
from app.exchange.base import AbstractExchange
class BybitAdapter(AbstractExchange):
"""Exchange adapter for Bybit."""
def get_name(self) -> str:
return "bybit"
def get_base_url(self) -> str:
return "https://api.bybit.com"
def get_ws_url(self) -> str:
return "wss://stream.bybit.com/v5/public/spot"
def _init_ccxt(self) -> ccxt.Exchange:
options: dict = {
"apiKey": self.api_key,
"secret": self.api_secret,
"enableRateLimit": True,
}
if self.testnet:
options["urls"] = {
"api": {
"public": "https://api-testnet.bybit.com",
"private": "https://api-testnet.bybit.com",
}
}
return ccxt.bybit(options)
+87
View File
@@ -0,0 +1,87 @@
from __future__ import annotations
import hashlib
import time
from typing import Optional
from app.exchange.base import AbstractExchange
from app.exchange.binance import BinanceAdapter
from app.exchange.bingx import BingXAdapter
from app.exchange.bybit import BybitAdapter
from app.exchange.gate import GateAdapter
from app.exchange.mexc import MEXCAdapter
class ExchangeFactory:
"""Factory for creating exchange adapter instances.
Implements the Singleton pattern via a global instance.
Caches adapter instances (TTL 5 min) to avoid repeated CCXT init (~20s).
"""
_adapters: dict[str, type[AbstractExchange]] = {
"binance": BinanceAdapter,
"bingx": BingXAdapter,
"bybit": BybitAdapter,
"gate": GateAdapter,
"mexc": MEXCAdapter,
}
def __init__(self) -> None:
self._cache: dict[str, tuple[AbstractExchange, float]] = {}
self._cache_ttl: float = 300.0 # 5 minutes
def register(self, name: str, adapter_cls: type[AbstractExchange]) -> None:
"""Register a new exchange adapter class."""
self._adapters[name] = adapter_cls
def _cache_key(self, name: str, api_key: str) -> str:
"""Generate a cache key from exchange name + API key."""
return f"{name}:{hashlib.sha256(api_key.encode()).hexdigest()[:16]}"
def create(
self,
name: str,
api_key: str = "",
api_secret: str = "",
testnet: bool = False,
) -> AbstractExchange:
"""Create (or retrieve cached) an exchange adapter instance by name.
Raises ValueError if the exchange is not registered.
"""
adapter_cls = self._adapters.get(name)
if adapter_cls is None:
raise ValueError(
f"Unknown exchange: {name!r}. "
f"Available exchanges: {', '.join(sorted(self._adapters))}"
)
# Return cached adapter if still fresh
if api_key:
key = self._cache_key(name, api_key)
cached = self._cache.get(key)
if cached is not None:
adapter, created_at = cached
if time.monotonic() - created_at < self._cache_ttl:
return adapter
# Expired — remove from cache
del self._cache[key]
adapter = adapter_cls(api_key=api_key, api_secret=api_secret, testnet=testnet)
# Warm up the CCXT client in the background — first balance call will trigger it
# but subsequent calls within TTL reuse the same instance
if api_key:
key = self._cache_key(name, api_key)
self._cache[key] = (adapter, time.monotonic())
return adapter
def get_available_exchanges(self) -> list[str]:
"""Return a list of registered exchange names."""
return list(self._adapters.keys())
# Global singleton factory instance
factory: ExchangeFactory = ExchangeFactory()
+26
View File
@@ -0,0 +1,26 @@
from __future__ import annotations
import ccxt
from app.exchange.base import AbstractExchange
class GateAdapter(AbstractExchange):
"""Exchange adapter for Gate.io."""
def get_name(self) -> str:
return "gate"
def get_base_url(self) -> str:
return "https://api.gateio.ws"
def get_ws_url(self) -> str:
return "wss://api.gateio.ws/ws/v4/"
def _init_ccxt(self) -> ccxt.Exchange:
options: dict = {
"apiKey": self.api_key,
"secret": self.api_secret,
"enableRateLimit": True,
}
return ccxt.gate(options)
+33
View File
@@ -0,0 +1,33 @@
from __future__ import annotations
import ccxt
from app.exchange.base import AbstractExchange
class MEXCAdapter(AbstractExchange):
"""Exchange adapter for MEXC."""
def get_name(self) -> str:
return "mexc"
def get_base_url(self) -> str:
return "https://api.mexc.com"
def get_ws_url(self) -> str:
return "wss://wbs.mexc.com/ws"
def _init_ccxt(self) -> ccxt.Exchange:
options: dict = {
"apiKey": self.api_key,
"secret": self.api_secret,
"enableRateLimit": True,
}
if self.testnet:
options["urls"] = {
"api": {
"public": "https://testnet-api.mexc.com",
"private": "https://testnet-api.mexc.com",
}
}
return ccxt.mexc(options)
+72
View File
@@ -0,0 +1,72 @@
from __future__ import annotations
import asyncio
import time
from typing import Optional
# Per-exchange rate limits (requests per second for standard API)
BINANCE_RATE_LIMIT: int = 10
BYBIT_RATE_LIMIT: int = 10
MEXC_RATE_LIMIT: int = 20
class RateLimiter:
"""Token bucket rate limiter for exchange API requests."""
def __init__(self, tokens_per_second: float, max_tokens: Optional[int] = None, name: str = "") -> None:
self.tokens_per_second = tokens_per_second
self.max_tokens = max_tokens if max_tokens is not None else int(tokens_per_second)
self.name = name
self._tokens: float = float(self.max_tokens)
self._last_refill: float = time.monotonic()
self._lock: asyncio.Lock = asyncio.Lock()
async def acquire(self) -> None:
"""Wait for a token to be available, blocking until one is."""
while True:
async with self._lock:
self._refill()
if self._tokens >= 1.0:
self._tokens -= 1.0
return
# How long until we have at least 1 token?
wait_time = (1.0 - self._tokens) / self.tokens_per_second
await asyncio.sleep(wait_time)
def _refill(self) -> None:
now = time.monotonic()
elapsed = now - self._last_refill
self._tokens = min(float(self.max_tokens), self._tokens + elapsed * self.tokens_per_second)
self._last_refill = now
async def __aenter__(self) -> "RateLimiter":
await self.acquire()
return self
async def __aexit__(
self,
exc_type: Optional[type[BaseException]],
exc_val: Optional[BaseException],
exc_tb: Optional[object],
) -> None:
pass
# Global rate limiter registry: singleton mapping exchange_name -> RateLimiter instance
_global_limiters: dict[str, RateLimiter] = {}
def GlobalRateLimiter(exchange_name: str) -> RateLimiter:
"""Get or create the singleton RateLimiter for an exchange."""
if exchange_name not in _global_limiters:
limit_map = {
"binance": BINANCE_RATE_LIMIT,
"bybit": BYBIT_RATE_LIMIT,
"mexc": MEXC_RATE_LIMIT,
}
tokens_per_second = limit_map.get(exchange_name, 10)
_global_limiters[exchange_name] = RateLimiter(
tokens_per_second=tokens_per_second,
name=exchange_name,
)
return _global_limiters[exchange_name]
+146
View File
@@ -0,0 +1,146 @@
from __future__ import annotations
from datetime import datetime
from decimal import Decimal
from typing import Optional
from pydantic import BaseModel
class CandleData(BaseModel):
symbol: str
exchange: str
timeframe: str
timestamp: datetime
open: Decimal
high: Decimal
low: Decimal
close: Decimal
volume: Decimal
class SymbolInfo(BaseModel):
symbol: str
base: str
quote: str
exchange: str
is_active: bool = True
class TickerData(BaseModel):
symbol: str
exchange: str
bid: Decimal
ask: Decimal
last: Decimal
volume_24h: Decimal
change_24h: Optional[Decimal] = None
timestamp: datetime
class CandleValidationError(ValueError):
"""Raised when OHLC data is inconsistent."""
def __init__(self, message: str, *, open: Decimal, high: Decimal, low: Decimal, close: Decimal) -> None:
self.open = open
self.high = high
self.low = low
self.close = close
super().__init__(message)
class OrderRequest(BaseModel):
"""Parameters for placing an order on an exchange."""
symbol: str
side: str # "buy" or "sell"
order_type: str = "market" # "market" or "limit"
amount: Decimal # base currency amount (e.g. BTC amount)
price: Optional[Decimal] = None # required for limit orders
reduce_only: bool = False
position_side: Optional[str] = None # "long" or "short" (for futures)
class OrderData(BaseModel):
"""Response from placing an order."""
exchange: str
symbol: str
order_id: str
client_order_id: Optional[str] = None
side: str
order_type: str
amount: Decimal
filled: Decimal
price: Optional[Decimal] = None
average: Optional[Decimal] = None
status: str # "open", "closed", "canceled", "rejected"
timestamp: datetime
raw: Optional[dict] = None
class BalanceData(BaseModel):
"""Account balance for a single asset."""
asset: str
free: Decimal
used: Decimal
total: Decimal
class BalanceResponse(BaseModel):
"""Full account balance snapshot."""
exchange: str
balances: list[BalanceData]
timestamp: datetime
class OpenOrderData(BaseModel):
"""An open order from the exchange."""
order_id: str
symbol: str
side: str # "buy" or "sell"
order_type: str # "market", "limit", etc.
amount: Decimal
filled: Decimal
price: Optional[Decimal] = None
average: Optional[Decimal] = None
status: str
timestamp: Optional[datetime] = None
class PositionData(BaseModel):
"""A futures/derivatives position."""
symbol: str
side: str # "long" or "short"
contracts: Decimal
entry_price: Optional[Decimal] = None
mark_price: Optional[Decimal] = None
unrealized_pnl: Optional[Decimal] = None
leverage: Optional[Decimal] = None
liquidation_price: Optional[Decimal] = None
percentage: Optional[Decimal] = None
class TradeData(BaseModel):
"""A filled trade from the exchange."""
trade_id: str
symbol: str
side: str # "buy" or "sell"
amount: Decimal
price: Decimal
cost: Optional[Decimal] = None
fee: Optional[Decimal] = None
fee_currency: Optional[str] = None
timestamp: Optional[datetime] = None
class CredentialSummaryResponse(BaseModel):
"""Aggregated summary for one credential."""
exchange_name: str
usdt_balance: Decimal
usdt_estimate: str # "≈ $1,234.56" for non-USDT assets
balance_count: int # number of tokens with balance > 0
open_orders_count: int
positions_count: int
balances: list[BalanceData]
open_orders: list[OpenOrderData]
positions: list[PositionData]
recent_trades: list[TradeData]
+100
View File
@@ -0,0 +1,100 @@
"""
One-shot script to force sync candle data from MEXC.
Inserts directly into candles_default partition.
"""
from __future__ import annotations
import asyncio
import logging
from decimal import Decimal
from sqlalchemy import and_, select, text
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import joinedload
from app.database import async_session_factory
from app.exchange.factory import factory as exchange_factory
from app.models.exchange import Exchange
from app.models.symbol import Symbol
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
TIMEFRAMES = ["15m", "30m", "1h", "4h", "1d"]
async def sync():
async with async_session_factory() as db:
# Get active MEXC symbols for our tokens
watchlist_symbols = ['BTC/USDT','ETH/USDT','SOL/USDT','NEAR/USDT','BNB/USDT','SQD/USDT','ZEN/USDT','KITE/USDT']
result = await db.execute(
select(Symbol)
.options(joinedload(Symbol.exchange))
.where(Symbol.symbol.in_(watchlist_symbols), Symbol.is_active == True)
)
symbols: list[Symbol] = list(result.scalars().all())
logger.info(f"Found {len(symbols)} active symbols to sync")
adapter = exchange_factory.create("mexc")
total_inserted = 0
for db_symbol in symbols:
for tf in TIMEFRAMES:
try:
candles = await adapter.fetch_ohlcv(
symbol=db_symbol.symbol,
timeframe=tf,
limit=200,
)
except Exception as e:
logger.error(f"Fetch error {db_symbol.symbol} {tf}: {e}")
continue
if not candles:
continue
inserted = 0
for c in candles:
try:
await db.execute(
text("""
INSERT INTO candles_default
(symbol_id, timeframe, timestamp, open, high, low, close, volume)
VALUES (:sid, :tf, :ts, :o, :h, :l, :c, :v)
ON CONFLICT (symbol_id, timeframe, timestamp) DO NOTHING
"""),
{
"sid": db_symbol.id,
"tf": c.timeframe,
"ts": c.timestamp,
"o": Decimal(str(c.open)),
"h": Decimal(str(c.high)),
"l": Decimal(str(c.low)),
"c": Decimal(str(c.close)),
"v": Decimal(str(c.volume)),
}
)
inserted += 1
except Exception as e:
logger.debug(f"Insert error: {e}")
await db.commit()
total_inserted += inserted
logger.info(f" ✅ {db_symbol.symbol:12s} {tf:4s}: {inserted} candles")
logger.info(f"\n✅ DONE! Total: {total_inserted} candles inserted")
# Verify
r = await db.execute(text("""
SELECT s.symbol, c.timeframe, MAX(c.timestamp) as ts
FROM candles_default c
JOIN symbols s ON s.id = c.symbol_id
WHERE s.symbol IN ('BTC/USDT','ETH/USDT','SOL/USDT')
GROUP BY s.symbol, c.timeframe
ORDER BY s.symbol
"""))
print("\n📊 Latest candles:")
for row in r:
print(f" {row[0]:12s} {row[1]:4s}: {row[2]}")
asyncio.run(sync())
+308
View File
@@ -0,0 +1,308 @@
from __future__ import annotations
import asyncio
import json
import logging
import time
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from fastapi import FastAPI, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from sqlalchemy import select, text
from app.config import settings
from app.core.exceptions import AppException
from app.core.middleware import RequestLoggingMiddleware
from pydantic import ValidationError
from app.database import async_session_factory, engine
from app.tasks.candle_fetcher import force_full_sync
from app.tasks.exchange_sync import sync_all_exchanges
from app.models import Exchange
from app.services.ws_push_service import setup_push_listener
# P3-4: Structured JSON logging
class JsonFormatter(logging.Formatter):
def format(self, record: logging.LogRecord) -> str:
log_entry = {
"timestamp": datetime.now(timezone.utc).isoformat(),
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
"module": record.module,
"line": record.lineno,
}
if record.exc_info and record.exc_info[0]:
log_entry["exception"] = self.formatException(record.exc_info)
return json.dumps(log_entry, default=str)
log_handler = logging.StreamHandler()
log_handler.setFormatter(JsonFormatter())
logging.getLogger().handlers = [log_handler]
logging.getLogger().setLevel(getattr(logging, settings.LOG_LEVEL.upper(), logging.INFO))
logger = logging.getLogger(__name__)
# Track application start time for uptime reporting
_start_time: float = time.time()
@asynccontextmanager
async def lifespan(app: FastAPI):
"""Application lifespan handler for startup and shutdown events."""
logger.info(
"Starting Trading Portal API on %s:%s",
settings.HOST,
settings.PORT,
)
logger.info("Log level: %s", settings.LOG_LEVEL)
# --- Auto-create tables (dev bootstrap, replaces Alembic) ---
try:
from app.database import Base, engine
from app import models # noqa: F401
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
logger.info("Database tables verified/created")
except Exception:
logger.exception("Table creation failed (non-fatal)")
# --- Startup: background tasks ---
# Candle fetcher: Binance USDT only (1,075 symbols), 4 TFs (15m/1h/4h/1d),
# batch of 25 symbols every 10 min. Full cycle ~7.2h, DB-safe.
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from app.tasks.candle_fetcher import setup_candle_scheduler
scheduler = setup_candle_scheduler(app)
app.state.scheduler = scheduler
scheduler.start()
logger.info("Background scheduler started with optimized candle fetcher")
# --- Schedule stale trade management (every 30 min) ---
try:
from apscheduler.triggers.interval import IntervalTrigger
async def _close_stale_trades():
from app.services.signal_service import close_stale_trades
await close_stale_trades()
scheduler.add_job(
_close_stale_trades,
IntervalTrigger(minutes=5),
id="close-stale-trades",
replace_existing=True,
misfire_grace_time=60,
)
logger.info("Stale trade management scheduled (every 5 min)")
# P3-5: Signal aging — expire old signals daily
async def _expire_signals():
from app.services.signal_service import expire_old_signals
await expire_old_signals(max_age_days=7)
scheduler.add_job(
_expire_signals,
IntervalTrigger(hours=6),
id="expire-old-signals",
replace_existing=True,
misfire_grace_time=300,
)
logger.info("Signal aging scheduled (every 6 hours)")
except Exception:
logger.exception("Failed to schedule stale trade management")
# --- Register WebSocket push listener ---
setup_push_listener(app)
# Exchange sync already done — skip on restart
# async def _background_exchange_sync():
# ...
logger.info("Exchange symbol sync skipped — already synced")
# --- Force full candle sync on startup (DISABLED — too heavy with 16K symbols) ---
# Background scheduler handles incremental sync every 2 min instead
# Uncomment when symbol count is manageable:
# async def _background_sync():
# try:
# await force_full_sync(app)
# except Exception:
# logger.exception("Force full sync failed (non-fatal)")
# asyncio.ensure_future(_background_sync())
logger.info("Full candle sync skipped — relying on incremental scheduler")
# --- Periodic win rate recalculation (every 6 hours) ---
async def _recalc_win_rates():
from app.services.signal_booster import compute_strategy_win_rates
while True:
await asyncio.sleep(21600) # 6 hours
try:
async with async_session_factory() as db:
await compute_strategy_win_rates(db)
logger.info("Win rates refreshed (periodic recalc)")
except Exception:
logger.exception("Win rate periodic refresh failed")
asyncio.ensure_future(_recalc_win_rates())
logger.info("Win rate recalc scheduled (every 6 hours)")
logger.info("Trading Portal API startup complete")
yield
# --- Shutdown ---
# P3-7: Graceful shutdown with DB cleanup
try:
scheduler.shutdown(wait=False)
logger.info("Background scheduler shutdown complete")
except Exception:
logger.exception("Scheduler shutdown failed (ignoring)")
# Close database engine
try:
await engine.dispose()
logger.info("Database engine disposed")
except Exception:
logger.exception("Engine disposal failed (ignoring)")
# Cancel remaining background tasks
tasks = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
for task in tasks:
task.cancel()
if tasks:
await asyncio.gather(*tasks, return_exceptions=True)
logger.info("Cancelled %d background tasks", len(tasks))
logger.info("Shutting down Trading Portal API")
app = FastAPI(
title="Trading Portal API",
description="Backend API for the Trading Portal application",
version="1.0.0",
lifespan=lifespan,
)
# ---------------------------------------------------------------------------
# CORS middleware
# ---------------------------------------------------------------------------
if settings.CORS_ORIGINS:
origins = [o.strip() for o in settings.CORS_ORIGINS.split(",") if o.strip()]
else:
origins = ["*"]
app.add_middleware(
CORSMiddleware,
allow_origins=origins,
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# ---------------------------------------------------------------------------
# Request logging middleware
# ---------------------------------------------------------------------------
app.add_middleware(RequestLoggingMiddleware)
# ---------------------------------------------------------------------------
# Exception handlers
# ---------------------------------------------------------------------------
@app.exception_handler(AppException)
async def app_exception_handler(request: Request, exc: AppException) -> JSONResponse:
"""Return a JSON response for application-level exceptions."""
return JSONResponse(
status_code=exc.status_code,
content={"detail": exc.detail, "code": exc.code},
)
@app.exception_handler(Exception)
async def unhandled_exception_handler(request: Request, exc: Exception) -> JSONResponse:
"""Catch any unhandled exception and return a generic 500 response."""
logger.exception("Unhandled exception: %s", str(exc))
return JSONResponse(
status_code=500,
content={"detail": "Internal server error", "code": "internal_error"},
)
@app.exception_handler(404)
async def not_found_handler(request: Request, exc) -> JSONResponse:
"""Return a JSON 404 for unmatched routes."""
return JSONResponse(
status_code=404,
content={"detail": "Not found", "code": "not_found"},
)
@app.exception_handler(ValidationError)
async def validation_error_handler(request: Request, exc: ValidationError) -> JSONResponse:
"""Return a JSON 422 with validation error details."""
return JSONResponse(
status_code=422,
content={
"detail": "Validation error",
"code": "validation_error",
"errors": exc.errors(),
},
)
# ---------------------------------------------------------------------------
# Health check
# ---------------------------------------------------------------------------
@app.get("/health")
async def health_check():
"""Return health status including version, uptime, and database connectivity.
NOTE: Does NOT test exchange connectivity to avoid thread pool exhaustion
and Docker health check timeouts.
"""
db_connected = False
db_latency = 0.0
try:
async with async_session_factory() as db_session:
start = time.time()
await db_session.execute(text("SELECT 1"))
db_latency = (time.time() - start) * 1000 # ms
db_connected = True
except Exception:
db_connected = False
uptime_seconds = time.time() - _start_time
health = {
"status": "healthy" if db_connected else "degraded",
"version": "1.0.0",
"uptime": round(uptime_seconds, 2),
"db_connected": db_connected,
"db_latency_ms": round(db_latency, 2),
}
return health
async def _test_exchange_connection(exchange_name: str) -> bool:
"""Try connecting to an exchange to verify it's reachable."""
try:
import asyncio
from app.exchange.factory import factory
adapter = factory.create(exchange_name)
loop = asyncio.get_running_loop()
await loop.run_in_executor(None, lambda: adapter.client.load_markets())
return True
except Exception:
return False
# ---------------------------------------------------------------------------
# V1 API routers
# ---------------------------------------------------------------------------
from app.api.v1.router import api_router
app.include_router(api_router)
# ---------------------------------------------------------------------------
# WebSocket router
# ---------------------------------------------------------------------------
from app.api.ws.candle_handler import router as ws_router
app.include_router(ws_router)
+165
View File
@@ -0,0 +1,165 @@
"""
Entrypoint: API-only service (no scheduler).
User-facing HTTP API — always fast, never blocked by candle fetching.
"""
from __future__ import annotations
import json
import logging
import time
from contextlib import asynccontextmanager
from datetime import datetime, timezone
from fastapi import FastAPI, Request
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse
from sqlalchemy import text
from pydantic import ValidationError
from app.config import settings
from app.core.exceptions import AppException
from app.core.middleware import RequestLoggingMiddleware
from app.database import async_session_factory, engine
# Structured JSON logging
class JsonFormatter(logging.Formatter):
def format(self, record: logging.LogRecord) -> str:
log_entry = {
"timestamp": datetime.now(timezone.utc).isoformat(),
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
"module": record.module,
"line": record.lineno,
}
if record.exc_info and record.exc_info[0]:
log_entry["exception"] = self.formatException(record.exc_info)
return json.dumps(log_entry, default=str)
log_handler = logging.StreamHandler()
log_handler.setFormatter(JsonFormatter())
logging.getLogger().handlers = [log_handler]
logging.getLogger().setLevel(getattr(logging, settings.LOG_LEVEL.upper(), logging.INFO))
logger = logging.getLogger(__name__)
_start_time: float = time.time()
@asynccontextmanager
async def lifespan(app: FastAPI):
logger.info("Starting Trading Portal API on %s:%s", settings.HOST, settings.PORT)
logger.info("Log level: %s", settings.LOG_LEVEL)
# Auto-create tables
try:
from app.database import Base, engine
from app import models # noqa: F401
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
logger.info("Database tables verified/created")
except Exception:
logger.exception("Table creation failed (non-fatal)")
# Register WebSocket push listener
from app.services.ws_push_service import setup_push_listener
setup_push_listener(app)
logger.info("Trading Portal API startup complete (API-only, no scheduler)")
yield
# Shutdown
try:
await engine.dispose()
logger.info("Database engine disposed")
except Exception:
logger.exception("Engine disposal failed")
app = FastAPI(
title="Trading Portal API",
description="User-facing API for Trading Portal",
version="2.0.0",
lifespan=lifespan,
)
# CORS
if settings.CORS_ORIGINS:
origins = [o.strip() for o in settings.CORS_ORIGINS.split(",") if o.strip()]
else:
origins = ["*"]
app.add_middleware(CORSMiddleware, allow_origins=origins, allow_credentials=True,
allow_methods=["*"], allow_headers=["*"])
app.add_middleware(RequestLoggingMiddleware)
# Exception handlers
@app.exception_handler(AppException)
async def app_exception_handler(request: Request, exc: AppException) -> JSONResponse:
return JSONResponse(status_code=exc.status_code,
content={"detail": exc.detail, "code": exc.code})
@app.exception_handler(Exception)
async def unhandled_exception_handler(request: Request, exc: Exception) -> JSONResponse:
logger.exception("Unhandled exception: %s", str(exc))
return JSONResponse(status_code=500,
content={"detail": "Internal server error", "code": "internal_error"})
@app.exception_handler(404)
async def not_found_handler(request: Request, exc) -> JSONResponse:
return JSONResponse(status_code=404,
content={"detail": "Not found", "code": "not_found"})
@app.exception_handler(ValidationError)
async def validation_error_handler(request: Request, exc: ValidationError) -> JSONResponse:
from uuid import UUID as _UUID
errors: list[dict] = []
for e in exc.errors():
# Convert non-JSON-serializable values (UUID, datetime, etc.)
safe = {}
for k, v in e.items():
if isinstance(v, _UUID):
safe[k] = str(v)
elif isinstance(v, (datetime,)):
safe[k] = v.isoformat()
else:
safe[k] = v
errors.append(safe)
return JSONResponse(status_code=422,
content={"detail": "Validation error", "code": "validation_error",
"errors": errors})
# Health check
@app.get("/health")
async def health_check():
db_connected = False
db_latency = 0.0
try:
async with async_session_factory() as db_session:
start = time.time()
await db_session.execute(text("SELECT 1"))
db_latency = (time.time() - start) * 1000
db_connected = True
except Exception:
db_connected = False
return {
"status": "healthy" if db_connected else "degraded",
"version": "2.0.0-api",
"uptime": round(time.time() - _start_time, 2),
"db_connected": db_connected,
"db_latency_ms": round(db_latency, 2),
}
# API routes
from app.api.v1.router import api_router
app.include_router(api_router)
from app.api.ws.candle_handler import router as ws_router
app.include_router(ws_router)
+183
View File
@@ -0,0 +1,183 @@
"""Entrypoint: Scheduler-only service (no HTTP server).
Runs candle fetcher, signal analysis, trade management, win rate recalc.
Heavy CPU work stays here, API stays fast.
"""
from __future__ import annotations
import asyncio
import json
import logging
import signal
import sys
from datetime import datetime, timezone
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from apscheduler.triggers.interval import IntervalTrigger
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
from app.config import settings
from app.database import Base
# ── Scheduler-specific engine: longer statement_timeout (30s) ──
# Unlike the API which needs fast-fail (<10s), the scheduler does heavy
# batch UPDATE queries (trade eviction executemany) that can take longer.
_scheduler_engine = create_async_engine(
settings.DATABASE_URL,
echo=(settings.LOG_LEVEL == "DEBUG"),
pool_pre_ping=True,
pool_size=20, # semaphore=8 limits concurrent analyses
max_overflow=5,
pool_timeout=5,
pool_recycle=600,
connect_args={
"server_settings": {
"idle_in_transaction_session_timeout": "60000",
"statement_timeout": "30000", # 30s vs API's 10s
}
},
)
_scheduler_session_factory = async_sessionmaker(
_scheduler_engine,
class_=AsyncSession,
expire_on_commit=False,
)
# Monkey-patch app.database so all downstream imports use scheduler engine
import app.database as _db_module
_db_module.engine = _scheduler_engine
_db_module.async_session_factory = _scheduler_session_factory
class JsonFormatter(logging.Formatter):
def format(self, record: logging.LogRecord) -> str:
log_entry = {
"timestamp": datetime.now(timezone.utc).isoformat(),
"level": record.levelname,
"logger": record.name,
"message": record.getMessage(),
"module": record.module,
"line": record.lineno,
}
if record.exc_info and record.exc_info[0]:
log_entry["exception"] = self.formatException(record.exc_info)
return json.dumps(log_entry, default=str)
log_handler = logging.StreamHandler()
log_handler.setFormatter(JsonFormatter())
logging.getLogger().handlers = [log_handler]
logging.getLogger().setLevel(getattr(logging, settings.LOG_LEVEL.upper(), logging.INFO))
logger = logging.getLogger(__name__)
async def _close_stale_trades():
from app.services.signal_service import close_stale_trades
await close_stale_trades()
async def _expire_signals():
from app.services.signal_service import expire_old_signals
await expire_old_signals(max_age_days=7)
async def _recalc_win_rates_loop():
"""Periodic win rate recalculation every 6 hours."""
from app.services.signal_booster import compute_strategy_win_rates
while True:
await asyncio.sleep(21600) # 6 hours
try:
async with async_session_factory() as db:
await compute_strategy_win_rates(db)
logger.info("Win rates refreshed (periodic recalc)")
except Exception:
logger.exception("Win rate periodic refresh failed")
async def main():
logger.info("=== Trading Portal Scheduler ===")
logger.info("Mode: scheduler-only (no HTTP server)")
# Auto-create tables (first run)
try:
from app.database import Base
from app import models # noqa: F401
async with _scheduler_engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
logger.info("Database tables verified/created")
except Exception:
logger.exception("Table creation failed (non-fatal)")
# Create scheduler
scheduler = AsyncIOScheduler()
# Candle fetcher — the core job (5-min interval)
from app.tasks.candle_fetcher import fetch_recent_candles, _TIMEFRAMES_OPTIMIZED
scheduler.add_job(
fetch_recent_candles,
trigger=IntervalTrigger(seconds=300),
args=[None, 2, _TIMEFRAMES_OPTIMIZED, 25],
id="fetch_candles_optimized",
name="Fetch trading candles (top 100 bases, 5 exchanges, 4TFs)",
replace_existing=True,
coalesce=True,
max_instances=1,
misfire_grace_time=600,
)
logger.info("Candle fetcher scheduled: every 5 min, 25/batch")
# Stale trade management (every 5 min)
scheduler.add_job(
_close_stale_trades,
IntervalTrigger(minutes=5),
id="close-stale-trades",
replace_existing=True,
)
# Real trade sync (every 5 min)
from app.services.trade_executor import sync_real_trades
scheduler.add_job(
sync_real_trades,
IntervalTrigger(minutes=5),
id="sync-real-trades",
replace_existing=True,
)
# Signal aging (every 6 hours)
scheduler.add_job(
_expire_signals,
IntervalTrigger(hours=6),
id="expire-old-signals",
replace_existing=True,
)
# Win rate recalc (6h loop)
asyncio.ensure_future(_recalc_win_rates_loop())
# Start
scheduler.start()
logger.info("Scheduler started with all jobs")
# Keep alive forever
stop_event = asyncio.Event()
def _shutdown(sig, frame):
logger.info("Received signal %s — shutting down", sig)
stop_event.set()
signal.signal(signal.SIGTERM, _shutdown)
signal.signal(signal.SIGINT, _shutdown)
await stop_event.wait()
# Graceful shutdown
scheduler.shutdown(wait=False)
await _scheduler_engine.dispose()
logger.info("Scheduler shutdown complete")
if __name__ == "__main__":
asyncio.run(main())
+24
View File
@@ -0,0 +1,24 @@
from app.models.user import User
from app.models.exchange import Exchange
from app.models.symbol import Symbol
from app.models.candle import Candle
from app.models.watchlist import Watchlist
from app.models.credential import ExchangeCredential
from app.models.refresh_token import RefreshToken
from app.models.real_trade import RealTrade
from app.models.signal import Signal, HypotheticalTrade
from app.models.audit_log import AuditLog
__all__ = [
"User",
"Exchange",
"Symbol",
"Candle",
"Watchlist",
"ExchangeCredential",
"RefreshToken",
"Signal",
"HypotheticalTrade",
"RealTrade",
"AuditLog",
]
+71
View File
@@ -0,0 +1,71 @@
"""AlertCondition ORM model — multi-condition user alerts."""
from __future__ import annotations
from datetime import datetime
import uuid
from sqlalchemy import (
Boolean,
DateTime,
ForeignKey,
Index,
Integer,
JSON,
String,
)
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column
from app.database import Base
class AlertCondition(Base):
"""A user-defined multi-condition alert.
Each alert consists of a name, a list of conditions (stored as JSON),
and notification preferences. Alerts are evaluated after each signal
generation; all active alerts whose conditions are satisfied trigger
a notification.
"""
__tablename__ = "alert_conditions"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
user_id: Mapped[uuid.UUID] = mapped_column(
PGUUID(as_uuid=True),
ForeignKey("users.id"),
nullable=False,
comment="The user who owns this alert",
)
name: Mapped[str] = mapped_column(
String(100), nullable=False, comment="Human-readable alert name"
)
conditions: Mapped[list] = mapped_column(
JSON,
nullable=False,
default=list,
comment="Array of condition objects, e.g. [{'indicator': 'rsi', 'operator': '>', 'value': 70, 'timeframe': '1h'}]",
)
notify_platform: Mapped[str] = mapped_column(
String(20),
nullable=False,
default="telegram",
comment="Notification channel: 'telegram', 'discord', or 'both'",
)
is_active: Mapped[bool] = mapped_column(
Boolean, default=True, nullable=False
)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, default=datetime.now
)
updated_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, default=datetime.now, onupdate=datetime.now
)
__table_args__ = (
Index("ix_alert_conditions_user_active", "user_id", "is_active"),
)
def __repr__(self) -> str:
return f"<AlertCondition id={self.id} name={self.name!r} user_id={self.user_id}>"
+48
View File
@@ -0,0 +1,48 @@
"""Audit Log model for tracking actions across the system."""
from __future__ import annotations
import uuid
from datetime import datetime
from sqlalchemy import JSON, DateTime, Index, Integer, String
from sqlalchemy.dialects.postgresql import UUID
from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.sql import func
from app.database import Base
class AuditLog(Base):
__tablename__ = "audit_logs"
id: Mapped[int] = mapped_column(
Integer, primary_key=True, autoincrement=True
)
user_id: Mapped[uuid.UUID | None] = mapped_column(
UUID(as_uuid=True), nullable=True, index=True
)
action: Mapped[str] = mapped_column(
String(50), nullable=False
)
resource: Mapped[str] = mapped_column(
String(100), nullable=False
)
details: Mapped[dict | None] = mapped_column(
JSON, nullable=True
)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
server_default=func.now(),
nullable=False,
)
__table_args__ = (
Index("ix_audit_logs_action_created_at", "action", "created_at"),
)
def __repr__(self) -> str:
return (
f"<AuditLog id={self.id} action={self.action!r} "
f"resource={self.resource!r}>"
)
+71
View File
@@ -0,0 +1,71 @@
from __future__ import annotations
from datetime import datetime
from decimal import Decimal
from sqlalchemy import (
ForeignKey,
Index,
Integer,
Numeric,
PrimaryKeyConstraint,
String,
desc,
)
from sqlalchemy.dialects.postgresql import TIMESTAMP
from sqlalchemy.orm import Mapped, mapped_column
from app.database import Base
class Candle(Base):
__tablename__ = "candles"
symbol_id: Mapped[int] = mapped_column(
Integer, ForeignKey("symbols.id"), nullable=False
)
timeframe: Mapped[str] = mapped_column(
String(10), nullable=False
)
timestamp: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), nullable=False
)
open: Mapped[Decimal] = mapped_column(
Numeric(20, 8), nullable=False
)
high: Mapped[Decimal] = mapped_column(
Numeric(20, 8), nullable=False
)
low: Mapped[Decimal] = mapped_column(
Numeric(20, 8), nullable=False
)
close: Mapped[Decimal] = mapped_column(
Numeric(20, 8), nullable=False
)
volume: Mapped[Decimal] = mapped_column(
Numeric(30, 8), nullable=False
)
__table_args__ = (
PrimaryKeyConstraint(
"symbol_id", "timeframe", "timestamp",
name="pk_candles",
),
Index(
"ix_candles_symbol_timeframe_ts_desc",
"symbol_id",
"timeframe",
desc("timestamp"),
),
{
"postgresql_partition_by": "RANGE (timestamp)",
"info": {"partitioned": True},
},
)
def __repr__(self) -> str:
return (
f"<Candle symbol_id={self.symbol_id} "
f"timeframe={self.timeframe!r} "
f"timestamp={self.timestamp}>"
)
+73
View File
@@ -0,0 +1,73 @@
from __future__ import annotations
import uuid
from datetime import datetime
from sqlalchemy import Boolean, ForeignKey, Integer, String
from sqlalchemy.dialects.postgresql import UUID, TIMESTAMP
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.sql import func
from app.database import Base
class ExchangeCredential(Base):
__tablename__ = "exchange_credentials"
id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True),
primary_key=True,
default=func.gen_random_uuid(),
)
user_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True),
ForeignKey("users.id", ondelete="CASCADE"),
nullable=False,
)
exchange_id: Mapped[int] = mapped_column(
Integer,
ForeignKey("exchanges.id"),
nullable=False,
)
api_key: Mapped[str] = mapped_column(
String(255), nullable=False
)
api_secret_enc: Mapped[str] = mapped_column(
String(512), nullable=False
)
api_secret_iv: Mapped[str] = mapped_column(
String(64), nullable=False
)
passphrase: Mapped[str | None] = mapped_column(
String(255), nullable=True
)
is_testnet: Mapped[bool] = mapped_column(
Boolean, default=False, nullable=False
)
is_active: Mapped[bool] = mapped_column(
Boolean, default=True, nullable=False
)
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
)
updated_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True),
server_default=func.now(),
onupdate=func.now(),
nullable=False,
)
# Relationships
user: Mapped["User"] = relationship(
"User", back_populates="credentials"
)
exchange: Mapped["Exchange"] = relationship(
"Exchange", back_populates="credentials"
)
def __repr__(self) -> str:
return (
f"<ExchangeCredential id={self.id} "
f"user_id={self.user_id} "
f"exchange_id={self.exchange_id}>"
)
+52
View File
@@ -0,0 +1,52 @@
from __future__ import annotations
from datetime import datetime
from sqlalchemy import Boolean, Index, Integer, String
from sqlalchemy.dialects.postgresql import TIMESTAMP
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.sql import func
from app.database import Base
class Exchange(Base):
__tablename__ = "exchanges"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
name: Mapped[str] = mapped_column(
String(50), unique=True, nullable=False
)
display_name: Mapped[str | None] = mapped_column(
String(100), nullable=True
)
base_url: Mapped[str | None] = mapped_column(
String(255), nullable=True
)
ws_url: Mapped[str | None] = mapped_column(
String(255), nullable=True
)
is_active: Mapped[bool] = mapped_column(
Boolean, default=True, nullable=False
)
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
)
# P2-8: Index for active exchanges
__table_args__ = (
Index("ix_exchanges_is_active", "is_active"),
)
# Relationships
symbols: Mapped[list["Symbol"]] = relationship(
"Symbol", back_populates="exchange", cascade="all, delete-orphan"
)
credentials: Mapped[list["ExchangeCredential"]] = relationship(
"ExchangeCredential",
back_populates="exchange",
cascade="all, delete-orphan",
)
def __repr__(self) -> str:
return f"<Exchange id={self.id} name={self.name!r}>"
+95
View File
@@ -0,0 +1,95 @@
"""RealTrade ORM model for actual exchange trade tracking."""
from __future__ import annotations
import uuid
from datetime import datetime
from decimal import Decimal
from sqlalchemy import (
DateTime,
ForeignKey,
Index,
Integer,
Numeric,
String,
)
from sqlalchemy.dialects.postgresql import UUID as PG_UUID
from sqlalchemy.orm import Mapped, mapped_column
from app.database import Base
class RealTrade(Base):
"""A real trade placed on a connected exchange and persisted locally."""
__tablename__ = "real_trades"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
user_id: Mapped[uuid.UUID] = mapped_column(
PG_UUID(as_uuid=True), ForeignKey("users.id"), nullable=False,
comment="FK to the user who placed this trade",
)
exchange: Mapped[str] = mapped_column(String(20), nullable=False)
symbol: Mapped[str] = mapped_column(String(50), nullable=False)
side: Mapped[str] = mapped_column(
String(10), nullable=False, comment="buy or sell"
)
order_type: Mapped[str] = mapped_column(
String(10), nullable=False, default="market",
comment="market or limit",
)
# Order details
amount: Mapped[Decimal] = mapped_column(Numeric(20, 8), nullable=False)
price: Mapped[Decimal | None] = mapped_column(
Numeric(20, 8), nullable=True,
comment="Limit price (null for market orders)",
)
filled_amount: Mapped[Decimal] = mapped_column(
Numeric(20, 8), nullable=False, default=Decimal("0"),
)
status: Mapped[str] = mapped_column(
String(20), nullable=False, default="open",
comment="open / filled / cancelled / rejected",
)
# P&L (computed when closed)
pnl: Mapped[Decimal | None] = mapped_column(
Numeric(20, 8), nullable=True,
comment="Realised P&L in quote currency",
)
pnl_percent: Mapped[Decimal | None] = mapped_column(
Numeric(10, 4), nullable=True,
comment="P&L as percentage",
)
# Exchange reference
order_id: Mapped[str | None] = mapped_column(
String(100), nullable=True,
comment="Exchange-side order ID",
)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False,
default=datetime.utcnow,
)
closed_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True,
)
__table_args__ = (
Index("ix_real_trades_user_id", "user_id"),
Index("ix_real_trades_symbol", "symbol"),
Index("ix_real_trades_status", "status"),
Index("ix_real_trades_created_at", "created_at"),
)
def __repr__(self) -> str:
return (
f"<RealTrade #{self.id} {self.side.upper()} "
f"{self.amount} {self.symbol} on {self.exchange} "
f"status={self.status}>"
)
+56
View File
@@ -0,0 +1,56 @@
from __future__ import annotations
import uuid
from datetime import datetime
from sqlalchemy import Boolean, ForeignKey, Index, String
from sqlalchemy.dialects.postgresql import INET, UUID, TIMESTAMP
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.sql import func
from app.database import Base
class RefreshToken(Base):
__tablename__ = "refresh_tokens"
id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True),
primary_key=True,
default=func.gen_random_uuid(),
)
user_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True),
ForeignKey("users.id", ondelete="CASCADE"),
nullable=False,
)
token_hash: Mapped[str] = mapped_column(
String(64), nullable=False
)
expires_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), nullable=False
)
revoked: Mapped[bool] = mapped_column(
Boolean, default=False, nullable=False
)
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
)
user_agent: Mapped[str | None] = mapped_column(
String(255), nullable=True
)
ip_address: Mapped[str | None] = mapped_column(
INET, nullable=True
)
__table_args__ = (
Index("ix_refresh_tokens_user_id", "user_id"),
)
# Relationships
user: Mapped["User"] = relationship(
"User", back_populates="refresh_tokens"
)
def __repr__(self) -> str:
return f"<RefreshToken id={self.id} user_id={self.user_id}>"
+98
View File
@@ -0,0 +1,98 @@
"""Signal and HypotheticalTrade ORM models for trading signal detection."""
from __future__ import annotations
from datetime import datetime
from decimal import Decimal
import uuid
from sqlalchemy import (
DateTime,
ForeignKey,
Index,
Integer,
Numeric,
String,
Text,
)
from sqlalchemy.dialects.postgresql import UUID as PGUUID
from sqlalchemy.orm import Mapped, mapped_column
from app.database import Base
class Signal(Base):
"""A detected trading signal based on Double Bollinger Bands + RSI analysis."""
__tablename__ = "signals"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
symbol: Mapped[str] = mapped_column(String(50), nullable=False, index=True)
exchange: Mapped[str] = mapped_column(String(20), nullable=False, default="mexc")
timeframe: Mapped[str] = mapped_column(String(10), nullable=False)
signal_type: Mapped[str] = mapped_column(String(30), nullable=False)
strength: Mapped[str] = mapped_column(String(20), nullable=False)
price: Mapped[Decimal] = mapped_column(Numeric(20, 8), nullable=False)
timestamp: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False
)
indicators_snapshot: Mapped[str | None] = mapped_column(Text, nullable=True)
status: Mapped[str] = mapped_column(String(20), nullable=False, default="ACTIVE")
note: Mapped[str | None] = mapped_column(Text, nullable=True)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, default=datetime.now
)
__table_args__ = (
Index("ix_signals_symbol_created", "symbol", "created_at"),
)
class HypotheticalTrade(Base):
"""A hypothetical (paper) trade automatically entered on signal generation.
When a STRONG_BUY or BUY signal fires, a LONG trade is opened.
When a STRONG_SELL or SELL signal fires, a SHORT trade is opened.
Trades are closed when an opposing signal fires or when stopped out.
"""
__tablename__ = "hypothetical_trades"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
user_id: Mapped[uuid.UUID | None] = mapped_column(
PGUUID(as_uuid=True), ForeignKey("users.id"), nullable=True,
comment="The user who owns this trade",
)
signal_id: Mapped[int | None] = mapped_column(
Integer, ForeignKey("signals.id"), nullable=True,
comment="The signal that opened this trade"
)
symbol: Mapped[str] = mapped_column(String(50), nullable=False)
exchange: Mapped[str] = mapped_column(String(20), nullable=False, default="mexc")
timeframe: Mapped[str] = mapped_column(String(10), nullable=False)
direction: Mapped[str] = mapped_column(String(10), nullable=False)
entry_price: Mapped[Decimal] = mapped_column(Numeric(20, 8), nullable=False)
entry_time: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False
)
entry_reason: Mapped[str | None] = mapped_column(String(30), nullable=True)
exit_price: Mapped[Decimal | None] = mapped_column(Numeric(20, 8), nullable=True)
exit_time: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True
)
exit_reason: Mapped[str | None] = mapped_column(String(30), nullable=True)
quantity: Mapped[Decimal] = mapped_column(Numeric(20, 8), nullable=False)
pnl: Mapped[Decimal | None] = mapped_column(Numeric(20, 8), nullable=True)
pnl_percent: Mapped[Decimal | None] = mapped_column(Numeric(14, 4), nullable=True)
status: Mapped[str] = mapped_column(String(10), nullable=False, default="OPEN")
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True), nullable=False, default=datetime.now
)
closed_at: Mapped[datetime | None] = mapped_column(
DateTime(timezone=True), nullable=True
)
__table_args__ = (
Index("ix_hyp_trades_symbol", "symbol"),
Index("ix_hyp_trades_status", "status"),
)
+46
View File
@@ -0,0 +1,46 @@
from __future__ import annotations
from sqlalchemy import Boolean, ForeignKey, Index, Integer, String, UniqueConstraint
from sqlalchemy.orm import Mapped, mapped_column, relationship
from app.database import Base
class Symbol(Base):
__tablename__ = "symbols"
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
exchange_id: Mapped[int] = mapped_column(
Integer, ForeignKey("exchanges.id"), nullable=False
)
symbol: Mapped[str] = mapped_column(
String(50), nullable=False
)
base: Mapped[str] = mapped_column(
String(20), nullable=False
)
quote: Mapped[str] = mapped_column(
String(20), nullable=False
)
is_active: Mapped[bool] = mapped_column(
Boolean, default=True, nullable=False
)
is_trading: Mapped[bool] = mapped_column(
Boolean, default=False, nullable=False
)
__table_args__ = (
UniqueConstraint("exchange_id", "symbol", name="uq_symbol_exchange_symbol"),
Index("ix_symbols_is_active", "is_active"), # P2-8
)
# Relationships
exchange: Mapped["Exchange"] = relationship(
"Exchange", back_populates="symbols"
)
watchlists: Mapped[list["Watchlist"]] = relationship(
"Watchlist", back_populates="symbol", cascade="all, delete-orphan"
)
def __repr__(self) -> str:
return f"<Symbol id={self.id} symbol={self.symbol!r}>"
+72
View File
@@ -0,0 +1,72 @@
from __future__ import annotations
import uuid
from datetime import datetime
from sqlalchemy import Boolean, JSON, String, UniqueConstraint
from sqlalchemy.dialects.postgresql import UUID, TIMESTAMP
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.sql import func
from app.database import Base
class User(Base):
__tablename__ = "users"
id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True),
primary_key=True,
default=func.gen_random_uuid(),
)
username: Mapped[str] = mapped_column(
String(50), unique=True, nullable=False
)
email: Mapped[str] = mapped_column(
String(255), unique=True, nullable=False
)
password_hash: Mapped[str] = mapped_column(
String(255), nullable=False
)
display_name: Mapped[str | None] = mapped_column(
String(100), nullable=True
)
is_active: Mapped[bool] = mapped_column(
Boolean, default=True, nullable=False
)
is_admin: Mapped[bool] = mapped_column(
Boolean, default=False, nullable=False
)
role: Mapped[str] = mapped_column(
String(20), default="trader", nullable=False
)
preferences: Mapped[dict | None] = mapped_column(
JSON, default={"default_exchange": "mexc", "default_timeframe": "1h"}, nullable=True
)
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
)
updated_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True),
server_default=func.now(),
onupdate=func.now(),
nullable=False,
)
# Relationships
watchlists: Mapped[list["Watchlist"]] = relationship(
"Watchlist", back_populates="user", cascade="all, delete-orphan"
)
credentials: Mapped[list["ExchangeCredential"]] = relationship(
"ExchangeCredential",
back_populates="user",
cascade="all, delete-orphan",
)
refresh_tokens: Mapped[list["RefreshToken"]] = relationship(
"RefreshToken",
back_populates="user",
cascade="all, delete-orphan",
)
def __repr__(self) -> str:
return f"<User id={self.id} username={self.username!r}>"
+61
View File
@@ -0,0 +1,61 @@
from __future__ import annotations
import uuid
from datetime import datetime
from sqlalchemy import ForeignKey, Integer, String, UniqueConstraint
from sqlalchemy.dialects.postgresql import UUID, TIMESTAMP
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.sql import func
from app.database import Base
class Watchlist(Base):
__tablename__ = "user_watchlists"
id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True),
primary_key=True,
default=func.gen_random_uuid(),
)
user_id: Mapped[uuid.UUID] = mapped_column(
UUID(as_uuid=True),
ForeignKey("users.id", ondelete="CASCADE"),
nullable=False,
)
symbol_id: Mapped[int] = mapped_column(
Integer,
ForeignKey("symbols.id", ondelete="CASCADE"),
nullable=False,
)
label: Mapped[str | None] = mapped_column(
String(50), nullable=True
)
sort_order: Mapped[int] = mapped_column(
Integer, default=0, nullable=False
)
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
)
__table_args__ = (
UniqueConstraint(
"user_id", "symbol_id", name="uq_watchlist_user_symbol"
),
)
# Relationships
user: Mapped["User"] = relationship(
"User", back_populates="watchlists"
)
symbol: Mapped["Symbol"] = relationship(
"Symbol", back_populates="watchlists"
)
def __repr__(self) -> str:
return (
f"<Watchlist id={self.id} "
f"user_id={self.user_id} "
f"symbol_id={self.symbol_id}>"
)
+106
View File
@@ -0,0 +1,106 @@
from app.schemas.auth import (
RegisterRequest,
LoginRequest,
TokenResponse,
RefreshRequest,
UserResponse,
UserSessionResponse,
LogoutRequest,
ChangePasswordRequest,
)
from app.schemas.user import (
UserUpdateRequest,
AdminCreateUserRequest,
AdminResetPasswordRequest,
AdminUserUpdateRequest,
)
from app.schemas.candle import (
CandleResponse,
CandleListResponse,
IndicatorResponse,
)
from app.schemas.symbol import (
SymbolResponse,
SymbolSearchResponse,
WatchlistResponse,
WatchlistCreateRequest,
WatchlistUpdateRequest,
)
from app.schemas.ws_message import (
WSSubscription,
WSCandleUpdate,
WSTickerUpdate,
WSConnectionStatus,
)
from app.schemas.credential import (
CredentialResponse,
CredentialCreateRequest,
CredentialUpdateRequest,
)
from app.schemas.exchange import (
ExchangeResponse,
ExchangeCreateRequest,
ExchangeUpdateRequest,
)
from app.schemas.health import (
HealthResponse,
DetailedHealthResponse,
DbHealth,
ExchangeHealth,
)
from app.schemas.strategy import (
StrategyEntry,
StrategyListResponse,
StrategyConfigRequest,
STRATEGY_NAMES,
STRATEGY_DISPLAY,
)
__all__ = [
# auth
"RegisterRequest",
"LoginRequest",
"TokenResponse",
"RefreshRequest",
"UserResponse",
"UserSessionResponse",
"LogoutRequest",
"ChangePasswordRequest",
# user
"UserUpdateRequest",
"AdminUserUpdateRequest",
# candle
"CandleResponse",
"CandleListResponse",
"IndicatorResponse",
# symbol
"SymbolResponse",
"SymbolSearchResponse",
"WatchlistResponse",
"WatchlistCreateRequest",
"WatchlistUpdateRequest",
# ws_message
"WSSubscription",
"WSCandleUpdate",
"WSTickerUpdate",
"WSConnectionStatus",
# credential
"CredentialResponse",
"CredentialCreateRequest",
"CredentialUpdateRequest",
# exchange
"ExchangeResponse",
"ExchangeCreateRequest",
"ExchangeUpdateRequest",
# health
"HealthResponse",
"DetailedHealthResponse",
"DbHealth",
"ExchangeHealth",
# strategy
"StrategyEntry",
"StrategyListResponse",
"StrategyConfigRequest",
"STRATEGY_NAMES",
"STRATEGY_DISPLAY",
]
+69
View File
@@ -0,0 +1,69 @@
from datetime import datetime
from typing import Optional
from uuid import UUID
from pydantic import BaseModel, field_validator
class RegisterRequest(BaseModel):
username: str
email: str
password: str
@field_validator("password")
@classmethod
def password_min_length(cls, v: str) -> str:
if len(v) < 8:
raise ValueError("password must be at least 8 characters")
return v
class LoginRequest(BaseModel):
username: str
password: str
class TokenResponse(BaseModel):
access_token: str
refresh_token: str
token_type: str = "bearer"
class RefreshRequest(BaseModel):
refresh_token: str
class UserResponse(BaseModel):
id: UUID
username: str
email: str
display_name: Optional[str] = None
is_active: bool = True
is_admin: bool = False
role: str = "trader"
preferences: Optional[dict] = None
created_at: datetime
class UserSessionResponse(BaseModel):
id: UUID
created_at: datetime
user_agent: Optional[str] = None
ip_address: Optional[str] = None
is_current: bool
class LogoutRequest(BaseModel):
refresh_token: str
class ChangePasswordRequest(BaseModel):
old_password: str
new_password: str
@field_validator("new_password")
@classmethod
def new_password_min_length(cls, v: str) -> str:
if len(v) < 8:
raise ValueError("new_password must be at least 8 characters")
return v
+29
View File
@@ -0,0 +1,29 @@
from datetime import datetime
from decimal import Decimal
from typing import Optional
from pydantic import BaseModel
class CandleResponse(BaseModel):
symbol_id: int
timeframe: str
timestamp: datetime
open: float
high: float
low: float
close: float
volume: float
class CandleListResponse(BaseModel):
candles: list[CandleResponse]
cursor: Optional[str] = None
has_more: bool
class IndicatorResponse(BaseModel):
symbol_id: int
timeframe: str
timestamp: datetime
indicators: dict
+34
View File
@@ -0,0 +1,34 @@
from typing import Optional
from uuid import UUID
from pydantic import BaseModel, field_validator
class CredentialResponse(BaseModel):
id: UUID
exchange_id: int
exchange_name: str
api_key: str
is_testnet: bool
is_active: bool
@field_validator("api_key")
@classmethod
def mask_api_key(cls, v: str) -> str:
"""Show full API key so users can copy it (middle chars left visible for reference)."""
return v
class CredentialCreateRequest(BaseModel):
exchange_id: int
api_key: str
api_secret: str
passphrase: Optional[str] = None
is_testnet: bool = False
class CredentialUpdateRequest(BaseModel):
api_key: Optional[str] = None
api_secret: Optional[str] = None
passphrase: Optional[str] = None
is_active: Optional[bool] = None
+24
View File
@@ -0,0 +1,24 @@
from typing import Optional
from pydantic import BaseModel
class ExchangeResponse(BaseModel):
id: int
name: str
display_name: str
is_active: bool = True
class ExchangeCreateRequest(BaseModel):
name: str
display_name: str
base_url: str
ws_url: str
class ExchangeUpdateRequest(BaseModel):
display_name: Optional[str] = None
base_url: Optional[str] = None
ws_url: Optional[str] = None
is_active: Optional[bool] = None
+31
View File
@@ -0,0 +1,31 @@
from datetime import datetime
from typing import Optional
from pydantic import BaseModel
class DbHealth(BaseModel):
connected: bool
latency_ms: float
pool_size: int
class ExchangeHealth(BaseModel):
exchange: str
connected: bool
last_sync: Optional[datetime] = None
class HealthResponse(BaseModel):
status: str
version: str
uptime: float
db_connected: bool
class DetailedHealthResponse(BaseModel):
status: str
version: str
uptime: float
db: DbHealth
exchange_connections: list[ExchangeHealth]
+78
View File
@@ -0,0 +1,78 @@
"""Pydantic schemas for real trade responses and requests."""
from __future__ import annotations
from datetime import datetime
from decimal import Decimal
from typing import Optional
from uuid import UUID
from pydantic import BaseModel, field_validator
class RealTradeResponse(BaseModel):
"""Single real trade record returned to the frontend."""
id: int
user_id: str
exchange: str
symbol: str
side: str
order_type: str
amount: float
price: Optional[float] = None
filled_amount: float
status: str
pnl: Optional[float] = None
pnl_percent: Optional[float] = None
order_id: Optional[str] = None
created_at: datetime
closed_at: Optional[datetime] = None
model_config = {"from_attributes": True}
@field_validator("user_id", mode="before")
@classmethod
def _coerce_user_id(cls, v: object) -> str:
if isinstance(v, UUID):
return str(v)
return str(v)
class RealTradeListResponse(BaseModel):
"""Paginated list of real trades with aggregate stats."""
trades: list[RealTradeResponse]
total: int
total_pnl: Optional[float] = None
win_rate: Optional[float] = None
class RealTradeCreateRequest(BaseModel):
"""Payload to create a real trade record (used after placing an order)."""
exchange: str
symbol: str
side: str
order_type: str = "market"
amount: float
price: Optional[float] = None
filled_amount: float = 0
status: str = "open"
order_id: Optional[str] = None
class WinRatePeriod(BaseModel):
"""Win-rate data for a single period (daily / weekly / monthly)."""
trades: int
wins: int
win_rate: float
class WinRateResponse(BaseModel):
"""Win-rate breakdown across daily, weekly, and monthly periods."""
daily: WinRatePeriod
weekly: WinRatePeriod
monthly: WinRatePeriod
+89
View File
@@ -0,0 +1,89 @@
"""Pydantic schemas for signal and trade responses."""
from __future__ import annotations
from datetime import datetime
from decimal import Decimal
from typing import Optional
from pydantic import BaseModel, model_validator
import json
class SignalResponse(BaseModel):
id: int
symbol: str
exchange: str
timeframe: str
signal_type: str
strength: str
price: float
timestamp: datetime
indicators_snapshot: Optional[str] = None
confidence: Optional[float] = None
status: str
note: Optional[str] = None
created_at: datetime
model_config = {"from_attributes": True}
@model_validator(mode="after")
def _extract_confidence(self) -> "SignalResponse":
if self.confidence is None and self.indicators_snapshot:
try:
snap = json.loads(self.indicators_snapshot)
self.confidence = snap.get("confidence")
except (json.JSONDecodeError, TypeError, AttributeError):
pass
return self
class TradeResponse(BaseModel):
id: int
signal_id: Optional[int] = None
symbol: str
exchange: str
timeframe: str
direction: str
entry_price: float
entry_time: datetime
entry_reason: Optional[str] = None
exit_price: Optional[float] = None
exit_time: Optional[datetime] = None
exit_reason: Optional[str] = None
quantity: float
pnl: Optional[float] = None
pnl_percent: Optional[float] = None
status: str
created_at: datetime
closed_at: Optional[datetime] = None
model_config = {"from_attributes": True}
class SignalListResponse(BaseModel):
signals: list[SignalResponse]
total: int
class TradeListResponse(BaseModel):
trades: list[TradeResponse]
total: int
total_pnl: Optional[float] = None
win_rate: Optional[float] = None
class ReviewResponse(BaseModel):
period: str # "weekly" or "monthly"
start_date: str
end_date: str
total_signals: int
total_trades: int
wins: int
losses: int
win_rate: float
total_pnl: float
best_trade: Optional[TradeResponse] = None
worst_trade: Optional[TradeResponse] = None
signals_by_type: dict[str, int]
symbol_performance: list[dict]
+66
View File
@@ -0,0 +1,66 @@
"""Pydantic schemas for per-user strategy configuration."""
from __future__ import annotations
from typing import Optional
from pydantic import BaseModel
# Available strategy names — full list of 13 voting algorithms
STRATEGY_NAMES = [
"double_bb_rsi",
"macd_crossover",
"supertrend",
"volume_breakout",
"ichimoku_cloud",
"divergence",
"smc",
"mtf",
"obv",
"stoch_rsi",
"mfi",
"fvg",
"candlestick",
]
STRATEGY_DISPLAY: dict[str, str] = {
"double_bb_rsi": "Double BB + RSI",
"macd_crossover": "MACD Crossover",
"supertrend": "SuperTrend",
"volume_breakout": "Volume Breakout",
"ichimoku_cloud": "Ichimoku Cloud",
"divergence": "Divergence",
"smc": "Market Structure (SMC)",
"mtf": "Multi-Timeframe",
"obv": "OBV Crossover",
"stoch_rsi": "Stochastic RSI",
"mfi": "Money Flow Index",
"fvg": "Fair Value Gap",
"candlestick": "Candlestick Patterns",
}
class StrategyEntry(BaseModel):
"""A single strategy with its enabled status."""
name: str
display_name: str
enabled: bool = False
class StrategyListResponse(BaseModel):
"""Response for GET /api/v1/strategies."""
strategies: list[StrategyEntry]
thresholds: dict[str, float] = {}
class StrategyConfigRequest(BaseModel):
"""Request body for PUT /api/v1/strategies.
Only provided fields will be updated.
"""
enabled_strategies: Optional[list[str]] = None
thresholds: Optional[dict[str, float]] = None
+37
View File
@@ -0,0 +1,37 @@
from typing import Optional
from uuid import UUID
from pydantic import BaseModel
class SymbolResponse(BaseModel):
id: int
exchange_id: int
symbol: str
base: str
quote: str
is_active: bool = True
class SymbolSearchResponse(BaseModel):
symbols: list[SymbolResponse]
class WatchlistResponse(BaseModel):
id: UUID
symbol_id: int
symbol: str
exchange: str
label: Optional[str] = None
sort_order: int
class WatchlistCreateRequest(BaseModel):
symbol_id: int
label: Optional[str] = None
sort_order: Optional[int] = None
class WatchlistUpdateRequest(BaseModel):
label: Optional[str] = None
sort_order: Optional[int] = None
+63
View File
@@ -0,0 +1,63 @@
from typing import Optional
import re
from pydantic import BaseModel, field_validator
class UserUpdateRequest(BaseModel):
display_name: Optional[str] = None
email: Optional[str] = None
preferences: Optional[dict] = None
_EMAIL_RE = re.compile(r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$")
def _validate_email(email: str) -> str:
if not _EMAIL_RE.match(email):
raise ValueError("invalid email format")
return email
def _validate_strong_password(password: str) -> str:
if len(password) < 8:
raise ValueError("password must be at least 8 characters")
if not re.search(r"\d", password):
raise ValueError("password must contain at least 1 number")
if not re.search(r"[^a-zA-Z0-9]", password):
raise ValueError("password must contain at least 1 special character")
return password
class AdminUserUpdateRequest(BaseModel):
is_active: Optional[bool] = None
is_admin: Optional[bool] = None
role: Optional[str] = None
email: Optional[str] = None
display_name: Optional[str] = None
class AdminCreateUserRequest(BaseModel):
username: str
email: str
password: str
display_name: Optional[str] = None
is_admin: bool = False
role: str = "trader"
@field_validator("email")
@classmethod
def email_must_be_valid(cls, v: str) -> str:
return _validate_email(v)
@field_validator("password")
@classmethod
def password_must_be_strong(cls, v: str) -> str:
return _validate_strong_password(v)
class AdminResetPasswordRequest(BaseModel):
new_password: str
@field_validator("new_password")
@classmethod
def new_password_must_be_strong(cls, v: str) -> str:
return _validate_strong_password(v)
+31
View File
@@ -0,0 +1,31 @@
from decimal import Decimal
from typing import Literal
from pydantic import BaseModel
from app.schemas.candle import CandleResponse
class WSSubscription(BaseModel):
symbol: str
timeframe: str
exchange: str
action: Literal["subscribe", "unsubscribe"]
class WSCandleUpdate(BaseModel):
type: Literal["candle"] = "candle"
data: CandleResponse
class WSTickerUpdate(BaseModel):
type: Literal["ticker"] = "ticker"
symbol: str
price: float
change_24h: float
volume: float
class WSConnectionStatus(BaseModel):
type: Literal["connection"] = "connection"
status: Literal["connected", "reconnecting", "disconnected"]
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")
View File
+398
View File
@@ -0,0 +1,398 @@
"""APScheduler background task that periodically fetches recent candles."""
from __future__ import annotations
import asyncio
import logging
from collections.abc import Awaitable, Callable
from datetime import datetime, timezone
from typing import Any
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from fastapi import FastAPI
from sqlalchemy import and_, select
from sqlalchemy.exc import InterfaceError
from sqlalchemy.orm import joinedload, selectinload
from sqlalchemy.ext.asyncio import AsyncSession
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.services.candle_service import candle_cache
logger = logging.getLogger(__name__)
# ── Rate-limit CCXT API calls (max 250 concurrent fetches) ──
_FETCH_SEMAPHORE = asyncio.Semaphore(250)
# ---------------------------------------------------------------------------
# After-fetch callbacks
# ---------------------------------------------------------------------------
# External services (e.g. WebSocket push) can register callbacks that are
# invoked after candles are successfully persisted to the database.
# Signature: async callback(exchange: str, symbol: str, timeframe: str, candle_data: dict)
_after_fetch_callbacks: list[
Callable[[str, str, str, dict[str, Any]], Awaitable[None]]
] = []
def register_after_fetch_callback(
callback: Callable[[str, str, str, dict[str, Any]], Awaitable[None]],
) -> None:
"""Register an async callback invoked after candles are saved.
The callback receives ``(exchange, symbol, timeframe, candle_data_dict)``
for every candle that was newly inserted or already existed (ON CONFLICT
DO NOTHING). Multiple callbacks are supported.
"""
_after_fetch_callbacks.append(callback)
logger.debug("Registered after-fetch callback: %s", callback.__name__)
# ---------------------------------------------------------------------------
# Timeframes we care about
# ---------------------------------------------------------------------------
_TIMEFRAMES_1M: list[str] = ["1m"]
_TIMEFRAMES_5M_PLUS: list[str] = ["5m", "15m", "30m", "1h", "4h", "1d", "1w", "1M"]
_TIMEFRAMES_OPTIMIZED: list[str] = ["15m", "30m", "1h", "4h", "1d", "1w", "1M"]
# Fetch top 100 trading symbols across ALL exchanges (is_trading flag)
# ~476 total symbols (100 bases × ~5 exchanges), batch=25 → ~100 min/cycle
# 7 timeframes × 25 symbols = 175 API calls/batch → well under semaphore=250
# Cache key prefix: "{exchange_name}:{symbol}:" — we append the timeframe later
def _cache_key_prefix(exchange_name: str, symbol: str) -> str:
return f"{exchange_name}:{symbol}:"
def _invalidate_candle_cache(exchange_name: str, symbol: str) -> None:
"""Remove all cached candle entries for a given exchange+symbol pair.
Builds exact keys for all known timeframes instead of scanning the
entire cache (which was O(cache_size × symbols) and CPU-heavy).
"""
prefix = _cache_key_prefix(exchange_name, symbol)
for tf in _TIMEFRAMES_1M + _TIMEFRAMES_5M_PLUS:
key = prefix + tf
candle_cache.pop(key, None)
# ====================================================================
# Core fetch function
# ====================================================================
async def fetch_recent_candles(
app: FastAPI,
fetch_limit: int = 2,
timeframes: list[str] | None = None,
max_symbols: int = 25, # 25 symbols/batch × 4TF = 100 API calls per 5-min tick
) -> None:
"""Fetch the latest candle(s) for active trading symbols from the DB.
Processes a batch of up to ``max_symbols`` per call, cycling through
symbols alphabetically so all trading pairs are eventually covered.
Uses is_trading=true flag — top 100 bases × ~5 exchanges ≈ 476 symbols.
Full cycle: ~100 minutes at 5-min interval.
"""
if timeframes is None:
timeframes = _TIMEFRAMES_OPTIMIZED # 15m, 1h, 4h, 1d
# Persist an offset counter via module-level list (mutable singleton)
# so the next call picks up where the last one left off.
if not hasattr(fetch_recent_candles, "_offset"):
fetch_recent_candles._offset = 0
# ── Step 1: Query symbols in a SHORT-lived session ──
# We must NOT hold the session during CCXT API calls (30-60s)
# because idle_in_transaction_session_timeout=60s kills idle connections.
async with async_session_factory() as query_db:
try:
query = (
select(Symbol)
.options(joinedload(Symbol.exchange))
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(
and_(
Symbol.is_trading == True, # noqa: E712
Symbol.is_active == True, # noqa: E712
Exchange.is_active == True, # noqa: E712
)
)
.order_by(Symbol.symbol)
)
result = await query_db.execute(query)
all_symbols: list[Symbol] = list(result.scalars().all())
except InterfaceError:
logger.warning("Symbol query failed — connection closed, skipping batch")
return
if not all_symbols:
logger.debug("No active trading symbols found — skipping candle fetch")
return
# Slice the batch using a rolling offset
total = len(all_symbols)
offset = fetch_recent_candles._offset
batch = all_symbols[offset:offset + max_symbols]
# Update / wrap the offset
fetch_recent_candles._offset = (offset + max_symbols) % total
logger.info(
"Fetching candles for %d/%d symbols (offset=%d, batch=%d-%d)",
len(batch), total, offset, offset + 1, offset + len(batch),
)
# ── Step 2: Fetch candles from CCXT WITHOUT holding DB session ──
exchange_map: dict[str, list[Symbol]] = {}
for sym in batch:
exchange_name = sym.exchange.name
exchange_map.setdefault(exchange_name, []).append(sym)
all_candle_values: list[dict[str, Any]] = []
new_candle_events: list[tuple[str, str, str, dict[str, Any]]] = []
for exchange_name, sym_list in exchange_map.items():
try:
adapter = exchange_factory.create(exchange_name)
except ValueError:
logger.warning("Unknown exchange %s — skipping", exchange_name)
continue
async def _fetch_one(sym: Symbol, tf: str):
async with _FETCH_SEMAPHORE:
try:
candles = await adapter.fetch_ohlcv(
symbol=sym.symbol,
timeframe=tf,
limit=fetch_limit,
)
return sym, tf, candles
except asyncio.CancelledError:
raise
except Exception as e:
err_str = str(e)
if "does not have market symbol" in err_str or "BadSymbol" in err_str:
logger.warning(
"Symbol %s not found on %s — marking inactive",
sym.symbol, exchange_name,
)
sym.is_active = False
else:
logger.exception(
"Failed to fetch %s %s on %s",
tf, sym.symbol, exchange_name,
)
return sym, tf, []
tasks = [_fetch_one(sym, tf) for sym in sym_list for tf in timeframes]
results = await asyncio.gather(*tasks, return_exceptions=True)
for result in results:
if isinstance(result, BaseException):
continue
sym, tf, candles = result
for c in candles:
all_candle_values.append(
{
"symbol_id": sym.id,
"timeframe": c.timeframe,
"timestamp": c.timestamp,
"open": c.open,
"high": c.high,
"low": c.low,
"close": c.close,
"volume": c.volume,
}
)
new_candle_events.append(
(
exchange_name,
sym.symbol,
c.timeframe,
{
"symbol": sym.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,
},
)
)
await asyncio.sleep(0.05)
# ── Step 3: Save candles in a FRESH, short-lived DB session ──
if all_candle_values:
async with async_session_factory() as save_db:
try:
from sqlalchemy.dialects.postgresql import insert as pg_insert
stmt = pg_insert(Candle).values(all_candle_values)
stmt = stmt.on_conflict_do_nothing(
index_elements=["symbol_id", "timeframe", "timestamp"]
)
await save_db.execute(stmt)
await save_db.commit()
logger.info(
"Fetched and stored %d recent candles across %d symbols (timeframes=%s)",
len(all_candle_values),
len(batch),
timeframes,
)
except InterfaceError:
logger.warning(
"Candle save: InterfaceError — connection already closed, skipping. "
"Data will be re-fetched next cycle."
)
except Exception:
logger.exception("Candle save: DB step failed")
try:
await save_db.rollback()
except InterfaceError:
pass
else:
logger.debug("No candle data fetched for this batch")
# 🔑 Session closed here — connection released back to pool!
# Non-DB operations below run without holding a pool connection.
if new_candle_events:
# --- Invoke after-fetch callbacks ---
if _after_fetch_callbacks:
for exchange_name, symbol_str, tf, candle_dict in new_candle_events:
for cb in _after_fetch_callbacks:
try:
await cb(exchange_name, symbol_str, tf, candle_dict)
except Exception:
logger.exception(
"After-fetch callback %s failed for %s:%s:%s",
cb.__name__,
exchange_name,
symbol_str,
tf,
)
# --- Batch signal analysis: one analysis per unique (exchange, symbol) ---
# Only analyze the best timeframe for trading (1h) to avoid thrashing
# when multiple TFs of the same symbol run concurrently.
processed_pairs: set[tuple[str, str]] = set()
for exchange_name, symbol_str, tf, _ in new_candle_events:
if tf != "1h":
continue # only 1h triggers trade signals
pair = (exchange_name, symbol_str)
if pair in processed_pairs:
continue
processed_pairs.add(pair)
if processed_pairs:
from app.services.signal_service import analyse_and_generate_signals
# 🔧 Optimized: semaphore 3 (was 8) — lower concurrency = lower CPU
# spike + less trade open/evict thrashing from concurrent symbol analysis.
_SIGNAL_SEMAPHORE = asyncio.Semaphore(3)
async def _analyse_one(ex_name: str, sym: str) -> None:
async with _SIGNAL_SEMAPHORE:
try:
await analyse_and_generate_signals(ex_name, sym, "1h")
except Exception:
logger.exception(
"Batch signal analysis failed for %s:%s",
ex_name, sym,
)
await asyncio.gather(
*(_analyse_one(*pair) for pair in processed_pairs),
return_exceptions=True,
)
logger.debug(
"Batch signal analysis: %d unique pairs processed",
len(processed_pairs),
)
# --- Invalidate cache ---
for sym in batch:
_invalidate_candle_cache(sym.exchange.name, sym.symbol)
# ====================================================================
# Scheduler setup
# ====================================================================
def setup_candle_scheduler(app: FastAPI) -> AsyncIOScheduler:
"""Create and configure an APScheduler ``AsyncIOScheduler``.
Jobs added:
- **1m timeframes**: ``fetch_recent_candles`` every 60 seconds.
- **5m+ timeframes**: ``fetch_recent_candles`` every 5 minutes.
The scheduler is started when the FastAPI application starts and
shut down when it stops (via the *lifespan* context manager).
"""
scheduler = AsyncIOScheduler()
# ─────────────────────────────────────────────────────────────────────
# DISABLED: 1m candles — user agreed not to fetch 1m (too heavy on DB)
# Kept as commented code for future reference.
# ─────────────────────────────────────────────────────────────────────
# scheduler.add_job(
# fetch_recent_candles,
# trigger="interval",
# seconds=120,
# args=[app, 2, _TIMEFRAMES_1M, 100],
# id="fetch_candles_1m",
# replace_existing=True,
# coalesce=True,
# max_instances=1,
# misfire_grace_time=300,
# name="Fetch 1m candles",
# )
# Every 5 minutes — 7 timeframes {15m,30m,1h,4h,1d,1w,1M}, 25 symbols/batch
# ~476 trading symbols (100 bases × ~5 exchanges) ÷ 25/batch × 5 min = ~100 min full cycle
# 7 TFs × 25 symbols = 175 API calls/batch → ~35 calls/min average (well under 250 semaphore)
scheduler.add_job(
fetch_recent_candles,
trigger="interval",
seconds=300,
args=[app, 2, _TIMEFRAMES_OPTIMIZED, 25],
id="fetch_candles_optimized",
replace_existing=True,
coalesce=True,
max_instances=1,
misfire_grace_time=600,
name="Fetch trading candles (top 100 bases, 5 exchanges, 7TFs: 15m,30m,1h,4h,1d,1w,1M)",
)
logger.info("Candle scheduler configured: top 100 bases, 5 exchanges, 7TFs, 25/batch, 5min")
# NOTE: Scheduler lifecycle is managed by main.py's lifespan handler.
# The deprecated @app.on_event() decorators do NOT fire when
# lifespan= is used in the FastAPI constructor, so we removed them.
# Call scheduler.start() in your lifespan startup block instead.
return scheduler
async def force_full_sync(app: FastAPI) -> None:
"""Run a one-time full historical candle sync on startup.
Fetches up to 500 candles per symbol/timeframe to backfill
missing data after an outage or initial deployment.
"""
logger.info("Starting one-time full candle sync (fetch_limit=500)...")
try:
await fetch_recent_candles(app, fetch_limit=500)
logger.info("One-time full candle sync completed")
except Exception:
logger.exception("One-time full candle sync failed (non-fatal)")
+126
View File
@@ -0,0 +1,126 @@
"""Symbol synchronisation tasks for keeping exchange market data up-to-date."""
from __future__ import annotations
import logging
from fastapi import FastAPI
from sqlalchemy import select
from sqlalchemy.dialects.postgresql import insert as pg_insert
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload
from app.database import async_session_factory
from app.exchange.factory import factory as exchange_factory
from app.models.exchange import Exchange
from app.models.symbol import Symbol
logger = logging.getLogger(__name__)
async def sync_exchange_symbols(
db: AsyncSession,
exchange_name: str,
) -> int:
"""Fetch all trading symbols from an exchange and upsert them into the DB.
Steps:
1. Look up the ``Exchange`` record by *exchange_name*.
2. Create an exchange adapter via ``ExchangeFactory``.
3. Call ``adapter.fetch_symbols()`` to retrieve all available symbols.
4. Upsert each symbol into the ``symbols`` table.
5. Return the total number of symbols upserted.
Raises
------
ValueError
If the exchange is unknown to the factory.
"""
# --- 1. Get exchange ---
result = await db.execute(
select(Exchange).where(Exchange.name == exchange_name)
)
exchange = result.scalar_one_or_none()
if exchange is None:
raise ValueError(
f"Exchange {exchange_name!r} not found in the database. "
"Create an Exchange record first."
)
# --- 2. Create adapter ---
adapter = exchange_factory.create(exchange_name)
# --- 3. Fetch symbols ---
symbols = await adapter.fetch_symbols()
if not symbols:
logger.warning("No symbols returned from %s", exchange_name)
return 0
# --- 4. Upsert symbols ---
values = [
{
"exchange_id": exchange.id,
"symbol": sym.symbol,
"base": sym.base,
"quote": sym.quote,
"is_active": sym.is_active,
}
for sym in symbols
]
stmt = pg_insert(Symbol).values(values)
stmt = stmt.on_conflict_do_update(
index_elements=["exchange_id", "symbol"],
set_={
"base": stmt.excluded.base,
"quote": stmt.excluded.quote,
"is_active": stmt.excluded.is_active,
},
)
await db.execute(stmt)
await db.commit()
logger.info(
"Synced %d symbols from %s",
len(values),
exchange_name,
)
return len(values)
async def sync_all_exchanges(app: FastAPI) -> dict[str, int]:
"""Sync symbols for every active exchange registered in the database.
Returns a dictionary mapping ``exchange_name`` to the number of
symbols synced.
"""
results: dict[str, int] = {}
async with async_session_factory() as db:
try:
exch_result = await db.execute(
select(Exchange).where(Exchange.is_active == True) # noqa: E712
)
exchanges: list[Exchange] = list(exch_result.scalars().all())
if not exchanges:
logger.info("No active exchanges found — nothing to sync")
return results
for exchange in exchanges:
try:
count = await sync_exchange_symbols(db, exchange.name)
results[exchange.name] = count
except Exception:
logger.exception(
"Failed to sync symbols for %s",
exchange.name,
)
results[exchange.name] = -1
return results
except Exception:
logger.exception("sync_all_exchanges failed")
await db.rollback()
return results
+163
View File
@@ -0,0 +1,163 @@
"""Stale data detection for ingested candle data."""
from __future__ import annotations
import logging
from datetime import datetime, timezone
from fastapi import FastAPI
from sqlalchemy import and_, func, select
from sqlalchemy.ext.asyncio import AsyncSession
from app.database import async_session_factory
from app.models.candle import Candle
from app.models.exchange import Exchange
from app.models.symbol import Symbol
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Timeframe string → seconds mapping
# ---------------------------------------------------------------------------
_TIMEFRAME_SECONDS: dict[str, int] = {
"1m": 60,
"3m": 180,
"5m": 300,
"15m": 900,
"30m": 1800,
"1h": 3600,
"2h": 7200,
"4h": 14400,
"6h": 21600,
"8h": 28800,
"12h": 43200,
"1d": 86400,
"3d": 259200,
"1w": 604800,
"1M": 2_592_000,
}
def _timeframe_to_seconds(tf: str) -> int:
"""Convert a timeframe string to seconds.
Falls back to 3600 (1 hour) for unknown timeframes.
"""
return _TIMEFRAME_SECONDS.get(tf, 3600)
# ====================================================================
# Public API
# ====================================================================
async def check_stale_candles(
db: AsyncSession,
app: FastAPI, # noqa: ARG001 — kept for consistent signature with other tasks
) -> list[dict]:
"""Scan all active symbols and detect stale candle data.
A candle is considered stale when the timestamp of the latest candle
plus *twice* the candle timeframe duration is still in the past
(i.e. ``latest_timestamp + 2 * timeframe_seconds < now``).
Returns
-------
list[dict]
Each entry contains:
- ``symbol`` — the trading pair (e.g. ``"BTC/USDT"``)
- ``exchange`` — the exchange name
- ``timeframe`` — the candle interval
- ``last_timestamp`` — the most recent candle's timestamp (ISO-8601)
- ``staleness_minutes`` — how many minutes behind expected
"""
stale_entries: list[dict] = []
try:
# --- Get all active symbols with their exchange info ---
result = await db.execute(
select(Symbol)
.join(Exchange, Exchange.id == Symbol.exchange_id)
.where(
and_(
Symbol.is_active == True, # noqa: E712
Exchange.is_active == True, # noqa: E712
)
)
)
symbols: list[Symbol] = list(result.scalars().all())
if not symbols:
logger.debug("No active symbols found — skipping stale check")
return []
now = datetime.now(tz=timezone.utc)
# Define the timeframes to check
timeframes_to_check = ["1m", "5m", "15m", "30m", "1h", "4h", "1d"]
for db_symbol in symbols:
for tf in timeframes_to_check:
tf_seconds = _timeframe_to_seconds(tf)
stale_threshold_seconds = 2 * tf_seconds
# Get the latest candle timestamp for this symbol + timeframe
ts_result = await db.execute(
select(func.max(Candle.timestamp)).where(
and_(
Candle.symbol_id == db_symbol.id,
Candle.timeframe == tf,
)
)
)
latest_ts: datetime | None = ts_result.scalar()
if latest_ts is None:
# No data at all — flag as stale
stale_entries.append(
{
"symbol": db_symbol.symbol,
"exchange": db_symbol.exchange.name,
"timeframe": tf,
"last_timestamp": None,
"staleness_minutes": None,
"reason": "no_data",
}
)
continue
# Ensure timezone-awareness
if latest_ts.tzinfo is None:
latest_ts = latest_ts.replace(tzinfo=timezone.utc)
# Expected latest timestamp
expected_latest = latest_ts.replace(tzinfo=timezone.utc) + (
__import__("datetime").timedelta(seconds=stale_threshold_seconds)
)
if expected_latest < now:
staleness_mins = (now - latest_ts).total_seconds() / 60.0
stale_entries.append(
{
"symbol": db_symbol.symbol,
"exchange": db_symbol.exchange.name,
"timeframe": tf,
"last_timestamp": latest_ts.isoformat(),
"staleness_minutes": round(staleness_mins, 1),
"reason": "stale",
}
)
if stale_entries:
logger.warning(
"Found %d stale candle entries across %d symbols",
len(stale_entries),
len(symbols),
)
else:
logger.info("All candles are up-to-date — no stale entries detected")
return stale_entries
except Exception:
logger.exception("check_stale_candles failed")
return stale_entries
+216
View File
@@ -0,0 +1,216 @@
"""
WebSocket Connection Manager for real-time candle data streaming.
Provides thread-safe subscription management and broadcasting to clients
subscribed to specific symbol / timeframe / exchange combinations.
The module exposes a singleton ``manager`` instance that should be imported
wherever WebSocket subscriptions or broadcasts are needed.
"""
from __future__ import annotations
import json
import logging
from datetime import datetime
from decimal import Decimal
from typing import Any
from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _subscription_key(exchange: str, symbol: str, timeframe: str) -> str:
"""Return the canonical subscription dict key."""
return f"{exchange}:{symbol}:{timeframe}"
def _json_safe(value: Any) -> Any:
"""Recursively convert non-JSON-safe types to serializable equivalents.
- ``Decimal`` → ``float``
- ``datetime`` → ISO-format ``str``
"""
if isinstance(value, Decimal):
return float(value)
if isinstance(value, datetime):
return value.isoformat()
if isinstance(value, dict):
return {k: _json_safe(v) for k, v in value.items()}
if isinstance(value, (list, tuple)):
return [_json_safe(v) for v in value]
return value
# ---------------------------------------------------------------------------
# ConnectionManager
# ---------------------------------------------------------------------------
class ConnectionManager:
"""Manages WebSocket connections and their subscriptions.
Provides O(1) lookups in both directions:
* ``_subscriptions``: ``{key: {websocket, ...}}``
* ``_client_subs``: ``{websocket: {key, ...}}``
All public methods are safe to call from concurrent asyncio tasks thanks
to an internal ``asyncio.Lock``.
"""
def __init__(self) -> None:
import asyncio
self._subscriptions: dict[str, set[WebSocket]] = {}
self._client_subs: dict[WebSocket, set[str]] = {}
self._lock = asyncio.Lock()
# ------------------------------------------------------------------
# Subscription management
# ------------------------------------------------------------------
async def subscribe(
self,
websocket: WebSocket,
symbol: str,
timeframe: str,
exchange: str,
) -> None:
"""Register *websocket* for candle updates on the given channel."""
key = _subscription_key(exchange, symbol, timeframe)
async with self._lock:
self._subscriptions.setdefault(key, set()).add(websocket)
self._client_subs.setdefault(websocket, set()).add(key)
logger.debug(
"WebSocket %s subscribed to %s (%d subscribers on channel)",
id(websocket),
key,
len(self._subscriptions[key]),
)
async def unsubscribe(
self,
websocket: WebSocket,
symbol: str,
timeframe: str,
exchange: str,
) -> None:
"""Remove *websocket* from the given channel."""
key = _subscription_key(exchange, symbol, timeframe)
async with self._lock:
self._subscriptions.get(key, set()).discard(websocket)
self._client_subs.get(websocket, set()).discard(key)
# Clean up empty buckets
if key in self._subscriptions and not self._subscriptions[key]:
del self._subscriptions[key]
async def unsubscribe_all(self, websocket: WebSocket) -> None:
"""Remove *websocket* from every channel it was subscribed to."""
async with self._lock:
keys = self._client_subs.pop(websocket, set())
for key in keys:
self._subscriptions.get(key, set()).discard(websocket)
if key in self._subscriptions and not self._subscriptions[key]:
del self._subscriptions[key]
logger.debug(
"WebSocket %s unsubscribed from %d channels",
id(websocket),
len(keys),
)
async def get_subscribers(
self,
symbol: str,
timeframe: str,
exchange: str,
) -> list[WebSocket]:
"""Return a snapshot list of websockets subscribed to the channel."""
key = _subscription_key(exchange, symbol, timeframe)
async with self._lock:
return list(self._subscriptions.get(key, set()))
# ------------------------------------------------------------------
# Broadcasting
# ------------------------------------------------------------------
async def broadcast(
self,
symbol: str,
timeframe: str,
exchange: str,
message: dict[str, Any],
) -> None:
"""Send *message* to every websocket subscribed to the channel.
Disconnected clients are detected during the send attempt and
automatically removed from all subscriptions.
Parameters
----------
symbol:
Trading pair, e.g. ``"BTC/USDT"``.
timeframe:
Candle interval, e.g. ``"1m"``, ``"1h"``.
exchange:
Exchange name, e.g. ``"mexc"``.
message:
Payload dict. ``Decimal`` and ``datetime`` values are
automatically converted to JSON-safe equivalents.
"""
key = _subscription_key(exchange, symbol, timeframe)
# --- 1. Snapshot subscribers under the lock ---
async with self._lock:
subscribers = list(self._subscriptions.get(key, set()))
if not subscribers:
return
# --- 2. Serialise once for all recipients ---
safe_message = _json_safe(message)
payload = json.dumps(safe_message, default=str)
# --- 3. Send to each subscriber, cleaning up on failure ---
dead: list[WebSocket] = []
for ws in subscribers:
try:
if ws.client_state == WebSocketState.DISCONNECTED:
dead.append(ws)
continue
await ws.send_text(payload)
except WebSocketDisconnect:
dead.append(ws)
except Exception:
logger.exception(
"Error sending WS message to %s on channel %s",
id(ws),
key,
)
dead.append(ws)
if dead:
async with self._lock:
for ws in dead:
keys = self._client_subs.pop(ws, set())
for k in keys:
self._subscriptions.get(k, set()).discard(ws)
if k in self._subscriptions and not self._subscriptions[k]:
del self._subscriptions[k]
logger.info(
"Cleaned up %d disconnected websocket(s) from channel %s",
len(dead),
key,
)
# ---------------------------------------------------------------------------
# Singleton
# ---------------------------------------------------------------------------
manager: ConnectionManager = ConnectionManager()
"""Module-level singleton ``ConnectionManager`` instance."""