diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 1c49ad51beb..52b814c788a 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3211,13 +3211,19 @@ async def global_spend_refresh(): "timeout": 6000, }, ) - await new_client.db.connect() - await _query_raw(new_client, sql_query) - verbose_proxy_logger.info("MonthlyGlobalSpend view refreshed") - return { - "message": "MonthlyGlobalSpend view refreshed", - "status": "success", - } + try: + await new_client.db.connect() + await _query_raw(new_client, sql_query) + verbose_proxy_logger.info("MonthlyGlobalSpend view refreshed") + return { + "message": "MonthlyGlobalSpend view refreshed", + "status": "success", + } + finally: + try: + await new_client.disconnect() + except Exception: + verbose_proxy_logger.exception("Failed to disconnect MonthlyGlobalSpend refresh client") except Exception as e: verbose_proxy_logger.exception("Failed to refresh materialized view - %s", e) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 19ceb3d3d1f..84b99a17932 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -5811,3 +5811,82 @@ def test_scoped_spend_report_range_at_max_allowed(client, monkeypatch): mock_prisma.db.query_raw.assert_awaited_once() finally: app.dependency_overrides.pop(ps.user_api_key_auth, None) +@pytest.mark.asyncio +async def test_global_spend_refresh_disconnects_client_on_success(monkeypatch): + from litellm.proxy.spend_tracking.spend_management_endpoints import ( + global_spend_refresh, + ) + + fake_singleton = MagicMock() + fake_singleton.db = MagicMock() + fake_singleton.db.query_raw = AsyncMock( + return_value=[{"relname": "MonthlyGlobalSpend", "relkind": "m"}] + ) + monkeypatch.setattr(ps, "prisma_client", fake_singleton) + monkeypatch.setenv("DATABASE_URL", "postgresql://localhost:5432/db") + + fake_client = MagicMock() + fake_client.db = MagicMock() + fake_client.db.connect = AsyncMock() + fake_client.db.query_raw = AsyncMock(return_value=None) + fake_client.disconnect = AsyncMock() + + with patch("litellm.proxy.utils.PrismaClient", return_value=fake_client): + result = await global_spend_refresh() + + assert result["status"] == "success" + fake_client.disconnect.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_global_spend_refresh_disconnects_client_on_failure(monkeypatch): + from litellm.proxy.spend_tracking.spend_management_endpoints import ( + global_spend_refresh, + ) + + fake_singleton = MagicMock() + fake_singleton.db = MagicMock() + fake_singleton.db.query_raw = AsyncMock( + return_value=[{"relname": "MonthlyGlobalSpend", "relkind": "m"}] + ) + monkeypatch.setattr(ps, "prisma_client", fake_singleton) + monkeypatch.setenv("DATABASE_URL", "postgresql://localhost:5432/db") + + fake_client = MagicMock() + fake_client.db = MagicMock() + fake_client.db.connect = AsyncMock() + fake_client.db.query_raw = AsyncMock(side_effect=Exception("refresh timed out")) + fake_client.disconnect = AsyncMock() + + with patch("litellm.proxy.utils.PrismaClient", return_value=fake_client): + result = await global_spend_refresh() + + assert result["status"] == "failure" + fake_client.disconnect.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_global_spend_refresh_swallows_disconnect_error(monkeypatch): + from litellm.proxy.spend_tracking.spend_management_endpoints import ( + global_spend_refresh, + ) + + fake_singleton = MagicMock() + fake_singleton.db = MagicMock() + fake_singleton.db.query_raw = AsyncMock( + return_value=[{"relname": "MonthlyGlobalSpend", "relkind": "m"}] + ) + monkeypatch.setattr(ps, "prisma_client", fake_singleton) + monkeypatch.setenv("DATABASE_URL", "postgresql://localhost:5432/db") + + fake_client = MagicMock() + fake_client.db = MagicMock() + fake_client.db.connect = AsyncMock() + fake_client.db.query_raw = AsyncMock(return_value=None) + fake_client.disconnect = AsyncMock(side_effect=Exception("disconnect failed")) + + with patch("litellm.proxy.utils.PrismaClient", return_value=fake_client): + result = await global_spend_refresh() + + assert result["status"] == "success" + fake_client.disconnect.assert_awaited_once()