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
+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))