from collections import defaultdict from typing import Any from fastapi import APIRouter, Depends from resonance_engine.adapters.db import db from resonance_engine.api import schemas, security from resonance_engine.dsp.enums import DSPClientName, DSPClientStatus, DSPId from resonance_engine.dsp.gateway import dsp_gateway from resonance_engine.dsp.models import DSPClient from resonance_engine.tasks.infra import get_recommended_concurrency from resonance_engine.tasks.models import CollectTask from resonance_engine.tasks.planner import COLD_REQUEST_LATENCY_S router = APIRouter( prefix="/dsp", tags=["DSP"], dependencies=[Depends(security.authenticate)], ) @router.get( "/clients", response_model=list[schemas.DSP], ) def dsp_list() -> Any: with db.autocommit(): clients = DSPClient.query.order_by(DSPClient.dsp_id, DSPClient.name).all() groups: defaultdict[DSPId, list[DSPClient]] = defaultdict(list) for client in clients: groups[client.dsp_id].append(client) return [ schemas.DSP( id=dsp_id, clients=[_client_schema(client) for client in clients] ) for dsp_id, clients in sorted(groups.items()) ] @router.post( "/clients/{client_name}/pause", response_model=schemas.DSPClient, ) def pause_dsp(client_name: DSPClientName) -> Any: with db.transaction(): client = DSPClient.query.where(DSPClient.name == client_name).one() client.status = DSPClientStatus.paused client.save() return _client_schema(client) @router.post( "/clients/{client_name}/resume", response_model=schemas.DSPClient, ) def resume_dsp(client_name: DSPClientName) -> Any: with db.transaction(): client = DSPClient.query.where(DSPClient.name == client_name).one() client.status = DSPClientStatus.active client.save() return _client_schema(client) @router.put( "/clients/{client_name}", response_model=schemas.DSPClient, ) def update_dsp(client_name: DSPClientName, body: schemas.UpdateDSPInput) -> Any: with db.transaction(): client = DSPClient.query.where(DSPClient.name == client_name).one() client.nominal_rps = body.nominal_rps client.save() return _client_schema(client) @router.get( "/clients/{name}/recommended-concurrency", response_model=schemas.RecommendedConcurrency, ) def recommended_concurrency(name: DSPClientName) -> Any: with db.autocommit(): client = DSPClient.query.where(DSPClient.name == name).one() signals = CollectTask.query.fanout_signals(client.name) observed_rps = signals.rps if signals else None latency_s = ( signals.latency_s if signals and signals.latency_s > 0 else COLD_REQUEST_LATENCY_S ) throttle_rate = signals.throttle_rate if signals else 0.0 recommended = get_recommended_concurrency( client.nominal_rps or 1, observed_rps=observed_rps, latency_s=latency_s, throttle_rate=throttle_rate, ) if observed_rps and observed_rps > 0: ran = observed_rps * latency_s basis = "backoff" if throttle_rate > 0 else "grow" else: ran = None basis = "cold" return schemas.RecommendedConcurrency( recommended=recommended, nominal_rps=client.nominal_rps, observed_rps=observed_rps, latency_s=latency_s, throttle_rate=throttle_rate, basis=basis, ran=ran, ) def _client_schema(client: DSPClient) -> schemas.DSPClient: return schemas.DSPClient( id=client.id, name=client.name, display_name=client.display_name, status=client.status, configured=dsp_gateway.is_configured(client.name), stats=dsp_gateway.stats(client.name), nominal_rps=client.nominal_rps, )