from __future__ import annotations from typing import Annotated from fastapi import APIRouter, Depends from pydantic import BaseModel, Field from ..config import get_settings from ..models.db import get_db from ..sim.ledger import Ledger from .auth import require_user router = APIRouter(prefix="/api/settings", tags=["settings"]) KEYS = ( "fee_rate", "exit_move_points", "rest_seconds", "max_rounds", "initial_equity", "perp_qty_eth", "option_qty_eth", ) class StrategySettingsBody(BaseModel): fee_rate: float | None = Field(default=None, ge=0, le=0.05) exit_move_points: float | None = Field(default=None, ge=1, le=500) rest_seconds: int | None = Field(default=None, ge=0, le=3600) max_rounds: int | None = Field(default=None, ge=1, le=20) initial_equity: float | None = Field(default=None, ge=1000) perp_qty_eth: float | None = Field(default=None, ge=0.01, le=100) option_qty_eth: float | None = Field(default=None, ge=0.01, le=100) def _read_settings() -> dict: db = get_db() s = get_settings() return { "fee_rate": float(db.get_setting("fee_rate", str(s.fee_rate)) or s.fee_rate), "exit_move_points": float( db.get_setting("exit_move_points", str(s.exit_move_points)) or s.exit_move_points ), "rest_seconds": int( float(db.get_setting("rest_seconds", str(s.rest_seconds)) or s.rest_seconds) ), "max_rounds": int( float(db.get_setting("max_rounds", str(s.max_rounds)) or s.max_rounds) ), "initial_equity": float( db.get_setting("initial_equity", str(s.initial_equity)) or s.initial_equity ), "perp_qty_eth": float( db.get_setting("perp_qty_eth", str(s.perp_qty_eth)) or s.perp_qty_eth ), "option_qty_eth": float( db.get_setting("option_qty_eth", str(s.option_qty_eth)) or s.option_qty_eth ), "ledger": Ledger(db).snapshot(), } @router.get("/strategy") async def get_strategy_settings(_user: Annotated[str, Depends(require_user)]) -> dict: return _read_settings() @router.put("/strategy") async def put_strategy_settings( body: StrategySettingsBody, _user: Annotated[str, Depends(require_user)], ) -> dict: db = get_db() data = body.model_dump(exclude_none=True) for k, v in data.items(): if k in KEYS: db.set_setting(k, str(v)) return _read_settings()