Initial commit: Trading Portal - FastAPI + React + PostgreSQL
This commit is contained in:
Executable
Executable
+21
@@ -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"}
|
||||
Executable
+360
@@ -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()
|
||||
Executable
+216
@@ -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)
|
||||
Executable
+410
@@ -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,
|
||||
}
|
||||
Executable
+116
@@ -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)
|
||||
Executable
+150
@@ -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"}
|
||||
Executable
+416
@@ -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)
|
||||
Executable
+228
@@ -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
|
||||
Executable
+476
@@ -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),
|
||||
}
|
||||
Executable
+114
@@ -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"}
|
||||
Executable
+155
@@ -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]}")
|
||||
Executable
+134
@@ -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)
|
||||
Executable
+56
@@ -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,
|
||||
}
|
||||
Executable
+71
@@ -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)
|
||||
Executable
+127
@@ -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)
|
||||
Executable
+181
@@ -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,
|
||||
)
|
||||
Executable
+196
@@ -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
|
||||
]
|
||||
Executable
Executable
+188
@@ -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))
|
||||
Reference in New Issue
Block a user