diff --git a/litellm/proxy/management_endpoints/roi_calculator_endpoints.py b/litellm/proxy/management_endpoints/roi_calculator_endpoints.py index 593d14bfaac..d8a1b781c54 100644 --- a/litellm/proxy/management_endpoints/roi_calculator_endpoints.py +++ b/litellm/proxy/management_endpoints/roi_calculator_endpoints.py @@ -6,6 +6,7 @@ from types import MappingProxyType from typing import Annotated, Final, Literal import httpx +from apscheduler.schedulers.asyncio import AsyncIOScheduler # pyright: ignore[reportMissingTypeStubs] # no upstream stubs from fastapi import APIRouter, Depends, FastAPI, HTTPException, Query from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError @@ -556,6 +557,17 @@ def _next_update(settings: ROISettings, status: ROISyncStatus, report: ROIReport return utc_anchor + timedelta(minutes=settings.update_interval_minutes) +def register_scheduled_sync(scheduler: AsyncIOScheduler) -> None: + scheduler.add_job( # pyright: ignore[reportUnknownMemberType] # APScheduler exposes untyped scheduling parameters + run_scheduled_sync, + "interval", + seconds=30, + id="roi_calculator_refresh", + max_instances=1, + replace_existing=True, + ) + + async def run_scheduled_sync() -> None: from litellm.proxy.proxy_server import prisma_client diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 915abf70b6f..83b439188fb 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1549,9 +1549,9 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: model_info_scheduler.start() if scheduler is not None and prisma_client is not None: - from litellm.proxy.management_endpoints.roi_calculator_endpoints import run_scheduled_sync + from litellm.proxy.management_endpoints.roi_calculator_endpoints import register_scheduled_sync - scheduler.add_job(run_scheduled_sync, "interval", seconds=30, id="roi_calculator_refresh", max_instances=1) + register_scheduled_sync(scheduler) # End of startup event yield diff --git a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py index eaf896831b1..df3a6256965 100644 --- a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py @@ -6,6 +6,7 @@ from types import MappingProxyType from typing import Final, cast import pytest +from apscheduler.schedulers.asyncio import AsyncIOScheduler from fastapi import FastAPI from fastapi.testclient import TestClient from pydantic import TypeAdapter @@ -16,7 +17,9 @@ from litellm.proxy.management_endpoints.roi_calculator_endpoints import ( _estimator_models_from_deployments, _next_update, get_roi_config_repository, + register_scheduled_sync, router, + run_scheduled_sync, ) from litellm.proxy.roi_calculator.estimator import estimator_options from litellm.proxy.roi_calculator.sample import sample_report @@ -25,6 +28,21 @@ from litellm.types.roi_calculator import ROIReport, ROISettings, ROISyncStatus _JSON_HEADERS: Final = MappingProxyType({"content-type": "application/json"}) +@pytest.mark.asyncio +async def test_repeated_startup_keeps_one_roi_schedule() -> None: + scheduler: Final = AsyncIOScheduler() + scheduler.start(paused=True) + try: + register_scheduled_sync(scheduler) + register_scheduled_sync(scheduler) + + jobs: Final = scheduler.get_jobs() + assert len(jobs) == 1 + assert jobs[0].func is run_scheduled_sync + finally: + scheduler.shutdown(wait=False) + + def _assert_json_round_trip(value: object) -> None: serialized: Final = json.dumps(value) decoded: Final[object] = cast(object, json.loads(serialized))