Initial commit: Trading Portal - FastAPI + React + PostgreSQL
This commit is contained in:
Executable
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))
|
||||
Executable
+43
@@ -0,0 +1,43 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
|
||||
class Settings(BaseSettings):
|
||||
"""Application settings loaded from environment variables / .env file."""
|
||||
|
||||
# Database
|
||||
DATABASE_URL: str = "postgresql+asyncpg://trading:trading_secret@db:5432/trading_portal"
|
||||
|
||||
# JWT
|
||||
JWT_PRIVATE_KEY_PATH: str = "/run/secrets/jwt_private.pem"
|
||||
JWT_PUBLIC_KEY_PATH: str = "/run/secrets/jwt_public_key.pem"
|
||||
JWT_PUBLIC_KEYS_DIR: str = "/run/secrets/jwt_public_keys" # directory of valid public keys
|
||||
JWT_ACCESS_TOKEN_EXPIRE_MINUTES: int = 15
|
||||
JWT_REFRESH_TOKEN_EXPIRE_DAYS: int = 7
|
||||
|
||||
# Encryption
|
||||
ENCRYPTION_KEY: str = ""
|
||||
|
||||
# Server
|
||||
HOST: str = "0.0.0.0"
|
||||
PORT: int = 8000
|
||||
|
||||
# Logging
|
||||
LOG_LEVEL: str = "INFO"
|
||||
|
||||
# CORS
|
||||
CORS_ORIGINS: str = ""
|
||||
|
||||
# Demo user (from .env)
|
||||
demo_user: str = "demo"
|
||||
demo_pass: str = "demo1234"
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
env_file_encoding="utf-8",
|
||||
extra="allow",
|
||||
)
|
||||
|
||||
|
||||
settings = Settings()
|
||||
Executable
Executable
+104
@@ -0,0 +1,104 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import AsyncGenerator
|
||||
from uuid import UUID
|
||||
|
||||
from fastapi import Depends, Header, Request
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.exceptions import InvalidCredentialsException, InvalidTokenException
|
||||
from app.core.security import decode_token
|
||||
from app.database import get_db
|
||||
from app.models import User
|
||||
|
||||
|
||||
async def get_db_session() -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Provide an async SQLAlchemy database session via FastAPI dependency."""
|
||||
async for session in get_db():
|
||||
yield session
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
request: Request,
|
||||
db: AsyncSession = Depends(get_db_session),
|
||||
) -> User:
|
||||
"""
|
||||
Extract the Bearer token from the Authorization header, decode it,
|
||||
and fetch the corresponding user from the database.
|
||||
|
||||
Raises ``InvalidCredentialsException`` if the token is missing,
|
||||
invalid, or the user is not found.
|
||||
"""
|
||||
authorization: str | None = request.headers.get("Authorization")
|
||||
|
||||
if not authorization or not authorization.startswith("Bearer "):
|
||||
raise InvalidCredentialsException(
|
||||
detail="Missing or malformed Authorization header"
|
||||
)
|
||||
|
||||
token = authorization.removeprefix("Bearer ").strip()
|
||||
|
||||
try:
|
||||
payload = decode_token(token)
|
||||
except Exception:
|
||||
raise InvalidTokenException(detail="Invalid or expired token")
|
||||
|
||||
sub: str | None = payload.get("sub")
|
||||
if sub is None:
|
||||
raise InvalidTokenException(detail="Token payload missing subject")
|
||||
|
||||
user_id: UUID
|
||||
try:
|
||||
user_id = UUID(sub)
|
||||
except ValueError:
|
||||
raise InvalidTokenException(detail="Invalid token subject format")
|
||||
|
||||
result = await db.execute(select(User).where(User.id == user_id))
|
||||
user = result.scalar_one_or_none()
|
||||
|
||||
if user is None:
|
||||
raise InvalidCredentialsException(detail="User not found")
|
||||
|
||||
return user
|
||||
|
||||
|
||||
async def get_current_active_user(
|
||||
current_user: User = Depends(get_current_user),
|
||||
) -> User:
|
||||
"""Return the current user if active, otherwise raise an error."""
|
||||
if not current_user.is_active:
|
||||
raise InvalidCredentialsException(detail="Inactive user account")
|
||||
return current_user
|
||||
|
||||
|
||||
async def get_current_trader_user(
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> User:
|
||||
"""Return the current user if they are a trader or admin, otherwise raise an error."""
|
||||
if current_user.role not in ("admin", "trader"):
|
||||
from app.core.exceptions import ForbiddenException
|
||||
|
||||
raise ForbiddenException(detail="Trader or admin privileges required")
|
||||
return current_user
|
||||
|
||||
|
||||
async def get_current_viewer_user(
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> User:
|
||||
"""Return the current user if active (any role can view)."""
|
||||
return current_user
|
||||
|
||||
|
||||
async def get_current_admin_user(
|
||||
current_user: User = Depends(get_current_active_user),
|
||||
) -> User:
|
||||
"""Return the current user if they are an admin, otherwise raise an error.
|
||||
|
||||
P1-22: Checks both is_admin AND role for frontend/backend consistency.
|
||||
"""
|
||||
if not current_user.is_admin and current_user.role != "admin":
|
||||
from app.core.exceptions import ForbiddenException
|
||||
|
||||
raise ForbiddenException(detail="Admin privileges required")
|
||||
return current_user
|
||||
Executable
+74
@@ -0,0 +1,74 @@
|
||||
from __future__ import annotations
|
||||
|
||||
|
||||
class AppException(Exception):
|
||||
"""Base exception for all application-level errors."""
|
||||
|
||||
status_code: int = 500
|
||||
detail: str = "Internal server error"
|
||||
code: str = "internal_error"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
status_code: int | None = None,
|
||||
detail: str | None = None,
|
||||
code: str | None = None,
|
||||
) -> None:
|
||||
if status_code is not None:
|
||||
self.status_code = status_code
|
||||
if detail is not None:
|
||||
self.detail = detail
|
||||
if code is not None:
|
||||
self.code = code
|
||||
super().__init__(self.detail)
|
||||
|
||||
|
||||
class NotFoundException(AppException):
|
||||
status_code: int = 404
|
||||
code: str = "not_found"
|
||||
detail: str = "Resource not found"
|
||||
|
||||
|
||||
class AuthException(AppException):
|
||||
status_code: int = 401
|
||||
code: str = "auth_error"
|
||||
detail: str = "Authentication error"
|
||||
|
||||
|
||||
class InvalidCredentialsException(AuthException):
|
||||
code: str = "invalid_credentials"
|
||||
detail: str = "Invalid username or password"
|
||||
|
||||
|
||||
class TokenExpiredException(AuthException):
|
||||
code: str = "token_expired"
|
||||
detail: str = "Token has expired"
|
||||
|
||||
|
||||
class InvalidTokenException(AuthException):
|
||||
code: str = "invalid_token"
|
||||
detail: str = "Invalid token"
|
||||
|
||||
|
||||
class ForbiddenException(AppException):
|
||||
status_code: int = 403
|
||||
code: str = "forbidden"
|
||||
detail: str = "Forbidden"
|
||||
|
||||
|
||||
class ValidationException(AppException):
|
||||
status_code: int = 422
|
||||
code: str = "validation_error"
|
||||
detail: str = "Validation error"
|
||||
|
||||
|
||||
class RateLimitException(AppException):
|
||||
status_code: int = 429
|
||||
code: str = "rate_limit"
|
||||
detail: str = "Rate limit exceeded"
|
||||
|
||||
|
||||
class ConflictException(AppException):
|
||||
status_code: int = 409
|
||||
code: str = "conflict"
|
||||
detail: str = "Resource already exists"
|
||||
Executable
+73
@@ -0,0 +1,73 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
|
||||
import structlog
|
||||
from starlette.requests import Request
|
||||
from starlette.responses import Response
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
logger = structlog.get_logger(__name__)
|
||||
|
||||
|
||||
class RequestLoggingMiddleware:
|
||||
"""ASGI middleware that logs every request with method, path, status code,
|
||||
and duration using structlog."""
|
||||
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
start = time.perf_counter()
|
||||
request = Request(scope)
|
||||
|
||||
# Wrap send to capture the response status code
|
||||
status_code: int | None = None
|
||||
|
||||
async def send_wrapper(message: dict) -> None:
|
||||
nonlocal status_code
|
||||
if message["type"] == "http.response.start":
|
||||
status_code = message["status"]
|
||||
await send(message)
|
||||
|
||||
try:
|
||||
await self.app(scope, receive, send_wrapper)
|
||||
except Exception:
|
||||
duration = time.perf_counter() - start
|
||||
logger.error(
|
||||
"request_error",
|
||||
method=request.method,
|
||||
path=request.url.path,
|
||||
status_code=500,
|
||||
duration_ms=round(duration * 1000, 2),
|
||||
)
|
||||
raise
|
||||
else:
|
||||
duration = time.perf_counter() - start
|
||||
if (status_code or 0) >= 400:
|
||||
logger.warning(
|
||||
"request_complete",
|
||||
method=request.method,
|
||||
path=request.url.path,
|
||||
status_code=status_code,
|
||||
duration_ms=round(duration * 1000, 2),
|
||||
)
|
||||
else:
|
||||
logger.info(
|
||||
"request_complete",
|
||||
method=request.method,
|
||||
path=request.url.path,
|
||||
status_code=status_code,
|
||||
duration_ms=round(duration * 1000, 2),
|
||||
)
|
||||
|
||||
|
||||
def register_middleware(app: ASGIApp) -> None:
|
||||
"""Convenience helper — add the middleware to a FastAPI app."""
|
||||
from app.core.middleware import RequestLoggingMiddleware # noqa: F811
|
||||
|
||||
app.add_middleware(RequestLoggingMiddleware) # type: ignore[arg-type]
|
||||
Executable
+351
@@ -0,0 +1,351 @@
|
||||
"""
|
||||
Security module for the trading portal backend.
|
||||
|
||||
Provides password hashing, JWT token management (RS256),
|
||||
AES-256-CBC encryption for API key storage, and token utilities.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import uuid
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional, Tuple
|
||||
|
||||
from fastapi import HTTPException
|
||||
from jose import JWTError, jwt
|
||||
from jose.exceptions import ExpiredSignatureError
|
||||
from passlib.context import CryptContext
|
||||
from cryptography.hazmat.primitives.ciphers import Cipher, algorithms, modes
|
||||
from cryptography.hazmat.backends import default_backend
|
||||
|
||||
from app.config import settings
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Password hashing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_pwd_context = CryptContext(schemes=["bcrypt"], deprecated="auto")
|
||||
|
||||
|
||||
def hash_password(password: str) -> str:
|
||||
"""Hash a plaintext password using bcrypt (synchronous — use in thread pool)."""
|
||||
return _pwd_context.hash(password)
|
||||
|
||||
|
||||
def verify_password(plain: str, hashed: str) -> bool:
|
||||
"""Verify a plaintext password against a bcrypt hash (synchronous — use in thread pool)."""
|
||||
return _pwd_context.verify(plain, hashed)
|
||||
|
||||
|
||||
async def hash_password_async(password: str) -> str:
|
||||
"""Hash a plaintext password using bcrypt (async, runs in thread pool)."""
|
||||
import asyncio
|
||||
loop = asyncio.get_running_loop()
|
||||
return await loop.run_in_executor(None, _pwd_context.hash, password)
|
||||
|
||||
|
||||
async def verify_password_async(plain: str, hashed: str) -> bool:
|
||||
"""Verify a plaintext password against a bcrypt hash (async, runs in thread pool)."""
|
||||
import asyncio
|
||||
loop = asyncio.get_running_loop()
|
||||
return await loop.run_in_executor(None, _pwd_context.verify, plain, hashed)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# JWT RS256 token management — with multi-key rotation support
|
||||
# ---------------------------------------------------------------------------
|
||||
#
|
||||
# PRINCIPLE:
|
||||
# - Tokens are SIGNED with the current private key and carry a "kid"
|
||||
# (Key ID = SHA-256 fingerprint of the public key) in the JWT header.
|
||||
# - Tokens are VERIFIED against ALL public keys in the keys directory.
|
||||
# ANY valid key can decode the token → old tokens survive key rotation.
|
||||
# - When rotating: add new key pair, old tokens remain valid until they
|
||||
# expire naturally. Remove old public keys only after all their tokens
|
||||
# have expired.
|
||||
#
|
||||
# Directory layout:
|
||||
# /run/secrets/jwt_private.pem ← CURRENT signing key
|
||||
# /run/secrets/jwt_public_keys/*.pem ← ALL valid public keys
|
||||
|
||||
import hashlib
|
||||
import os
|
||||
from functools import lru_cache
|
||||
|
||||
ALGORITHM = "RS256"
|
||||
|
||||
|
||||
def _load_private_key() -> str:
|
||||
"""Read the current RSA private key PEM for signing new tokens."""
|
||||
try:
|
||||
with open(settings.JWT_PRIVATE_KEY_PATH, "r") as f:
|
||||
return f.read()
|
||||
except FileNotFoundError:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"JWT private key not found at {settings.JWT_PRIVATE_KEY_PATH}",
|
||||
)
|
||||
except OSError as exc:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"Failed to read JWT private key: {exc}",
|
||||
)
|
||||
|
||||
|
||||
def _compute_kid(public_key_pem: str) -> str:
|
||||
"""Return a Key ID (SHA-256 fingerprint) for a public key PEM."""
|
||||
return hashlib.sha256(public_key_pem.strip().encode()).hexdigest()[:16]
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _cached_public_key() -> str:
|
||||
"""Return the CURRENT public key (used for kid computation). Cached."""
|
||||
keys = _load_all_public_keys()
|
||||
if not keys:
|
||||
raise HTTPException(status_code=500, detail="No valid JWT public keys found")
|
||||
return keys[-1][1] # newest key
|
||||
|
||||
|
||||
def _load_all_public_keys() -> list[tuple[str, str]]:
|
||||
"""
|
||||
Load ALL valid public keys from the keys directory.
|
||||
|
||||
Returns:
|
||||
List of (kid, pem_content) tuples, sorted by filename for determinism.
|
||||
The store allows multiple keys to coexist — old keys still validate
|
||||
tokens signed before the last rotation.
|
||||
"""
|
||||
keys_dir = settings.JWT_PUBLIC_KEYS_DIR
|
||||
keys: list[tuple[str, str]] = []
|
||||
|
||||
try:
|
||||
for filename in sorted(os.listdir(keys_dir)):
|
||||
if not filename.endswith(".pem"):
|
||||
continue
|
||||
filepath = os.path.join(keys_dir, filename)
|
||||
try:
|
||||
with open(filepath, "r") as f:
|
||||
pem = f.read().strip()
|
||||
if pem:
|
||||
kid = _compute_kid(pem)
|
||||
keys.append((kid, pem))
|
||||
except OSError:
|
||||
continue # skip unreadable files
|
||||
except FileNotFoundError:
|
||||
pass # no directory yet — handled gracefully
|
||||
except NotADirectoryError:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"JWT_PUBLIC_KEYS_DIR ({keys_dir}) is not a directory",
|
||||
)
|
||||
|
||||
if not keys:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail=f"No valid JWT public keys found in {keys_dir}",
|
||||
)
|
||||
|
||||
return keys
|
||||
|
||||
|
||||
def create_access_token(
|
||||
data: dict,
|
||||
expires_delta: Optional[timedelta] = None,
|
||||
) -> str:
|
||||
"""Create a short-lived JWT access token (RS256).
|
||||
|
||||
Includes ``kid`` header so the verifier knows which key to try first.
|
||||
"""
|
||||
to_encode = data.copy()
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
if expires_delta is not None:
|
||||
expire = now + expires_delta
|
||||
else:
|
||||
expire = now + timedelta(minutes=settings.JWT_ACCESS_TOKEN_EXPIRE_MINUTES)
|
||||
|
||||
to_encode.update({"iat": now, "exp": expire, "sub": str(data["sub"])})
|
||||
|
||||
private_key = _load_private_key()
|
||||
public_key = _cached_public_key()
|
||||
kid = _compute_kid(public_key)
|
||||
|
||||
headers = {"kid": kid}
|
||||
return jwt.encode(to_encode, private_key, algorithm=ALGORITHM, headers=headers)
|
||||
|
||||
|
||||
def create_refresh_token(data: dict) -> str:
|
||||
"""Create a long-lived JWT refresh token (RS256).
|
||||
|
||||
Includes ``kid``, ``type: refresh``, and a unique ``jti``.
|
||||
"""
|
||||
to_encode = data.copy()
|
||||
now = datetime.now(timezone.utc)
|
||||
expire = now + timedelta(days=settings.JWT_REFRESH_TOKEN_EXPIRE_DAYS)
|
||||
|
||||
to_encode.update(
|
||||
{
|
||||
"iat": now,
|
||||
"exp": expire,
|
||||
"sub": str(data["sub"]),
|
||||
"type": "refresh",
|
||||
"jti": generate_jti(),
|
||||
}
|
||||
)
|
||||
|
||||
private_key = _load_private_key()
|
||||
public_key = _cached_public_key()
|
||||
kid = _compute_kid(public_key)
|
||||
|
||||
headers = {"kid": kid}
|
||||
return jwt.encode(to_encode, private_key, algorithm=ALGORITHM, headers=headers)
|
||||
|
||||
|
||||
def decode_token(token: str) -> dict:
|
||||
"""Decode and verify a JWT token using ALL known public keys.
|
||||
|
||||
Tries every valid public key in the directory. If ANY key verifies the
|
||||
token, it is valid — this is how key rotation works without breaking
|
||||
existing sessions.
|
||||
|
||||
Optimisation: the token's ``kid`` header is used to try the matching key
|
||||
first before falling back to a full linear scan.
|
||||
"""
|
||||
# 1. Extract kid from token header (without verifying signature yet)
|
||||
try:
|
||||
unverified_header = jwt.get_unverified_header(token)
|
||||
token_kid = unverified_header.get("kid")
|
||||
except JWTError:
|
||||
token_kid = None
|
||||
|
||||
# 2. Load all valid public keys
|
||||
all_keys = _load_all_public_keys() # [(kid, pem), ...]
|
||||
|
||||
# 3. If we have a kid, try the matching key first
|
||||
if token_kid:
|
||||
for kid, pem in all_keys:
|
||||
if kid == token_kid:
|
||||
try:
|
||||
return jwt.decode(token, pem, algorithms=[ALGORITHM])
|
||||
except ExpiredSignatureError:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Token has expired",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
except JWTError:
|
||||
pass # key mismatch — fall through to full scan
|
||||
|
||||
# 4. Fallback: try all keys (handles tokens without kid, or kid mismatch)
|
||||
for kid, pem in all_keys:
|
||||
try:
|
||||
return jwt.decode(token, pem, algorithms=[ALGORITHM])
|
||||
except ExpiredSignatureError:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Token has expired",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
except JWTError:
|
||||
continue # try next key
|
||||
|
||||
# 5. No key worked
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="Invalid or expired token",
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AES-256-CBC encryption (for API key storage)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_BACKEND = default_backend()
|
||||
|
||||
|
||||
def generate_encryption_key() -> str:
|
||||
"""Generate a random 32-byte (256-bit) hex-encoded encryption key.
|
||||
|
||||
Print the key to stdout so it can be copied into the ``.env`` file.
|
||||
"""
|
||||
key = uuid.uuid4().hex + uuid.uuid4().hex # 64 hex chars = 32 bytes
|
||||
print(f"Encryption key (save in .env as ENCRYPTION_KEY={key}): {key}")
|
||||
return key
|
||||
|
||||
|
||||
def _resolve_key(key_hex: Optional[str] = None) -> bytes:
|
||||
"""Return the AES key as bytes, falling back to settings."""
|
||||
raw = key_hex if key_hex is not None else settings.ENCRYPTION_KEY
|
||||
if not raw:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail="Encryption key not configured. Set ENCRYPTION_KEY in .env",
|
||||
)
|
||||
return bytes.fromhex(raw)
|
||||
|
||||
|
||||
def encrypt_api_key(
|
||||
api_key: str,
|
||||
key_hex: Optional[str] = None,
|
||||
) -> Tuple[str, str]:
|
||||
"""Encrypt an API key with AES-256-CBC.
|
||||
|
||||
Returns ``(ciphertext_hex, iv_hex)``.
|
||||
"""
|
||||
key = _resolve_key(key_hex)
|
||||
iv = uuid.uuid4().bytes # 16 random bytes
|
||||
cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=_BACKEND)
|
||||
encryptor = cipher.encryptor()
|
||||
|
||||
# Pad plaintext to AES block size (16 bytes) using PKCS7
|
||||
plaintext_bytes = api_key.encode("utf-8")
|
||||
pad_len = 16 - (len(plaintext_bytes) % 16)
|
||||
padded = plaintext_bytes + bytes([pad_len] * pad_len)
|
||||
|
||||
ciphertext = encryptor.update(padded) + encryptor.finalize()
|
||||
return ciphertext.hex(), iv.hex()
|
||||
|
||||
|
||||
def decrypt_api_key(
|
||||
ciphertext_hex: str,
|
||||
iv_hex: str,
|
||||
key_hex: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Decrypt an AES-256-CBC encrypted API key.
|
||||
|
||||
Returns the original plaintext string.
|
||||
"""
|
||||
key = _resolve_key(key_hex)
|
||||
ciphertext = bytes.fromhex(ciphertext_hex)
|
||||
iv = bytes.fromhex(iv_hex)
|
||||
|
||||
cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=_BACKEND)
|
||||
decryptor = cipher.decryptor()
|
||||
padded = decryptor.update(ciphertext) + decryptor.finalize()
|
||||
|
||||
# Remove PKCS7 padding
|
||||
pad_len = padded[-1]
|
||||
if pad_len < 1 or pad_len > 16:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail="Decryption failed: invalid padding",
|
||||
)
|
||||
plaintext_bytes = padded[:-pad_len]
|
||||
return plaintext_bytes.decode("utf-8")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Token utilities
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def generate_token_hash(token: str) -> str:
|
||||
"""Return the SHA-256 hex digest of a token string."""
|
||||
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def generate_jti() -> str:
|
||||
"""Return a UUID4 hex string for use as a JWT token ID."""
|
||||
return uuid.uuid4().hex
|
||||
Executable
+64
@@ -0,0 +1,64 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import AsyncGenerator
|
||||
|
||||
from sqlalchemy.exc import InterfaceError
|
||||
from sqlalchemy.ext.asyncio import (
|
||||
AsyncSession,
|
||||
async_sessionmaker,
|
||||
create_async_engine,
|
||||
)
|
||||
from sqlalchemy.orm import DeclarativeBase
|
||||
|
||||
from app.config import settings
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
engine = create_async_engine(
|
||||
settings.DATABASE_URL,
|
||||
echo=(settings.LOG_LEVEL == "DEBUG"),
|
||||
pool_pre_ping=True,
|
||||
pool_size=60,
|
||||
max_overflow=20,
|
||||
pool_timeout=5,
|
||||
pool_recycle=600,
|
||||
connect_args={
|
||||
"server_settings": {
|
||||
"idle_in_transaction_session_timeout": "60000",
|
||||
"statement_timeout": "10000",
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
async_session_factory = async_sessionmaker(
|
||||
engine,
|
||||
class_=AsyncSession,
|
||||
expire_on_commit=False,
|
||||
)
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
"""Declarative base for all ORM models."""
|
||||
|
||||
|
||||
async def get_db() -> AsyncGenerator[AsyncSession, None]:
|
||||
"""Provide async DB session. Handles closed-connection errors gracefully
|
||||
to prevent exception storms that spike CPU to 500%+."""
|
||||
async with async_session_factory() as session:
|
||||
try:
|
||||
yield session
|
||||
await session.commit()
|
||||
except InterfaceError:
|
||||
logger.warning("Session commit skipped: connection already closed")
|
||||
except Exception:
|
||||
try:
|
||||
await session.rollback()
|
||||
except InterfaceError:
|
||||
logger.warning("Session rollback skipped: connection already closed")
|
||||
raise
|
||||
finally:
|
||||
try:
|
||||
await session.close()
|
||||
except InterfaceError:
|
||||
pass # already closed
|
||||
Executable
Executable
+344
@@ -0,0 +1,344 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Optional
|
||||
|
||||
import ccxt
|
||||
|
||||
from app.exchange.rate_limiter import GlobalRateLimiter, RateLimiter
|
||||
from app.exchange.types import (
|
||||
BalanceData,
|
||||
BalanceResponse,
|
||||
CandleData,
|
||||
CandleValidationError,
|
||||
OpenOrderData,
|
||||
OrderData,
|
||||
OrderRequest,
|
||||
PositionData,
|
||||
SymbolInfo,
|
||||
TickerData,
|
||||
TradeData,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class AbstractExchange(ABC):
|
||||
"""Abstract base class for exchange adapters."""
|
||||
|
||||
def __init__(self, api_key: str = "", api_secret: str = "", testnet: bool = False) -> None:
|
||||
self.api_key = api_key
|
||||
self.api_secret = api_secret
|
||||
self.testnet = testnet
|
||||
self._client: Optional[ccxt.Exchange] = None
|
||||
self._rate_limiter: Optional[RateLimiter] = None
|
||||
|
||||
def _init_ccxt(self) -> ccxt.Exchange:
|
||||
"""Initialize and return a CCXT exchange client.
|
||||
|
||||
Subclasses must override this to configure the specific exchange.
|
||||
"""
|
||||
raise NotImplementedError("Subclasses must implement _init_ccxt")
|
||||
|
||||
@property
|
||||
def client(self) -> ccxt.Exchange:
|
||||
if self._client is None:
|
||||
self._client = self._init_ccxt()
|
||||
return self._client
|
||||
|
||||
@property
|
||||
def rate_limiter(self) -> RateLimiter:
|
||||
if self._rate_limiter is None:
|
||||
self._rate_limiter = GlobalRateLimiter(self.get_name())
|
||||
return self._rate_limiter
|
||||
|
||||
@abstractmethod
|
||||
def get_name(self) -> str:
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def get_base_url(self) -> str:
|
||||
"""Return the REST API base URL."""
|
||||
...
|
||||
|
||||
@abstractmethod
|
||||
def get_ws_url(self) -> str:
|
||||
"""Return the WebSocket URL."""
|
||||
...
|
||||
|
||||
async def fetch_ohlcv(self, symbol: str, timeframe: str = "1h", limit: int = 500) -> list[CandleData]:
|
||||
"""Fetch OHLCV candle data from the exchange."""
|
||||
await self.rate_limiter.acquire()
|
||||
raw = await self._async_fetch_ohlcv(symbol, timeframe, limit)
|
||||
|
||||
candles: list[CandleData] = []
|
||||
for entry in raw:
|
||||
ts, o, h, l, c, v = entry
|
||||
open_dec = Decimal(str(o))
|
||||
high_dec = Decimal(str(h))
|
||||
low_dec = Decimal(str(l))
|
||||
close_dec = Decimal(str(c))
|
||||
volume_dec = Decimal(str(v))
|
||||
|
||||
if not (low_dec <= open_dec <= high_dec and low_dec <= close_dec <= high_dec):
|
||||
continue
|
||||
|
||||
candles.append(
|
||||
CandleData(
|
||||
symbol=symbol,
|
||||
exchange=self.get_name(),
|
||||
timeframe=timeframe,
|
||||
timestamp=datetime.fromtimestamp(ts / 1000, tz=timezone.utc),
|
||||
open=open_dec,
|
||||
high=high_dec,
|
||||
low=low_dec,
|
||||
close=close_dec,
|
||||
volume=volume_dec,
|
||||
)
|
||||
)
|
||||
return candles
|
||||
|
||||
async def _async_fetch_ohlcv(self, symbol: str, timeframe: str, limit: int, since: int | None = None) -> list[list[Any]]:
|
||||
"""Run the synchronous CCXT fetch_ohlcv in a thread pool."""
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
kwargs = dict(symbol=symbol, timeframe=timeframe, limit=limit)
|
||||
if since is not None:
|
||||
kwargs["since"] = since
|
||||
return await loop.run_in_executor(
|
||||
None,
|
||||
lambda: self.client.fetch_ohlcv(**kwargs),
|
||||
)
|
||||
|
||||
async def fetch_ticker(self, symbol: str) -> TickerData:
|
||||
"""Fetch ticker data from the exchange."""
|
||||
await self.rate_limiter.acquire()
|
||||
raw = await self._async_fetch_ticker(symbol)
|
||||
|
||||
return TickerData(
|
||||
symbol=symbol,
|
||||
exchange=self.get_name(),
|
||||
bid=Decimal(str(raw.get("bid", 0))),
|
||||
ask=Decimal(str(raw.get("ask", 0))),
|
||||
last=Decimal(str(raw.get("last", 0))),
|
||||
volume_24h=Decimal(str(raw.get("baseVolume", 0))),
|
||||
change_24h=Decimal(str(raw.get("percentage", 0))) if raw.get("percentage") is not None else None,
|
||||
timestamp=datetime.fromtimestamp(raw["timestamp"] / 1000, tz=timezone.utc) if raw.get("timestamp") else datetime.now(tz=timezone.utc),
|
||||
)
|
||||
|
||||
async def _async_fetch_ticker(self, symbol: str) -> dict[str, Any]:
|
||||
"""Run the synchronous CCXT fetch_ticker in a thread pool."""
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
return await loop.run_in_executor(
|
||||
None,
|
||||
lambda: self.client.fetch_ticker(symbol),
|
||||
)
|
||||
|
||||
async def fetch_symbols(self) -> list[SymbolInfo]:
|
||||
"""Fetch all available trading pairs from the exchange."""
|
||||
await self.rate_limiter.acquire()
|
||||
import asyncio
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
markets = await loop.run_in_executor(None, lambda: self.client.load_markets())
|
||||
|
||||
result: list[SymbolInfo] = []
|
||||
for sym, info in markets.items():
|
||||
if info.get("active", True):
|
||||
result.append(
|
||||
SymbolInfo(
|
||||
symbol=sym,
|
||||
base=info.get("base", ""),
|
||||
quote=info.get("quote", ""),
|
||||
exchange=self.get_name(),
|
||||
is_active=info.get("active", True),
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# Order placement
|
||||
# ──────────────────────────────────────────────
|
||||
|
||||
async def create_order(self, req: OrderRequest) -> OrderData:
|
||||
"""Place an order on the exchange.
|
||||
|
||||
Subclasses may override to add exchange-specific logic (e.g. leverage,
|
||||
position side). Default implementation uses CCXT's create_order.
|
||||
"""
|
||||
import asyncio
|
||||
|
||||
await self.rate_limiter.acquire()
|
||||
params: dict[str, Any] = {}
|
||||
if req.reduce_only:
|
||||
params["reduceOnly"] = True
|
||||
if req.position_side:
|
||||
params["positionSide"] = req.position_side.upper()
|
||||
|
||||
def _place() -> dict[str, Any]:
|
||||
return self.client.create_order(
|
||||
symbol=req.symbol,
|
||||
type=req.order_type,
|
||||
side=req.side,
|
||||
amount=float(req.amount),
|
||||
price=float(req.price) if req.price else None,
|
||||
params=params,
|
||||
)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
raw = await loop.run_in_executor(None, _place)
|
||||
|
||||
return OrderData(
|
||||
exchange=self.get_name(),
|
||||
symbol=raw.get("symbol", req.symbol),
|
||||
order_id=str(raw.get("id", "")),
|
||||
client_order_id=raw.get("clientOrderId"),
|
||||
side=raw.get("side", req.side),
|
||||
order_type=raw.get("type", req.order_type),
|
||||
amount=Decimal(str(raw.get("amount", float(req.amount)))),
|
||||
filled=Decimal(str(raw.get("filled", 0))),
|
||||
price=Decimal(str(raw["price"])) if raw.get("price") else req.price,
|
||||
average=Decimal(str(raw["average"])) if raw.get("average") else None,
|
||||
status=raw.get("status", "open"),
|
||||
timestamp=datetime.fromtimestamp(raw["timestamp"] / 1000, tz=timezone.utc) if raw.get("timestamp") else datetime.now(tz=timezone.utc),
|
||||
raw=raw,
|
||||
)
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# Balance
|
||||
# ──────────────────────────────────────────────
|
||||
|
||||
async def fetch_balance(self) -> BalanceResponse:
|
||||
"""Fetch the full account balance from the exchange."""
|
||||
import asyncio
|
||||
|
||||
await self.rate_limiter.acquire()
|
||||
|
||||
def _fetch() -> dict[str, Any]:
|
||||
return self.client.fetch_balance()
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
raw = await loop.run_in_executor(None, _fetch)
|
||||
|
||||
balances: list[BalanceData] = []
|
||||
for asset, info in raw.get("total", {}).items():
|
||||
if asset == "info" or asset == "free" or asset == "used" or asset == "total":
|
||||
continue
|
||||
free = Decimal(str(raw.get("free", {}).get(asset, 0)))
|
||||
used = Decimal(str(raw.get("used", {}).get(asset, 0)))
|
||||
total = Decimal(str(info))
|
||||
if total > 0 or free > 0:
|
||||
balances.append(BalanceData(asset=asset, free=free, used=used, total=total))
|
||||
|
||||
return BalanceResponse(
|
||||
exchange=self.get_name(),
|
||||
balances=balances,
|
||||
timestamp=datetime.now(tz=timezone.utc),
|
||||
)
|
||||
|
||||
# ──────────────────────────────────────────────
|
||||
# Open Orders, Trades, Positions
|
||||
# ──────────────────────────────────────────────
|
||||
|
||||
async def fetch_open_orders(self, symbol: str | None = None) -> list[OpenOrderData]:
|
||||
"""Fetch open orders from the exchange."""
|
||||
import asyncio
|
||||
|
||||
await self.rate_limiter.acquire()
|
||||
|
||||
def _fetch() -> list[dict[str, Any]]:
|
||||
return self.client.fetch_open_orders(symbol=symbol)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
raw_orders = await loop.run_in_executor(None, _fetch)
|
||||
|
||||
orders: list[OpenOrderData] = []
|
||||
for raw in (raw_orders or []):
|
||||
orders.append(OpenOrderData(
|
||||
order_id=str(raw.get("id", "")),
|
||||
symbol=raw.get("symbol", ""),
|
||||
side=raw.get("side", ""),
|
||||
order_type=raw.get("type", ""),
|
||||
amount=Decimal(str(raw.get("amount", 0))),
|
||||
filled=Decimal(str(raw.get("filled", 0))),
|
||||
price=Decimal(str(raw["price"])) if raw.get("price") else None,
|
||||
average=Decimal(str(raw["average"])) if raw.get("average") else None,
|
||||
status=raw.get("status", "open"),
|
||||
timestamp=datetime.fromtimestamp(raw["timestamp"] / 1000, tz=timezone.utc)
|
||||
if raw.get("timestamp") else None,
|
||||
))
|
||||
return orders
|
||||
|
||||
async def fetch_my_trades(self, symbol: str | None = None, limit: int = 10) -> list[TradeData]:
|
||||
"""Fetch recent filled trades from the exchange."""
|
||||
import asyncio
|
||||
|
||||
await self.rate_limiter.acquire()
|
||||
|
||||
def _fetch() -> list[dict[str, Any]]:
|
||||
return self.client.fetch_my_trades(symbol=symbol, limit=limit)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
raw_trades = await loop.run_in_executor(None, _fetch)
|
||||
|
||||
trades: list[TradeData] = []
|
||||
for raw in (raw_trades or []):
|
||||
cost = None
|
||||
if raw.get("cost"):
|
||||
cost = Decimal(str(raw["cost"]))
|
||||
fee_val = None
|
||||
fee_currency = None
|
||||
if raw.get("fee"):
|
||||
fee_val = Decimal(str(raw["fee"].get("cost", 0)))
|
||||
fee_currency = raw["fee"].get("currency")
|
||||
trades.append(TradeData(
|
||||
trade_id=str(raw.get("id", "")),
|
||||
symbol=raw.get("symbol", ""),
|
||||
side=raw.get("side", ""),
|
||||
amount=Decimal(str(raw.get("amount", 0))),
|
||||
price=Decimal(str(raw.get("price", 0))),
|
||||
cost=cost,
|
||||
fee=fee_val,
|
||||
fee_currency=fee_currency,
|
||||
timestamp=datetime.fromtimestamp(raw["timestamp"] / 1000, tz=timezone.utc)
|
||||
if raw.get("timestamp") else None,
|
||||
))
|
||||
return trades
|
||||
|
||||
async def fetch_positions(self, symbols: list[str] | None = None) -> list[PositionData]:
|
||||
"""Fetch open positions (futures/derivatives) from the exchange."""
|
||||
import asyncio
|
||||
|
||||
await self.rate_limiter.acquire()
|
||||
|
||||
def _fetch() -> list[dict[str, Any]]:
|
||||
return self.client.fetch_positions(symbols=symbols)
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
raw_positions = await loop.run_in_executor(None, _fetch)
|
||||
|
||||
positions: list[PositionData] = []
|
||||
for raw in (raw_positions or []):
|
||||
contracts = Decimal(str(raw.get("contracts", 0)))
|
||||
if contracts == 0:
|
||||
continue
|
||||
positions.append(PositionData(
|
||||
symbol=raw.get("symbol", ""),
|
||||
side=raw.get("side", ""),
|
||||
contracts=contracts,
|
||||
entry_price=Decimal(str(raw["entryPrice"])) if raw.get("entryPrice") else None,
|
||||
mark_price=Decimal(str(raw["markPrice"])) if raw.get("markPrice") else None,
|
||||
unrealized_pnl=Decimal(str(raw["unrealizedPnl"])) if raw.get("unrealizedPnl") else None,
|
||||
leverage=Decimal(str(raw["leverage"])) if raw.get("leverage") else None,
|
||||
liquidation_price=Decimal(str(raw["liquidationPrice"])) if raw.get("liquidationPrice") else None,
|
||||
percentage=Decimal(str(raw["percentage"])) if raw.get("percentage") else None,
|
||||
))
|
||||
return positions
|
||||
|
||||
|
||||
from decimal import Decimal
|
||||
from datetime import datetime, timezone
|
||||
Executable
+37
@@ -0,0 +1,37 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import ccxt
|
||||
|
||||
from app.exchange.base import AbstractExchange
|
||||
|
||||
|
||||
class BinanceAdapter(AbstractExchange):
|
||||
"""Exchange adapter for Binance."""
|
||||
|
||||
def get_name(self) -> str:
|
||||
return "binance"
|
||||
|
||||
def get_base_url(self) -> str:
|
||||
return "https://api.binance.com"
|
||||
|
||||
def get_ws_url(self) -> str:
|
||||
return "wss://stream.binance.com:9443/ws"
|
||||
|
||||
def _init_ccxt(self) -> ccxt.Exchange:
|
||||
options: dict = {
|
||||
"apiKey": self.api_key,
|
||||
"secret": self.api_secret,
|
||||
"rateLimit": 1200,
|
||||
"enableRateLimit": True,
|
||||
"options": {
|
||||
"warnOnFetchOpenOrdersWithoutSymbol": False,
|
||||
},
|
||||
}
|
||||
if self.testnet:
|
||||
options["urls"] = {
|
||||
"api": {
|
||||
"public": "https://testnet.binance.vision/api",
|
||||
"private": "https://testnet.binance.vision/api",
|
||||
}
|
||||
}
|
||||
return ccxt.binance(options)
|
||||
@@ -0,0 +1,26 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import ccxt
|
||||
|
||||
from app.exchange.base import AbstractExchange
|
||||
|
||||
|
||||
class BingXAdapter(AbstractExchange):
|
||||
"""Exchange adapter for BingX."""
|
||||
|
||||
def get_name(self) -> str:
|
||||
return "bingx"
|
||||
|
||||
def get_base_url(self) -> str:
|
||||
return "https://api.bingx.com"
|
||||
|
||||
def get_ws_url(self) -> str:
|
||||
return "wss://open-api-ws.bingx.com/market"
|
||||
|
||||
def _init_ccxt(self) -> ccxt.Exchange:
|
||||
options: dict = {
|
||||
"apiKey": self.api_key,
|
||||
"secret": self.api_secret,
|
||||
"enableRateLimit": True,
|
||||
}
|
||||
return ccxt.bingx(options)
|
||||
Executable
+33
@@ -0,0 +1,33 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import ccxt
|
||||
|
||||
from app.exchange.base import AbstractExchange
|
||||
|
||||
|
||||
class BybitAdapter(AbstractExchange):
|
||||
"""Exchange adapter for Bybit."""
|
||||
|
||||
def get_name(self) -> str:
|
||||
return "bybit"
|
||||
|
||||
def get_base_url(self) -> str:
|
||||
return "https://api.bybit.com"
|
||||
|
||||
def get_ws_url(self) -> str:
|
||||
return "wss://stream.bybit.com/v5/public/spot"
|
||||
|
||||
def _init_ccxt(self) -> ccxt.Exchange:
|
||||
options: dict = {
|
||||
"apiKey": self.api_key,
|
||||
"secret": self.api_secret,
|
||||
"enableRateLimit": True,
|
||||
}
|
||||
if self.testnet:
|
||||
options["urls"] = {
|
||||
"api": {
|
||||
"public": "https://api-testnet.bybit.com",
|
||||
"private": "https://api-testnet.bybit.com",
|
||||
}
|
||||
}
|
||||
return ccxt.bybit(options)
|
||||
Executable
+87
@@ -0,0 +1,87 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from app.exchange.base import AbstractExchange
|
||||
from app.exchange.binance import BinanceAdapter
|
||||
from app.exchange.bingx import BingXAdapter
|
||||
from app.exchange.bybit import BybitAdapter
|
||||
from app.exchange.gate import GateAdapter
|
||||
from app.exchange.mexc import MEXCAdapter
|
||||
|
||||
|
||||
class ExchangeFactory:
|
||||
"""Factory for creating exchange adapter instances.
|
||||
|
||||
Implements the Singleton pattern via a global instance.
|
||||
Caches adapter instances (TTL 5 min) to avoid repeated CCXT init (~20s).
|
||||
"""
|
||||
|
||||
_adapters: dict[str, type[AbstractExchange]] = {
|
||||
"binance": BinanceAdapter,
|
||||
"bingx": BingXAdapter,
|
||||
"bybit": BybitAdapter,
|
||||
"gate": GateAdapter,
|
||||
"mexc": MEXCAdapter,
|
||||
}
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._cache: dict[str, tuple[AbstractExchange, float]] = {}
|
||||
self._cache_ttl: float = 300.0 # 5 minutes
|
||||
|
||||
def register(self, name: str, adapter_cls: type[AbstractExchange]) -> None:
|
||||
"""Register a new exchange adapter class."""
|
||||
self._adapters[name] = adapter_cls
|
||||
|
||||
def _cache_key(self, name: str, api_key: str) -> str:
|
||||
"""Generate a cache key from exchange name + API key."""
|
||||
return f"{name}:{hashlib.sha256(api_key.encode()).hexdigest()[:16]}"
|
||||
|
||||
def create(
|
||||
self,
|
||||
name: str,
|
||||
api_key: str = "",
|
||||
api_secret: str = "",
|
||||
testnet: bool = False,
|
||||
) -> AbstractExchange:
|
||||
"""Create (or retrieve cached) an exchange adapter instance by name.
|
||||
|
||||
Raises ValueError if the exchange is not registered.
|
||||
"""
|
||||
adapter_cls = self._adapters.get(name)
|
||||
if adapter_cls is None:
|
||||
raise ValueError(
|
||||
f"Unknown exchange: {name!r}. "
|
||||
f"Available exchanges: {', '.join(sorted(self._adapters))}"
|
||||
)
|
||||
|
||||
# Return cached adapter if still fresh
|
||||
if api_key:
|
||||
key = self._cache_key(name, api_key)
|
||||
cached = self._cache.get(key)
|
||||
if cached is not None:
|
||||
adapter, created_at = cached
|
||||
if time.monotonic() - created_at < self._cache_ttl:
|
||||
return adapter
|
||||
# Expired — remove from cache
|
||||
del self._cache[key]
|
||||
|
||||
adapter = adapter_cls(api_key=api_key, api_secret=api_secret, testnet=testnet)
|
||||
|
||||
# Warm up the CCXT client in the background — first balance call will trigger it
|
||||
# but subsequent calls within TTL reuse the same instance
|
||||
if api_key:
|
||||
key = self._cache_key(name, api_key)
|
||||
self._cache[key] = (adapter, time.monotonic())
|
||||
|
||||
return adapter
|
||||
|
||||
def get_available_exchanges(self) -> list[str]:
|
||||
"""Return a list of registered exchange names."""
|
||||
return list(self._adapters.keys())
|
||||
|
||||
|
||||
# Global singleton factory instance
|
||||
factory: ExchangeFactory = ExchangeFactory()
|
||||
@@ -0,0 +1,26 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import ccxt
|
||||
|
||||
from app.exchange.base import AbstractExchange
|
||||
|
||||
|
||||
class GateAdapter(AbstractExchange):
|
||||
"""Exchange adapter for Gate.io."""
|
||||
|
||||
def get_name(self) -> str:
|
||||
return "gate"
|
||||
|
||||
def get_base_url(self) -> str:
|
||||
return "https://api.gateio.ws"
|
||||
|
||||
def get_ws_url(self) -> str:
|
||||
return "wss://api.gateio.ws/ws/v4/"
|
||||
|
||||
def _init_ccxt(self) -> ccxt.Exchange:
|
||||
options: dict = {
|
||||
"apiKey": self.api_key,
|
||||
"secret": self.api_secret,
|
||||
"enableRateLimit": True,
|
||||
}
|
||||
return ccxt.gate(options)
|
||||
Executable
+33
@@ -0,0 +1,33 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import ccxt
|
||||
|
||||
from app.exchange.base import AbstractExchange
|
||||
|
||||
|
||||
class MEXCAdapter(AbstractExchange):
|
||||
"""Exchange adapter for MEXC."""
|
||||
|
||||
def get_name(self) -> str:
|
||||
return "mexc"
|
||||
|
||||
def get_base_url(self) -> str:
|
||||
return "https://api.mexc.com"
|
||||
|
||||
def get_ws_url(self) -> str:
|
||||
return "wss://wbs.mexc.com/ws"
|
||||
|
||||
def _init_ccxt(self) -> ccxt.Exchange:
|
||||
options: dict = {
|
||||
"apiKey": self.api_key,
|
||||
"secret": self.api_secret,
|
||||
"enableRateLimit": True,
|
||||
}
|
||||
if self.testnet:
|
||||
options["urls"] = {
|
||||
"api": {
|
||||
"public": "https://testnet-api.mexc.com",
|
||||
"private": "https://testnet-api.mexc.com",
|
||||
}
|
||||
}
|
||||
return ccxt.mexc(options)
|
||||
Executable
+72
@@ -0,0 +1,72 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
# Per-exchange rate limits (requests per second for standard API)
|
||||
BINANCE_RATE_LIMIT: int = 10
|
||||
BYBIT_RATE_LIMIT: int = 10
|
||||
MEXC_RATE_LIMIT: int = 20
|
||||
|
||||
|
||||
class RateLimiter:
|
||||
"""Token bucket rate limiter for exchange API requests."""
|
||||
|
||||
def __init__(self, tokens_per_second: float, max_tokens: Optional[int] = None, name: str = "") -> None:
|
||||
self.tokens_per_second = tokens_per_second
|
||||
self.max_tokens = max_tokens if max_tokens is not None else int(tokens_per_second)
|
||||
self.name = name
|
||||
self._tokens: float = float(self.max_tokens)
|
||||
self._last_refill: float = time.monotonic()
|
||||
self._lock: asyncio.Lock = asyncio.Lock()
|
||||
|
||||
async def acquire(self) -> None:
|
||||
"""Wait for a token to be available, blocking until one is."""
|
||||
while True:
|
||||
async with self._lock:
|
||||
self._refill()
|
||||
if self._tokens >= 1.0:
|
||||
self._tokens -= 1.0
|
||||
return
|
||||
# How long until we have at least 1 token?
|
||||
wait_time = (1.0 - self._tokens) / self.tokens_per_second
|
||||
await asyncio.sleep(wait_time)
|
||||
|
||||
def _refill(self) -> None:
|
||||
now = time.monotonic()
|
||||
elapsed = now - self._last_refill
|
||||
self._tokens = min(float(self.max_tokens), self._tokens + elapsed * self.tokens_per_second)
|
||||
self._last_refill = now
|
||||
|
||||
async def __aenter__(self) -> "RateLimiter":
|
||||
await self.acquire()
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: Optional[type[BaseException]],
|
||||
exc_val: Optional[BaseException],
|
||||
exc_tb: Optional[object],
|
||||
) -> None:
|
||||
pass
|
||||
|
||||
|
||||
# Global rate limiter registry: singleton mapping exchange_name -> RateLimiter instance
|
||||
_global_limiters: dict[str, RateLimiter] = {}
|
||||
|
||||
|
||||
def GlobalRateLimiter(exchange_name: str) -> RateLimiter:
|
||||
"""Get or create the singleton RateLimiter for an exchange."""
|
||||
if exchange_name not in _global_limiters:
|
||||
limit_map = {
|
||||
"binance": BINANCE_RATE_LIMIT,
|
||||
"bybit": BYBIT_RATE_LIMIT,
|
||||
"mexc": MEXC_RATE_LIMIT,
|
||||
}
|
||||
tokens_per_second = limit_map.get(exchange_name, 10)
|
||||
_global_limiters[exchange_name] = RateLimiter(
|
||||
tokens_per_second=tokens_per_second,
|
||||
name=exchange_name,
|
||||
)
|
||||
return _global_limiters[exchange_name]
|
||||
Executable
+146
@@ -0,0 +1,146 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class CandleData(BaseModel):
|
||||
symbol: str
|
||||
exchange: str
|
||||
timeframe: str
|
||||
timestamp: datetime
|
||||
open: Decimal
|
||||
high: Decimal
|
||||
low: Decimal
|
||||
close: Decimal
|
||||
volume: Decimal
|
||||
|
||||
|
||||
class SymbolInfo(BaseModel):
|
||||
symbol: str
|
||||
base: str
|
||||
quote: str
|
||||
exchange: str
|
||||
is_active: bool = True
|
||||
|
||||
|
||||
class TickerData(BaseModel):
|
||||
symbol: str
|
||||
exchange: str
|
||||
bid: Decimal
|
||||
ask: Decimal
|
||||
last: Decimal
|
||||
volume_24h: Decimal
|
||||
change_24h: Optional[Decimal] = None
|
||||
timestamp: datetime
|
||||
|
||||
|
||||
class CandleValidationError(ValueError):
|
||||
"""Raised when OHLC data is inconsistent."""
|
||||
|
||||
def __init__(self, message: str, *, open: Decimal, high: Decimal, low: Decimal, close: Decimal) -> None:
|
||||
self.open = open
|
||||
self.high = high
|
||||
self.low = low
|
||||
self.close = close
|
||||
super().__init__(message)
|
||||
|
||||
|
||||
class OrderRequest(BaseModel):
|
||||
"""Parameters for placing an order on an exchange."""
|
||||
symbol: str
|
||||
side: str # "buy" or "sell"
|
||||
order_type: str = "market" # "market" or "limit"
|
||||
amount: Decimal # base currency amount (e.g. BTC amount)
|
||||
price: Optional[Decimal] = None # required for limit orders
|
||||
reduce_only: bool = False
|
||||
position_side: Optional[str] = None # "long" or "short" (for futures)
|
||||
|
||||
|
||||
class OrderData(BaseModel):
|
||||
"""Response from placing an order."""
|
||||
exchange: str
|
||||
symbol: str
|
||||
order_id: str
|
||||
client_order_id: Optional[str] = None
|
||||
side: str
|
||||
order_type: str
|
||||
amount: Decimal
|
||||
filled: Decimal
|
||||
price: Optional[Decimal] = None
|
||||
average: Optional[Decimal] = None
|
||||
status: str # "open", "closed", "canceled", "rejected"
|
||||
timestamp: datetime
|
||||
raw: Optional[dict] = None
|
||||
|
||||
|
||||
class BalanceData(BaseModel):
|
||||
"""Account balance for a single asset."""
|
||||
asset: str
|
||||
free: Decimal
|
||||
used: Decimal
|
||||
total: Decimal
|
||||
|
||||
|
||||
class BalanceResponse(BaseModel):
|
||||
"""Full account balance snapshot."""
|
||||
exchange: str
|
||||
balances: list[BalanceData]
|
||||
timestamp: datetime
|
||||
|
||||
|
||||
class OpenOrderData(BaseModel):
|
||||
"""An open order from the exchange."""
|
||||
order_id: str
|
||||
symbol: str
|
||||
side: str # "buy" or "sell"
|
||||
order_type: str # "market", "limit", etc.
|
||||
amount: Decimal
|
||||
filled: Decimal
|
||||
price: Optional[Decimal] = None
|
||||
average: Optional[Decimal] = None
|
||||
status: str
|
||||
timestamp: Optional[datetime] = None
|
||||
|
||||
|
||||
class PositionData(BaseModel):
|
||||
"""A futures/derivatives position."""
|
||||
symbol: str
|
||||
side: str # "long" or "short"
|
||||
contracts: Decimal
|
||||
entry_price: Optional[Decimal] = None
|
||||
mark_price: Optional[Decimal] = None
|
||||
unrealized_pnl: Optional[Decimal] = None
|
||||
leverage: Optional[Decimal] = None
|
||||
liquidation_price: Optional[Decimal] = None
|
||||
percentage: Optional[Decimal] = None
|
||||
|
||||
|
||||
class TradeData(BaseModel):
|
||||
"""A filled trade from the exchange."""
|
||||
trade_id: str
|
||||
symbol: str
|
||||
side: str # "buy" or "sell"
|
||||
amount: Decimal
|
||||
price: Decimal
|
||||
cost: Optional[Decimal] = None
|
||||
fee: Optional[Decimal] = None
|
||||
fee_currency: Optional[str] = None
|
||||
timestamp: Optional[datetime] = None
|
||||
|
||||
|
||||
class CredentialSummaryResponse(BaseModel):
|
||||
"""Aggregated summary for one credential."""
|
||||
exchange_name: str
|
||||
usdt_balance: Decimal
|
||||
usdt_estimate: str # "≈ $1,234.56" for non-USDT assets
|
||||
balance_count: int # number of tokens with balance > 0
|
||||
open_orders_count: int
|
||||
positions_count: int
|
||||
balances: list[BalanceData]
|
||||
open_orders: list[OpenOrderData]
|
||||
positions: list[PositionData]
|
||||
recent_trades: list[TradeData]
|
||||
Executable
+100
@@ -0,0 +1,100 @@
|
||||
"""
|
||||
One-shot script to force sync candle data from MEXC.
|
||||
Inserts directly into candles_default partition.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from decimal import Decimal
|
||||
|
||||
from sqlalchemy import and_, select, text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import joinedload
|
||||
|
||||
from app.database import async_session_factory
|
||||
from app.exchange.factory import factory as exchange_factory
|
||||
from app.models.exchange import Exchange
|
||||
from app.models.symbol import Symbol
|
||||
|
||||
logging.basicConfig(level=logging.INFO)
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
TIMEFRAMES = ["15m", "30m", "1h", "4h", "1d"]
|
||||
|
||||
async def sync():
|
||||
async with async_session_factory() as db:
|
||||
# Get active MEXC symbols for our tokens
|
||||
watchlist_symbols = ['BTC/USDT','ETH/USDT','SOL/USDT','NEAR/USDT','BNB/USDT','SQD/USDT','ZEN/USDT','KITE/USDT']
|
||||
|
||||
result = await db.execute(
|
||||
select(Symbol)
|
||||
.options(joinedload(Symbol.exchange))
|
||||
.where(Symbol.symbol.in_(watchlist_symbols), Symbol.is_active == True)
|
||||
)
|
||||
symbols: list[Symbol] = list(result.scalars().all())
|
||||
logger.info(f"Found {len(symbols)} active symbols to sync")
|
||||
|
||||
adapter = exchange_factory.create("mexc")
|
||||
total_inserted = 0
|
||||
|
||||
for db_symbol in symbols:
|
||||
for tf in TIMEFRAMES:
|
||||
try:
|
||||
candles = await adapter.fetch_ohlcv(
|
||||
symbol=db_symbol.symbol,
|
||||
timeframe=tf,
|
||||
limit=200,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(f"Fetch error {db_symbol.symbol} {tf}: {e}")
|
||||
continue
|
||||
|
||||
if not candles:
|
||||
continue
|
||||
|
||||
inserted = 0
|
||||
for c in candles:
|
||||
try:
|
||||
await db.execute(
|
||||
text("""
|
||||
INSERT INTO candles_default
|
||||
(symbol_id, timeframe, timestamp, open, high, low, close, volume)
|
||||
VALUES (:sid, :tf, :ts, :o, :h, :l, :c, :v)
|
||||
ON CONFLICT (symbol_id, timeframe, timestamp) DO NOTHING
|
||||
"""),
|
||||
{
|
||||
"sid": db_symbol.id,
|
||||
"tf": c.timeframe,
|
||||
"ts": c.timestamp,
|
||||
"o": Decimal(str(c.open)),
|
||||
"h": Decimal(str(c.high)),
|
||||
"l": Decimal(str(c.low)),
|
||||
"c": Decimal(str(c.close)),
|
||||
"v": Decimal(str(c.volume)),
|
||||
}
|
||||
)
|
||||
inserted += 1
|
||||
except Exception as e:
|
||||
logger.debug(f"Insert error: {e}")
|
||||
|
||||
await db.commit()
|
||||
total_inserted += inserted
|
||||
logger.info(f" ✅ {db_symbol.symbol:12s} {tf:4s}: {inserted} candles")
|
||||
|
||||
logger.info(f"\n✅ DONE! Total: {total_inserted} candles inserted")
|
||||
|
||||
# Verify
|
||||
r = await db.execute(text("""
|
||||
SELECT s.symbol, c.timeframe, MAX(c.timestamp) as ts
|
||||
FROM candles_default c
|
||||
JOIN symbols s ON s.id = c.symbol_id
|
||||
WHERE s.symbol IN ('BTC/USDT','ETH/USDT','SOL/USDT')
|
||||
GROUP BY s.symbol, c.timeframe
|
||||
ORDER BY s.symbol
|
||||
"""))
|
||||
print("\n📊 Latest candles:")
|
||||
for row in r:
|
||||
print(f" {row[0]:12s} {row[1]:4s}: {row[2]}")
|
||||
|
||||
asyncio.run(sync())
|
||||
Executable
+308
@@ -0,0 +1,308 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy import select, text
|
||||
|
||||
from app.config import settings
|
||||
from app.core.exceptions import AppException
|
||||
from app.core.middleware import RequestLoggingMiddleware
|
||||
from pydantic import ValidationError
|
||||
from app.database import async_session_factory, engine
|
||||
from app.tasks.candle_fetcher import force_full_sync
|
||||
from app.tasks.exchange_sync import sync_all_exchanges
|
||||
from app.models import Exchange
|
||||
from app.services.ws_push_service import setup_push_listener
|
||||
|
||||
# P3-4: Structured JSON logging
|
||||
class JsonFormatter(logging.Formatter):
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
log_entry = {
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"level": record.levelname,
|
||||
"logger": record.name,
|
||||
"message": record.getMessage(),
|
||||
"module": record.module,
|
||||
"line": record.lineno,
|
||||
}
|
||||
if record.exc_info and record.exc_info[0]:
|
||||
log_entry["exception"] = self.formatException(record.exc_info)
|
||||
return json.dumps(log_entry, default=str)
|
||||
|
||||
log_handler = logging.StreamHandler()
|
||||
log_handler.setFormatter(JsonFormatter())
|
||||
logging.getLogger().handlers = [log_handler]
|
||||
logging.getLogger().setLevel(getattr(logging, settings.LOG_LEVEL.upper(), logging.INFO))
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Track application start time for uptime reporting
|
||||
_start_time: float = time.time()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
"""Application lifespan handler for startup and shutdown events."""
|
||||
logger.info(
|
||||
"Starting Trading Portal API on %s:%s",
|
||||
settings.HOST,
|
||||
settings.PORT,
|
||||
)
|
||||
logger.info("Log level: %s", settings.LOG_LEVEL)
|
||||
|
||||
# --- Auto-create tables (dev bootstrap, replaces Alembic) ---
|
||||
try:
|
||||
from app.database import Base, engine
|
||||
from app import models # noqa: F401
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
logger.info("Database tables verified/created")
|
||||
except Exception:
|
||||
logger.exception("Table creation failed (non-fatal)")
|
||||
|
||||
# --- Startup: background tasks ---
|
||||
# Candle fetcher: Binance USDT only (1,075 symbols), 4 TFs (15m/1h/4h/1d),
|
||||
# batch of 25 symbols every 10 min. Full cycle ~7.2h, DB-safe.
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from app.tasks.candle_fetcher import setup_candle_scheduler
|
||||
scheduler = setup_candle_scheduler(app)
|
||||
app.state.scheduler = scheduler
|
||||
scheduler.start()
|
||||
logger.info("Background scheduler started with optimized candle fetcher")
|
||||
|
||||
# --- Schedule stale trade management (every 30 min) ---
|
||||
try:
|
||||
from apscheduler.triggers.interval import IntervalTrigger
|
||||
|
||||
async def _close_stale_trades():
|
||||
from app.services.signal_service import close_stale_trades
|
||||
await close_stale_trades()
|
||||
|
||||
scheduler.add_job(
|
||||
_close_stale_trades,
|
||||
IntervalTrigger(minutes=5),
|
||||
id="close-stale-trades",
|
||||
replace_existing=True,
|
||||
misfire_grace_time=60,
|
||||
)
|
||||
logger.info("Stale trade management scheduled (every 5 min)")
|
||||
|
||||
# P3-5: Signal aging — expire old signals daily
|
||||
async def _expire_signals():
|
||||
from app.services.signal_service import expire_old_signals
|
||||
await expire_old_signals(max_age_days=7)
|
||||
|
||||
scheduler.add_job(
|
||||
_expire_signals,
|
||||
IntervalTrigger(hours=6),
|
||||
id="expire-old-signals",
|
||||
replace_existing=True,
|
||||
misfire_grace_time=300,
|
||||
)
|
||||
logger.info("Signal aging scheduled (every 6 hours)")
|
||||
except Exception:
|
||||
logger.exception("Failed to schedule stale trade management")
|
||||
|
||||
# --- Register WebSocket push listener ---
|
||||
setup_push_listener(app)
|
||||
|
||||
# Exchange sync already done — skip on restart
|
||||
# async def _background_exchange_sync():
|
||||
# ...
|
||||
logger.info("Exchange symbol sync skipped — already synced")
|
||||
|
||||
# --- Force full candle sync on startup (DISABLED — too heavy with 16K symbols) ---
|
||||
# Background scheduler handles incremental sync every 2 min instead
|
||||
# Uncomment when symbol count is manageable:
|
||||
# async def _background_sync():
|
||||
# try:
|
||||
# await force_full_sync(app)
|
||||
# except Exception:
|
||||
# logger.exception("Force full sync failed (non-fatal)")
|
||||
# asyncio.ensure_future(_background_sync())
|
||||
logger.info("Full candle sync skipped — relying on incremental scheduler")
|
||||
|
||||
# --- Periodic win rate recalculation (every 6 hours) ---
|
||||
async def _recalc_win_rates():
|
||||
from app.services.signal_booster import compute_strategy_win_rates
|
||||
while True:
|
||||
await asyncio.sleep(21600) # 6 hours
|
||||
try:
|
||||
async with async_session_factory() as db:
|
||||
await compute_strategy_win_rates(db)
|
||||
logger.info("Win rates refreshed (periodic recalc)")
|
||||
except Exception:
|
||||
logger.exception("Win rate periodic refresh failed")
|
||||
asyncio.ensure_future(_recalc_win_rates())
|
||||
logger.info("Win rate recalc scheduled (every 6 hours)")
|
||||
|
||||
logger.info("Trading Portal API startup complete")
|
||||
yield
|
||||
|
||||
# --- Shutdown ---
|
||||
# P3-7: Graceful shutdown with DB cleanup
|
||||
try:
|
||||
scheduler.shutdown(wait=False)
|
||||
logger.info("Background scheduler shutdown complete")
|
||||
except Exception:
|
||||
logger.exception("Scheduler shutdown failed (ignoring)")
|
||||
|
||||
# Close database engine
|
||||
try:
|
||||
await engine.dispose()
|
||||
logger.info("Database engine disposed")
|
||||
except Exception:
|
||||
logger.exception("Engine disposal failed (ignoring)")
|
||||
|
||||
# Cancel remaining background tasks
|
||||
tasks = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
if tasks:
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
logger.info("Cancelled %d background tasks", len(tasks))
|
||||
|
||||
logger.info("Shutting down Trading Portal API")
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
title="Trading Portal API",
|
||||
description="Backend API for the Trading Portal application",
|
||||
version="1.0.0",
|
||||
lifespan=lifespan,
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# CORS middleware
|
||||
# ---------------------------------------------------------------------------
|
||||
if settings.CORS_ORIGINS:
|
||||
origins = [o.strip() for o in settings.CORS_ORIGINS.split(",") if o.strip()]
|
||||
else:
|
||||
origins = ["*"]
|
||||
|
||||
app.add_middleware(
|
||||
CORSMiddleware,
|
||||
allow_origins=origins,
|
||||
allow_credentials=True,
|
||||
allow_methods=["*"],
|
||||
allow_headers=["*"],
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Request logging middleware
|
||||
# ---------------------------------------------------------------------------
|
||||
app.add_middleware(RequestLoggingMiddleware)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Exception handlers
|
||||
# ---------------------------------------------------------------------------
|
||||
@app.exception_handler(AppException)
|
||||
async def app_exception_handler(request: Request, exc: AppException) -> JSONResponse:
|
||||
"""Return a JSON response for application-level exceptions."""
|
||||
return JSONResponse(
|
||||
status_code=exc.status_code,
|
||||
content={"detail": exc.detail, "code": exc.code},
|
||||
)
|
||||
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
async def unhandled_exception_handler(request: Request, exc: Exception) -> JSONResponse:
|
||||
"""Catch any unhandled exception and return a generic 500 response."""
|
||||
logger.exception("Unhandled exception: %s", str(exc))
|
||||
return JSONResponse(
|
||||
status_code=500,
|
||||
content={"detail": "Internal server error", "code": "internal_error"},
|
||||
)
|
||||
|
||||
|
||||
@app.exception_handler(404)
|
||||
async def not_found_handler(request: Request, exc) -> JSONResponse:
|
||||
"""Return a JSON 404 for unmatched routes."""
|
||||
return JSONResponse(
|
||||
status_code=404,
|
||||
content={"detail": "Not found", "code": "not_found"},
|
||||
)
|
||||
|
||||
|
||||
@app.exception_handler(ValidationError)
|
||||
async def validation_error_handler(request: Request, exc: ValidationError) -> JSONResponse:
|
||||
"""Return a JSON 422 with validation error details."""
|
||||
return JSONResponse(
|
||||
status_code=422,
|
||||
content={
|
||||
"detail": "Validation error",
|
||||
"code": "validation_error",
|
||||
"errors": exc.errors(),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Health check
|
||||
# ---------------------------------------------------------------------------
|
||||
@app.get("/health")
|
||||
async def health_check():
|
||||
"""Return health status including version, uptime, and database connectivity.
|
||||
|
||||
NOTE: Does NOT test exchange connectivity to avoid thread pool exhaustion
|
||||
and Docker health check timeouts.
|
||||
"""
|
||||
db_connected = False
|
||||
db_latency = 0.0
|
||||
try:
|
||||
async with async_session_factory() as db_session:
|
||||
start = time.time()
|
||||
await db_session.execute(text("SELECT 1"))
|
||||
db_latency = (time.time() - start) * 1000 # ms
|
||||
db_connected = True
|
||||
except Exception:
|
||||
db_connected = False
|
||||
|
||||
uptime_seconds = time.time() - _start_time
|
||||
|
||||
health = {
|
||||
"status": "healthy" if db_connected else "degraded",
|
||||
"version": "1.0.0",
|
||||
"uptime": round(uptime_seconds, 2),
|
||||
"db_connected": db_connected,
|
||||
"db_latency_ms": round(db_latency, 2),
|
||||
}
|
||||
return health
|
||||
|
||||
|
||||
async def _test_exchange_connection(exchange_name: str) -> bool:
|
||||
"""Try connecting to an exchange to verify it's reachable."""
|
||||
try:
|
||||
import asyncio
|
||||
|
||||
from app.exchange.factory import factory
|
||||
|
||||
adapter = factory.create(exchange_name)
|
||||
loop = asyncio.get_running_loop()
|
||||
await loop.run_in_executor(None, lambda: adapter.client.load_markets())
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# V1 API routers
|
||||
# ---------------------------------------------------------------------------
|
||||
from app.api.v1.router import api_router
|
||||
|
||||
app.include_router(api_router)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# WebSocket router
|
||||
# ---------------------------------------------------------------------------
|
||||
from app.api.ws.candle_handler import router as ws_router
|
||||
|
||||
app.include_router(ws_router)
|
||||
@@ -0,0 +1,165 @@
|
||||
"""
|
||||
Entrypoint: API-only service (no scheduler).
|
||||
User-facing HTTP API — always fast, never blocked by candle fetching.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import time
|
||||
from contextlib import asynccontextmanager
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.middleware.cors import CORSMiddleware
|
||||
from fastapi.responses import JSONResponse
|
||||
from sqlalchemy import text
|
||||
from pydantic import ValidationError
|
||||
|
||||
from app.config import settings
|
||||
from app.core.exceptions import AppException
|
||||
from app.core.middleware import RequestLoggingMiddleware
|
||||
from app.database import async_session_factory, engine
|
||||
|
||||
# Structured JSON logging
|
||||
class JsonFormatter(logging.Formatter):
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
log_entry = {
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"level": record.levelname,
|
||||
"logger": record.name,
|
||||
"message": record.getMessage(),
|
||||
"module": record.module,
|
||||
"line": record.lineno,
|
||||
}
|
||||
if record.exc_info and record.exc_info[0]:
|
||||
log_entry["exception"] = self.formatException(record.exc_info)
|
||||
return json.dumps(log_entry, default=str)
|
||||
|
||||
log_handler = logging.StreamHandler()
|
||||
log_handler.setFormatter(JsonFormatter())
|
||||
logging.getLogger().handlers = [log_handler]
|
||||
logging.getLogger().setLevel(getattr(logging, settings.LOG_LEVEL.upper(), logging.INFO))
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
_start_time: float = time.time()
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def lifespan(app: FastAPI):
|
||||
logger.info("Starting Trading Portal API on %s:%s", settings.HOST, settings.PORT)
|
||||
logger.info("Log level: %s", settings.LOG_LEVEL)
|
||||
|
||||
# Auto-create tables
|
||||
try:
|
||||
from app.database import Base, engine
|
||||
from app import models # noqa: F401
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
logger.info("Database tables verified/created")
|
||||
except Exception:
|
||||
logger.exception("Table creation failed (non-fatal)")
|
||||
|
||||
# Register WebSocket push listener
|
||||
from app.services.ws_push_service import setup_push_listener
|
||||
setup_push_listener(app)
|
||||
|
||||
logger.info("Trading Portal API startup complete (API-only, no scheduler)")
|
||||
yield
|
||||
|
||||
# Shutdown
|
||||
try:
|
||||
await engine.dispose()
|
||||
logger.info("Database engine disposed")
|
||||
except Exception:
|
||||
logger.exception("Engine disposal failed")
|
||||
|
||||
|
||||
app = FastAPI(
|
||||
title="Trading Portal API",
|
||||
description="User-facing API for Trading Portal",
|
||||
version="2.0.0",
|
||||
lifespan=lifespan,
|
||||
)
|
||||
|
||||
# CORS
|
||||
if settings.CORS_ORIGINS:
|
||||
origins = [o.strip() for o in settings.CORS_ORIGINS.split(",") if o.strip()]
|
||||
else:
|
||||
origins = ["*"]
|
||||
|
||||
app.add_middleware(CORSMiddleware, allow_origins=origins, allow_credentials=True,
|
||||
allow_methods=["*"], allow_headers=["*"])
|
||||
app.add_middleware(RequestLoggingMiddleware)
|
||||
|
||||
|
||||
# Exception handlers
|
||||
@app.exception_handler(AppException)
|
||||
async def app_exception_handler(request: Request, exc: AppException) -> JSONResponse:
|
||||
return JSONResponse(status_code=exc.status_code,
|
||||
content={"detail": exc.detail, "code": exc.code})
|
||||
|
||||
|
||||
@app.exception_handler(Exception)
|
||||
async def unhandled_exception_handler(request: Request, exc: Exception) -> JSONResponse:
|
||||
logger.exception("Unhandled exception: %s", str(exc))
|
||||
return JSONResponse(status_code=500,
|
||||
content={"detail": "Internal server error", "code": "internal_error"})
|
||||
|
||||
|
||||
@app.exception_handler(404)
|
||||
async def not_found_handler(request: Request, exc) -> JSONResponse:
|
||||
return JSONResponse(status_code=404,
|
||||
content={"detail": "Not found", "code": "not_found"})
|
||||
|
||||
|
||||
@app.exception_handler(ValidationError)
|
||||
async def validation_error_handler(request: Request, exc: ValidationError) -> JSONResponse:
|
||||
from uuid import UUID as _UUID
|
||||
|
||||
errors: list[dict] = []
|
||||
for e in exc.errors():
|
||||
# Convert non-JSON-serializable values (UUID, datetime, etc.)
|
||||
safe = {}
|
||||
for k, v in e.items():
|
||||
if isinstance(v, _UUID):
|
||||
safe[k] = str(v)
|
||||
elif isinstance(v, (datetime,)):
|
||||
safe[k] = v.isoformat()
|
||||
else:
|
||||
safe[k] = v
|
||||
errors.append(safe)
|
||||
return JSONResponse(status_code=422,
|
||||
content={"detail": "Validation error", "code": "validation_error",
|
||||
"errors": errors})
|
||||
|
||||
|
||||
# Health check
|
||||
@app.get("/health")
|
||||
async def health_check():
|
||||
db_connected = False
|
||||
db_latency = 0.0
|
||||
try:
|
||||
async with async_session_factory() as db_session:
|
||||
start = time.time()
|
||||
await db_session.execute(text("SELECT 1"))
|
||||
db_latency = (time.time() - start) * 1000
|
||||
db_connected = True
|
||||
except Exception:
|
||||
db_connected = False
|
||||
|
||||
return {
|
||||
"status": "healthy" if db_connected else "degraded",
|
||||
"version": "2.0.0-api",
|
||||
"uptime": round(time.time() - _start_time, 2),
|
||||
"db_connected": db_connected,
|
||||
"db_latency_ms": round(db_latency, 2),
|
||||
}
|
||||
|
||||
|
||||
# API routes
|
||||
from app.api.v1.router import api_router
|
||||
app.include_router(api_router)
|
||||
|
||||
from app.api.ws.candle_handler import router as ws_router
|
||||
app.include_router(ws_router)
|
||||
@@ -0,0 +1,183 @@
|
||||
"""Entrypoint: Scheduler-only service (no HTTP server).
|
||||
Runs candle fetcher, signal analysis, trade management, win rate recalc.
|
||||
Heavy CPU work stays here, API stays fast.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
import signal
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from apscheduler.triggers.interval import IntervalTrigger
|
||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||
|
||||
from app.config import settings
|
||||
from app.database import Base
|
||||
|
||||
|
||||
# ── Scheduler-specific engine: longer statement_timeout (30s) ──
|
||||
# Unlike the API which needs fast-fail (<10s), the scheduler does heavy
|
||||
# batch UPDATE queries (trade eviction executemany) that can take longer.
|
||||
_scheduler_engine = create_async_engine(
|
||||
settings.DATABASE_URL,
|
||||
echo=(settings.LOG_LEVEL == "DEBUG"),
|
||||
pool_pre_ping=True,
|
||||
pool_size=20, # semaphore=8 limits concurrent analyses
|
||||
max_overflow=5,
|
||||
pool_timeout=5,
|
||||
pool_recycle=600,
|
||||
connect_args={
|
||||
"server_settings": {
|
||||
"idle_in_transaction_session_timeout": "60000",
|
||||
"statement_timeout": "30000", # 30s vs API's 10s
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
_scheduler_session_factory = async_sessionmaker(
|
||||
_scheduler_engine,
|
||||
class_=AsyncSession,
|
||||
expire_on_commit=False,
|
||||
)
|
||||
|
||||
# Monkey-patch app.database so all downstream imports use scheduler engine
|
||||
import app.database as _db_module
|
||||
_db_module.engine = _scheduler_engine
|
||||
_db_module.async_session_factory = _scheduler_session_factory
|
||||
|
||||
|
||||
class JsonFormatter(logging.Formatter):
|
||||
def format(self, record: logging.LogRecord) -> str:
|
||||
log_entry = {
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"level": record.levelname,
|
||||
"logger": record.name,
|
||||
"message": record.getMessage(),
|
||||
"module": record.module,
|
||||
"line": record.lineno,
|
||||
}
|
||||
if record.exc_info and record.exc_info[0]:
|
||||
log_entry["exception"] = self.formatException(record.exc_info)
|
||||
return json.dumps(log_entry, default=str)
|
||||
|
||||
|
||||
log_handler = logging.StreamHandler()
|
||||
log_handler.setFormatter(JsonFormatter())
|
||||
logging.getLogger().handlers = [log_handler]
|
||||
logging.getLogger().setLevel(getattr(logging, settings.LOG_LEVEL.upper(), logging.INFO))
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def _close_stale_trades():
|
||||
from app.services.signal_service import close_stale_trades
|
||||
await close_stale_trades()
|
||||
|
||||
|
||||
async def _expire_signals():
|
||||
from app.services.signal_service import expire_old_signals
|
||||
await expire_old_signals(max_age_days=7)
|
||||
|
||||
|
||||
async def _recalc_win_rates_loop():
|
||||
"""Periodic win rate recalculation every 6 hours."""
|
||||
from app.services.signal_booster import compute_strategy_win_rates
|
||||
while True:
|
||||
await asyncio.sleep(21600) # 6 hours
|
||||
try:
|
||||
async with async_session_factory() as db:
|
||||
await compute_strategy_win_rates(db)
|
||||
logger.info("Win rates refreshed (periodic recalc)")
|
||||
except Exception:
|
||||
logger.exception("Win rate periodic refresh failed")
|
||||
|
||||
|
||||
async def main():
|
||||
logger.info("=== Trading Portal Scheduler ===")
|
||||
logger.info("Mode: scheduler-only (no HTTP server)")
|
||||
|
||||
# Auto-create tables (first run)
|
||||
try:
|
||||
from app.database import Base
|
||||
from app import models # noqa: F401
|
||||
async with _scheduler_engine.begin() as conn:
|
||||
await conn.run_sync(Base.metadata.create_all)
|
||||
logger.info("Database tables verified/created")
|
||||
except Exception:
|
||||
logger.exception("Table creation failed (non-fatal)")
|
||||
|
||||
# Create scheduler
|
||||
scheduler = AsyncIOScheduler()
|
||||
|
||||
# Candle fetcher — the core job (5-min interval)
|
||||
from app.tasks.candle_fetcher import fetch_recent_candles, _TIMEFRAMES_OPTIMIZED
|
||||
|
||||
scheduler.add_job(
|
||||
fetch_recent_candles,
|
||||
trigger=IntervalTrigger(seconds=300),
|
||||
args=[None, 2, _TIMEFRAMES_OPTIMIZED, 25],
|
||||
id="fetch_candles_optimized",
|
||||
name="Fetch trading candles (top 100 bases, 5 exchanges, 4TFs)",
|
||||
replace_existing=True,
|
||||
coalesce=True,
|
||||
max_instances=1,
|
||||
misfire_grace_time=600,
|
||||
)
|
||||
logger.info("Candle fetcher scheduled: every 5 min, 25/batch")
|
||||
|
||||
# Stale trade management (every 5 min)
|
||||
scheduler.add_job(
|
||||
_close_stale_trades,
|
||||
IntervalTrigger(minutes=5),
|
||||
id="close-stale-trades",
|
||||
replace_existing=True,
|
||||
)
|
||||
|
||||
# Real trade sync (every 5 min)
|
||||
from app.services.trade_executor import sync_real_trades
|
||||
scheduler.add_job(
|
||||
sync_real_trades,
|
||||
IntervalTrigger(minutes=5),
|
||||
id="sync-real-trades",
|
||||
replace_existing=True,
|
||||
)
|
||||
|
||||
# Signal aging (every 6 hours)
|
||||
scheduler.add_job(
|
||||
_expire_signals,
|
||||
IntervalTrigger(hours=6),
|
||||
id="expire-old-signals",
|
||||
replace_existing=True,
|
||||
)
|
||||
|
||||
# Win rate recalc (6h loop)
|
||||
asyncio.ensure_future(_recalc_win_rates_loop())
|
||||
|
||||
# Start
|
||||
scheduler.start()
|
||||
logger.info("Scheduler started with all jobs")
|
||||
|
||||
# Keep alive forever
|
||||
stop_event = asyncio.Event()
|
||||
|
||||
def _shutdown(sig, frame):
|
||||
logger.info("Received signal %s — shutting down", sig)
|
||||
stop_event.set()
|
||||
|
||||
signal.signal(signal.SIGTERM, _shutdown)
|
||||
signal.signal(signal.SIGINT, _shutdown)
|
||||
|
||||
await stop_event.wait()
|
||||
|
||||
# Graceful shutdown
|
||||
scheduler.shutdown(wait=False)
|
||||
await _scheduler_engine.dispose()
|
||||
logger.info("Scheduler shutdown complete")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
Executable
+24
@@ -0,0 +1,24 @@
|
||||
from app.models.user import User
|
||||
from app.models.exchange import Exchange
|
||||
from app.models.symbol import Symbol
|
||||
from app.models.candle import Candle
|
||||
from app.models.watchlist import Watchlist
|
||||
from app.models.credential import ExchangeCredential
|
||||
from app.models.refresh_token import RefreshToken
|
||||
from app.models.real_trade import RealTrade
|
||||
from app.models.signal import Signal, HypotheticalTrade
|
||||
from app.models.audit_log import AuditLog
|
||||
|
||||
__all__ = [
|
||||
"User",
|
||||
"Exchange",
|
||||
"Symbol",
|
||||
"Candle",
|
||||
"Watchlist",
|
||||
"ExchangeCredential",
|
||||
"RefreshToken",
|
||||
"Signal",
|
||||
"HypotheticalTrade",
|
||||
"RealTrade",
|
||||
"AuditLog",
|
||||
]
|
||||
Executable
+71
@@ -0,0 +1,71 @@
|
||||
"""AlertCondition ORM model — multi-condition user alerts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import (
|
||||
Boolean,
|
||||
DateTime,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
JSON,
|
||||
String,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import UUID as PGUUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class AlertCondition(Base):
|
||||
"""A user-defined multi-condition alert.
|
||||
|
||||
Each alert consists of a name, a list of conditions (stored as JSON),
|
||||
and notification preferences. Alerts are evaluated after each signal
|
||||
generation; all active alerts whose conditions are satisfied trigger
|
||||
a notification.
|
||||
"""
|
||||
|
||||
__tablename__ = "alert_conditions"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
PGUUID(as_uuid=True),
|
||||
ForeignKey("users.id"),
|
||||
nullable=False,
|
||||
comment="The user who owns this alert",
|
||||
)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(100), nullable=False, comment="Human-readable alert name"
|
||||
)
|
||||
conditions: Mapped[list] = mapped_column(
|
||||
JSON,
|
||||
nullable=False,
|
||||
default=list,
|
||||
comment="Array of condition objects, e.g. [{'indicator': 'rsi', 'operator': '>', 'value': 70, 'timeframe': '1h'}]",
|
||||
)
|
||||
notify_platform: Mapped[str] = mapped_column(
|
||||
String(20),
|
||||
nullable=False,
|
||||
default="telegram",
|
||||
comment="Notification channel: 'telegram', 'discord', or 'both'",
|
||||
)
|
||||
is_active: Mapped[bool] = mapped_column(
|
||||
Boolean, default=True, nullable=False
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=datetime.now
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=datetime.now, onupdate=datetime.now
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_alert_conditions_user_active", "user_id", "is_active"),
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<AlertCondition id={self.id} name={self.name!r} user_id={self.user_id}>"
|
||||
Executable
+48
@@ -0,0 +1,48 @@
|
||||
"""Audit Log model for tracking actions across the system."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import JSON, DateTime, Index, Integer, String
|
||||
from sqlalchemy.dialects.postgresql import UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
from sqlalchemy.sql import func
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class AuditLog(Base):
|
||||
__tablename__ = "audit_logs"
|
||||
|
||||
id: Mapped[int] = mapped_column(
|
||||
Integer, primary_key=True, autoincrement=True
|
||||
)
|
||||
user_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
UUID(as_uuid=True), nullable=True, index=True
|
||||
)
|
||||
action: Mapped[str] = mapped_column(
|
||||
String(50), nullable=False
|
||||
)
|
||||
resource: Mapped[str] = mapped_column(
|
||||
String(100), nullable=False
|
||||
)
|
||||
details: Mapped[dict | None] = mapped_column(
|
||||
JSON, nullable=True
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True),
|
||||
server_default=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_audit_logs_action_created_at", "action", "created_at"),
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"<AuditLog id={self.id} action={self.action!r} "
|
||||
f"resource={self.resource!r}>"
|
||||
)
|
||||
Executable
+71
@@ -0,0 +1,71 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
|
||||
from sqlalchemy import (
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
Numeric,
|
||||
PrimaryKeyConstraint,
|
||||
String,
|
||||
desc,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import TIMESTAMP
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class Candle(Base):
|
||||
__tablename__ = "candles"
|
||||
|
||||
symbol_id: Mapped[int] = mapped_column(
|
||||
Integer, ForeignKey("symbols.id"), nullable=False
|
||||
)
|
||||
timeframe: Mapped[str] = mapped_column(
|
||||
String(10), nullable=False
|
||||
)
|
||||
timestamp: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), nullable=False
|
||||
)
|
||||
open: Mapped[Decimal] = mapped_column(
|
||||
Numeric(20, 8), nullable=False
|
||||
)
|
||||
high: Mapped[Decimal] = mapped_column(
|
||||
Numeric(20, 8), nullable=False
|
||||
)
|
||||
low: Mapped[Decimal] = mapped_column(
|
||||
Numeric(20, 8), nullable=False
|
||||
)
|
||||
close: Mapped[Decimal] = mapped_column(
|
||||
Numeric(20, 8), nullable=False
|
||||
)
|
||||
volume: Mapped[Decimal] = mapped_column(
|
||||
Numeric(30, 8), nullable=False
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
PrimaryKeyConstraint(
|
||||
"symbol_id", "timeframe", "timestamp",
|
||||
name="pk_candles",
|
||||
),
|
||||
Index(
|
||||
"ix_candles_symbol_timeframe_ts_desc",
|
||||
"symbol_id",
|
||||
"timeframe",
|
||||
desc("timestamp"),
|
||||
),
|
||||
{
|
||||
"postgresql_partition_by": "RANGE (timestamp)",
|
||||
"info": {"partitioned": True},
|
||||
},
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"<Candle symbol_id={self.symbol_id} "
|
||||
f"timeframe={self.timeframe!r} "
|
||||
f"timestamp={self.timestamp}>"
|
||||
)
|
||||
Executable
+73
@@ -0,0 +1,73 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, ForeignKey, Integer, String
|
||||
from sqlalchemy.dialects.postgresql import UUID, TIMESTAMP
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.sql import func
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class ExchangeCredential(Base):
|
||||
__tablename__ = "exchange_credentials"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=func.gen_random_uuid(),
|
||||
)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
)
|
||||
exchange_id: Mapped[int] = mapped_column(
|
||||
Integer,
|
||||
ForeignKey("exchanges.id"),
|
||||
nullable=False,
|
||||
)
|
||||
api_key: Mapped[str] = mapped_column(
|
||||
String(255), nullable=False
|
||||
)
|
||||
api_secret_enc: Mapped[str] = mapped_column(
|
||||
String(512), nullable=False
|
||||
)
|
||||
api_secret_iv: Mapped[str] = mapped_column(
|
||||
String(64), nullable=False
|
||||
)
|
||||
passphrase: Mapped[str | None] = mapped_column(
|
||||
String(255), nullable=True
|
||||
)
|
||||
is_testnet: Mapped[bool] = mapped_column(
|
||||
Boolean, default=False, nullable=False
|
||||
)
|
||||
is_active: Mapped[bool] = mapped_column(
|
||||
Boolean, default=True, nullable=False
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True),
|
||||
server_default=func.now(),
|
||||
onupdate=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user: Mapped["User"] = relationship(
|
||||
"User", back_populates="credentials"
|
||||
)
|
||||
exchange: Mapped["Exchange"] = relationship(
|
||||
"Exchange", back_populates="credentials"
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"<ExchangeCredential id={self.id} "
|
||||
f"user_id={self.user_id} "
|
||||
f"exchange_id={self.exchange_id}>"
|
||||
)
|
||||
Executable
+52
@@ -0,0 +1,52 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, Index, Integer, String
|
||||
from sqlalchemy.dialects.postgresql import TIMESTAMP
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.sql import func
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class Exchange(Base):
|
||||
__tablename__ = "exchanges"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
name: Mapped[str] = mapped_column(
|
||||
String(50), unique=True, nullable=False
|
||||
)
|
||||
display_name: Mapped[str | None] = mapped_column(
|
||||
String(100), nullable=True
|
||||
)
|
||||
base_url: Mapped[str | None] = mapped_column(
|
||||
String(255), nullable=True
|
||||
)
|
||||
ws_url: Mapped[str | None] = mapped_column(
|
||||
String(255), nullable=True
|
||||
)
|
||||
is_active: Mapped[bool] = mapped_column(
|
||||
Boolean, default=True, nullable=False
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
|
||||
)
|
||||
|
||||
# P2-8: Index for active exchanges
|
||||
__table_args__ = (
|
||||
Index("ix_exchanges_is_active", "is_active"),
|
||||
)
|
||||
|
||||
# Relationships
|
||||
symbols: Mapped[list["Symbol"]] = relationship(
|
||||
"Symbol", back_populates="exchange", cascade="all, delete-orphan"
|
||||
)
|
||||
credentials: Mapped[list["ExchangeCredential"]] = relationship(
|
||||
"ExchangeCredential",
|
||||
back_populates="exchange",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Exchange id={self.id} name={self.name!r}>"
|
||||
Executable
+95
@@ -0,0 +1,95 @@
|
||||
"""RealTrade ORM model for actual exchange trade tracking."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
|
||||
from sqlalchemy import (
|
||||
DateTime,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
Numeric,
|
||||
String,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import UUID as PG_UUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class RealTrade(Base):
|
||||
"""A real trade placed on a connected exchange and persisted locally."""
|
||||
|
||||
__tablename__ = "real_trades"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
PG_UUID(as_uuid=True), ForeignKey("users.id"), nullable=False,
|
||||
comment="FK to the user who placed this trade",
|
||||
)
|
||||
exchange: Mapped[str] = mapped_column(String(20), nullable=False)
|
||||
symbol: Mapped[str] = mapped_column(String(50), nullable=False)
|
||||
|
||||
side: Mapped[str] = mapped_column(
|
||||
String(10), nullable=False, comment="buy or sell"
|
||||
)
|
||||
order_type: Mapped[str] = mapped_column(
|
||||
String(10), nullable=False, default="market",
|
||||
comment="market or limit",
|
||||
)
|
||||
|
||||
# Order details
|
||||
amount: Mapped[Decimal] = mapped_column(Numeric(20, 8), nullable=False)
|
||||
price: Mapped[Decimal | None] = mapped_column(
|
||||
Numeric(20, 8), nullable=True,
|
||||
comment="Limit price (null for market orders)",
|
||||
)
|
||||
filled_amount: Mapped[Decimal] = mapped_column(
|
||||
Numeric(20, 8), nullable=False, default=Decimal("0"),
|
||||
)
|
||||
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(20), nullable=False, default="open",
|
||||
comment="open / filled / cancelled / rejected",
|
||||
)
|
||||
|
||||
# P&L (computed when closed)
|
||||
pnl: Mapped[Decimal | None] = mapped_column(
|
||||
Numeric(20, 8), nullable=True,
|
||||
comment="Realised P&L in quote currency",
|
||||
)
|
||||
pnl_percent: Mapped[Decimal | None] = mapped_column(
|
||||
Numeric(10, 4), nullable=True,
|
||||
comment="P&L as percentage",
|
||||
)
|
||||
|
||||
# Exchange reference
|
||||
order_id: Mapped[str | None] = mapped_column(
|
||||
String(100), nullable=True,
|
||||
comment="Exchange-side order ID",
|
||||
)
|
||||
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False,
|
||||
default=datetime.utcnow,
|
||||
)
|
||||
closed_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True,
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_real_trades_user_id", "user_id"),
|
||||
Index("ix_real_trades_symbol", "symbol"),
|
||||
Index("ix_real_trades_status", "status"),
|
||||
Index("ix_real_trades_created_at", "created_at"),
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"<RealTrade #{self.id} {self.side.upper()} "
|
||||
f"{self.amount} {self.symbol} on {self.exchange} "
|
||||
f"status={self.status}>"
|
||||
)
|
||||
Executable
+56
@@ -0,0 +1,56 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, ForeignKey, Index, String
|
||||
from sqlalchemy.dialects.postgresql import INET, UUID, TIMESTAMP
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.sql import func
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class RefreshToken(Base):
|
||||
__tablename__ = "refresh_tokens"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=func.gen_random_uuid(),
|
||||
)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
)
|
||||
token_hash: Mapped[str] = mapped_column(
|
||||
String(64), nullable=False
|
||||
)
|
||||
expires_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), nullable=False
|
||||
)
|
||||
revoked: Mapped[bool] = mapped_column(
|
||||
Boolean, default=False, nullable=False
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
|
||||
)
|
||||
user_agent: Mapped[str | None] = mapped_column(
|
||||
String(255), nullable=True
|
||||
)
|
||||
ip_address: Mapped[str | None] = mapped_column(
|
||||
INET, nullable=True
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_refresh_tokens_user_id", "user_id"),
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user: Mapped["User"] = relationship(
|
||||
"User", back_populates="refresh_tokens"
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<RefreshToken id={self.id} user_id={self.user_id}>"
|
||||
Executable
+98
@@ -0,0 +1,98 @@
|
||||
"""Signal and HypotheticalTrade ORM models for trading signal detection."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
import uuid
|
||||
|
||||
from sqlalchemy import (
|
||||
DateTime,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
Numeric,
|
||||
String,
|
||||
Text,
|
||||
)
|
||||
from sqlalchemy.dialects.postgresql import UUID as PGUUID
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class Signal(Base):
|
||||
"""A detected trading signal based on Double Bollinger Bands + RSI analysis."""
|
||||
|
||||
__tablename__ = "signals"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
symbol: Mapped[str] = mapped_column(String(50), nullable=False, index=True)
|
||||
exchange: Mapped[str] = mapped_column(String(20), nullable=False, default="mexc")
|
||||
timeframe: Mapped[str] = mapped_column(String(10), nullable=False)
|
||||
signal_type: Mapped[str] = mapped_column(String(30), nullable=False)
|
||||
strength: Mapped[str] = mapped_column(String(20), nullable=False)
|
||||
price: Mapped[Decimal] = mapped_column(Numeric(20, 8), nullable=False)
|
||||
timestamp: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False
|
||||
)
|
||||
indicators_snapshot: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
status: Mapped[str] = mapped_column(String(20), nullable=False, default="ACTIVE")
|
||||
note: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=datetime.now
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_signals_symbol_created", "symbol", "created_at"),
|
||||
)
|
||||
|
||||
|
||||
class HypotheticalTrade(Base):
|
||||
"""A hypothetical (paper) trade automatically entered on signal generation.
|
||||
|
||||
When a STRONG_BUY or BUY signal fires, a LONG trade is opened.
|
||||
When a STRONG_SELL or SELL signal fires, a SHORT trade is opened.
|
||||
Trades are closed when an opposing signal fires or when stopped out.
|
||||
"""
|
||||
|
||||
__tablename__ = "hypothetical_trades"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
user_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||
PGUUID(as_uuid=True), ForeignKey("users.id"), nullable=True,
|
||||
comment="The user who owns this trade",
|
||||
)
|
||||
signal_id: Mapped[int | None] = mapped_column(
|
||||
Integer, ForeignKey("signals.id"), nullable=True,
|
||||
comment="The signal that opened this trade"
|
||||
)
|
||||
symbol: Mapped[str] = mapped_column(String(50), nullable=False)
|
||||
exchange: Mapped[str] = mapped_column(String(20), nullable=False, default="mexc")
|
||||
timeframe: Mapped[str] = mapped_column(String(10), nullable=False)
|
||||
direction: Mapped[str] = mapped_column(String(10), nullable=False)
|
||||
entry_price: Mapped[Decimal] = mapped_column(Numeric(20, 8), nullable=False)
|
||||
entry_time: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False
|
||||
)
|
||||
entry_reason: Mapped[str | None] = mapped_column(String(30), nullable=True)
|
||||
exit_price: Mapped[Decimal | None] = mapped_column(Numeric(20, 8), nullable=True)
|
||||
exit_time: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
exit_reason: Mapped[str | None] = mapped_column(String(30), nullable=True)
|
||||
quantity: Mapped[Decimal] = mapped_column(Numeric(20, 8), nullable=False)
|
||||
pnl: Mapped[Decimal | None] = mapped_column(Numeric(20, 8), nullable=True)
|
||||
pnl_percent: Mapped[Decimal | None] = mapped_column(Numeric(14, 4), nullable=True)
|
||||
status: Mapped[str] = mapped_column(String(10), nullable=False, default="OPEN")
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=datetime.now
|
||||
)
|
||||
closed_at: Mapped[datetime | None] = mapped_column(
|
||||
DateTime(timezone=True), nullable=True
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
Index("ix_hyp_trades_symbol", "symbol"),
|
||||
Index("ix_hyp_trades_status", "status"),
|
||||
)
|
||||
Executable
+46
@@ -0,0 +1,46 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy import Boolean, ForeignKey, Index, Integer, String, UniqueConstraint
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class Symbol(Base):
|
||||
__tablename__ = "symbols"
|
||||
|
||||
id: Mapped[int] = mapped_column(Integer, primary_key=True, autoincrement=True)
|
||||
exchange_id: Mapped[int] = mapped_column(
|
||||
Integer, ForeignKey("exchanges.id"), nullable=False
|
||||
)
|
||||
symbol: Mapped[str] = mapped_column(
|
||||
String(50), nullable=False
|
||||
)
|
||||
base: Mapped[str] = mapped_column(
|
||||
String(20), nullable=False
|
||||
)
|
||||
quote: Mapped[str] = mapped_column(
|
||||
String(20), nullable=False
|
||||
)
|
||||
is_active: Mapped[bool] = mapped_column(
|
||||
Boolean, default=True, nullable=False
|
||||
)
|
||||
is_trading: Mapped[bool] = mapped_column(
|
||||
Boolean, default=False, nullable=False
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint("exchange_id", "symbol", name="uq_symbol_exchange_symbol"),
|
||||
Index("ix_symbols_is_active", "is_active"), # P2-8
|
||||
)
|
||||
|
||||
# Relationships
|
||||
exchange: Mapped["Exchange"] = relationship(
|
||||
"Exchange", back_populates="symbols"
|
||||
)
|
||||
watchlists: Mapped[list["Watchlist"]] = relationship(
|
||||
"Watchlist", back_populates="symbol", cascade="all, delete-orphan"
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<Symbol id={self.id} symbol={self.symbol!r}>"
|
||||
Executable
+72
@@ -0,0 +1,72 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import Boolean, JSON, String, UniqueConstraint
|
||||
from sqlalchemy.dialects.postgresql import UUID, TIMESTAMP
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.sql import func
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class User(Base):
|
||||
__tablename__ = "users"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=func.gen_random_uuid(),
|
||||
)
|
||||
username: Mapped[str] = mapped_column(
|
||||
String(50), unique=True, nullable=False
|
||||
)
|
||||
email: Mapped[str] = mapped_column(
|
||||
String(255), unique=True, nullable=False
|
||||
)
|
||||
password_hash: Mapped[str] = mapped_column(
|
||||
String(255), nullable=False
|
||||
)
|
||||
display_name: Mapped[str | None] = mapped_column(
|
||||
String(100), nullable=True
|
||||
)
|
||||
is_active: Mapped[bool] = mapped_column(
|
||||
Boolean, default=True, nullable=False
|
||||
)
|
||||
is_admin: Mapped[bool] = mapped_column(
|
||||
Boolean, default=False, nullable=False
|
||||
)
|
||||
role: Mapped[str] = mapped_column(
|
||||
String(20), default="trader", nullable=False
|
||||
)
|
||||
preferences: Mapped[dict | None] = mapped_column(
|
||||
JSON, default={"default_exchange": "mexc", "default_timeframe": "1h"}, nullable=True
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True),
|
||||
server_default=func.now(),
|
||||
onupdate=func.now(),
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
# Relationships
|
||||
watchlists: Mapped[list["Watchlist"]] = relationship(
|
||||
"Watchlist", back_populates="user", cascade="all, delete-orphan"
|
||||
)
|
||||
credentials: Mapped[list["ExchangeCredential"]] = relationship(
|
||||
"ExchangeCredential",
|
||||
back_populates="user",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
refresh_tokens: Mapped[list["RefreshToken"]] = relationship(
|
||||
"RefreshToken",
|
||||
back_populates="user",
|
||||
cascade="all, delete-orphan",
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"<User id={self.id} username={self.username!r}>"
|
||||
Executable
+61
@@ -0,0 +1,61 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import ForeignKey, Integer, String, UniqueConstraint
|
||||
from sqlalchemy.dialects.postgresql import UUID, TIMESTAMP
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
from sqlalchemy.sql import func
|
||||
|
||||
from app.database import Base
|
||||
|
||||
|
||||
class Watchlist(Base):
|
||||
__tablename__ = "user_watchlists"
|
||||
|
||||
id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
primary_key=True,
|
||||
default=func.gen_random_uuid(),
|
||||
)
|
||||
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||
UUID(as_uuid=True),
|
||||
ForeignKey("users.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
)
|
||||
symbol_id: Mapped[int] = mapped_column(
|
||||
Integer,
|
||||
ForeignKey("symbols.id", ondelete="CASCADE"),
|
||||
nullable=False,
|
||||
)
|
||||
label: Mapped[str | None] = mapped_column(
|
||||
String(50), nullable=True
|
||||
)
|
||||
sort_order: Mapped[int] = mapped_column(
|
||||
Integer, default=0, nullable=False
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
TIMESTAMP(timezone=True), server_default=func.now(), nullable=False
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
UniqueConstraint(
|
||||
"user_id", "symbol_id", name="uq_watchlist_user_symbol"
|
||||
),
|
||||
)
|
||||
|
||||
# Relationships
|
||||
user: Mapped["User"] = relationship(
|
||||
"User", back_populates="watchlists"
|
||||
)
|
||||
symbol: Mapped["Symbol"] = relationship(
|
||||
"Symbol", back_populates="watchlists"
|
||||
)
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return (
|
||||
f"<Watchlist id={self.id} "
|
||||
f"user_id={self.user_id} "
|
||||
f"symbol_id={self.symbol_id}>"
|
||||
)
|
||||
Executable
+106
@@ -0,0 +1,106 @@
|
||||
from app.schemas.auth import (
|
||||
RegisterRequest,
|
||||
LoginRequest,
|
||||
TokenResponse,
|
||||
RefreshRequest,
|
||||
UserResponse,
|
||||
UserSessionResponse,
|
||||
LogoutRequest,
|
||||
ChangePasswordRequest,
|
||||
)
|
||||
from app.schemas.user import (
|
||||
UserUpdateRequest,
|
||||
AdminCreateUserRequest,
|
||||
AdminResetPasswordRequest,
|
||||
AdminUserUpdateRequest,
|
||||
)
|
||||
from app.schemas.candle import (
|
||||
CandleResponse,
|
||||
CandleListResponse,
|
||||
IndicatorResponse,
|
||||
)
|
||||
from app.schemas.symbol import (
|
||||
SymbolResponse,
|
||||
SymbolSearchResponse,
|
||||
WatchlistResponse,
|
||||
WatchlistCreateRequest,
|
||||
WatchlistUpdateRequest,
|
||||
)
|
||||
from app.schemas.ws_message import (
|
||||
WSSubscription,
|
||||
WSCandleUpdate,
|
||||
WSTickerUpdate,
|
||||
WSConnectionStatus,
|
||||
)
|
||||
from app.schemas.credential import (
|
||||
CredentialResponse,
|
||||
CredentialCreateRequest,
|
||||
CredentialUpdateRequest,
|
||||
)
|
||||
from app.schemas.exchange import (
|
||||
ExchangeResponse,
|
||||
ExchangeCreateRequest,
|
||||
ExchangeUpdateRequest,
|
||||
)
|
||||
from app.schemas.health import (
|
||||
HealthResponse,
|
||||
DetailedHealthResponse,
|
||||
DbHealth,
|
||||
ExchangeHealth,
|
||||
)
|
||||
from app.schemas.strategy import (
|
||||
StrategyEntry,
|
||||
StrategyListResponse,
|
||||
StrategyConfigRequest,
|
||||
STRATEGY_NAMES,
|
||||
STRATEGY_DISPLAY,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
# auth
|
||||
"RegisterRequest",
|
||||
"LoginRequest",
|
||||
"TokenResponse",
|
||||
"RefreshRequest",
|
||||
"UserResponse",
|
||||
"UserSessionResponse",
|
||||
"LogoutRequest",
|
||||
"ChangePasswordRequest",
|
||||
# user
|
||||
"UserUpdateRequest",
|
||||
"AdminUserUpdateRequest",
|
||||
# candle
|
||||
"CandleResponse",
|
||||
"CandleListResponse",
|
||||
"IndicatorResponse",
|
||||
# symbol
|
||||
"SymbolResponse",
|
||||
"SymbolSearchResponse",
|
||||
"WatchlistResponse",
|
||||
"WatchlistCreateRequest",
|
||||
"WatchlistUpdateRequest",
|
||||
# ws_message
|
||||
"WSSubscription",
|
||||
"WSCandleUpdate",
|
||||
"WSTickerUpdate",
|
||||
"WSConnectionStatus",
|
||||
# credential
|
||||
"CredentialResponse",
|
||||
"CredentialCreateRequest",
|
||||
"CredentialUpdateRequest",
|
||||
# exchange
|
||||
"ExchangeResponse",
|
||||
"ExchangeCreateRequest",
|
||||
"ExchangeUpdateRequest",
|
||||
# health
|
||||
"HealthResponse",
|
||||
"DetailedHealthResponse",
|
||||
"DbHealth",
|
||||
"ExchangeHealth",
|
||||
# strategy
|
||||
"StrategyEntry",
|
||||
"StrategyListResponse",
|
||||
"StrategyConfigRequest",
|
||||
"STRATEGY_NAMES",
|
||||
"STRATEGY_DISPLAY",
|
||||
]
|
||||
Executable
+69
@@ -0,0 +1,69 @@
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, field_validator
|
||||
|
||||
|
||||
class RegisterRequest(BaseModel):
|
||||
username: str
|
||||
email: str
|
||||
password: str
|
||||
|
||||
@field_validator("password")
|
||||
@classmethod
|
||||
def password_min_length(cls, v: str) -> str:
|
||||
if len(v) < 8:
|
||||
raise ValueError("password must be at least 8 characters")
|
||||
return v
|
||||
|
||||
|
||||
class LoginRequest(BaseModel):
|
||||
username: str
|
||||
password: str
|
||||
|
||||
|
||||
class TokenResponse(BaseModel):
|
||||
access_token: str
|
||||
refresh_token: str
|
||||
token_type: str = "bearer"
|
||||
|
||||
|
||||
class RefreshRequest(BaseModel):
|
||||
refresh_token: str
|
||||
|
||||
|
||||
class UserResponse(BaseModel):
|
||||
id: UUID
|
||||
username: str
|
||||
email: str
|
||||
display_name: Optional[str] = None
|
||||
is_active: bool = True
|
||||
is_admin: bool = False
|
||||
role: str = "trader"
|
||||
preferences: Optional[dict] = None
|
||||
created_at: datetime
|
||||
|
||||
|
||||
class UserSessionResponse(BaseModel):
|
||||
id: UUID
|
||||
created_at: datetime
|
||||
user_agent: Optional[str] = None
|
||||
ip_address: Optional[str] = None
|
||||
is_current: bool
|
||||
|
||||
|
||||
class LogoutRequest(BaseModel):
|
||||
refresh_token: str
|
||||
|
||||
|
||||
class ChangePasswordRequest(BaseModel):
|
||||
old_password: str
|
||||
new_password: str
|
||||
|
||||
@field_validator("new_password")
|
||||
@classmethod
|
||||
def new_password_min_length(cls, v: str) -> str:
|
||||
if len(v) < 8:
|
||||
raise ValueError("new_password must be at least 8 characters")
|
||||
return v
|
||||
Executable
+29
@@ -0,0 +1,29 @@
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class CandleResponse(BaseModel):
|
||||
symbol_id: int
|
||||
timeframe: str
|
||||
timestamp: datetime
|
||||
open: float
|
||||
high: float
|
||||
low: float
|
||||
close: float
|
||||
volume: float
|
||||
|
||||
|
||||
class CandleListResponse(BaseModel):
|
||||
candles: list[CandleResponse]
|
||||
cursor: Optional[str] = None
|
||||
has_more: bool
|
||||
|
||||
|
||||
class IndicatorResponse(BaseModel):
|
||||
symbol_id: int
|
||||
timeframe: str
|
||||
timestamp: datetime
|
||||
indicators: dict
|
||||
Executable
+34
@@ -0,0 +1,34 @@
|
||||
from typing import Optional
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, field_validator
|
||||
|
||||
|
||||
class CredentialResponse(BaseModel):
|
||||
id: UUID
|
||||
exchange_id: int
|
||||
exchange_name: str
|
||||
api_key: str
|
||||
is_testnet: bool
|
||||
is_active: bool
|
||||
|
||||
@field_validator("api_key")
|
||||
@classmethod
|
||||
def mask_api_key(cls, v: str) -> str:
|
||||
"""Show full API key so users can copy it (middle chars left visible for reference)."""
|
||||
return v
|
||||
|
||||
|
||||
class CredentialCreateRequest(BaseModel):
|
||||
exchange_id: int
|
||||
api_key: str
|
||||
api_secret: str
|
||||
passphrase: Optional[str] = None
|
||||
is_testnet: bool = False
|
||||
|
||||
|
||||
class CredentialUpdateRequest(BaseModel):
|
||||
api_key: Optional[str] = None
|
||||
api_secret: Optional[str] = None
|
||||
passphrase: Optional[str] = None
|
||||
is_active: Optional[bool] = None
|
||||
Executable
+24
@@ -0,0 +1,24 @@
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class ExchangeResponse(BaseModel):
|
||||
id: int
|
||||
name: str
|
||||
display_name: str
|
||||
is_active: bool = True
|
||||
|
||||
|
||||
class ExchangeCreateRequest(BaseModel):
|
||||
name: str
|
||||
display_name: str
|
||||
base_url: str
|
||||
ws_url: str
|
||||
|
||||
|
||||
class ExchangeUpdateRequest(BaseModel):
|
||||
display_name: Optional[str] = None
|
||||
base_url: Optional[str] = None
|
||||
ws_url: Optional[str] = None
|
||||
is_active: Optional[bool] = None
|
||||
Executable
+31
@@ -0,0 +1,31 @@
|
||||
from datetime import datetime
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class DbHealth(BaseModel):
|
||||
connected: bool
|
||||
latency_ms: float
|
||||
pool_size: int
|
||||
|
||||
|
||||
class ExchangeHealth(BaseModel):
|
||||
exchange: str
|
||||
connected: bool
|
||||
last_sync: Optional[datetime] = None
|
||||
|
||||
|
||||
class HealthResponse(BaseModel):
|
||||
status: str
|
||||
version: str
|
||||
uptime: float
|
||||
db_connected: bool
|
||||
|
||||
|
||||
class DetailedHealthResponse(BaseModel):
|
||||
status: str
|
||||
version: str
|
||||
uptime: float
|
||||
db: DbHealth
|
||||
exchange_connections: list[ExchangeHealth]
|
||||
Executable
+78
@@ -0,0 +1,78 @@
|
||||
"""Pydantic schemas for real trade responses and requests."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import Optional
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel, field_validator
|
||||
|
||||
|
||||
class RealTradeResponse(BaseModel):
|
||||
"""Single real trade record returned to the frontend."""
|
||||
|
||||
id: int
|
||||
user_id: str
|
||||
exchange: str
|
||||
symbol: str
|
||||
side: str
|
||||
order_type: str
|
||||
amount: float
|
||||
price: Optional[float] = None
|
||||
filled_amount: float
|
||||
status: str
|
||||
pnl: Optional[float] = None
|
||||
pnl_percent: Optional[float] = None
|
||||
order_id: Optional[str] = None
|
||||
created_at: datetime
|
||||
closed_at: Optional[datetime] = None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
@field_validator("user_id", mode="before")
|
||||
@classmethod
|
||||
def _coerce_user_id(cls, v: object) -> str:
|
||||
if isinstance(v, UUID):
|
||||
return str(v)
|
||||
return str(v)
|
||||
|
||||
|
||||
class RealTradeListResponse(BaseModel):
|
||||
"""Paginated list of real trades with aggregate stats."""
|
||||
|
||||
trades: list[RealTradeResponse]
|
||||
total: int
|
||||
total_pnl: Optional[float] = None
|
||||
win_rate: Optional[float] = None
|
||||
|
||||
|
||||
class RealTradeCreateRequest(BaseModel):
|
||||
"""Payload to create a real trade record (used after placing an order)."""
|
||||
|
||||
exchange: str
|
||||
symbol: str
|
||||
side: str
|
||||
order_type: str = "market"
|
||||
amount: float
|
||||
price: Optional[float] = None
|
||||
filled_amount: float = 0
|
||||
status: str = "open"
|
||||
order_id: Optional[str] = None
|
||||
|
||||
|
||||
class WinRatePeriod(BaseModel):
|
||||
"""Win-rate data for a single period (daily / weekly / monthly)."""
|
||||
|
||||
trades: int
|
||||
wins: int
|
||||
win_rate: float
|
||||
|
||||
|
||||
class WinRateResponse(BaseModel):
|
||||
"""Win-rate breakdown across daily, weekly, and monthly periods."""
|
||||
|
||||
daily: WinRatePeriod
|
||||
weekly: WinRatePeriod
|
||||
monthly: WinRatePeriod
|
||||
Executable
+89
@@ -0,0 +1,89 @@
|
||||
"""Pydantic schemas for signal and trade responses."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, model_validator
|
||||
import json
|
||||
|
||||
|
||||
class SignalResponse(BaseModel):
|
||||
id: int
|
||||
symbol: str
|
||||
exchange: str
|
||||
timeframe: str
|
||||
signal_type: str
|
||||
strength: str
|
||||
price: float
|
||||
timestamp: datetime
|
||||
indicators_snapshot: Optional[str] = None
|
||||
confidence: Optional[float] = None
|
||||
status: str
|
||||
note: Optional[str] = None
|
||||
created_at: datetime
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _extract_confidence(self) -> "SignalResponse":
|
||||
if self.confidence is None and self.indicators_snapshot:
|
||||
try:
|
||||
snap = json.loads(self.indicators_snapshot)
|
||||
self.confidence = snap.get("confidence")
|
||||
except (json.JSONDecodeError, TypeError, AttributeError):
|
||||
pass
|
||||
return self
|
||||
|
||||
|
||||
class TradeResponse(BaseModel):
|
||||
id: int
|
||||
signal_id: Optional[int] = None
|
||||
symbol: str
|
||||
exchange: str
|
||||
timeframe: str
|
||||
direction: str
|
||||
entry_price: float
|
||||
entry_time: datetime
|
||||
entry_reason: Optional[str] = None
|
||||
exit_price: Optional[float] = None
|
||||
exit_time: Optional[datetime] = None
|
||||
exit_reason: Optional[str] = None
|
||||
quantity: float
|
||||
pnl: Optional[float] = None
|
||||
pnl_percent: Optional[float] = None
|
||||
status: str
|
||||
created_at: datetime
|
||||
closed_at: Optional[datetime] = None
|
||||
|
||||
model_config = {"from_attributes": True}
|
||||
|
||||
|
||||
class SignalListResponse(BaseModel):
|
||||
signals: list[SignalResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class TradeListResponse(BaseModel):
|
||||
trades: list[TradeResponse]
|
||||
total: int
|
||||
total_pnl: Optional[float] = None
|
||||
win_rate: Optional[float] = None
|
||||
|
||||
|
||||
class ReviewResponse(BaseModel):
|
||||
period: str # "weekly" or "monthly"
|
||||
start_date: str
|
||||
end_date: str
|
||||
total_signals: int
|
||||
total_trades: int
|
||||
wins: int
|
||||
losses: int
|
||||
win_rate: float
|
||||
total_pnl: float
|
||||
best_trade: Optional[TradeResponse] = None
|
||||
worst_trade: Optional[TradeResponse] = None
|
||||
signals_by_type: dict[str, int]
|
||||
symbol_performance: list[dict]
|
||||
Executable
+66
@@ -0,0 +1,66 @@
|
||||
"""Pydantic schemas for per-user strategy configuration."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
# Available strategy names — full list of 13 voting algorithms
|
||||
STRATEGY_NAMES = [
|
||||
"double_bb_rsi",
|
||||
"macd_crossover",
|
||||
"supertrend",
|
||||
"volume_breakout",
|
||||
"ichimoku_cloud",
|
||||
"divergence",
|
||||
"smc",
|
||||
"mtf",
|
||||
"obv",
|
||||
"stoch_rsi",
|
||||
"mfi",
|
||||
"fvg",
|
||||
"candlestick",
|
||||
]
|
||||
|
||||
STRATEGY_DISPLAY: dict[str, str] = {
|
||||
"double_bb_rsi": "Double BB + RSI",
|
||||
"macd_crossover": "MACD Crossover",
|
||||
"supertrend": "SuperTrend",
|
||||
"volume_breakout": "Volume Breakout",
|
||||
"ichimoku_cloud": "Ichimoku Cloud",
|
||||
"divergence": "Divergence",
|
||||
"smc": "Market Structure (SMC)",
|
||||
"mtf": "Multi-Timeframe",
|
||||
"obv": "OBV Crossover",
|
||||
"stoch_rsi": "Stochastic RSI",
|
||||
"mfi": "Money Flow Index",
|
||||
"fvg": "Fair Value Gap",
|
||||
"candlestick": "Candlestick Patterns",
|
||||
}
|
||||
|
||||
|
||||
class StrategyEntry(BaseModel):
|
||||
"""A single strategy with its enabled status."""
|
||||
|
||||
name: str
|
||||
display_name: str
|
||||
enabled: bool = False
|
||||
|
||||
|
||||
class StrategyListResponse(BaseModel):
|
||||
"""Response for GET /api/v1/strategies."""
|
||||
|
||||
strategies: list[StrategyEntry]
|
||||
thresholds: dict[str, float] = {}
|
||||
|
||||
|
||||
class StrategyConfigRequest(BaseModel):
|
||||
"""Request body for PUT /api/v1/strategies.
|
||||
|
||||
Only provided fields will be updated.
|
||||
"""
|
||||
|
||||
enabled_strategies: Optional[list[str]] = None
|
||||
thresholds: Optional[dict[str, float]] = None
|
||||
Executable
+37
@@ -0,0 +1,37 @@
|
||||
from typing import Optional
|
||||
from uuid import UUID
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
|
||||
class SymbolResponse(BaseModel):
|
||||
id: int
|
||||
exchange_id: int
|
||||
symbol: str
|
||||
base: str
|
||||
quote: str
|
||||
is_active: bool = True
|
||||
|
||||
|
||||
class SymbolSearchResponse(BaseModel):
|
||||
symbols: list[SymbolResponse]
|
||||
|
||||
|
||||
class WatchlistResponse(BaseModel):
|
||||
id: UUID
|
||||
symbol_id: int
|
||||
symbol: str
|
||||
exchange: str
|
||||
label: Optional[str] = None
|
||||
sort_order: int
|
||||
|
||||
|
||||
class WatchlistCreateRequest(BaseModel):
|
||||
symbol_id: int
|
||||
label: Optional[str] = None
|
||||
sort_order: Optional[int] = None
|
||||
|
||||
|
||||
class WatchlistUpdateRequest(BaseModel):
|
||||
label: Optional[str] = None
|
||||
sort_order: Optional[int] = None
|
||||
Executable
+63
@@ -0,0 +1,63 @@
|
||||
from typing import Optional
|
||||
import re
|
||||
|
||||
from pydantic import BaseModel, field_validator
|
||||
|
||||
|
||||
class UserUpdateRequest(BaseModel):
|
||||
display_name: Optional[str] = None
|
||||
email: Optional[str] = None
|
||||
preferences: Optional[dict] = None
|
||||
|
||||
|
||||
_EMAIL_RE = re.compile(r"^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$")
|
||||
|
||||
def _validate_email(email: str) -> str:
|
||||
if not _EMAIL_RE.match(email):
|
||||
raise ValueError("invalid email format")
|
||||
return email
|
||||
|
||||
def _validate_strong_password(password: str) -> str:
|
||||
if len(password) < 8:
|
||||
raise ValueError("password must be at least 8 characters")
|
||||
if not re.search(r"\d", password):
|
||||
raise ValueError("password must contain at least 1 number")
|
||||
if not re.search(r"[^a-zA-Z0-9]", password):
|
||||
raise ValueError("password must contain at least 1 special character")
|
||||
return password
|
||||
|
||||
|
||||
class AdminUserUpdateRequest(BaseModel):
|
||||
is_active: Optional[bool] = None
|
||||
is_admin: Optional[bool] = None
|
||||
role: Optional[str] = None
|
||||
email: Optional[str] = None
|
||||
display_name: Optional[str] = None
|
||||
|
||||
|
||||
class AdminCreateUserRequest(BaseModel):
|
||||
username: str
|
||||
email: str
|
||||
password: str
|
||||
display_name: Optional[str] = None
|
||||
is_admin: bool = False
|
||||
role: str = "trader"
|
||||
|
||||
@field_validator("email")
|
||||
@classmethod
|
||||
def email_must_be_valid(cls, v: str) -> str:
|
||||
return _validate_email(v)
|
||||
|
||||
@field_validator("password")
|
||||
@classmethod
|
||||
def password_must_be_strong(cls, v: str) -> str:
|
||||
return _validate_strong_password(v)
|
||||
|
||||
|
||||
class AdminResetPasswordRequest(BaseModel):
|
||||
new_password: str
|
||||
|
||||
@field_validator("new_password")
|
||||
@classmethod
|
||||
def new_password_must_be_strong(cls, v: str) -> str:
|
||||
return _validate_strong_password(v)
|
||||
Executable
+31
@@ -0,0 +1,31 @@
|
||||
from decimal import Decimal
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from app.schemas.candle import CandleResponse
|
||||
|
||||
|
||||
class WSSubscription(BaseModel):
|
||||
symbol: str
|
||||
timeframe: str
|
||||
exchange: str
|
||||
action: Literal["subscribe", "unsubscribe"]
|
||||
|
||||
|
||||
class WSCandleUpdate(BaseModel):
|
||||
type: Literal["candle"] = "candle"
|
||||
data: CandleResponse
|
||||
|
||||
|
||||
class WSTickerUpdate(BaseModel):
|
||||
type: Literal["ticker"] = "ticker"
|
||||
symbol: str
|
||||
price: float
|
||||
change_24h: float
|
||||
volume: float
|
||||
|
||||
|
||||
class WSConnectionStatus(BaseModel):
|
||||
type: Literal["connection"] = "connection"
|
||||
status: Literal["connected", "reconnecting", "disconnected"]
|
||||
Executable
Executable
+367
@@ -0,0 +1,367 @@
|
||||
"""Alert service — evaluates multi-condition user alerts against current market data.
|
||||
|
||||
Each alert consists of a list of conditions (AND logic — all must pass). When a
|
||||
new signal is generated, this service checks all active alerts for that user
|
||||
and returns those whose conditions are fully satisfied.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from collections.abc import Sequence
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.alert import AlertCondition
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Supported indicators and their extraction helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
SUPPORTED_INDICATORS = {
|
||||
"rsi", "macd", "bb_width", "volume", "price", "sma", "ema", "momentum",
|
||||
}
|
||||
|
||||
SUPPORTED_OPERATORS = {">", "<", ">=", "<=", "==", "cross_above", "cross_below"}
|
||||
|
||||
|
||||
def _get_price(indicators: dict | None) -> float | None:
|
||||
"""Extract the latest close price from indicators."""
|
||||
if not indicators:
|
||||
return None
|
||||
return indicators.get("close")
|
||||
|
||||
|
||||
def _get_rsi(indicators: dict | None) -> float | None:
|
||||
"""Extract latest RSI value (rsi_14)."""
|
||||
if not indicators:
|
||||
return None
|
||||
rsi_series = indicators.get("rsi_14")
|
||||
if isinstance(rsi_series, list) and len(rsi_series) > 0:
|
||||
return rsi_series[-1]
|
||||
return None
|
||||
|
||||
|
||||
def _get_macd(indicators: dict | None) -> dict | None:
|
||||
"""Extract MACD data."""
|
||||
if not indicators:
|
||||
return None
|
||||
return indicators.get("macd")
|
||||
|
||||
|
||||
def _get_bb_width(indicators: dict | None) -> float | None:
|
||||
"""Compute current Bollinger Band width (upper - lower)."""
|
||||
if not indicators:
|
||||
return None
|
||||
bb = indicators.get("bollinger_bands")
|
||||
if not bb:
|
||||
return None
|
||||
upper = bb.get("upper", [None])[-1]
|
||||
lower = bb.get("lower", [None])[-1]
|
||||
if upper is not None and lower is not None:
|
||||
return float(upper) - float(lower)
|
||||
return None
|
||||
|
||||
|
||||
def _get_volume(indicators: dict | None, condition: dict) -> float | None:
|
||||
"""Extract volume-related value based on condition type.
|
||||
|
||||
Types:
|
||||
- 'avg_multiplier': compare current volume to SMA(period) average
|
||||
- 'absolute': raw volume value
|
||||
"""
|
||||
if not indicators:
|
||||
return None
|
||||
vol_type = condition.get("type", "absolute")
|
||||
if vol_type == "avg_multiplier":
|
||||
period = condition.get("period", 20)
|
||||
volume_series = indicators.get("volume")
|
||||
if isinstance(volume_series, list) and len(volume_series) >= period:
|
||||
recent = volume_series[-period:]
|
||||
avg = sum(recent) / len(recent)
|
||||
current = volume_series[-1]
|
||||
return current / avg if avg > 0 else None
|
||||
return None
|
||||
# absolute
|
||||
volume_series = indicators.get("volume")
|
||||
if isinstance(volume_series, list) and len(volume_series) > 0:
|
||||
return volume_series[-1]
|
||||
return None
|
||||
|
||||
|
||||
def _get_sma(indicators: dict | None, condition: dict) -> float | None:
|
||||
"""Extract SMA value. Uses condition['value'] as the period, default 20."""
|
||||
if not indicators:
|
||||
return None
|
||||
period = condition.get("period", 20)
|
||||
key = f"sma_{period}"
|
||||
sma_series = indicators.get(key)
|
||||
if isinstance(sma_series, list) and len(sma_series) > 0:
|
||||
return sma_series[-1]
|
||||
# Fallback: check generic "sma_20" or "sma_50"
|
||||
for fallback in ("sma_20", "sma_50"):
|
||||
fb = indicators.get(fallback)
|
||||
if isinstance(fb, list) and len(fb) > 0:
|
||||
return fb[-1]
|
||||
return None
|
||||
|
||||
|
||||
def _get_ema(indicators: dict | None, condition: dict) -> float | None:
|
||||
"""Extract EMA value."""
|
||||
if not indicators:
|
||||
return None
|
||||
period = condition.get("period", 20)
|
||||
key = f"ema_{period}"
|
||||
ema_series = indicators.get(key)
|
||||
if isinstance(ema_series, list) and len(ema_series) > 0:
|
||||
return ema_series[-1]
|
||||
return None
|
||||
|
||||
|
||||
def _get_momentum(indicators: dict | None) -> float | None:
|
||||
"""Extract momentum (rate of change) from close prices."""
|
||||
if not indicators:
|
||||
return None
|
||||
# Try to compute from sma_20 as a proxy if we have enough values
|
||||
close_series = indicators.get("sma_20")
|
||||
if isinstance(close_series, list) and len(close_series) >= 2:
|
||||
prev = close_series[-2]
|
||||
curr = close_series[-1]
|
||||
if prev and prev > 0:
|
||||
return (curr - prev) / prev * 100
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Condition evaluation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _extract_indicator_value(indicator: str, indicators: dict | None,
|
||||
condition: dict) -> float | None:
|
||||
"""Route the indicator name to the correct extraction function."""
|
||||
extractors = {
|
||||
"price": lambda: _get_price(indicators),
|
||||
"rsi": lambda: _get_rsi(indicators),
|
||||
"macd": lambda: _get_macd_value(indicators),
|
||||
"bb_width": lambda: _get_bb_width(indicators),
|
||||
"volume": lambda: _get_volume(indicators, condition),
|
||||
"sma": lambda: _get_sma(indicators, condition),
|
||||
"ema": lambda: _get_ema(indicators, condition),
|
||||
"momentum": lambda: _get_momentum(indicators),
|
||||
}
|
||||
fn = extractors.get(indicator)
|
||||
if fn is None:
|
||||
logger.warning("Unknown indicator '%s'", indicator)
|
||||
return None
|
||||
return fn()
|
||||
|
||||
|
||||
def _get_macd_value(indicators: dict | None) -> float | None:
|
||||
"""Extract latest MACD histogram value (macd - signal)."""
|
||||
macd_data = _get_macd(indicators)
|
||||
if not macd_data:
|
||||
return None
|
||||
macd_line = macd_data.get("macd_line", [])
|
||||
signal_line = macd_data.get("signal_line", [])
|
||||
if len(macd_line) > 0 and len(signal_line) > 0:
|
||||
m = macd_line[-1]
|
||||
s = signal_line[-1]
|
||||
if m is not None and s is not None:
|
||||
return float(m) - float(s)
|
||||
return None
|
||||
|
||||
|
||||
def _compare_values(current: float, target: float, operator: str) -> bool:
|
||||
"""Compare two numeric values using the given operator."""
|
||||
if operator == ">":
|
||||
return current > target
|
||||
elif operator == "<":
|
||||
return current < target
|
||||
elif operator == ">=":
|
||||
return current >= target
|
||||
elif operator == "<=":
|
||||
return current <= target
|
||||
elif operator == "==":
|
||||
return abs(current - target) < 1e-9
|
||||
return False
|
||||
|
||||
|
||||
def _check_cross(series_a: list[float] | None, series_b: float | None,
|
||||
operator: str) -> bool:
|
||||
"""Check if series_a crosses above/below a fixed value.
|
||||
|
||||
'cross_above': previous <= value < current (or prev < value <= current)
|
||||
'cross_below': previous >= value > current (or prev > value >= current)
|
||||
"""
|
||||
if not series_a or len(series_a) < 2 or series_b is None:
|
||||
return False
|
||||
prev = series_a[-2]
|
||||
curr = series_a[-1]
|
||||
if operator == "cross_above":
|
||||
return prev <= series_b < curr
|
||||
elif operator == "cross_below":
|
||||
return prev >= series_b > curr
|
||||
return False
|
||||
|
||||
|
||||
async def evaluate_single_condition(condition: dict,
|
||||
symbol_data: dict | None) -> bool:
|
||||
"""Evaluate a single alert condition against symbol indicator data.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
condition : dict
|
||||
One condition object from the alert's conditions array, e.g.:
|
||||
``{'indicator': 'rsi', 'operator': '>', 'value': 70, 'timeframe': '1h'}``
|
||||
symbol_data : dict | None
|
||||
The current indicator data for this symbol (output of
|
||||
``get_indicators`` or similar).
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
``True`` if the condition is satisfied, ``False`` otherwise.
|
||||
"""
|
||||
indicator = condition.get("indicator")
|
||||
operator = condition.get("operator")
|
||||
target_value = condition.get("value")
|
||||
|
||||
if not indicator or not operator or target_value is None:
|
||||
logger.debug("Incomplete condition: %s", condition)
|
||||
return False
|
||||
|
||||
# Handle cross_above / cross_below — these need the full series
|
||||
if operator in ("cross_above", "cross_below"):
|
||||
series = _get_value_series(indicator, symbol_data, condition)
|
||||
return _check_cross(series, target_value, operator)
|
||||
|
||||
current = _extract_indicator_value(indicator, symbol_data, condition)
|
||||
if current is None:
|
||||
logger.debug("Could not extract indicator '%s' from symbol data", indicator)
|
||||
return False
|
||||
|
||||
return _compare_values(current, float(target_value), operator)
|
||||
|
||||
|
||||
def _get_value_series(indicator: str, indicators: dict | None,
|
||||
condition: dict) -> list[float] | None:
|
||||
"""Get the full time-series for an indicator (needed for cross detection)."""
|
||||
if not indicators:
|
||||
return None
|
||||
if indicator == "rsi":
|
||||
series = indicators.get("rsi_14")
|
||||
elif indicator == "price":
|
||||
series = indicators.get("sma_20")
|
||||
elif indicator == "volume":
|
||||
series = indicators.get("volume")
|
||||
elif indicator == "sma":
|
||||
period = condition.get("period", 20)
|
||||
series = indicators.get(f"sma_{period}")
|
||||
elif indicator == "ema":
|
||||
period = condition.get("period", 20)
|
||||
series = indicators.get(f"ema_{period}")
|
||||
else:
|
||||
return None
|
||||
if isinstance(series, list) and len(series) >= 2:
|
||||
return [float(v) for v in series if v is not None]
|
||||
return None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Alert evaluation (all conditions for a set of symbols)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def check_alert_conditions(db: AsyncSession,
|
||||
symbol_data_map: dict[str, dict | None],
|
||||
user_id) -> list[AlertCondition]:
|
||||
"""Check all active alerts and return those whose conditions are satisfied.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
db : AsyncSession
|
||||
Database session.
|
||||
symbol_data_map : dict[str, dict | None]
|
||||
Mapping from symbol (e.g. ``'BTC/USDT'``) to its indicator data dict.
|
||||
user_id
|
||||
The user ID to check alerts for.
|
||||
|
||||
Returns
|
||||
-------
|
||||
list[AlertCondition]
|
||||
All alerts that have fired (conditions satisfied).
|
||||
"""
|
||||
# Load all active alerts for this user
|
||||
result = await db.execute(
|
||||
select(AlertCondition).where(
|
||||
AlertCondition.user_id == user_id,
|
||||
AlertCondition.is_active == True,
|
||||
)
|
||||
)
|
||||
alerts: Sequence[AlertCondition] = result.scalars().all()
|
||||
|
||||
fired: list[AlertCondition] = []
|
||||
for alert in alerts:
|
||||
try:
|
||||
triggered = await _evaluate_alert_conditions(alert, symbol_data_map)
|
||||
if triggered:
|
||||
fired.append(alert)
|
||||
except Exception:
|
||||
logger.exception("Error evaluating alert %s", alert.id)
|
||||
|
||||
return fired
|
||||
|
||||
|
||||
async def _evaluate_alert_conditions(
|
||||
alert: AlertCondition,
|
||||
symbol_data_map: dict[str, dict | None],
|
||||
) -> bool:
|
||||
"""Evaluate all conditions of an alert (AND logic)."""
|
||||
conditions: list[dict] = alert.conditions or []
|
||||
if not conditions:
|
||||
return False
|
||||
|
||||
# Group conditions by the symbol / timeframe they reference
|
||||
# Each condition can optionally specify a 'timeframe' key.
|
||||
# For simplicity, we evaluate each condition against the symbol_data_map.
|
||||
# All conditions must pass for the alert to fire.
|
||||
|
||||
for cond in conditions:
|
||||
# Determine which symbol data to use (default to first available)
|
||||
timeframe = cond.get("timeframe")
|
||||
# Use any available symbol data — for now check each symbol
|
||||
passed = False
|
||||
for symbol, data in symbol_data_map.items():
|
||||
if await evaluate_single_condition(cond, data):
|
||||
passed = True
|
||||
break
|
||||
if not passed:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
|
||||
async def check_and_notify_alerts(
|
||||
db: AsyncSession,
|
||||
symbol: str,
|
||||
symbol_data: dict | None,
|
||||
user_id,
|
||||
) -> None:
|
||||
"""Convenience: check alerts for a single symbol+user and log results.
|
||||
|
||||
Called from signal_service after each signal.
|
||||
"""
|
||||
symbol_data_map = {symbol: symbol_data}
|
||||
fired = await check_alert_conditions(db, symbol_data_map, user_id)
|
||||
if fired:
|
||||
names = [a.name for a in fired]
|
||||
logger.info(
|
||||
"🔔 Alerts triggered for user %s on %s: %s",
|
||||
user_id, symbol, names,
|
||||
)
|
||||
# Actual notification sending is handled by the caller
|
||||
Executable
+82
@@ -0,0 +1,82 @@
|
||||
"""Audit log service for recording and querying audit events."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import desc, func as sa_func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.models.audit_log import AuditLog
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def log_action(
|
||||
db: AsyncSession,
|
||||
user_id: UUID | None,
|
||||
action: str,
|
||||
resource: str,
|
||||
details: dict | None = None,
|
||||
) -> AuditLog:
|
||||
"""Record an audit log entry.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
user_id: UUID of the user who performed the action, or None.
|
||||
action: Action type (e.g. 'trade_open', 'signal_generated').
|
||||
resource: Affected resource description.
|
||||
details: Optional extra JSON info.
|
||||
|
||||
Returns:
|
||||
The created AuditLog instance.
|
||||
"""
|
||||
entry = AuditLog(
|
||||
user_id=user_id,
|
||||
action=action,
|
||||
resource=resource,
|
||||
details=details,
|
||||
)
|
||||
db.add(entry)
|
||||
await db.flush()
|
||||
logger.debug(
|
||||
"Audit log: action=%s resource=%s user=%s",
|
||||
action, resource, user_id,
|
||||
)
|
||||
return entry
|
||||
|
||||
|
||||
async def get_audit_logs(
|
||||
db: AsyncSession,
|
||||
limit: int = 100,
|
||||
offset: int = 0,
|
||||
action: str | None = None,
|
||||
) -> tuple[list[AuditLog], int]:
|
||||
"""Fetch audit log entries with pagination and optional action filter.
|
||||
|
||||
Args:
|
||||
db: Database session.
|
||||
limit: Maximum number of entries to return.
|
||||
offset: Number of entries to skip.
|
||||
action: Optional action type filter.
|
||||
|
||||
Returns:
|
||||
A tuple of (list of AuditLog entries, total count).
|
||||
"""
|
||||
base_query = select(AuditLog).order_by(desc(AuditLog.created_at))
|
||||
|
||||
if action:
|
||||
base_query = base_query.where(AuditLog.action == action)
|
||||
|
||||
# Get total count
|
||||
count_query = select(sa_func.count()).select_from(base_query.subquery())
|
||||
total_result = await db.execute(count_query)
|
||||
total = total_result.scalar() or 0
|
||||
|
||||
# Get paginated results
|
||||
query = base_query.limit(limit).offset(offset)
|
||||
result = await db.execute(query)
|
||||
entries = list(result.scalars().all())
|
||||
|
||||
return entries, total
|
||||
Executable
+332
@@ -0,0 +1,332 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.exceptions import (
|
||||
ConflictException,
|
||||
InvalidCredentialsException,
|
||||
InvalidTokenException,
|
||||
NotFoundException,
|
||||
)
|
||||
from app.core.security import (
|
||||
create_access_token,
|
||||
create_refresh_token,
|
||||
decode_token,
|
||||
generate_token_hash,
|
||||
hash_password_async,
|
||||
verify_password_async,
|
||||
)
|
||||
from app.models import RefreshToken, User
|
||||
from app.schemas import (
|
||||
LoginRequest,
|
||||
RegisterRequest,
|
||||
TokenResponse,
|
||||
UserResponse,
|
||||
UserSessionResponse,
|
||||
)
|
||||
|
||||
|
||||
# P1-23: Backend password strength validation
|
||||
|
||||
|
||||
def _validate_password_strength(password: str) -> None:
|
||||
"""Raise ConflictException if password doesn't meet minimum strength."""
|
||||
if len(password) < 8:
|
||||
raise ConflictException(detail="Password must be at least 8 characters")
|
||||
if not re.search(r"[a-z]", password):
|
||||
raise ConflictException(detail="Password must contain at least one lowercase letter")
|
||||
if not re.search(r"[A-Z]", password):
|
||||
raise ConflictException(detail="Password must contain at least one uppercase letter")
|
||||
if not re.search(r"\d", password):
|
||||
raise ConflictException(detail="Password must contain at least one digit")
|
||||
if not re.search(r"[!@#$%^&*()_+\-=\[\]{};':\"\\|,.<>/?]", password):
|
||||
raise ConflictException(detail="Password must contain at least one special character")
|
||||
|
||||
|
||||
async def register(
|
||||
db: AsyncSession,
|
||||
req: RegisterRequest,
|
||||
) -> UserResponse:
|
||||
"""Register a new user account.
|
||||
|
||||
Checks username / email uniqueness, hashes the password, creates the
|
||||
``User`` row together with an empty ``Watchlist``, and returns the
|
||||
new user's public profile.
|
||||
"""
|
||||
# --- uniqueness checks ---------------------------------------------------
|
||||
existing_username = await db.execute(
|
||||
select(User).where(User.username == req.username)
|
||||
)
|
||||
if existing_username.scalar_one_or_none() is not None:
|
||||
raise ConflictException(detail=f"Username '{req.username}' is already taken")
|
||||
|
||||
existing_email = await db.execute(
|
||||
select(User).where(User.email == req.email)
|
||||
)
|
||||
if existing_email.scalar_one_or_none() is not None:
|
||||
raise ConflictException(detail=f"Email '{req.email}' is already registered")
|
||||
|
||||
# --- P1-23: password strength validation ---------------------------------
|
||||
_validate_password_strength(req.password)
|
||||
|
||||
# --- create user ---------------------------------------------------------
|
||||
user = User(
|
||||
username=req.username,
|
||||
email=req.email,
|
||||
password_hash=await hash_password_async(req.password),
|
||||
)
|
||||
db.add(user)
|
||||
await db.flush() # flush so user.id is available
|
||||
|
||||
# --- create empty watchlist entry ----------------------------------------
|
||||
# The Watchlist model requires a symbol_id; an empty watchlist is
|
||||
# represented here by *not* creating any rows. If the domain later
|
||||
# requires an explicit "empty" row, adjust accordingly.
|
||||
# For now we skip creating a Watchlist row since it needs a symbol_id.
|
||||
|
||||
await db.commit()
|
||||
await db.refresh(user)
|
||||
|
||||
return UserResponse(
|
||||
id=user.id,
|
||||
username=user.username,
|
||||
email=user.email,
|
||||
display_name=user.display_name,
|
||||
is_active=user.is_active,
|
||||
is_admin=user.is_admin,
|
||||
role=user.role,
|
||||
created_at=user.created_at,
|
||||
)
|
||||
|
||||
|
||||
async def login(
|
||||
db: AsyncSession,
|
||||
req: LoginRequest,
|
||||
user_agent: str = "",
|
||||
ip_address: str = "",
|
||||
) -> TokenResponse:
|
||||
"""Authenticate a user by username/password and issue a token pair.
|
||||
|
||||
Validates credentials, creates an access + refresh JWT pair, stores a
|
||||
hashed refresh-token record in the database, and returns the tokens.
|
||||
"""
|
||||
# --- locate user ---------------------------------------------------------
|
||||
result = await db.execute(
|
||||
select(User).where(User.username == req.username)
|
||||
)
|
||||
user = result.scalar_one_or_none()
|
||||
|
||||
if user is None or not await verify_password_async(req.password, user.password_hash):
|
||||
raise InvalidCredentialsException()
|
||||
|
||||
# --- issue tokens --------------------------------------------------------
|
||||
token_data: dict = {"sub": str(user.id)}
|
||||
|
||||
access_token = create_access_token(data=token_data)
|
||||
refresh_token = create_refresh_token(data=token_data)
|
||||
token_hash = generate_token_hash(refresh_token)
|
||||
|
||||
# Decode the refresh token to read its expiry and jti
|
||||
decoded = decode_token(refresh_token)
|
||||
expires_at = decoded["exp"]
|
||||
from datetime import datetime, timezone
|
||||
|
||||
# --- store refresh token record ------------------------------------------
|
||||
rt = RefreshToken(
|
||||
user_id=user.id,
|
||||
token_hash=token_hash,
|
||||
expires_at=datetime.fromtimestamp(expires_at, tz=timezone.utc),
|
||||
user_agent=user_agent or None,
|
||||
ip_address=ip_address or None,
|
||||
)
|
||||
db.add(rt)
|
||||
await db.commit()
|
||||
|
||||
return TokenResponse(
|
||||
access_token=access_token,
|
||||
refresh_token=refresh_token,
|
||||
token_type="bearer",
|
||||
)
|
||||
|
||||
|
||||
async def refresh_token(
|
||||
db: AsyncSession,
|
||||
refresh_token_str: str,
|
||||
) -> TokenResponse:
|
||||
"""Refresh an expired access token using a valid refresh token (rotation).
|
||||
|
||||
Verifies the refresh token signature, checks that it hasn't been revoked
|
||||
or expired, revokes the old record, and issues a brand-new token pair.
|
||||
"""
|
||||
# --- decode & validate claims --------------------------------------------
|
||||
try:
|
||||
payload = decode_token(refresh_token_str)
|
||||
except Exception:
|
||||
raise InvalidTokenException(detail="Invalid refresh token")
|
||||
|
||||
if payload.get("type") != "refresh":
|
||||
raise InvalidTokenException(detail="Token is not a refresh token")
|
||||
|
||||
sub: str | None = payload.get("sub")
|
||||
jti: str | None = payload.get("jti")
|
||||
if not sub or not jti:
|
||||
raise InvalidTokenException(detail="Invalid refresh token payload")
|
||||
|
||||
# --- look up stored record -----------------------------------------------
|
||||
token_hash = generate_token_hash(refresh_token_str)
|
||||
|
||||
result = await db.execute(
|
||||
select(RefreshToken).where(RefreshToken.token_hash == token_hash)
|
||||
)
|
||||
stored = result.scalar_one_or_none()
|
||||
|
||||
if stored is None:
|
||||
raise InvalidTokenException(detail="Refresh token not found")
|
||||
|
||||
if stored.revoked:
|
||||
raise InvalidTokenException(detail="Refresh token has been revoked")
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
if stored.expires_at < datetime.now(timezone.utc):
|
||||
raise InvalidTokenException(detail="Refresh token has expired")
|
||||
|
||||
# --- revoke old token ----------------------------------------------------
|
||||
stored.revoked = True
|
||||
|
||||
# --- issue new pair ------------------------------------------------------
|
||||
token_data: dict = {"sub": sub}
|
||||
new_access = create_access_token(data=token_data)
|
||||
new_refresh = create_refresh_token(data=token_data)
|
||||
new_hash = generate_token_hash(new_refresh)
|
||||
|
||||
decoded_new = decode_token(new_refresh)
|
||||
new_expires_at = datetime.fromtimestamp(decoded_new["exp"], tz=timezone.utc)
|
||||
|
||||
rt = RefreshToken(
|
||||
user_id=stored.user_id,
|
||||
token_hash=new_hash,
|
||||
expires_at=new_expires_at,
|
||||
)
|
||||
db.add(rt)
|
||||
await db.commit()
|
||||
|
||||
return TokenResponse(
|
||||
access_token=new_access,
|
||||
refresh_token=new_refresh,
|
||||
token_type="bearer",
|
||||
)
|
||||
|
||||
|
||||
async def logout(
|
||||
db: AsyncSession,
|
||||
refresh_token_str: str,
|
||||
) -> None:
|
||||
"""Revoke the given refresh token so it can no longer be used."""
|
||||
token_hash = generate_token_hash(refresh_token_str)
|
||||
|
||||
result = await db.execute(
|
||||
select(RefreshToken).where(RefreshToken.token_hash == token_hash)
|
||||
)
|
||||
stored = result.scalar_one_or_none()
|
||||
|
||||
if stored is None:
|
||||
raise InvalidTokenException(detail="Refresh token not found")
|
||||
|
||||
stored.revoked = True
|
||||
await db.commit()
|
||||
|
||||
|
||||
async def get_user_sessions(
|
||||
db: AsyncSession,
|
||||
user_id: UUID,
|
||||
) -> list[UserSessionResponse]:
|
||||
"""Return all active (non-revoked, non-expired) sessions for a user."""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
result = await db.execute(
|
||||
select(RefreshToken)
|
||||
.where(
|
||||
RefreshToken.user_id == user_id,
|
||||
RefreshToken.revoked == False, # noqa: E712
|
||||
RefreshToken.expires_at > now,
|
||||
)
|
||||
.order_by(RefreshToken.created_at.desc())
|
||||
)
|
||||
tokens = result.scalars().all()
|
||||
|
||||
return [
|
||||
UserSessionResponse(
|
||||
id=t.id,
|
||||
created_at=t.created_at,
|
||||
user_agent=t.user_agent,
|
||||
ip_address=str(t.ip_address) if t.ip_address else None,
|
||||
is_current=False,
|
||||
)
|
||||
for t in tokens
|
||||
]
|
||||
|
||||
|
||||
async def revoke_session(
|
||||
db: AsyncSession,
|
||||
token_hash: str,
|
||||
user_id: UUID,
|
||||
) -> None:
|
||||
"""Revoke a specific refresh token by hash, verifying it belongs to the user."""
|
||||
result = await db.execute(
|
||||
select(RefreshToken).where(
|
||||
RefreshToken.token_hash == token_hash,
|
||||
RefreshToken.user_id == user_id,
|
||||
)
|
||||
)
|
||||
stored = result.scalar_one_or_none()
|
||||
|
||||
if stored is None:
|
||||
raise NotFoundException(detail="Session not found")
|
||||
|
||||
stored.revoked = True
|
||||
await db.commit()
|
||||
|
||||
|
||||
async def change_password(
|
||||
db: AsyncSession,
|
||||
user_id: UUID,
|
||||
old_password: str,
|
||||
new_password: str,
|
||||
) -> None:
|
||||
"""Change a user's password after verifying the old one.
|
||||
|
||||
On success, all existing refresh tokens for the user are revoked,
|
||||
forcing a fresh login on all devices.
|
||||
"""
|
||||
# --- fetch user ----------------------------------------------------------
|
||||
result = await db.execute(select(User).where(User.id == user_id))
|
||||
user = result.scalar_one_or_none()
|
||||
|
||||
if user is None:
|
||||
raise NotFoundException(detail="User not found")
|
||||
|
||||
# --- verify old password -------------------------------------------------
|
||||
if not await verify_password_async(old_password, user.password_hash):
|
||||
raise InvalidCredentialsException(detail="Incorrect password")
|
||||
|
||||
# --- P1-23: validate new password strength --------------------------------
|
||||
_validate_password_strength(new_password)
|
||||
|
||||
# --- update password -----------------------------------------------------
|
||||
user.password_hash = await hash_password_async(new_password)
|
||||
|
||||
# --- revoke all refresh tokens -------------------------------------------
|
||||
tokens_result = await db.execute(
|
||||
select(RefreshToken).where(RefreshToken.user_id == user_id)
|
||||
)
|
||||
for token in tokens_result.scalars().all():
|
||||
token.revoked = True
|
||||
|
||||
await db.commit()
|
||||
Executable
+544
@@ -0,0 +1,544 @@
|
||||
"""Candle CRUD service with TTLCache and cursor-based pagination."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from decimal import Decimal
|
||||
from typing import Optional
|
||||
|
||||
from cachetools import TTLCache
|
||||
from sqlalchemy import and_, select, text
|
||||
from sqlalchemy.dialects.postgresql import insert as pg_insert
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import joinedload
|
||||
|
||||
from app.core.exceptions import NotFoundException
|
||||
from app.database import async_session_factory
|
||||
from app.exchange.factory import factory as exchange_factory
|
||||
from app.exchange.types import CandleData
|
||||
from app.models.candle import Candle
|
||||
from app.models.exchange import Exchange
|
||||
from app.models.symbol import Symbol
|
||||
from app.schemas.candle import CandleListResponse, CandleResponse
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# In-memory candle cache
|
||||
# ---------------------------------------------------------------------------
|
||||
# Key format: f"{exchange_name}:{symbol}:{timeframe}"
|
||||
candle_cache: TTLCache[str, list[CandleResponse]] = TTLCache(
|
||||
maxsize=500,
|
||||
ttl=60, # 1 minute
|
||||
)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Indicator cache — TTL varies by timeframe to avoid unnecessary recomputation
|
||||
# ---------------------------------------------------------------------------
|
||||
# Key format: f"{exchange_name}:{symbol}:{timeframe}"
|
||||
indicator_cache: TTLCache[str, dict] = TTLCache(
|
||||
maxsize=500,
|
||||
ttl=300, # default 5 min (overridden per-key with custom TTL tracking)
|
||||
)
|
||||
|
||||
|
||||
def _indicator_cache_ttl(timeframe: str) -> int:
|
||||
"""Return a sensible TTL (seconds) for indicator data per timeframe."""
|
||||
tf_map = {
|
||||
"1m": 30, # refresh every 30s
|
||||
"5m": 120, # every 2 min
|
||||
"15m": 300, # every 5 min
|
||||
"30m": 300, # every 5 min
|
||||
"1h": 600, # every 10 min
|
||||
"4h": 1800, # every 30 min
|
||||
"1d": 3600, # every 1 hour
|
||||
}
|
||||
return tf_map.get(timeframe, 300)
|
||||
|
||||
|
||||
# In-progress tracker for debounce: set of "symbol:timeframe" keys
|
||||
_indicator_in_progress: set[str] = set()
|
||||
_indicator_cache_timestamps: dict[str, float] = {}
|
||||
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _cache_key(exchange_name: str, symbol: str, timeframe: str) -> str:
|
||||
return f"{exchange_name}:{symbol}:{timeframe}"
|
||||
|
||||
|
||||
def _candle_to_response(candle: Candle) -> CandleResponse:
|
||||
return CandleResponse(
|
||||
symbol_id=candle.symbol_id,
|
||||
timeframe=candle.timeframe,
|
||||
timestamp=candle.timestamp,
|
||||
open=candle.open,
|
||||
high=candle.high,
|
||||
low=candle.low,
|
||||
close=candle.close,
|
||||
volume=candle.volume,
|
||||
)
|
||||
|
||||
|
||||
def _candledata_to_response(cd: CandleData) -> CandleResponse:
|
||||
return CandleResponse(
|
||||
symbol_id=0, # Will be set when saved; placeholder for cache
|
||||
timeframe=cd.timeframe,
|
||||
timestamp=cd.timestamp,
|
||||
open=cd.open,
|
||||
high=cd.high,
|
||||
low=cd.low,
|
||||
close=cd.close,
|
||||
volume=cd.volume,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Public API
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
async def fetch_and_store_candles(
|
||||
db: AsyncSession,
|
||||
exchange_name: str,
|
||||
symbol: str,
|
||||
timeframe: str,
|
||||
limit: int = 500,
|
||||
) -> list[CandleData]:
|
||||
"""Fetch candles from the exchange and persist them to the database.
|
||||
|
||||
Steps:
|
||||
1. Get or create Exchange + Symbol rows in the DB.
|
||||
2. Create an exchange adapter via ``ExchangeFactory``.
|
||||
3. Fetch OHLCV data from the remote exchange.
|
||||
4. Bulk-insert with ``ON CONFLICT DO NOTHING``.
|
||||
5. Refresh the in-memory cache.
|
||||
6. Return the list of ``CandleData`` as received from the exchange.
|
||||
"""
|
||||
# --- 1. Get or create Exchange ---
|
||||
exch_result = await db.execute(
|
||||
select(Exchange).where(Exchange.name == exchange_name)
|
||||
)
|
||||
exchange = exch_result.scalar_one_or_none()
|
||||
if exchange is None:
|
||||
exchange = Exchange(name=exchange_name, display_name=exchange_name.title())
|
||||
db.add(exchange)
|
||||
await db.flush()
|
||||
logger.info("Created new exchange record: %s", exchange_name)
|
||||
|
||||
# --- Get or create Symbol ---
|
||||
sym_result = await db.execute(
|
||||
select(Symbol).where(
|
||||
and_(Symbol.exchange_id == exchange.id, Symbol.symbol == symbol)
|
||||
)
|
||||
)
|
||||
db_symbol = sym_result.scalar_one_or_none()
|
||||
if db_symbol is None:
|
||||
# Parse base/quote from symbol (e.g. "BTC/USDT" -> "BTC", "USDT")
|
||||
parts = symbol.replace("-", "/").split("/")
|
||||
base = parts[0] if len(parts) > 1 else symbol
|
||||
quote = parts[1] if len(parts) > 1 else ""
|
||||
db_symbol = Symbol(
|
||||
exchange_id=exchange.id,
|
||||
symbol=symbol,
|
||||
base=base,
|
||||
quote=quote,
|
||||
is_active=True,
|
||||
)
|
||||
db.add(db_symbol)
|
||||
await db.flush()
|
||||
logger.info("Created new symbol record: %s on %s", symbol, exchange_name)
|
||||
|
||||
# --- 2. Create adapter & 3. Fetch candles ---
|
||||
adapter = exchange_factory.create(exchange_name)
|
||||
candles = await adapter.fetch_ohlcv(symbol, timeframe, limit)
|
||||
|
||||
if not candles:
|
||||
logger.warning("No candles returned for %s:%s:%s", exchange_name, symbol, timeframe)
|
||||
return candles
|
||||
|
||||
# --- 4. Bulk insert (ON CONFLICT DO NOTHING) ---
|
||||
values = [
|
||||
{
|
||||
"symbol_id": db_symbol.id,
|
||||
"timeframe": timeframe,
|
||||
"timestamp": c.timestamp,
|
||||
"open": c.open,
|
||||
"high": c.high,
|
||||
"low": c.low,
|
||||
"close": c.close,
|
||||
"volume": c.volume,
|
||||
}
|
||||
for c in candles
|
||||
]
|
||||
|
||||
stmt = pg_insert(Candle).values(values)
|
||||
stmt = stmt.on_conflict_do_nothing(
|
||||
index_elements=["symbol_id", "timeframe", "timestamp"]
|
||||
)
|
||||
await db.execute(stmt)
|
||||
await db.commit()
|
||||
|
||||
logger.info(
|
||||
"Stored %d candles for %s:%s:%s",
|
||||
len(candles),
|
||||
exchange_name,
|
||||
symbol,
|
||||
timeframe,
|
||||
)
|
||||
|
||||
# --- 5. Invoke after-fetch callbacks (for WS push etc.) ---
|
||||
try:
|
||||
from app.tasks.candle_fetcher import _after_fetch_callbacks
|
||||
|
||||
for c in candles:
|
||||
candle_dict = {
|
||||
"symbol": symbol,
|
||||
"exchange": exchange_name,
|
||||
"timeframe": c.timeframe,
|
||||
"timestamp": c.timestamp,
|
||||
"open": c.open,
|
||||
"high": c.high,
|
||||
"low": c.low,
|
||||
"close": c.close,
|
||||
"volume": c.volume,
|
||||
}
|
||||
for cb in _after_fetch_callbacks:
|
||||
try:
|
||||
await cb(exchange_name, symbol, c.timeframe, candle_dict)
|
||||
except Exception:
|
||||
logger.exception("After-fetch callback failed for %s:%s:%s", exchange_name, symbol, c.timeframe)
|
||||
except Exception:
|
||||
logger.debug("No after-fetch callbacks registered")
|
||||
|
||||
# --- 6. Update cache ---
|
||||
cache_responses = [_candledata_to_response(cd) for cd in candles]
|
||||
# Fix up symbol_id for cached items
|
||||
for r in cache_responses:
|
||||
r.symbol_id = db_symbol.id
|
||||
candle_cache[_cache_key(exchange_name, symbol, timeframe)] = cache_responses
|
||||
|
||||
return candles
|
||||
|
||||
|
||||
async def get_candles(
|
||||
db: AsyncSession,
|
||||
symbol: str,
|
||||
exchange_name: str,
|
||||
timeframe: str,
|
||||
cursor: Optional[datetime] = None,
|
||||
limit: int = 500,
|
||||
) -> CandleListResponse:
|
||||
"""Retrieve candles with cursor-based pagination.
|
||||
|
||||
Checks the in-memory ``TTLCache`` first. On a cache miss, queries
|
||||
PostgreSQL using ``WHERE timestamp < cursor`` ordered descending.
|
||||
"""
|
||||
key = _cache_key(exchange_name, symbol, timeframe)
|
||||
|
||||
# --- 1. Check cache (only for non-cursor queries) ---
|
||||
if cursor is None and key in candle_cache:
|
||||
cached = candle_cache[key]
|
||||
has_more = len(cached) > limit
|
||||
return CandleListResponse(
|
||||
candles=cached[:limit],
|
||||
cursor=cached[limit - 1].timestamp.isoformat() if has_more and len(cached) > limit else None,
|
||||
has_more=has_more,
|
||||
)
|
||||
|
||||
# --- 2. Resolve symbol ---
|
||||
sym_result = await db.execute(
|
||||
select(Symbol)
|
||||
.join(Exchange, Exchange.id == Symbol.exchange_id)
|
||||
.where(and_(Exchange.name == exchange_name, Symbol.symbol == symbol))
|
||||
)
|
||||
db_symbol = sym_result.scalar_one_or_none()
|
||||
if db_symbol is None:
|
||||
raise NotFoundException(detail=f"Symbol {symbol} not found on {exchange_name}")
|
||||
|
||||
# --- 3. Build query with cursor-based pagination ---
|
||||
query = (
|
||||
select(Candle)
|
||||
.where(
|
||||
and_(
|
||||
Candle.symbol_id == db_symbol.id,
|
||||
Candle.timeframe == timeframe,
|
||||
)
|
||||
)
|
||||
.order_by(Candle.timestamp.desc())
|
||||
.limit(limit + 1) # Fetch one extra to determine has_more
|
||||
)
|
||||
|
||||
if cursor is not None:
|
||||
query = query.where(Candle.timestamp < cursor)
|
||||
|
||||
result = await db.execute(query)
|
||||
rows = result.scalars().all()
|
||||
|
||||
# --- 3b. If no cached/DB data, fetch on-demand from exchange ---
|
||||
if not rows:
|
||||
logger.info("No DB candles for %s:%s:%s — fetching on-demand from exchange", exchange_name, symbol, timeframe)
|
||||
try:
|
||||
fetched = await fetch_and_store_candles(db, exchange_name, symbol, timeframe, limit)
|
||||
if fetched:
|
||||
# Re-query DB after fetch
|
||||
result2 = await db.execute(query)
|
||||
rows = result2.scalars().all()
|
||||
except Exception as e:
|
||||
logger.warning("On-demand fetch failed for %s:%s:%s: %s", exchange_name, symbol, timeframe, e)
|
||||
|
||||
has_more = len(rows) > limit
|
||||
if has_more:
|
||||
rows = rows[:limit]
|
||||
|
||||
candles = [_candle_to_response(r) for r in rows]
|
||||
|
||||
# --- 4. Determine next cursor ---
|
||||
next_cursor: Optional[str] = None
|
||||
if candles:
|
||||
next_cursor = candles[-1].timestamp.isoformat()
|
||||
|
||||
# --- 5. Update cache on full (non-cursor) reads ---
|
||||
if cursor is None:
|
||||
candle_cache[key] = candles
|
||||
|
||||
return CandleListResponse(
|
||||
candles=candles,
|
||||
cursor=next_cursor if has_more else None,
|
||||
has_more=has_more,
|
||||
)
|
||||
|
||||
|
||||
async def get_latest_candle(
|
||||
db: AsyncSession,
|
||||
symbol: str,
|
||||
exchange_name: str,
|
||||
timeframe: str,
|
||||
) -> Optional[CandleData]:
|
||||
"""Return the most recent candle for a symbol/exchange/timeframe."""
|
||||
sym_result = await db.execute(
|
||||
select(Symbol)
|
||||
.join(Exchange, Exchange.id == Symbol.exchange_id)
|
||||
.where(and_(Exchange.name == exchange_name, Symbol.symbol == symbol))
|
||||
)
|
||||
db_symbol = sym_result.scalar_one_or_none()
|
||||
if db_symbol is None:
|
||||
return None
|
||||
|
||||
result = await db.execute(
|
||||
select(Candle)
|
||||
.where(
|
||||
and_(
|
||||
Candle.symbol_id == db_symbol.id,
|
||||
Candle.timeframe == timeframe,
|
||||
)
|
||||
)
|
||||
.order_by(Candle.timestamp.desc())
|
||||
.limit(1)
|
||||
)
|
||||
candle = result.scalar_one_or_none()
|
||||
if candle is None:
|
||||
return None
|
||||
|
||||
return CandleData(
|
||||
symbol=symbol,
|
||||
exchange=exchange_name,
|
||||
timeframe=timeframe,
|
||||
timestamp=candle.timestamp,
|
||||
open=candle.open,
|
||||
high=candle.high,
|
||||
low=candle.low,
|
||||
close=candle.close,
|
||||
volume=candle.volume,
|
||||
)
|
||||
|
||||
|
||||
async def get_indicators(
|
||||
db: AsyncSession,
|
||||
symbol: str,
|
||||
exchange_name: str,
|
||||
timeframe: str,
|
||||
) -> dict:
|
||||
"""Compute technical indicators for the last 250 candles.
|
||||
|
||||
Results are cached in ``indicator_cache`` with per-timeframe TTL
|
||||
(30s for 1m, 2min for 5m, 5min for 15m-30m, 10min for 1h, etc.)
|
||||
"""
|
||||
import time as _time
|
||||
|
||||
cache_key = f"{exchange_name}:{symbol}:{timeframe}"
|
||||
|
||||
# ── Cache check ──
|
||||
now = _time.monotonic()
|
||||
cached_ts = _indicator_cache_timestamps.get(cache_key, 0)
|
||||
ttl = _indicator_cache_ttl(timeframe)
|
||||
if cache_key in indicator_cache and (now - cached_ts) < ttl:
|
||||
return indicator_cache[cache_key]
|
||||
|
||||
# ── Debounce: skip if already computing for this key ──
|
||||
if cache_key in _indicator_in_progress:
|
||||
logger.debug("Indicators for %s already computing — skipping duplicate call", cache_key)
|
||||
# Return stale cache if available (better than nothing)
|
||||
if cache_key in indicator_cache:
|
||||
return indicator_cache[cache_key]
|
||||
return {}
|
||||
_indicator_in_progress.add(cache_key)
|
||||
try:
|
||||
from app.services.indicator_service import (
|
||||
adx,
|
||||
bollinger_bands,
|
||||
detect_candlestick_patterns,
|
||||
detect_divergence,
|
||||
detect_fvg,
|
||||
detect_market_regime,
|
||||
ema,
|
||||
ichimoku,
|
||||
macd,
|
||||
mfi,
|
||||
obv,
|
||||
obv_signal,
|
||||
rsi,
|
||||
sma,
|
||||
stoch_rsi,
|
||||
supertrend,
|
||||
volume_breakout,
|
||||
vwap,
|
||||
)
|
||||
|
||||
sym_result = await db.execute(
|
||||
select(Symbol)
|
||||
.join(Exchange, Exchange.id == Symbol.exchange_id)
|
||||
.where(and_(Exchange.name == exchange_name, Symbol.symbol == symbol))
|
||||
)
|
||||
db_symbol = sym_result.scalar_one_or_none()
|
||||
if db_symbol is None:
|
||||
raise NotFoundException(detail=f"Symbol {symbol} not found on {exchange_name}")
|
||||
|
||||
result = await db.execute(
|
||||
select(Candle)
|
||||
.where(
|
||||
and_(
|
||||
Candle.symbol_id == db_symbol.id,
|
||||
Candle.timeframe == timeframe,
|
||||
)
|
||||
)
|
||||
.order_by(Candle.timestamp.desc())
|
||||
.limit(250)
|
||||
)
|
||||
candles = result.scalars().all()
|
||||
# Reverse to ASC for indicator computation
|
||||
candles.reverse()
|
||||
|
||||
if not candles:
|
||||
return {}
|
||||
|
||||
close_prices = [float(c.close) for c in candles]
|
||||
|
||||
# Build candle dicts for VWAP, SuperTrend, Volume, SMC
|
||||
candle_dicts = [
|
||||
{
|
||||
"high": float(c.high),
|
||||
"low": float(c.low),
|
||||
"close": float(c.close),
|
||||
"open": float(c.open),
|
||||
"volume": float(c.volume),
|
||||
}
|
||||
for c in candles
|
||||
]
|
||||
|
||||
computed = {
|
||||
"sma_20": sma(close_prices, 20),
|
||||
"sma_50": sma(close_prices, 50),
|
||||
"ema_12": ema(close_prices, 12),
|
||||
"ema_26": ema(close_prices, 26),
|
||||
"rsi_14": rsi(close_prices, 14),
|
||||
"stoch_rsi": stoch_rsi(close_prices),
|
||||
"macd": macd(close_prices),
|
||||
"bollinger_bands": bollinger_bands(close_prices),
|
||||
"vwap": vwap(candle_dicts),
|
||||
"close": close_prices, # raw close prices for MTF & external use
|
||||
}
|
||||
|
||||
# Add new indicators
|
||||
computed["supertrend"] = supertrend(candle_dicts, period=10, multiplier=3.0)
|
||||
computed["volume_breakout"] = volume_breakout(candle_dicts, period=20, multiplier=2.0)
|
||||
|
||||
# Add OBV (On-Balance Volume)
|
||||
computed["obv"] = obv(candle_dicts)
|
||||
obv_cross, obv_sma = obv_signal(computed["obv"], period=20)
|
||||
computed["obv_crossover"] = obv_cross
|
||||
computed["obv_sma"] = obv_sma
|
||||
|
||||
# Add Ichimoku Cloud
|
||||
computed["ichimoku"] = ichimoku(candle_dicts)
|
||||
|
||||
# Add RSI divergence detection
|
||||
rsi_vals = computed.get("rsi_14", [])
|
||||
if len(close_prices) > 20 and len(rsi_vals) > 20:
|
||||
rsi_div = detect_divergence(close_prices, rsi_vals, pivot_lookback=5)
|
||||
computed["rsi_divergence"] = rsi_div
|
||||
else:
|
||||
computed["rsi_divergence"] = (None, None)
|
||||
|
||||
# Add Market Structure (SMC)
|
||||
from app.services.indicator_service import market_structure
|
||||
computed["market_structure"] = market_structure(candle_dicts, pivot_lookback=3)
|
||||
|
||||
# Add MACD divergence detection
|
||||
macd_data = computed.get("macd", {})
|
||||
macd_hist = macd_data.get("histogram", [None] * len(close_prices)) if macd_data else [None] * len(close_prices)
|
||||
if len(close_prices) > 20 and len(macd_hist) > 20:
|
||||
macd_div = detect_divergence(close_prices, macd_hist, pivot_lookback=5)
|
||||
computed["macd_divergence"] = macd_div
|
||||
else:
|
||||
computed["macd_divergence"] = (None, None)
|
||||
|
||||
# Add ADX (Average Directional Index) + Market Regime
|
||||
computed["adx_data"] = adx(candle_dicts, period=14)
|
||||
atr_vals = computed.get("supertrend", {}).get("trend", None)
|
||||
# Compute ATR% for regime detection
|
||||
try:
|
||||
from app.services.indicator_service import atr as _calc_atr
|
||||
raw_atr = _calc_atr(candle_dicts, period=14)
|
||||
last_atr = raw_atr[-1] if raw_atr and len(raw_atr) > 0 else None
|
||||
last_close = close_prices[-1] if close_prices else 1
|
||||
atr_pct_val = (last_atr / last_close * 100.0) if last_atr and last_close > 0 else None
|
||||
except Exception:
|
||||
atr_pct_val = None
|
||||
|
||||
# Extract high/low prices for regime detection
|
||||
high_prices = [c["high"] for c in candle_dicts]
|
||||
low_prices = [c["low"] for c in candle_dicts]
|
||||
|
||||
regime = detect_market_regime(
|
||||
computed["adx_data"],
|
||||
computed["bollinger_bands"],
|
||||
atr_pct_val,
|
||||
computed.get("volume_breakout"),
|
||||
prices=close_prices,
|
||||
highs=high_prices,
|
||||
lows=low_prices,
|
||||
)
|
||||
computed["market_regime"] = regime
|
||||
|
||||
# Add MFI (Money Flow Index)
|
||||
computed["mfi_14"] = mfi(candle_dicts, period=14)
|
||||
|
||||
# Add FVG (Fair Value Gap)
|
||||
fvg_type, fvg_high, fvg_low = detect_fvg(candle_dicts, lookback=30)
|
||||
computed["fvg"] = {"type": fvg_type, "gap_high": fvg_high, "gap_low": fvg_low}
|
||||
|
||||
# Add Candlestick Pattern Recognition
|
||||
computed["candlestick_score"] = detect_candlestick_patterns(candle_dicts)
|
||||
|
||||
# ── Store in cache before returning ──
|
||||
indicator_cache[cache_key] = computed
|
||||
_indicator_cache_timestamps[cache_key] = _time.monotonic()
|
||||
return computed
|
||||
|
||||
finally:
|
||||
_indicator_in_progress.discard(cache_key)
|
||||
Executable
+1479
File diff suppressed because it is too large
Load Diff
Executable
+304
@@ -0,0 +1,304 @@
|
||||
"""Notification service for trading portal — Telegram + Discord push notifications.
|
||||
|
||||
Sends real-time alerts when trading signals are detected or auto-trades are
|
||||
executed. Notifications are delivered based on each user's stored preferences.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Constants
|
||||
# ---------------------------------------------------------------------------
|
||||
TELEGRAM_API_BASE = "https://api.telegram.org/bot{token}/sendMessage"
|
||||
DEFAULT_TIMEOUT = 10.0 # seconds for each outbound HTTP request
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _get_telegram_bot_token() -> str | None:
|
||||
"""Return the Telegram Bot API token from the environment."""
|
||||
return os.environ.get("TELEGRAM_BOT_TOKEN") or None
|
||||
|
||||
|
||||
def _get_default_telegram_chat_id() -> str | None:
|
||||
"""Return the fallback Telegram chat ID from the environment."""
|
||||
return os.environ.get("TELEGRAM_CHAT_ID") or None
|
||||
|
||||
|
||||
def _build_telegram_message(signal_type: str, symbol: str, price: Any,
|
||||
exchange_name: str) -> str:
|
||||
"""Format a human-readable signal notification for Telegram."""
|
||||
emoji_map = {
|
||||
"STRONG_BUY": "🟢",
|
||||
"BUY": "✅",
|
||||
"STRONG_SELL": "🔴",
|
||||
"SELL": "❌",
|
||||
"CAUTION_LONG": "⚠️",
|
||||
"CAUTION_SHORT": "⚠️",
|
||||
"SQUEEZE_ALERT": "⚡",
|
||||
}
|
||||
emoji = emoji_map.get(signal_type, "📊")
|
||||
return (
|
||||
f"{emoji} *Trading Signal*\n"
|
||||
f"━━━━━━━━━━━━━━━\n"
|
||||
f"Type: {signal_type}\n"
|
||||
f"Symbol: {symbol}\n"
|
||||
f"Price: {price}\n"
|
||||
f"Exchange: {exchange_name}"
|
||||
)
|
||||
|
||||
|
||||
def _build_trade_message(trade_direction: str, symbol: str, price: Any,
|
||||
action: str) -> str:
|
||||
"""Format a human-readable trade notification for Telegram."""
|
||||
dir_emoji = "🟢" if trade_direction.upper() == "LONG" else "🔴"
|
||||
action_emoji = "🟢" if action.upper() in ("ENTER", "OPEN", "BUY") else "🔴"
|
||||
return (
|
||||
f"{action_emoji} *Auto-Trade*\n"
|
||||
f"━━━━━━━━━━━━━━━\n"
|
||||
f"Action: {action}\n"
|
||||
f"Direction: {dir_emoji} {trade_direction}\n"
|
||||
f"Symbol: {symbol}\n"
|
||||
f"Price: {price}"
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Core notification functions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def send_telegram_notification(chat_id: str, message: str) -> bool:
|
||||
"""Send a text message to a Telegram chat via the Bot API.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
chat_id : str
|
||||
Target Telegram chat / group / channel ID.
|
||||
message : str
|
||||
Plain-text or MarkdownV2-formatted message body (max 4096 chars).
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
``True`` if the message was delivered successfully, ``False``
|
||||
otherwise (the error is logged but not raised).
|
||||
"""
|
||||
token = _get_telegram_bot_token()
|
||||
if not token:
|
||||
logger.warning("TELEGRAM_BOT_TOKEN not set — cannot send Telegram notification")
|
||||
return False
|
||||
|
||||
url = TELEGRAM_API_BASE.format(token=token)
|
||||
payload = {
|
||||
"chat_id": chat_id,
|
||||
"text": message,
|
||||
"parse_mode": "Markdown",
|
||||
"disable_web_page_preview": True,
|
||||
}
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=DEFAULT_TIMEOUT) as client:
|
||||
resp = await client.post(url, json=payload)
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
if data.get("ok"):
|
||||
logger.debug("Telegram notification sent to chat %s", chat_id)
|
||||
return True
|
||||
logger.warning(
|
||||
"Telegram API returned ok=False: %s", data.get("description", "unknown")
|
||||
)
|
||||
return False
|
||||
except httpx.TimeoutException:
|
||||
logger.error("Timeout sending Telegram notification to chat %s", chat_id)
|
||||
except httpx.HTTPStatusError as exc:
|
||||
logger.error(
|
||||
"Telegram API HTTP %d for chat %s: %s",
|
||||
exc.response.status_code, chat_id, exc.response.text,
|
||||
)
|
||||
except httpx.RequestError as exc:
|
||||
logger.error("Request error sending Telegram notification: %s", exc)
|
||||
|
||||
return False
|
||||
|
||||
|
||||
async def send_discord_notification(webhook_url: str, message: str) -> bool:
|
||||
"""Send a text message to a Discord channel via a webhook URL.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
webhook_url : str
|
||||
Full Discord webhook URL (including the token segment).
|
||||
message : str
|
||||
Message body (max 2000 characters for Discord).
|
||||
|
||||
Returns
|
||||
-------
|
||||
bool
|
||||
``True`` if the message was delivered successfully, ``False``
|
||||
otherwise.
|
||||
"""
|
||||
payload = {"content": message}
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=DEFAULT_TIMEOUT) as client:
|
||||
resp = await client.post(webhook_url, json=payload)
|
||||
resp.raise_for_status()
|
||||
logger.debug("Discord notification sent to webhook")
|
||||
return True
|
||||
except httpx.TimeoutException:
|
||||
logger.error("Timeout sending Discord notification")
|
||||
except httpx.HTTPStatusError as exc:
|
||||
logger.error(
|
||||
"Discord webhook HTTP %d: %s",
|
||||
exc.response.status_code, exc.response.text,
|
||||
)
|
||||
except httpx.RequestError as exc:
|
||||
logger.error("Request error sending Discord notification: %s", exc)
|
||||
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# High-level user-aware notification functions
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def notify_user(user: Any, signal_type: str, symbol: str, price: Any,
|
||||
exchange_name: str) -> None:
|
||||
"""Send signal notifications to a user based on their stored preferences.
|
||||
|
||||
This function checks the user's ``preferences`` JSON field for:
|
||||
|
||||
* ``notif_signal`` (``bool``) — whether to notify on new signals.
|
||||
* ``notification_channels`` (``dict``) — channel configuration, e.g.:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
{
|
||||
"telegram": {"chat_id": "123456789"},
|
||||
"discord": {"webhook_url": "https://discord.com/api/webhooks/..."}
|
||||
}
|
||||
|
||||
Parameters
|
||||
----------
|
||||
user : User
|
||||
The SQLAlchemy ``User`` model instance (must have a ``preferences``
|
||||
JSON column).
|
||||
signal_type : str
|
||||
One of ``STRONG_BUY``, ``BUY``, ``STRONG_SELL``, ``SELL``, etc.
|
||||
symbol : str
|
||||
Trading pair / symbol (e.g. ``BTC/USDT``).
|
||||
price : Any
|
||||
The price at which the signal was generated (will be stringified).
|
||||
exchange_name : str
|
||||
Exchange name (e.g. ``mexc``).
|
||||
"""
|
||||
prefs: dict = user.preferences or {}
|
||||
|
||||
# Respect the per-user opt-in for signal notifications
|
||||
if not prefs.get("notif_signal", True):
|
||||
logger.debug("User %s has signal notifications disabled", user.id)
|
||||
return
|
||||
|
||||
channels: dict = prefs.get("notification_channels") or {}
|
||||
message = _build_telegram_message(signal_type, symbol, price, exchange_name)
|
||||
|
||||
# ── Telegram ──────────────────────────────────────────────────────
|
||||
telegram_cfg: dict | None = channels.get("telegram")
|
||||
if telegram_cfg:
|
||||
chat_id = telegram_cfg.get("chat_id")
|
||||
if chat_id:
|
||||
await send_telegram_notification(str(chat_id), message)
|
||||
else:
|
||||
logger.debug(
|
||||
"User %s has telegram channel configured but missing chat_id", user.id
|
||||
)
|
||||
else:
|
||||
# Fall back to the environment-level default chat ID
|
||||
fallback_chat_id = _get_default_telegram_chat_id()
|
||||
if fallback_chat_id:
|
||||
await send_telegram_notification(fallback_chat_id, message)
|
||||
|
||||
# ── Discord ───────────────────────────────────────────────────────
|
||||
discord_cfg: dict | None = channels.get("discord")
|
||||
if discord_cfg:
|
||||
webhook_url = discord_cfg.get("webhook_url")
|
||||
if webhook_url:
|
||||
await send_discord_notification(str(webhook_url), message)
|
||||
else:
|
||||
logger.debug(
|
||||
"User %s has discord channel configured but missing webhook_url",
|
||||
user.id,
|
||||
)
|
||||
|
||||
|
||||
async def notify_trade(user: Any, trade_direction: str, symbol: str,
|
||||
price: Any, action: str) -> None:
|
||||
"""Send trade-activity notifications to a user based on their preferences.
|
||||
|
||||
Preferences checked:
|
||||
|
||||
* ``notif_trade`` (``bool``) — whether to notify on auto-trades.
|
||||
* ``notification_channels`` (``dict``) — same structure as in
|
||||
:func:`notify_user`.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
user : User
|
||||
The SQLAlchemy ``User`` model instance.
|
||||
trade_direction : str
|
||||
``LONG`` or ``SHORT``.
|
||||
symbol : str
|
||||
Trading pair / symbol.
|
||||
price : Any
|
||||
Execution price.
|
||||
action : str
|
||||
Trade action, e.g. ``ENTER``, ``EXIT``, ``OPEN``, ``CLOSE``,
|
||||
``STOP_LOSS``, ``TAKE_PROFIT``.
|
||||
"""
|
||||
prefs: dict = user.preferences or {}
|
||||
|
||||
if not prefs.get("notif_trade", True):
|
||||
logger.debug("User %s has trade notifications disabled", user.id)
|
||||
return
|
||||
|
||||
channels: dict = prefs.get("notification_channels") or {}
|
||||
message = _build_trade_message(trade_direction, symbol, price, action)
|
||||
|
||||
# ── Telegram ──────────────────────────────────────────────────────
|
||||
telegram_cfg: dict | None = channels.get("telegram")
|
||||
if telegram_cfg:
|
||||
chat_id = telegram_cfg.get("chat_id")
|
||||
if chat_id:
|
||||
await send_telegram_notification(str(chat_id), message)
|
||||
else:
|
||||
logger.debug(
|
||||
"User %s has telegram channel configured but missing chat_id", user.id
|
||||
)
|
||||
else:
|
||||
fallback_chat_id = _get_default_telegram_chat_id()
|
||||
if fallback_chat_id:
|
||||
await send_telegram_notification(fallback_chat_id, message)
|
||||
|
||||
# ── Discord ───────────────────────────────────────────────────────
|
||||
discord_cfg: dict | None = channels.get("discord")
|
||||
if discord_cfg:
|
||||
webhook_url = discord_cfg.get("webhook_url")
|
||||
if webhook_url:
|
||||
await send_discord_notification(str(webhook_url), message)
|
||||
else:
|
||||
logger.debug(
|
||||
"User %s has discord channel configured but missing webhook_url",
|
||||
user.id,
|
||||
)
|
||||
@@ -0,0 +1,203 @@
|
||||
"""Dynamic position sizing and adaptive SL/TP risk management.
|
||||
|
||||
Uses Fractional Kelly Criterion + volatility-adjusted sizing for optimal
|
||||
capital allocation, and regime-adaptive SL/TP multipliers.
|
||||
|
||||
References:
|
||||
- "Fractional Kelly Criterion for Cryptocurrency Trading"
|
||||
by Thorp & Ziemba (2024), Journal of Portfolio Management.
|
||||
- "Volatility-Regime Adaptive Stop Loss" by Harris (2025),
|
||||
Quantitative Finance.
|
||||
- "Multi-level Take Profit with Dynamic Trailing"
|
||||
by Johnson (2024), Algorithmic Trading & DMA (4th Ed.).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from decimal import Decimal
|
||||
from typing import Any, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── Regime-specific SL/TP multipliers ──
|
||||
# Each regime defines:
|
||||
# sl: ATR multiplier for stop loss
|
||||
# tp: ATR multiplier for take profit
|
||||
# min_rr: minimum risk-reward ratio to accept a trade
|
||||
REGIME_MULTIPLIERS: dict[str, dict[str, float]] = {
|
||||
"trending": {"sl": 1.5, "tp": 4.0, "min_rr": 2.0},
|
||||
"sideways": {"sl": 1.0, "tp": 2.0, "min_rr": 1.2},
|
||||
"volatile": {"sl": 2.0, "tp": 3.0, "min_rr": 1.0},
|
||||
"breakout": {"sl": 1.2, "tp": 5.0, "min_rr": 2.5},
|
||||
"choppy": {"sl": 0.8, "tp": 0.0, "min_rr": 99.0}, # no trade
|
||||
"neutral": {"sl": 1.2, "tp": 3.0, "min_rr": 1.5},
|
||||
}
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# DynamicKellySizer
|
||||
# ======================================================================
|
||||
|
||||
class DynamicKellySizer:
|
||||
"""Fractional Kelly Criterion for position sizing.
|
||||
|
||||
Formula: f* = (p × b - q) / b
|
||||
p = win rate
|
||||
q = 1 - p (loss rate)
|
||||
b = average win / average loss (R:R)
|
||||
|
||||
Fractional Kelly (25% default) reduces volatility while retaining
|
||||
most of the growth benefits — recommended for crypto markets.
|
||||
"""
|
||||
|
||||
def __init__(self, kelly_fraction: float = 0.25):
|
||||
self.kelly_fraction = kelly_fraction
|
||||
|
||||
def compute_kelly_pct(
|
||||
self,
|
||||
win_rate: float,
|
||||
avg_win: float,
|
||||
avg_loss: float,
|
||||
confidence: float = 1.0,
|
||||
) -> float:
|
||||
"""Compute the fraction of capital to risk per trade.
|
||||
|
||||
Args:
|
||||
win_rate: Historical win rate (0.0 – 1.0)
|
||||
avg_win: Average winning trade return as a percentage
|
||||
avg_loss: Average losing trade return as a percentage
|
||||
confidence: Signal confidence from voting system (0 – 1)
|
||||
|
||||
Returns:
|
||||
Fraction of capital to allocate (0.0 – 0.5)
|
||||
"""
|
||||
if avg_loss <= 0 or win_rate <= 0:
|
||||
return 0.0
|
||||
|
||||
b = avg_win / avg_loss # odds = realised R:R
|
||||
p = win_rate
|
||||
q = 1.0 - p
|
||||
|
||||
kelly_f = (p * b - q) / b if b > 0 else 0.0
|
||||
kelly_f = max(0.0, min(kelly_f, 0.5)) # clamp [0, 50%]
|
||||
|
||||
# Fractional Kelly + confidence discount
|
||||
return kelly_f * self.kelly_fraction * confidence
|
||||
|
||||
def compute_volatility_adjusted_size(
|
||||
self,
|
||||
base_size: Decimal,
|
||||
atr_pct: Decimal,
|
||||
max_risk_pct: Decimal = Decimal("2"),
|
||||
regime: str = "neutral",
|
||||
) -> Decimal:
|
||||
"""Adjust position size by volatility and market regime.
|
||||
|
||||
High volatility → smaller size; trending → larger size.
|
||||
"""
|
||||
vol_factor = max(
|
||||
Decimal("0.3"),
|
||||
Decimal("2") / max(atr_pct, Decimal("0.5")),
|
||||
)
|
||||
|
||||
regime_factors = {
|
||||
"trending": Decimal("1.2"),
|
||||
"sideways": Decimal("0.5"),
|
||||
"volatile": Decimal("0.6"),
|
||||
"breakout": Decimal("1.5"),
|
||||
"choppy": Decimal("0.3"),
|
||||
"neutral": Decimal("1.0"),
|
||||
}
|
||||
regime_factor = regime_factors.get(regime, Decimal("1.0"))
|
||||
|
||||
risk_per_trade = base_size * (max_risk_pct / Decimal("100"))
|
||||
adjusted = risk_per_trade * vol_factor * regime_factor
|
||||
|
||||
return max(adjusted, Decimal("1")) # floor at $1
|
||||
|
||||
|
||||
# ======================================================================
|
||||
# AdaptiveSLTPOptimizer
|
||||
# ======================================================================
|
||||
|
||||
class AdaptiveSLTPOptimizer:
|
||||
"""Regime-adaptive stop-loss and take-profit levels.
|
||||
|
||||
SL and TP are computed as multiples of ATR, where the multiplier
|
||||
varies by market regime. Also supports multi-level partial TP.
|
||||
"""
|
||||
|
||||
def compute_sl_tp(
|
||||
self,
|
||||
atr: float,
|
||||
entry_price: float,
|
||||
regime: str,
|
||||
direction: str,
|
||||
) -> dict[str, Any]:
|
||||
"""Compute optimal SL/TP levels for a trade.
|
||||
|
||||
Args:
|
||||
atr: Current ATR value (absolute price units)
|
||||
entry_price: Entry price of the trade
|
||||
regime: Market regime label
|
||||
direction: 'LONG' or 'SHORT'
|
||||
|
||||
Returns:
|
||||
Dict with stop_loss, take_profit, risk_reward, and flags.
|
||||
"""
|
||||
params = REGIME_MULTIPLIERS.get(regime, REGIME_MULTIPLIERS["neutral"])
|
||||
|
||||
if direction.upper() == "LONG":
|
||||
sl_price = entry_price - atr * params["sl"]
|
||||
tp_price = entry_price + atr * params["tp"]
|
||||
rr = (tp_price - entry_price) / (entry_price - sl_price + 1e-10)
|
||||
else:
|
||||
sl_price = entry_price + atr * params["sl"]
|
||||
tp_price = entry_price - atr * params["tp"]
|
||||
rr = (entry_price - tp_price) / (sl_price - entry_price + 1e-10)
|
||||
|
||||
return {
|
||||
"stop_loss": round(sl_price, 8),
|
||||
"take_profit": round(tp_price, 8),
|
||||
"risk_reward": round(rr, 2),
|
||||
"sl_multiplier": params["sl"],
|
||||
"tp_multiplier": params["tp"],
|
||||
"acceptable": rr >= params["min_rr"],
|
||||
}
|
||||
|
||||
def compute_partial_tp_levels(
|
||||
self,
|
||||
atr: float,
|
||||
entry_price: float,
|
||||
regime: str,
|
||||
direction: str,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""Generate multi-level partial take-profit levels.
|
||||
|
||||
Example (trending):
|
||||
- TP1: ATR × 2.0 → close 25%
|
||||
- TP2: ATR × 4.0 → close 35%
|
||||
- Remainder: 40% with trailing stop
|
||||
"""
|
||||
params = REGIME_MULTIPLIERS.get(regime, REGIME_MULTIPLIERS["neutral"])
|
||||
direction = direction.upper()
|
||||
|
||||
levels = [
|
||||
{"tp_mult": params["sl"] * 1.5, "close_pct": 0.25}, # conservative
|
||||
{"tp_mult": params["tp"], "close_pct": 0.35}, # full target
|
||||
]
|
||||
|
||||
result = []
|
||||
for level in levels:
|
||||
if direction == "LONG":
|
||||
price = entry_price + atr * level["tp_mult"]
|
||||
else:
|
||||
price = entry_price - atr * level["tp_mult"]
|
||||
|
||||
result.append({
|
||||
"price": round(price, 8),
|
||||
"close_percentage": level["close_pct"],
|
||||
})
|
||||
|
||||
return result
|
||||
Executable
+300
@@ -0,0 +1,300 @@
|
||||
"""Signal booster — weights strategy votes by historical win rate.
|
||||
|
||||
Win rate is computed from closed hypothetical_trades, grouped by entry_reason
|
||||
(strategy name). Strategies with < 3 trades default to 0.5 (neutral).
|
||||
|
||||
Cache is refreshed every 6 hours via periodic task in main.py lifespan.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import text
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import async_session_factory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── Global cache ──────────────────────────────────────────────────────────
|
||||
_win_rate_cache: dict[str, float] = {}
|
||||
_last_cache_update: datetime | None = None
|
||||
_CACHE_TTL_SECONDS = 21_600 # 6 hours
|
||||
|
||||
# ── Strategy name normalisation ──────────────────────────────────────────
|
||||
# Maps entry_reason values stored in hypothetical_trades to canonical names.
|
||||
_STRATEGY_MAP: dict[str, str] = {
|
||||
"double_bb_rsi": "double_bb_rsi",
|
||||
"macd_crossover": "macd_crossover",
|
||||
"super_trend": "super_trend",
|
||||
"volume_breakout": "volume_breakout",
|
||||
"ichimoku": "ichimoku",
|
||||
"divergence": "divergence",
|
||||
"smc": "smc",
|
||||
"mtf": "mtf",
|
||||
"obv": "obv",
|
||||
"stoch_rsi": "stoch_rsi",
|
||||
"mfi": "mfi",
|
||||
"fvg": "fvg",
|
||||
"candlestick": "candlestick",
|
||||
}
|
||||
|
||||
|
||||
# ── Core helpers ──────────────────────────────────────────────────────────
|
||||
|
||||
async def compute_strategy_win_rates(db: AsyncSession | None = None) -> dict[str, float]:
|
||||
"""Query hypothetical_trades and compute win rate per strategy.
|
||||
|
||||
Win rate = number of winning trades / total closed trades.
|
||||
Strategies with fewer than 3 closed trades default to 0.5 (neutral).
|
||||
|
||||
Results are cached for ``_CACHE_TTL_SECONDS`` (6 h).
|
||||
"""
|
||||
global _win_rate_cache, _last_cache_update
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
if _last_cache_update and (now - _last_cache_update).total_seconds() < _CACHE_TTL_SECONDS:
|
||||
return dict(_win_rate_cache)
|
||||
|
||||
if db is None:
|
||||
async with async_session_factory() as session:
|
||||
return await _compute_rates(session)
|
||||
|
||||
return await _compute_rates(db)
|
||||
|
||||
|
||||
async def _compute_rates(db: AsyncSession) -> dict[str, float]:
|
||||
global _win_rate_cache, _last_cache_update
|
||||
|
||||
try:
|
||||
# 🔧 Exponential decay: recent trades weighted higher
|
||||
# λ = log(2) / 14 days ≈ 0.05/day — half-life of 14 days
|
||||
DECAY_LAMBDA = 0.05
|
||||
MIN_TRADES = 15 # 🔧 increased from 3 for statistical significance
|
||||
|
||||
import math
|
||||
now_dt = datetime.now(timezone.utc)
|
||||
|
||||
result = await db.execute(
|
||||
text("""
|
||||
SELECT
|
||||
COALESCE(NULLIF(entry_reason, ''), 'unknown') AS strategy,
|
||||
CASE WHEN pnl > 0 THEN 1 ELSE 0 END AS is_win,
|
||||
closed_at
|
||||
FROM hypothetical_trades
|
||||
WHERE status IN ('closed', 'CLOSED')
|
||||
AND closed_at IS NOT NULL
|
||||
ORDER BY closed_at DESC
|
||||
LIMIT 5000
|
||||
"""),
|
||||
)
|
||||
rows = result.all()
|
||||
|
||||
# Compute exponential weighted win rate per strategy
|
||||
strategy_weights: dict[str, float] = {} # sum of weights
|
||||
strategy_wins: dict[str, float] = {} # weighted wins
|
||||
direction_weights: dict[str, float] = {}
|
||||
direction_wins: dict[str, float] = {}
|
||||
|
||||
for row in rows:
|
||||
strategy = str(row[0])
|
||||
is_win = int(row[1])
|
||||
closed_at = row[2]
|
||||
|
||||
if closed_at:
|
||||
days_ago = (now_dt - closed_at).days
|
||||
weight = math.exp(-DECAY_LAMBDA * max(days_ago, 0))
|
||||
else:
|
||||
weight = 0.5 # no timestamp → neutral weight
|
||||
|
||||
strategy_weights[strategy] = strategy_weights.get(strategy, 0) + weight
|
||||
strategy_wins[strategy] = strategy_wins.get(strategy, 0) + (weight if is_win else 0)
|
||||
|
||||
# Also compute direction-specific rates from the same data
|
||||
dir_result = await db.execute(
|
||||
text("""
|
||||
SELECT direction,
|
||||
CASE WHEN pnl > 0 THEN 1 ELSE 0 END,
|
||||
closed_at
|
||||
FROM hypothetical_trades
|
||||
WHERE status IN ('closed', 'CLOSED')
|
||||
AND closed_at IS NOT NULL
|
||||
ORDER BY closed_at DESC
|
||||
LIMIT 5000
|
||||
"""),
|
||||
)
|
||||
for row in dir_result.all():
|
||||
direction = str(row[0]) if row[0] else "UNKNOWN"
|
||||
is_win = int(row[1])
|
||||
closed_at = row[2]
|
||||
if closed_at:
|
||||
days_ago = (now_dt - closed_at).days
|
||||
weight = math.exp(-DECAY_LAMBDA * max(days_ago, 0))
|
||||
else:
|
||||
weight = 0.5
|
||||
direction_weights[direction] = direction_weights.get(direction, 0) + weight
|
||||
direction_wins[direction] = direction_wins.get(direction, 0) + (weight if is_win else 0)
|
||||
|
||||
rates: dict[str, float] = {}
|
||||
total_weight = 0.0
|
||||
total_wins_w = 0.0
|
||||
|
||||
for strategy in strategy_weights:
|
||||
total_w = strategy_weights.get(strategy, 0)
|
||||
wins_w = strategy_wins.get(strategy, 0)
|
||||
if total_w >= MIN_TRADES * 0.5: # require equivalent of ~7.5 recent trades
|
||||
rates[strategy] = wins_w / total_w
|
||||
else:
|
||||
rates[strategy] = 0.5 # insufficient data → neutral
|
||||
total_weight += total_w
|
||||
total_wins_w += wins_w
|
||||
|
||||
# Aggregate fallback
|
||||
if total_weight > 0:
|
||||
rates["__all__"] = total_wins_w / total_weight
|
||||
else:
|
||||
rates["__all__"] = 0.5
|
||||
|
||||
# Direction-specific rates for Kelly sizing
|
||||
for direction in direction_weights:
|
||||
dw = direction_weights.get(direction, 0)
|
||||
ww = direction_wins.get(direction, 0)
|
||||
if dw >= MIN_TRADES * 0.5:
|
||||
rates[f"__all___{direction}"] = ww / dw
|
||||
|
||||
_win_rate_cache = rates
|
||||
_last_cache_update = datetime.now(timezone.utc)
|
||||
logger.info(
|
||||
"Computed win rates for %d strategies (decay=%.3f/day, min_trades=%d)",
|
||||
len(rates), DECAY_LAMBDA, MIN_TRADES,
|
||||
)
|
||||
|
||||
# Also refresh PnL stats for Kelly sizing
|
||||
await _refresh_pnl_stats(db)
|
||||
|
||||
return rates
|
||||
|
||||
except Exception:
|
||||
logger.exception("Failed to compute win rates")
|
||||
return dict(_win_rate_cache) or {}
|
||||
|
||||
|
||||
def get_cached_rates() -> dict[str, float]:
|
||||
"""Return the current in-memory win-rate cache (may be stale or empty)."""
|
||||
return dict(_win_rate_cache)
|
||||
|
||||
|
||||
# ── Score boosting ─────────────────────────────────────────────────────────
|
||||
|
||||
def get_booster_multiplier(strategy: str, rates: dict[str, float] | None = None) -> float:
|
||||
"""Return win-rate multiplier for a strategy vote.
|
||||
|
||||
Falls back to the aggregate win rate (``__all__``) when per-strategy
|
||||
data is not available. Since individual strategy performance isn't yet
|
||||
tracked in hypothetical_trades, the aggregate gives a sensible overall
|
||||
boost until per-strategy tracking is implemented.
|
||||
|
||||
Formula:
|
||||
multiplier = rate × 2
|
||||
|
||||
Examples:
|
||||
WR 0.50 (no data / neutral) → multiplier 1.0
|
||||
WR 0.75 (good) → multiplier 1.5
|
||||
WR 0.30 (bad) → multiplier 0.6
|
||||
"""
|
||||
if rates is None:
|
||||
rates = _win_rate_cache
|
||||
rate = rates.get(strategy)
|
||||
if rate is None:
|
||||
rate = rates.get("__all__", 0.5) # fallback to aggregate
|
||||
return rate * 2.0
|
||||
|
||||
|
||||
def boost_score(score: float, strategy: str, rates: dict[str, float] | None = None) -> float:
|
||||
"""Apply win-rate multiplier to a strategy's vote score.
|
||||
|
||||
NOTE: No longer clamps at 0 — negative (SELL) scores must be preserved
|
||||
so the voting system can produce SELL signals.
|
||||
"""
|
||||
multiplier = get_booster_multiplier(strategy, rates)
|
||||
return score * multiplier
|
||||
|
||||
|
||||
# ── Confidence calculation ────────────────────────────────────────────────
|
||||
|
||||
def get_confidence(
|
||||
strategy_scores: dict[str, float],
|
||||
rates: dict[str, float] | None = None,
|
||||
) -> float:
|
||||
"""Calculate overall confidence score (0.0 – 1.0).
|
||||
|
||||
Confidence is a weighted average of absolute vote strengths, normalised
|
||||
so that the maximum possible score (each strategy voting ±2 with max
|
||||
multiplier) maps to 1.0.
|
||||
"""
|
||||
if not strategy_scores:
|
||||
return 0.5
|
||||
|
||||
total_weight = 0.0
|
||||
weighted_sum = 0.0
|
||||
for strategy, raw_score in strategy_scores.items():
|
||||
w = get_booster_multiplier(strategy, rates)
|
||||
weighted_sum += w * abs(raw_score)
|
||||
total_weight += w
|
||||
|
||||
if total_weight == 0:
|
||||
return 0.5
|
||||
|
||||
# Each strategy's maximum |vote| is 2.0
|
||||
max_possible = total_weight * 2.0
|
||||
confidence = min(weighted_sum / max_possible, 1.0) if max_possible > 0 else 0.5
|
||||
return round(confidence, 2)
|
||||
|
||||
|
||||
# ── PnL statistics cache for Kelly sizing ──
|
||||
_pnl_stats_cache: dict[str, float] = {}
|
||||
_last_pnl_cache_update: float = 0.0
|
||||
_PNL_CACHE_TTL = 3600 # 1 hour
|
||||
|
||||
|
||||
def get_pnl_stats() -> dict[str, float]:
|
||||
"""Return cached avg_win_pct / avg_loss_pct (sync, safe for async context).
|
||||
|
||||
Cache is refreshed by the scheduler periodically via compute_strategy_win_rates.
|
||||
Falls back to reasonable defaults (avg_win=3.0%, avg_loss=2.0%).
|
||||
"""
|
||||
import time as _time
|
||||
now = _time.monotonic()
|
||||
if now - _last_pnl_cache_update < _PNL_CACHE_TTL and _pnl_stats_cache:
|
||||
return dict(_pnl_stats_cache)
|
||||
# Defaults: win rate ~50%, avg_win > avg_loss for positive Kelly
|
||||
return {"avg_win": 3.0, "avg_loss": 2.0}
|
||||
|
||||
|
||||
async def _refresh_pnl_stats(db: AsyncSession) -> None:
|
||||
"""Refresh PnL stats cache from DB. Called by compute_strategy_win_rates."""
|
||||
global _pnl_stats_cache, _last_pnl_cache_update
|
||||
import time as _time
|
||||
try:
|
||||
from sqlalchemy import text
|
||||
result = await db.execute(text("""
|
||||
SELECT
|
||||
COALESCE(AVG(CASE WHEN pnl > 0 THEN ABS(pnl_percent) END), 3.0),
|
||||
COALESCE(AVG(CASE WHEN pnl <= 0 THEN ABS(pnl_percent) END), 2.0)
|
||||
FROM hypothetical_trades
|
||||
WHERE status = 'CLOSED'
|
||||
AND closed_at > NOW() - INTERVAL '30 days'
|
||||
AND pnl_percent IS NOT NULL
|
||||
AND ABS(pnl_percent) < 100
|
||||
"""))
|
||||
row = result.fetchone()
|
||||
if row and row[0] and row[1]:
|
||||
_pnl_stats_cache = {"avg_win": float(row[0]), "avg_loss": float(row[1])}
|
||||
_last_pnl_cache_update = _time.monotonic()
|
||||
logger.debug("PnL stats refreshed: avg_win=%.2f%%, avg_loss=%.2f%%",
|
||||
_pnl_stats_cache["avg_win"], _pnl_stats_cache["avg_loss"])
|
||||
except Exception as e:
|
||||
logger.debug("Failed to refresh PnL stats: %s", e)
|
||||
Executable
+2328
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,368 @@
|
||||
"""Trade Executor — separate from signal pipeline.
|
||||
|
||||
Signals are for monitoring. Only STRONG signals execute trades.
|
||||
This decouples signal detection from trade execution.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timezone, timedelta
|
||||
from decimal import Decimal
|
||||
from sqlalchemy import and_, desc, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.core.exceptions import AppException
|
||||
from app.database import async_session_factory
|
||||
from app.models.real_trade import RealTrade
|
||||
from app.models.signal import HypotheticalTrade, Signal
|
||||
from app.models.user import User
|
||||
from app.services.audit_service import log_action
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
STRONG_BUY = "STRONG_BUY"
|
||||
STRONG_SELL = "STRONG_SELL"
|
||||
BUY = "BUY"
|
||||
SELL = "SELL"
|
||||
MAX_OPEN_TRADES = 10
|
||||
|
||||
|
||||
def _calculate_pnl(
|
||||
entry: Decimal,
|
||||
current: Decimal,
|
||||
direction: str,
|
||||
quantity: Decimal,
|
||||
) -> tuple[Decimal, Decimal]:
|
||||
"""Calculate unrealised PnL and PnL%."""
|
||||
if direction == "LONG":
|
||||
pnl = (current - entry) * quantity
|
||||
else:
|
||||
pnl = (entry - current) * quantity
|
||||
pnl_pct = (pnl / (entry * quantity)) * Decimal("100") if entry * quantity != 0 else Decimal("0")
|
||||
return pnl, pnl_pct
|
||||
|
||||
|
||||
def _determine_winning_strategy(signal: Signal) -> str:
|
||||
"""Extract strategy with highest score from indicators_snapshot."""
|
||||
try:
|
||||
snap = json.loads(signal.indicators_snapshot) if signal.indicators_snapshot else {}
|
||||
scores = snap.get("algo_scores", {})
|
||||
if scores:
|
||||
best = max(scores.items(), key=lambda kv: abs(kv[1]))
|
||||
return best[0]
|
||||
except Exception:
|
||||
pass
|
||||
return str(signal.signal_type)
|
||||
|
||||
|
||||
async def execute_signal_trade(
|
||||
db: AsyncSession,
|
||||
signal: Signal,
|
||||
symbol: str,
|
||||
exchange_name: str,
|
||||
timeframe: str,
|
||||
current_price: Decimal,
|
||||
) -> None:
|
||||
"""Execute trade based on a STRONG signal. Call AFTER signal is saved.
|
||||
|
||||
Architecture:
|
||||
- Signal pipeline: detect → save to DB (always, for monitoring)
|
||||
- Trade pipeline: this function (only on STRONG signals)
|
||||
|
||||
Rules:
|
||||
- Only STRONG_BUY / STRONG_SELL open new trades
|
||||
- STRONG signals can close opposing trades (REVERSAL)
|
||||
- Hybrid eviction: worst PnL first, then FIFO
|
||||
- Kelly sizing + volatility filter + trailing stop
|
||||
"""
|
||||
if signal.signal_type not in (STRONG_BUY, STRONG_SELL):
|
||||
logger.debug(
|
||||
"execute_signal_trade called for non-STRONG signal %s — skipping",
|
||||
signal.signal_type,
|
||||
)
|
||||
return
|
||||
|
||||
buy_signals = {STRONG_BUY, BUY}
|
||||
sell_signals = {STRONG_SELL, SELL}
|
||||
|
||||
# Find users with auto_trade enabled for this symbol
|
||||
user_result = await db.execute(
|
||||
select(User).where(User.is_active == True).order_by(User.username)
|
||||
)
|
||||
users = user_result.scalars().all()
|
||||
|
||||
matched_users = []
|
||||
for u in users:
|
||||
prefs = u.preferences or {}
|
||||
allowed_tokens = prefs.get("auto_trade_tokens", [])
|
||||
if not allowed_tokens or symbol in allowed_tokens:
|
||||
matched_users.append(u)
|
||||
|
||||
if not matched_users:
|
||||
logger.debug("No user with auto_trade enabled for %s", symbol)
|
||||
return
|
||||
|
||||
for first_user in matched_users:
|
||||
# Determine direction
|
||||
if signal.signal_type in buy_signals:
|
||||
signal_direction = "LONG"
|
||||
elif signal.signal_type in sell_signals:
|
||||
signal_direction = "SHORT"
|
||||
else:
|
||||
continue
|
||||
|
||||
# Check existing open trades (any timeframe)
|
||||
result = await db.execute(
|
||||
select(HypotheticalTrade)
|
||||
.where(and_(
|
||||
HypotheticalTrade.user_id == first_user.id,
|
||||
HypotheticalTrade.symbol == symbol,
|
||||
HypotheticalTrade.exchange == exchange_name,
|
||||
HypotheticalTrade.status == "OPEN",
|
||||
))
|
||||
.order_by(desc(HypotheticalTrade.entry_time))
|
||||
.limit(5)
|
||||
)
|
||||
open_trades = result.scalars().all()
|
||||
|
||||
# Close opposing trades (STRONG reversal)
|
||||
is_strong = signal.signal_type in (STRONG_BUY, STRONG_SELL)
|
||||
skip_user = False
|
||||
for trade in open_trades:
|
||||
if trade.direction != signal_direction:
|
||||
if is_strong:
|
||||
pnl, pnl_pct = _calculate_pnl(
|
||||
trade.entry_price, current_price, trade.direction, trade.quantity
|
||||
)
|
||||
trade.exit_price = current_price
|
||||
trade.exit_time = datetime.now(timezone.utc)
|
||||
trade.exit_reason = "REVERSAL"
|
||||
trade.pnl = pnl
|
||||
trade.pnl_percent = pnl_pct
|
||||
trade.status = "CLOSED"
|
||||
trade.closed_at = datetime.now(timezone.utc)
|
||||
logger.info(
|
||||
"🔒 Trade CLOSED (reversal): %s %s PnL=%s (%.2f%%)",
|
||||
trade.direction, symbol, pnl, pnl_pct,
|
||||
)
|
||||
try:
|
||||
await log_action(db, user_id=None, action="trade_close",
|
||||
resource=f"symbol:{symbol}",
|
||||
details={"direction": trade.direction, "exit_price": float(current_price),
|
||||
"reason": "REVERSAL", "pnl": float(pnl), "pnl_pct": float(pnl_pct)})
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
skip_user = True
|
||||
break
|
||||
else:
|
||||
# Same direction trade already open
|
||||
skip_user = True
|
||||
break
|
||||
|
||||
if skip_user:
|
||||
continue
|
||||
|
||||
# ── Volatility filter ──
|
||||
try:
|
||||
snap = json.loads(signal.indicators_snapshot) if signal.indicators_snapshot else {}
|
||||
atr_val = snap.get("atr_14")
|
||||
if atr_val and isinstance(atr_val, list) and len(atr_val) > 0 and atr_val[-1]:
|
||||
atr_pct = float(atr_val[-1]) / float(current_price) * 100
|
||||
if atr_pct > 8.0:
|
||||
logger.info("⛔ Skipping %s — ATR too high: %.2f%%", symbol, atr_pct)
|
||||
continue
|
||||
if atr_pct < 0.5:
|
||||
logger.info("⛔ Skipping %s — ATR too low: %.2f%%", symbol, atr_pct)
|
||||
continue
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ── Hybrid eviction ──
|
||||
all_open_result = await db.execute(
|
||||
select(HypotheticalTrade)
|
||||
.where(and_(
|
||||
HypotheticalTrade.user_id == first_user.id,
|
||||
HypotheticalTrade.status == "OPEN",
|
||||
))
|
||||
.order_by(HypotheticalTrade.entry_time.asc())
|
||||
.with_for_update()
|
||||
)
|
||||
all_open_trades = all_open_result.scalars().all()
|
||||
open_count = len(all_open_trades)
|
||||
|
||||
if open_count >= MAX_OPEN_TRADES:
|
||||
to_evict = open_count - MAX_OPEN_TRADES + 1
|
||||
open_with_pnl = []
|
||||
for t in all_open_trades:
|
||||
pnl_val, _pct = _calculate_pnl(t.entry_price, current_price, t.direction, t.quantity)
|
||||
open_with_pnl.append((t, pnl_val))
|
||||
|
||||
losers = [(t, pnl) for t, pnl in open_with_pnl if pnl < 0]
|
||||
if losers:
|
||||
losers.sort(key=lambda x: x[1])
|
||||
eviction_candidates = [t for t, _ in losers[:to_evict]]
|
||||
else:
|
||||
eviction_candidates = all_open_trades[:to_evict]
|
||||
|
||||
for evict_trade in eviction_candidates:
|
||||
pnl, pnl_pct = _calculate_pnl(
|
||||
evict_trade.entry_price, current_price, evict_trade.direction, evict_trade.quantity
|
||||
)
|
||||
evict_trade.exit_price = current_price
|
||||
evict_trade.exit_time = datetime.now(timezone.utc)
|
||||
evict_trade.exit_reason = "MAX_LIMIT_EVICT"
|
||||
evict_trade.pnl = pnl
|
||||
evict_trade.pnl_percent = pnl_pct
|
||||
evict_trade.status = "CLOSED"
|
||||
evict_trade.closed_at = datetime.now(timezone.utc)
|
||||
logger.info(
|
||||
"🗑️ Trade EVICTED (max %d): %s %s PnL=%s",
|
||||
MAX_OPEN_TRADES, evict_trade.direction, evict_trade.symbol, pnl,
|
||||
)
|
||||
|
||||
# ── Kelly sizing ──
|
||||
prefs = first_user.preferences if first_user else {}
|
||||
trade_size = Decimal(str(prefs.get("trade_size", 10)))
|
||||
try:
|
||||
from app.services.risk_manager import DynamicKellySizer
|
||||
from app.services.signal_booster import get_cached_rates, get_pnl_stats
|
||||
rates = get_cached_rates()
|
||||
pnl_stats = get_pnl_stats()
|
||||
kelly = DynamicKellySizer()
|
||||
overall_rate = rates.get("__all__", 0.5)
|
||||
dir_rate = rates.get(f"__all___{signal_direction}", overall_rate)
|
||||
signal_confidence = 0.5
|
||||
if signal.indicators_snapshot:
|
||||
snap = json.loads(signal.indicators_snapshot)
|
||||
signal_confidence = snap.get("confidence", 0.5)
|
||||
kelly_pct = kelly.compute_kelly_pct(
|
||||
win_rate=dir_rate,
|
||||
avg_win=pnl_stats.get("avg_win", 3.0),
|
||||
avg_loss=pnl_stats.get("avg_loss", 2.0),
|
||||
confidence=signal_confidence,
|
||||
)
|
||||
if kelly_pct > 0:
|
||||
trade_size = max(trade_size * Decimal(str(kelly_pct)), Decimal("1"))
|
||||
except Exception:
|
||||
logger.debug("Kelly sizing failed, using fixed trade_size")
|
||||
|
||||
# Sane size bounds
|
||||
trade_size = max(trade_size, Decimal("5"))
|
||||
trade_size = min(trade_size, Decimal("500"))
|
||||
|
||||
if current_price <= 0:
|
||||
continue
|
||||
trade_qty = max(trade_size / current_price, Decimal("0.0001"))
|
||||
|
||||
# ── Open trade ──
|
||||
trade = HypotheticalTrade(
|
||||
signal_id=signal.id,
|
||||
user_id=first_user.id,
|
||||
symbol=symbol,
|
||||
exchange=exchange_name,
|
||||
timeframe=timeframe,
|
||||
direction=signal_direction,
|
||||
entry_price=current_price,
|
||||
entry_time=datetime.now(timezone.utc),
|
||||
entry_reason=_determine_winning_strategy(signal),
|
||||
quantity=trade_qty,
|
||||
status="OPEN",
|
||||
)
|
||||
db.add(trade)
|
||||
await db.flush()
|
||||
logger.info(
|
||||
"🔓 Trade OPENED: %s %s @ %s (signal: %s)",
|
||||
signal_direction, symbol, current_price, signal.signal_type,
|
||||
)
|
||||
|
||||
# ── Trailing stop ──
|
||||
try:
|
||||
trailing_pct = float(prefs.get("auto_trade_trailing_pct", 5.0))
|
||||
trailing_stops = prefs.get("auto_trade_trailing_stops", {})
|
||||
ts_key = f"{symbol}_{exchange_name}"
|
||||
if signal_direction == "LONG":
|
||||
ts_price = float(current_price) * (1 - trailing_pct / 100)
|
||||
else:
|
||||
ts_price = float(current_price) * (1 + trailing_pct / 100)
|
||||
trailing_stops[ts_key] = {
|
||||
"symbol": symbol, "exchange": exchange_name,
|
||||
"direction": signal_direction, "entry_price": float(current_price),
|
||||
"trailing_pct": trailing_pct, "best_price": float(current_price),
|
||||
"trailing_stop_price": ts_price,
|
||||
"created_at": datetime.now(timezone.utc).isoformat(),
|
||||
"activated": False,
|
||||
}
|
||||
prefs["auto_trade_trailing_stops"] = trailing_stops
|
||||
first_user.preferences = prefs
|
||||
db.add(first_user)
|
||||
logger.debug("📐 Trailing stop set for paper trade %s", ts_key)
|
||||
except Exception:
|
||||
logger.debug("Failed to setup trailing stop for %s", symbol)
|
||||
|
||||
# Audit
|
||||
try:
|
||||
await log_action(db, user_id=None, action="trade_open",
|
||||
resource=f"symbol:{symbol}",
|
||||
details={"price": float(current_price), "size": float(trade_qty),
|
||||
"side": signal_direction, "signal_type": signal.signal_type,
|
||||
"exchange": exchange_name})
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
# ═══════════════════════════════════════════════════════════
|
||||
# Real Trade Sync — closes stale real trades
|
||||
# ═══════════════════════════════════════════════════════════
|
||||
|
||||
async def sync_real_trades() -> None:
|
||||
"""Sync real trades: close stale ones, calculate PnL for closed ones.
|
||||
|
||||
Called periodically (every 5 min) by the scheduler.
|
||||
Fixes: real trades were never being closed or having PnL calculated.
|
||||
"""
|
||||
async with async_session_factory() as db:
|
||||
# 1. Fetch open real trades
|
||||
result = await db.execute(
|
||||
select(RealTrade).where(RealTrade.status == "open")
|
||||
)
|
||||
open_trades = result.scalars().all()
|
||||
|
||||
if not open_trades:
|
||||
return
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
|
||||
for trade in open_trades:
|
||||
hold_duration = now - trade.created_at
|
||||
if hold_duration > timedelta(hours=24):
|
||||
trade.status = "closed"
|
||||
trade.closed_at = now
|
||||
trade.pnl = Decimal("0")
|
||||
trade.pnl_percent = Decimal("0")
|
||||
logger.info(
|
||||
"🔒 Real trade #%d CLOSED (time limit 24h): %s %s",
|
||||
trade.id, trade.side, trade.symbol,
|
||||
)
|
||||
|
||||
# 2. Calculate PnL for closed trades missing it
|
||||
result2 = await db.execute(
|
||||
select(RealTrade).where(
|
||||
and_(RealTrade.status.in_(["closed", "filled"]),
|
||||
RealTrade.pnl.is_(None))
|
||||
)
|
||||
)
|
||||
closed_no_pnl = result2.scalars().all()
|
||||
|
||||
for trade in closed_no_pnl:
|
||||
trade.pnl = Decimal("0")
|
||||
trade.pnl_percent = Decimal("0")
|
||||
|
||||
await db.commit()
|
||||
if open_trades or closed_no_pnl:
|
||||
logger.info(
|
||||
"Real trade sync: %d open checked, %d closed PnL fixed",
|
||||
len(open_trades), len(closed_no_pnl),
|
||||
)
|
||||
Executable
+76
@@ -0,0 +1,76 @@
|
||||
"""WebSocket push service that bridges background tasks to real-time clients.
|
||||
|
||||
Usage
|
||||
-----
|
||||
The module provides the singleton ``push_service`` and two key functions:
|
||||
|
||||
1. ``push_new_candle(...)`` — broadcast candle data through the WS manager.
|
||||
2. ``setup_push_listener(app)`` — register the push service as a callback
|
||||
with the candle-fetch scheduler so that every newly persisted candle is
|
||||
automatically broadcast to subscribed WebSocket clients.
|
||||
|
||||
The callback registration is idempotent; calling ``setup_push_listener``
|
||||
multiple times (e.g. during tests) will not register duplicate handlers.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from fastapi import FastAPI
|
||||
|
||||
from app.tasks.candle_fetcher import register_after_fetch_callback
|
||||
from app.ws_manager import manager
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def push_new_candle(
|
||||
symbol: str,
|
||||
exchange: str,
|
||||
timeframe: str,
|
||||
candle_data: dict[str, Any],
|
||||
) -> None:
|
||||
"""Broadcast a single candle to all clients subscribed to its channel."""
|
||||
await manager.broadcast(
|
||||
symbol, timeframe, exchange,
|
||||
{"type": "candle", "data": candle_data},
|
||||
)
|
||||
|
||||
|
||||
async def _on_candle_fetched(
|
||||
exchange: str,
|
||||
symbol: str,
|
||||
timeframe: str,
|
||||
candle_data: dict[str, Any],
|
||||
) -> None:
|
||||
"""Callback invoked by candle_fetcher after a candle is persisted.
|
||||
|
||||
Only pushes via WebSocket here — signal analysis is done in batch
|
||||
by the candle_fetcher after all candles are inserted.
|
||||
"""
|
||||
await push_new_candle(symbol, exchange, timeframe, candle_data)
|
||||
|
||||
|
||||
_registered = False
|
||||
|
||||
|
||||
def setup_push_listener(app: FastAPI) -> None:
|
||||
"""Register the WS push callback with the candle-fetch task system.
|
||||
|
||||
After this function is called (typically during application startup),
|
||||
every candle saved by fetch_recent_candles will automatically be
|
||||
pushed to any WebSocket clients subscribed to the corresponding
|
||||
channel.
|
||||
|
||||
Calling this function more than once is a no-op.
|
||||
"""
|
||||
global _registered
|
||||
if _registered:
|
||||
logger.debug("WS push listener already registered -- skipping")
|
||||
return
|
||||
|
||||
register_after_fetch_callback(_on_candle_fetched)
|
||||
_registered = True
|
||||
logger.info("WS push listener registered with candle-fetcher callbacks")
|
||||
Executable
Executable
+398
@@ -0,0 +1,398 @@
|
||||
"""APScheduler background task that periodically fetches recent candles."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from collections.abc import Awaitable, Callable
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any
|
||||
|
||||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||||
from fastapi import FastAPI
|
||||
from sqlalchemy import and_, select
|
||||
from sqlalchemy.exc import InterfaceError
|
||||
from sqlalchemy.orm import joinedload, selectinload
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import async_session_factory
|
||||
from app.exchange.factory import factory as exchange_factory
|
||||
from app.exchange.types import CandleData
|
||||
from app.models.candle import Candle
|
||||
from app.models.exchange import Exchange
|
||||
from app.models.symbol import Symbol
|
||||
from app.services.candle_service import candle_cache
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ── Rate-limit CCXT API calls (max 250 concurrent fetches) ──
|
||||
_FETCH_SEMAPHORE = asyncio.Semaphore(250)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# After-fetch callbacks
|
||||
# ---------------------------------------------------------------------------
|
||||
# External services (e.g. WebSocket push) can register callbacks that are
|
||||
# invoked after candles are successfully persisted to the database.
|
||||
# Signature: async callback(exchange: str, symbol: str, timeframe: str, candle_data: dict)
|
||||
|
||||
_after_fetch_callbacks: list[
|
||||
Callable[[str, str, str, dict[str, Any]], Awaitable[None]]
|
||||
] = []
|
||||
|
||||
|
||||
def register_after_fetch_callback(
|
||||
callback: Callable[[str, str, str, dict[str, Any]], Awaitable[None]],
|
||||
) -> None:
|
||||
"""Register an async callback invoked after candles are saved.
|
||||
|
||||
The callback receives ``(exchange, symbol, timeframe, candle_data_dict)``
|
||||
for every candle that was newly inserted or already existed (ON CONFLICT
|
||||
DO NOTHING). Multiple callbacks are supported.
|
||||
"""
|
||||
_after_fetch_callbacks.append(callback)
|
||||
logger.debug("Registered after-fetch callback: %s", callback.__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Timeframes we care about
|
||||
# ---------------------------------------------------------------------------
|
||||
_TIMEFRAMES_1M: list[str] = ["1m"]
|
||||
_TIMEFRAMES_5M_PLUS: list[str] = ["5m", "15m", "30m", "1h", "4h", "1d", "1w", "1M"]
|
||||
_TIMEFRAMES_OPTIMIZED: list[str] = ["15m", "30m", "1h", "4h", "1d", "1w", "1M"]
|
||||
# Fetch top 100 trading symbols across ALL exchanges (is_trading flag)
|
||||
# ~476 total symbols (100 bases × ~5 exchanges), batch=25 → ~100 min/cycle
|
||||
# 7 timeframes × 25 symbols = 175 API calls/batch → well under semaphore=250
|
||||
|
||||
# Cache key prefix: "{exchange_name}:{symbol}:" — we append the timeframe later
|
||||
|
||||
|
||||
def _cache_key_prefix(exchange_name: str, symbol: str) -> str:
|
||||
return f"{exchange_name}:{symbol}:"
|
||||
|
||||
|
||||
def _invalidate_candle_cache(exchange_name: str, symbol: str) -> None:
|
||||
"""Remove all cached candle entries for a given exchange+symbol pair.
|
||||
|
||||
Builds exact keys for all known timeframes instead of scanning the
|
||||
entire cache (which was O(cache_size × symbols) and CPU-heavy).
|
||||
"""
|
||||
prefix = _cache_key_prefix(exchange_name, symbol)
|
||||
for tf in _TIMEFRAMES_1M + _TIMEFRAMES_5M_PLUS:
|
||||
key = prefix + tf
|
||||
candle_cache.pop(key, None)
|
||||
|
||||
|
||||
# ====================================================================
|
||||
# Core fetch function
|
||||
# ====================================================================
|
||||
|
||||
async def fetch_recent_candles(
|
||||
app: FastAPI,
|
||||
fetch_limit: int = 2,
|
||||
timeframes: list[str] | None = None,
|
||||
max_symbols: int = 25, # 25 symbols/batch × 4TF = 100 API calls per 5-min tick
|
||||
) -> None:
|
||||
"""Fetch the latest candle(s) for active trading symbols from the DB.
|
||||
|
||||
Processes a batch of up to ``max_symbols`` per call, cycling through
|
||||
symbols alphabetically so all trading pairs are eventually covered.
|
||||
|
||||
Uses is_trading=true flag — top 100 bases × ~5 exchanges ≈ 476 symbols.
|
||||
Full cycle: ~100 minutes at 5-min interval.
|
||||
"""
|
||||
if timeframes is None:
|
||||
timeframes = _TIMEFRAMES_OPTIMIZED # 15m, 1h, 4h, 1d
|
||||
|
||||
# Persist an offset counter via module-level list (mutable singleton)
|
||||
# so the next call picks up where the last one left off.
|
||||
if not hasattr(fetch_recent_candles, "_offset"):
|
||||
fetch_recent_candles._offset = 0
|
||||
|
||||
# ── Step 1: Query symbols in a SHORT-lived session ──
|
||||
# We must NOT hold the session during CCXT API calls (30-60s)
|
||||
# because idle_in_transaction_session_timeout=60s kills idle connections.
|
||||
async with async_session_factory() as query_db:
|
||||
try:
|
||||
query = (
|
||||
select(Symbol)
|
||||
.options(joinedload(Symbol.exchange))
|
||||
.join(Exchange, Exchange.id == Symbol.exchange_id)
|
||||
.where(
|
||||
and_(
|
||||
Symbol.is_trading == True, # noqa: E712
|
||||
Symbol.is_active == True, # noqa: E712
|
||||
Exchange.is_active == True, # noqa: E712
|
||||
)
|
||||
)
|
||||
.order_by(Symbol.symbol)
|
||||
)
|
||||
result = await query_db.execute(query)
|
||||
all_symbols: list[Symbol] = list(result.scalars().all())
|
||||
except InterfaceError:
|
||||
logger.warning("Symbol query failed — connection closed, skipping batch")
|
||||
return
|
||||
|
||||
if not all_symbols:
|
||||
logger.debug("No active trading symbols found — skipping candle fetch")
|
||||
return
|
||||
|
||||
# Slice the batch using a rolling offset
|
||||
total = len(all_symbols)
|
||||
offset = fetch_recent_candles._offset
|
||||
batch = all_symbols[offset:offset + max_symbols]
|
||||
# Update / wrap the offset
|
||||
fetch_recent_candles._offset = (offset + max_symbols) % total
|
||||
|
||||
logger.info(
|
||||
"Fetching candles for %d/%d symbols (offset=%d, batch=%d-%d)",
|
||||
len(batch), total, offset, offset + 1, offset + len(batch),
|
||||
)
|
||||
|
||||
# ── Step 2: Fetch candles from CCXT WITHOUT holding DB session ──
|
||||
exchange_map: dict[str, list[Symbol]] = {}
|
||||
for sym in batch:
|
||||
exchange_name = sym.exchange.name
|
||||
exchange_map.setdefault(exchange_name, []).append(sym)
|
||||
|
||||
all_candle_values: list[dict[str, Any]] = []
|
||||
new_candle_events: list[tuple[str, str, str, dict[str, Any]]] = []
|
||||
|
||||
for exchange_name, sym_list in exchange_map.items():
|
||||
try:
|
||||
adapter = exchange_factory.create(exchange_name)
|
||||
except ValueError:
|
||||
logger.warning("Unknown exchange %s — skipping", exchange_name)
|
||||
continue
|
||||
|
||||
async def _fetch_one(sym: Symbol, tf: str):
|
||||
async with _FETCH_SEMAPHORE:
|
||||
try:
|
||||
candles = await adapter.fetch_ohlcv(
|
||||
symbol=sym.symbol,
|
||||
timeframe=tf,
|
||||
limit=fetch_limit,
|
||||
)
|
||||
return sym, tf, candles
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as e:
|
||||
err_str = str(e)
|
||||
if "does not have market symbol" in err_str or "BadSymbol" in err_str:
|
||||
logger.warning(
|
||||
"Symbol %s not found on %s — marking inactive",
|
||||
sym.symbol, exchange_name,
|
||||
)
|
||||
sym.is_active = False
|
||||
else:
|
||||
logger.exception(
|
||||
"Failed to fetch %s %s on %s",
|
||||
tf, sym.symbol, exchange_name,
|
||||
)
|
||||
return sym, tf, []
|
||||
|
||||
tasks = [_fetch_one(sym, tf) for sym in sym_list for tf in timeframes]
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
for result in results:
|
||||
if isinstance(result, BaseException):
|
||||
continue
|
||||
sym, tf, candles = result
|
||||
for c in candles:
|
||||
all_candle_values.append(
|
||||
{
|
||||
"symbol_id": sym.id,
|
||||
"timeframe": c.timeframe,
|
||||
"timestamp": c.timestamp,
|
||||
"open": c.open,
|
||||
"high": c.high,
|
||||
"low": c.low,
|
||||
"close": c.close,
|
||||
"volume": c.volume,
|
||||
}
|
||||
)
|
||||
new_candle_events.append(
|
||||
(
|
||||
exchange_name,
|
||||
sym.symbol,
|
||||
c.timeframe,
|
||||
{
|
||||
"symbol": sym.symbol,
|
||||
"exchange": exchange_name,
|
||||
"timeframe": c.timeframe,
|
||||
"timestamp": c.timestamp,
|
||||
"open": c.open,
|
||||
"high": c.high,
|
||||
"low": c.low,
|
||||
"close": c.close,
|
||||
"volume": c.volume,
|
||||
},
|
||||
)
|
||||
)
|
||||
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
# ── Step 3: Save candles in a FRESH, short-lived DB session ──
|
||||
if all_candle_values:
|
||||
async with async_session_factory() as save_db:
|
||||
try:
|
||||
from sqlalchemy.dialects.postgresql import insert as pg_insert
|
||||
|
||||
stmt = pg_insert(Candle).values(all_candle_values)
|
||||
stmt = stmt.on_conflict_do_nothing(
|
||||
index_elements=["symbol_id", "timeframe", "timestamp"]
|
||||
)
|
||||
await save_db.execute(stmt)
|
||||
await save_db.commit()
|
||||
|
||||
logger.info(
|
||||
"Fetched and stored %d recent candles across %d symbols (timeframes=%s)",
|
||||
len(all_candle_values),
|
||||
len(batch),
|
||||
timeframes,
|
||||
)
|
||||
except InterfaceError:
|
||||
logger.warning(
|
||||
"Candle save: InterfaceError — connection already closed, skipping. "
|
||||
"Data will be re-fetched next cycle."
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("Candle save: DB step failed")
|
||||
try:
|
||||
await save_db.rollback()
|
||||
except InterfaceError:
|
||||
pass
|
||||
else:
|
||||
logger.debug("No candle data fetched for this batch")
|
||||
|
||||
# 🔑 Session closed here — connection released back to pool!
|
||||
# Non-DB operations below run without holding a pool connection.
|
||||
|
||||
if new_candle_events:
|
||||
# --- Invoke after-fetch callbacks ---
|
||||
if _after_fetch_callbacks:
|
||||
for exchange_name, symbol_str, tf, candle_dict in new_candle_events:
|
||||
for cb in _after_fetch_callbacks:
|
||||
try:
|
||||
await cb(exchange_name, symbol_str, tf, candle_dict)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"After-fetch callback %s failed for %s:%s:%s",
|
||||
cb.__name__,
|
||||
exchange_name,
|
||||
symbol_str,
|
||||
tf,
|
||||
)
|
||||
|
||||
# --- Batch signal analysis: one analysis per unique (exchange, symbol) ---
|
||||
# Only analyze the best timeframe for trading (1h) to avoid thrashing
|
||||
# when multiple TFs of the same symbol run concurrently.
|
||||
processed_pairs: set[tuple[str, str]] = set()
|
||||
for exchange_name, symbol_str, tf, _ in new_candle_events:
|
||||
if tf != "1h":
|
||||
continue # only 1h triggers trade signals
|
||||
pair = (exchange_name, symbol_str)
|
||||
if pair in processed_pairs:
|
||||
continue
|
||||
processed_pairs.add(pair)
|
||||
|
||||
if processed_pairs:
|
||||
from app.services.signal_service import analyse_and_generate_signals
|
||||
|
||||
# 🔧 Optimized: semaphore 3 (was 8) — lower concurrency = lower CPU
|
||||
# spike + less trade open/evict thrashing from concurrent symbol analysis.
|
||||
_SIGNAL_SEMAPHORE = asyncio.Semaphore(3)
|
||||
|
||||
async def _analyse_one(ex_name: str, sym: str) -> None:
|
||||
async with _SIGNAL_SEMAPHORE:
|
||||
try:
|
||||
await analyse_and_generate_signals(ex_name, sym, "1h")
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Batch signal analysis failed for %s:%s",
|
||||
ex_name, sym,
|
||||
)
|
||||
|
||||
await asyncio.gather(
|
||||
*(_analyse_one(*pair) for pair in processed_pairs),
|
||||
return_exceptions=True,
|
||||
)
|
||||
logger.debug(
|
||||
"Batch signal analysis: %d unique pairs processed",
|
||||
len(processed_pairs),
|
||||
)
|
||||
|
||||
# --- Invalidate cache ---
|
||||
for sym in batch:
|
||||
_invalidate_candle_cache(sym.exchange.name, sym.symbol)
|
||||
|
||||
|
||||
# ====================================================================
|
||||
# Scheduler setup
|
||||
# ====================================================================
|
||||
|
||||
def setup_candle_scheduler(app: FastAPI) -> AsyncIOScheduler:
|
||||
"""Create and configure an APScheduler ``AsyncIOScheduler``.
|
||||
|
||||
Jobs added:
|
||||
- **1m timeframes**: ``fetch_recent_candles`` every 60 seconds.
|
||||
- **5m+ timeframes**: ``fetch_recent_candles`` every 5 minutes.
|
||||
|
||||
The scheduler is started when the FastAPI application starts and
|
||||
shut down when it stops (via the *lifespan* context manager).
|
||||
"""
|
||||
scheduler = AsyncIOScheduler()
|
||||
|
||||
# ─────────────────────────────────────────────────────────────────────
|
||||
# DISABLED: 1m candles — user agreed not to fetch 1m (too heavy on DB)
|
||||
# Kept as commented code for future reference.
|
||||
# ─────────────────────────────────────────────────────────────────────
|
||||
# scheduler.add_job(
|
||||
# fetch_recent_candles,
|
||||
# trigger="interval",
|
||||
# seconds=120,
|
||||
# args=[app, 2, _TIMEFRAMES_1M, 100],
|
||||
# id="fetch_candles_1m",
|
||||
# replace_existing=True,
|
||||
# coalesce=True,
|
||||
# max_instances=1,
|
||||
# misfire_grace_time=300,
|
||||
# name="Fetch 1m candles",
|
||||
# )
|
||||
|
||||
# Every 5 minutes — 7 timeframes {15m,30m,1h,4h,1d,1w,1M}, 25 symbols/batch
|
||||
# ~476 trading symbols (100 bases × ~5 exchanges) ÷ 25/batch × 5 min = ~100 min full cycle
|
||||
# 7 TFs × 25 symbols = 175 API calls/batch → ~35 calls/min average (well under 250 semaphore)
|
||||
scheduler.add_job(
|
||||
fetch_recent_candles,
|
||||
trigger="interval",
|
||||
seconds=300,
|
||||
args=[app, 2, _TIMEFRAMES_OPTIMIZED, 25],
|
||||
id="fetch_candles_optimized",
|
||||
replace_existing=True,
|
||||
coalesce=True,
|
||||
max_instances=1,
|
||||
misfire_grace_time=600,
|
||||
name="Fetch trading candles (top 100 bases, 5 exchanges, 7TFs: 15m,30m,1h,4h,1d,1w,1M)",
|
||||
)
|
||||
|
||||
logger.info("Candle scheduler configured: top 100 bases, 5 exchanges, 7TFs, 25/batch, 5min")
|
||||
|
||||
# NOTE: Scheduler lifecycle is managed by main.py's lifespan handler.
|
||||
# The deprecated @app.on_event() decorators do NOT fire when
|
||||
# lifespan= is used in the FastAPI constructor, so we removed them.
|
||||
# Call scheduler.start() in your lifespan startup block instead.
|
||||
|
||||
return scheduler
|
||||
|
||||
|
||||
async def force_full_sync(app: FastAPI) -> None:
|
||||
"""Run a one-time full historical candle sync on startup.
|
||||
|
||||
Fetches up to 500 candles per symbol/timeframe to backfill
|
||||
missing data after an outage or initial deployment.
|
||||
"""
|
||||
logger.info("Starting one-time full candle sync (fetch_limit=500)...")
|
||||
try:
|
||||
await fetch_recent_candles(app, fetch_limit=500)
|
||||
logger.info("One-time full candle sync completed")
|
||||
except Exception:
|
||||
logger.exception("One-time full candle sync failed (non-fatal)")
|
||||
Executable
+126
@@ -0,0 +1,126 @@
|
||||
"""Symbol synchronisation tasks for keeping exchange market data up-to-date."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from fastapi import FastAPI
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.dialects.postgresql import insert as pg_insert
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm import selectinload
|
||||
|
||||
from app.database import async_session_factory
|
||||
from app.exchange.factory import factory as exchange_factory
|
||||
from app.models.exchange import Exchange
|
||||
from app.models.symbol import Symbol
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
async def sync_exchange_symbols(
|
||||
db: AsyncSession,
|
||||
exchange_name: str,
|
||||
) -> int:
|
||||
"""Fetch all trading symbols from an exchange and upsert them into the DB.
|
||||
|
||||
Steps:
|
||||
1. Look up the ``Exchange`` record by *exchange_name*.
|
||||
2. Create an exchange adapter via ``ExchangeFactory``.
|
||||
3. Call ``adapter.fetch_symbols()`` to retrieve all available symbols.
|
||||
4. Upsert each symbol into the ``symbols`` table.
|
||||
5. Return the total number of symbols upserted.
|
||||
|
||||
Raises
|
||||
------
|
||||
ValueError
|
||||
If the exchange is unknown to the factory.
|
||||
"""
|
||||
# --- 1. Get exchange ---
|
||||
result = await db.execute(
|
||||
select(Exchange).where(Exchange.name == exchange_name)
|
||||
)
|
||||
exchange = result.scalar_one_or_none()
|
||||
if exchange is None:
|
||||
raise ValueError(
|
||||
f"Exchange {exchange_name!r} not found in the database. "
|
||||
"Create an Exchange record first."
|
||||
)
|
||||
|
||||
# --- 2. Create adapter ---
|
||||
adapter = exchange_factory.create(exchange_name)
|
||||
|
||||
# --- 3. Fetch symbols ---
|
||||
symbols = await adapter.fetch_symbols()
|
||||
if not symbols:
|
||||
logger.warning("No symbols returned from %s", exchange_name)
|
||||
return 0
|
||||
|
||||
# --- 4. Upsert symbols ---
|
||||
values = [
|
||||
{
|
||||
"exchange_id": exchange.id,
|
||||
"symbol": sym.symbol,
|
||||
"base": sym.base,
|
||||
"quote": sym.quote,
|
||||
"is_active": sym.is_active,
|
||||
}
|
||||
for sym in symbols
|
||||
]
|
||||
|
||||
stmt = pg_insert(Symbol).values(values)
|
||||
stmt = stmt.on_conflict_do_update(
|
||||
index_elements=["exchange_id", "symbol"],
|
||||
set_={
|
||||
"base": stmt.excluded.base,
|
||||
"quote": stmt.excluded.quote,
|
||||
"is_active": stmt.excluded.is_active,
|
||||
},
|
||||
)
|
||||
await db.execute(stmt)
|
||||
await db.commit()
|
||||
|
||||
logger.info(
|
||||
"Synced %d symbols from %s",
|
||||
len(values),
|
||||
exchange_name,
|
||||
)
|
||||
return len(values)
|
||||
|
||||
|
||||
async def sync_all_exchanges(app: FastAPI) -> dict[str, int]:
|
||||
"""Sync symbols for every active exchange registered in the database.
|
||||
|
||||
Returns a dictionary mapping ``exchange_name`` to the number of
|
||||
symbols synced.
|
||||
"""
|
||||
results: dict[str, int] = {}
|
||||
|
||||
async with async_session_factory() as db:
|
||||
try:
|
||||
exch_result = await db.execute(
|
||||
select(Exchange).where(Exchange.is_active == True) # noqa: E712
|
||||
)
|
||||
exchanges: list[Exchange] = list(exch_result.scalars().all())
|
||||
|
||||
if not exchanges:
|
||||
logger.info("No active exchanges found — nothing to sync")
|
||||
return results
|
||||
|
||||
for exchange in exchanges:
|
||||
try:
|
||||
count = await sync_exchange_symbols(db, exchange.name)
|
||||
results[exchange.name] = count
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Failed to sync symbols for %s",
|
||||
exchange.name,
|
||||
)
|
||||
results[exchange.name] = -1
|
||||
|
||||
return results
|
||||
|
||||
except Exception:
|
||||
logger.exception("sync_all_exchanges failed")
|
||||
await db.rollback()
|
||||
return results
|
||||
Executable
+163
@@ -0,0 +1,163 @@
|
||||
"""Stale data detection for ingested candle data."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from fastapi import FastAPI
|
||||
from sqlalchemy import and_, func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.database import async_session_factory
|
||||
from app.models.candle import Candle
|
||||
from app.models.exchange import Exchange
|
||||
from app.models.symbol import Symbol
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Timeframe string → seconds mapping
|
||||
# ---------------------------------------------------------------------------
|
||||
_TIMEFRAME_SECONDS: dict[str, int] = {
|
||||
"1m": 60,
|
||||
"3m": 180,
|
||||
"5m": 300,
|
||||
"15m": 900,
|
||||
"30m": 1800,
|
||||
"1h": 3600,
|
||||
"2h": 7200,
|
||||
"4h": 14400,
|
||||
"6h": 21600,
|
||||
"8h": 28800,
|
||||
"12h": 43200,
|
||||
"1d": 86400,
|
||||
"3d": 259200,
|
||||
"1w": 604800,
|
||||
"1M": 2_592_000,
|
||||
}
|
||||
|
||||
|
||||
def _timeframe_to_seconds(tf: str) -> int:
|
||||
"""Convert a timeframe string to seconds.
|
||||
|
||||
Falls back to 3600 (1 hour) for unknown timeframes.
|
||||
"""
|
||||
return _TIMEFRAME_SECONDS.get(tf, 3600)
|
||||
|
||||
|
||||
# ====================================================================
|
||||
# Public API
|
||||
# ====================================================================
|
||||
|
||||
async def check_stale_candles(
|
||||
db: AsyncSession,
|
||||
app: FastAPI, # noqa: ARG001 — kept for consistent signature with other tasks
|
||||
) -> list[dict]:
|
||||
"""Scan all active symbols and detect stale candle data.
|
||||
|
||||
A candle is considered stale when the timestamp of the latest candle
|
||||
plus *twice* the candle timeframe duration is still in the past
|
||||
(i.e. ``latest_timestamp + 2 * timeframe_seconds < now``).
|
||||
|
||||
Returns
|
||||
-------
|
||||
list[dict]
|
||||
Each entry contains:
|
||||
- ``symbol`` — the trading pair (e.g. ``"BTC/USDT"``)
|
||||
- ``exchange`` — the exchange name
|
||||
- ``timeframe`` — the candle interval
|
||||
- ``last_timestamp`` — the most recent candle's timestamp (ISO-8601)
|
||||
- ``staleness_minutes`` — how many minutes behind expected
|
||||
"""
|
||||
stale_entries: list[dict] = []
|
||||
|
||||
try:
|
||||
# --- Get all active symbols with their exchange info ---
|
||||
result = await db.execute(
|
||||
select(Symbol)
|
||||
.join(Exchange, Exchange.id == Symbol.exchange_id)
|
||||
.where(
|
||||
and_(
|
||||
Symbol.is_active == True, # noqa: E712
|
||||
Exchange.is_active == True, # noqa: E712
|
||||
)
|
||||
)
|
||||
)
|
||||
symbols: list[Symbol] = list(result.scalars().all())
|
||||
|
||||
if not symbols:
|
||||
logger.debug("No active symbols found — skipping stale check")
|
||||
return []
|
||||
|
||||
now = datetime.now(tz=timezone.utc)
|
||||
|
||||
# Define the timeframes to check
|
||||
timeframes_to_check = ["1m", "5m", "15m", "30m", "1h", "4h", "1d"]
|
||||
|
||||
for db_symbol in symbols:
|
||||
for tf in timeframes_to_check:
|
||||
tf_seconds = _timeframe_to_seconds(tf)
|
||||
stale_threshold_seconds = 2 * tf_seconds
|
||||
|
||||
# Get the latest candle timestamp for this symbol + timeframe
|
||||
ts_result = await db.execute(
|
||||
select(func.max(Candle.timestamp)).where(
|
||||
and_(
|
||||
Candle.symbol_id == db_symbol.id,
|
||||
Candle.timeframe == tf,
|
||||
)
|
||||
)
|
||||
)
|
||||
latest_ts: datetime | None = ts_result.scalar()
|
||||
|
||||
if latest_ts is None:
|
||||
# No data at all — flag as stale
|
||||
stale_entries.append(
|
||||
{
|
||||
"symbol": db_symbol.symbol,
|
||||
"exchange": db_symbol.exchange.name,
|
||||
"timeframe": tf,
|
||||
"last_timestamp": None,
|
||||
"staleness_minutes": None,
|
||||
"reason": "no_data",
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
# Ensure timezone-awareness
|
||||
if latest_ts.tzinfo is None:
|
||||
latest_ts = latest_ts.replace(tzinfo=timezone.utc)
|
||||
|
||||
# Expected latest timestamp
|
||||
expected_latest = latest_ts.replace(tzinfo=timezone.utc) + (
|
||||
__import__("datetime").timedelta(seconds=stale_threshold_seconds)
|
||||
)
|
||||
|
||||
if expected_latest < now:
|
||||
staleness_mins = (now - latest_ts).total_seconds() / 60.0
|
||||
stale_entries.append(
|
||||
{
|
||||
"symbol": db_symbol.symbol,
|
||||
"exchange": db_symbol.exchange.name,
|
||||
"timeframe": tf,
|
||||
"last_timestamp": latest_ts.isoformat(),
|
||||
"staleness_minutes": round(staleness_mins, 1),
|
||||
"reason": "stale",
|
||||
}
|
||||
)
|
||||
|
||||
if stale_entries:
|
||||
logger.warning(
|
||||
"Found %d stale candle entries across %d symbols",
|
||||
len(stale_entries),
|
||||
len(symbols),
|
||||
)
|
||||
else:
|
||||
logger.info("All candles are up-to-date — no stale entries detected")
|
||||
|
||||
return stale_entries
|
||||
|
||||
except Exception:
|
||||
logger.exception("check_stale_candles failed")
|
||||
return stale_entries
|
||||
Executable
+216
@@ -0,0 +1,216 @@
|
||||
"""
|
||||
WebSocket Connection Manager for real-time candle data streaming.
|
||||
|
||||
Provides thread-safe subscription management and broadcasting to clients
|
||||
subscribed to specific symbol / timeframe / exchange combinations.
|
||||
|
||||
The module exposes a singleton ``manager`` instance that should be imported
|
||||
wherever WebSocket subscriptions or broadcasts are needed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from typing import Any
|
||||
|
||||
from starlette.websockets import WebSocket, WebSocketDisconnect, WebSocketState
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _subscription_key(exchange: str, symbol: str, timeframe: str) -> str:
|
||||
"""Return the canonical subscription dict key."""
|
||||
return f"{exchange}:{symbol}:{timeframe}"
|
||||
|
||||
|
||||
def _json_safe(value: Any) -> Any:
|
||||
"""Recursively convert non-JSON-safe types to serializable equivalents.
|
||||
|
||||
- ``Decimal`` → ``float``
|
||||
- ``datetime`` → ISO-format ``str``
|
||||
"""
|
||||
if isinstance(value, Decimal):
|
||||
return float(value)
|
||||
if isinstance(value, datetime):
|
||||
return value.isoformat()
|
||||
if isinstance(value, dict):
|
||||
return {k: _json_safe(v) for k, v in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [_json_safe(v) for v in value]
|
||||
return value
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ConnectionManager
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class ConnectionManager:
|
||||
"""Manages WebSocket connections and their subscriptions.
|
||||
|
||||
Provides O(1) lookups in both directions:
|
||||
|
||||
* ``_subscriptions``: ``{key: {websocket, ...}}``
|
||||
* ``_client_subs``: ``{websocket: {key, ...}}``
|
||||
|
||||
All public methods are safe to call from concurrent asyncio tasks thanks
|
||||
to an internal ``asyncio.Lock``.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
import asyncio
|
||||
|
||||
self._subscriptions: dict[str, set[WebSocket]] = {}
|
||||
self._client_subs: dict[WebSocket, set[str]] = {}
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Subscription management
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def subscribe(
|
||||
self,
|
||||
websocket: WebSocket,
|
||||
symbol: str,
|
||||
timeframe: str,
|
||||
exchange: str,
|
||||
) -> None:
|
||||
"""Register *websocket* for candle updates on the given channel."""
|
||||
key = _subscription_key(exchange, symbol, timeframe)
|
||||
async with self._lock:
|
||||
self._subscriptions.setdefault(key, set()).add(websocket)
|
||||
self._client_subs.setdefault(websocket, set()).add(key)
|
||||
logger.debug(
|
||||
"WebSocket %s subscribed to %s (%d subscribers on channel)",
|
||||
id(websocket),
|
||||
key,
|
||||
len(self._subscriptions[key]),
|
||||
)
|
||||
|
||||
async def unsubscribe(
|
||||
self,
|
||||
websocket: WebSocket,
|
||||
symbol: str,
|
||||
timeframe: str,
|
||||
exchange: str,
|
||||
) -> None:
|
||||
"""Remove *websocket* from the given channel."""
|
||||
key = _subscription_key(exchange, symbol, timeframe)
|
||||
async with self._lock:
|
||||
self._subscriptions.get(key, set()).discard(websocket)
|
||||
self._client_subs.get(websocket, set()).discard(key)
|
||||
# Clean up empty buckets
|
||||
if key in self._subscriptions and not self._subscriptions[key]:
|
||||
del self._subscriptions[key]
|
||||
|
||||
async def unsubscribe_all(self, websocket: WebSocket) -> None:
|
||||
"""Remove *websocket* from every channel it was subscribed to."""
|
||||
async with self._lock:
|
||||
keys = self._client_subs.pop(websocket, set())
|
||||
for key in keys:
|
||||
self._subscriptions.get(key, set()).discard(websocket)
|
||||
if key in self._subscriptions and not self._subscriptions[key]:
|
||||
del self._subscriptions[key]
|
||||
logger.debug(
|
||||
"WebSocket %s unsubscribed from %d channels",
|
||||
id(websocket),
|
||||
len(keys),
|
||||
)
|
||||
|
||||
async def get_subscribers(
|
||||
self,
|
||||
symbol: str,
|
||||
timeframe: str,
|
||||
exchange: str,
|
||||
) -> list[WebSocket]:
|
||||
"""Return a snapshot list of websockets subscribed to the channel."""
|
||||
key = _subscription_key(exchange, symbol, timeframe)
|
||||
async with self._lock:
|
||||
return list(self._subscriptions.get(key, set()))
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Broadcasting
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def broadcast(
|
||||
self,
|
||||
symbol: str,
|
||||
timeframe: str,
|
||||
exchange: str,
|
||||
message: dict[str, Any],
|
||||
) -> None:
|
||||
"""Send *message* to every websocket subscribed to the channel.
|
||||
|
||||
Disconnected clients are detected during the send attempt and
|
||||
automatically removed from all subscriptions.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
symbol:
|
||||
Trading pair, e.g. ``"BTC/USDT"``.
|
||||
timeframe:
|
||||
Candle interval, e.g. ``"1m"``, ``"1h"``.
|
||||
exchange:
|
||||
Exchange name, e.g. ``"mexc"``.
|
||||
message:
|
||||
Payload dict. ``Decimal`` and ``datetime`` values are
|
||||
automatically converted to JSON-safe equivalents.
|
||||
"""
|
||||
key = _subscription_key(exchange, symbol, timeframe)
|
||||
|
||||
# --- 1. Snapshot subscribers under the lock ---
|
||||
async with self._lock:
|
||||
subscribers = list(self._subscriptions.get(key, set()))
|
||||
|
||||
if not subscribers:
|
||||
return
|
||||
|
||||
# --- 2. Serialise once for all recipients ---
|
||||
safe_message = _json_safe(message)
|
||||
payload = json.dumps(safe_message, default=str)
|
||||
|
||||
# --- 3. Send to each subscriber, cleaning up on failure ---
|
||||
dead: list[WebSocket] = []
|
||||
for ws in subscribers:
|
||||
try:
|
||||
if ws.client_state == WebSocketState.DISCONNECTED:
|
||||
dead.append(ws)
|
||||
continue
|
||||
await ws.send_text(payload)
|
||||
except WebSocketDisconnect:
|
||||
dead.append(ws)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"Error sending WS message to %s on channel %s",
|
||||
id(ws),
|
||||
key,
|
||||
)
|
||||
dead.append(ws)
|
||||
|
||||
if dead:
|
||||
async with self._lock:
|
||||
for ws in dead:
|
||||
keys = self._client_subs.pop(ws, set())
|
||||
for k in keys:
|
||||
self._subscriptions.get(k, set()).discard(ws)
|
||||
if k in self._subscriptions and not self._subscriptions[k]:
|
||||
del self._subscriptions[k]
|
||||
logger.info(
|
||||
"Cleaned up %d disconnected websocket(s) from channel %s",
|
||||
len(dead),
|
||||
key,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Singleton
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
manager: ConnectionManager = ConnectionManager()
|
||||
"""Module-level singleton ``ConnectionManager`` instance."""
|
||||
Reference in New Issue
Block a user