mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
* fix(proxy): mark session/SSO/SAML cookies Secure behind a TLS-terminating reverse proxy litellm only sees a plain-HTTP hop when TLS terminates at a reverse proxy, so cookie Secure attributes previously derived from (or defaulted without regard to) the literal request scheme could be dropped in production. The token session cookie set by every login path never carried Secure/HttpOnly/ SameSite at all. Adds IPAddressUtils.is_request_https, a single trust-aware resolver used by every cookie-setting call site: PROXY_BASE_URL, then X-Forwarded-Proto only from a configured trusted proxy (general_settings.use_x_forwarded_for + mcp_trusted_proxy_ranges), then the literal scheme. An unconfigured or untrusted caller cannot spoof the header to force Secure on. Resolves LIT-6748 * fix(proxy): make the shared session-cookie helper public, type new test helpers set_session_token_cookie is imported across modules (ui_sso.py -> proxy_server.py), so the leading underscore was misleading and breached basedpyright's reportPrivateUsage budget with zero headroom. Also adds missing parameter/return type annotations to the new test helper functions per repo convention.
12509 lines
492 KiB
Python
12509 lines
492 KiB
Python
import asyncio
|
||
import contextlib
|
||
import importlib
|
||
import json
|
||
import os
|
||
import re
|
||
import socket
|
||
import subprocess
|
||
import types
|
||
from datetime import datetime, timedelta, timezone
|
||
from pathlib import Path
|
||
from unittest import mock
|
||
from unittest.mock import AsyncMock, MagicMock, mock_open, patch
|
||
|
||
import click
|
||
import httpx
|
||
import pytest
|
||
import yaml
|
||
from fastapi import FastAPI
|
||
from fastapi.staticfiles import StaticFiles
|
||
from fastapi.testclient import TestClient
|
||
|
||
|
||
import litellm
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
from litellm.caching.caching import RedisCache
|
||
from litellm.caching.redis_cluster_cache import RedisClusterCache
|
||
from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||
from litellm.proxy.proxy_server import app, initialize
|
||
from litellm.utils import _invalidate_model_cost_lowercase_map
|
||
|
||
example_embedding_result = {
|
||
"object": "list",
|
||
"data": [
|
||
{
|
||
"object": "embedding",
|
||
"index": 0,
|
||
"embedding": [
|
||
-0.006929283495992422,
|
||
-0.005336422007530928,
|
||
-4.547132266452536e-05,
|
||
-0.024047505110502243,
|
||
-0.006929283495992422,
|
||
-0.005336422007530928,
|
||
-4.547132266452536e-05,
|
||
-0.024047505110502243,
|
||
-0.006929283495992422,
|
||
-0.005336422007530928,
|
||
-4.547132266452536e-05,
|
||
-0.024047505110502243,
|
||
],
|
||
}
|
||
],
|
||
"model": "text-embedding-3-small",
|
||
"usage": {"prompt_tokens": 5, "total_tokens": 5},
|
||
}
|
||
|
||
|
||
def mock_patch_aembedding():
|
||
return mock.patch(
|
||
"litellm.proxy.proxy_server.llm_router.aembedding",
|
||
return_value=example_embedding_result,
|
||
)
|
||
|
||
|
||
@pytest.fixture(scope="function")
|
||
def client_no_auth():
|
||
# Assuming litellm.proxy.proxy_server is an object
|
||
from litellm.proxy.proxy_server import cleanup_router_config_variables
|
||
|
||
cleanup_router_config_variables()
|
||
filepath = os.path.dirname(os.path.abspath(__file__))
|
||
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
|
||
# initialize can get run in parallel, it sets specific variables for the fast api app, sinc eit gets run in parallel different tests use the wrong variables
|
||
asyncio.run(initialize(config=config_fp, debug=True))
|
||
return TestClient(app)
|
||
|
||
|
||
def test_cors_exposes_cache_key_header_to_browser_js():
|
||
from fastapi.middleware.cors import CORSMiddleware
|
||
|
||
from litellm.constants import LITELLM_UI_ALLOW_HEADERS
|
||
|
||
cors_middleware = next(m for m in app.user_middleware if m.cls is CORSMiddleware)
|
||
assert cors_middleware.kwargs["expose_headers"] is LITELLM_UI_ALLOW_HEADERS
|
||
assert "x-litellm-cache-key" in cors_middleware.kwargs["expose_headers"]
|
||
|
||
|
||
def test_login_v2_returns_redirect_url_and_sets_cookie(monkeypatch):
|
||
mock_login_result = {"user_id": "test-user"}
|
||
mock_prisma_client = MagicMock()
|
||
mock_authenticate_user = AsyncMock(return_value=mock_login_result)
|
||
mock_create_ui_token_object = MagicMock(return_value={"user_id": "test-user"})
|
||
mock_jwt_encode = MagicMock(return_value="signed-token")
|
||
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
mock_authenticate_user,
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.create_ui_token_object",
|
||
mock_create_ui_token_object,
|
||
)
|
||
monkeypatch.setattr("jwt.encode", mock_jwt_encode)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
|
||
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v2/login",
|
||
json={"username": "alice", "password": "secret"},
|
||
)
|
||
|
||
assert response.status_code == 200
|
||
assert response.json() == {
|
||
"redirect_url": "http://testserver/ui?login=success",
|
||
"token": "signed-token",
|
||
}
|
||
assert response.cookies.get("token") == "signed-token"
|
||
|
||
mock_authenticate_user.assert_awaited_once_with(
|
||
username="alice",
|
||
password="secret",
|
||
master_key="test-master-key",
|
||
prisma_client=mock_prisma_client,
|
||
general_settings={},
|
||
)
|
||
mock_create_ui_token_object.assert_called_once_with(
|
||
login_result=mock_login_result,
|
||
general_settings={},
|
||
premium_user=False,
|
||
)
|
||
mock_jwt_encode.assert_called_once()
|
||
payload, secret = mock_jwt_encode.call_args.args
|
||
# The UI session token carries a bounded-lifetime `exp` claim (dynamic timestamp), alongside
|
||
# the user_id; assert its presence rather than an exact expiry value.
|
||
assert payload["user_id"] == "test-user"
|
||
assert isinstance(payload.get("exp"), int) and payload["exp"] > 0
|
||
assert set(payload.keys()) == {"user_id", "exp"}
|
||
assert secret == "test-master-key"
|
||
assert mock_jwt_encode.call_args.kwargs == {"algorithm": "HS256"}
|
||
|
||
|
||
def _mock_login_v2_deps(monkeypatch: pytest.MonkeyPatch) -> None:
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
AsyncMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.create_ui_token_object",
|
||
MagicMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock())
|
||
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
|
||
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
|
||
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
|
||
|
||
|
||
def test_login_v2_sets_secure_cookie_over_direct_https(monkeypatch):
|
||
"""Regression: the token cookie previously carried no Secure/HttpOnly/SameSite
|
||
attributes at all, so it was always sent over plain HTTP."""
|
||
_mock_login_v2_deps(monkeypatch)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||
|
||
client = TestClient(app, base_url="https://testserver")
|
||
response = client.post("/v2/login", json={"username": "alice", "password": "secret"})
|
||
|
||
assert response.status_code == 200
|
||
cookie = response.headers.get("set-cookie")
|
||
assert "Secure" in cookie
|
||
assert "HttpOnly" not in cookie # deliberate: the dashboard reads this cookie via JS
|
||
assert "samesite=lax" in cookie.lower()
|
||
|
||
|
||
def test_login_v2_does_not_set_secure_cookie_over_direct_http(monkeypatch):
|
||
_mock_login_v2_deps(monkeypatch)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||
|
||
client = TestClient(app, base_url="http://testserver")
|
||
response = client.post("/v2/login", json={"username": "alice", "password": "secret"})
|
||
|
||
assert response.status_code == 200
|
||
assert "Secure" not in response.headers.get("set-cookie")
|
||
|
||
|
||
def test_login_v2_sets_secure_cookie_behind_trusted_tls_terminating_proxy(monkeypatch):
|
||
"""THE regression: litellm only sees a plain-HTTP hop when TLS terminates at a
|
||
reverse proxy, but the token cookie must still be Secure when the direct peer is
|
||
a configured trusted proxy reporting X-Forwarded-Proto: https."""
|
||
_mock_login_v2_deps(monkeypatch)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"use_x_forwarded_for": True, "mcp_trusted_proxy_ranges": ["10.0.0.0/8"]},
|
||
)
|
||
|
||
client = TestClient(app, base_url="http://testserver", client=("10.0.0.5", 50000))
|
||
response = client.post(
|
||
"/v2/login",
|
||
json={"username": "alice", "password": "secret"},
|
||
headers={"X-Forwarded-Proto": "https"},
|
||
)
|
||
|
||
assert response.status_code == 200
|
||
assert "Secure" in response.headers.get("set-cookie")
|
||
|
||
|
||
def test_login_v2_returns_json_on_proxy_exception(monkeypatch):
|
||
"""Test that /v2/login returns JSON error when ProxyException is raised"""
|
||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_authenticate_user = AsyncMock(
|
||
side_effect=ProxyException(
|
||
message="Invalid credentials",
|
||
type=ProxyErrorTypes.auth_error,
|
||
param="password",
|
||
code=401,
|
||
)
|
||
)
|
||
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
mock_authenticate_user,
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v2/login",
|
||
json={"username": "alice", "password": "wrong"},
|
||
)
|
||
|
||
assert response.status_code == 401
|
||
assert response.headers["content-type"] == "application/json"
|
||
data = response.json()
|
||
assert "error" in data
|
||
assert data["error"]["message"] == "Invalid credentials"
|
||
assert data["error"]["type"] == "auth_error"
|
||
|
||
|
||
def test_login_v2_returns_json_on_http_exception(monkeypatch):
|
||
"""Test that /v2/login converts HTTPException to JSON error response"""
|
||
from fastapi import HTTPException
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_authenticate_user = AsyncMock(side_effect=HTTPException(status_code=401, detail="Unauthorized"))
|
||
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
mock_authenticate_user,
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v2/login",
|
||
json={"username": "alice", "password": "secret"},
|
||
)
|
||
|
||
assert response.status_code == 401
|
||
assert response.headers["content-type"] == "application/json"
|
||
data = response.json()
|
||
assert "error" in data
|
||
assert isinstance(data["error"], dict)
|
||
|
||
|
||
def test_login_v2_returns_json_on_unexpected_exception(monkeypatch):
|
||
"""Test that /v2/login returns JSON error when unexpected exception occurs"""
|
||
mock_prisma_client = MagicMock()
|
||
mock_authenticate_user = AsyncMock(side_effect=ValueError("Unexpected error"))
|
||
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
mock_authenticate_user,
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v2/login",
|
||
json={"username": "alice", "password": "secret"},
|
||
)
|
||
|
||
assert response.status_code == 500
|
||
assert response.headers["content-type"] == "application/json"
|
||
data = response.json()
|
||
assert "error" in data
|
||
assert isinstance(data["error"], dict)
|
||
assert "Unexpected error" in data["error"]["message"]
|
||
|
||
|
||
def test_login_v2_returns_json_on_invalid_json_body(monkeypatch):
|
||
"""Test that /v2/login returns JSON error when request body is invalid JSON"""
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v2/login",
|
||
content="invalid json",
|
||
headers={"Content-Type": "application/json"},
|
||
)
|
||
|
||
assert response.status_code == 500
|
||
assert response.headers["content-type"] == "application/json"
|
||
data = response.json()
|
||
assert "error" in data
|
||
assert isinstance(data["error"], dict)
|
||
|
||
|
||
def test_login_v3_rejected_without_control_plane_url(monkeypatch):
|
||
"""v3/login returns 404 when control_plane_url is not configured."""
|
||
mock_prisma_client = MagicMock()
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v3/login",
|
||
json={"username": "alice", "password": "secret"},
|
||
)
|
||
|
||
assert response.status_code == 404
|
||
assert "control_plane_url" in response.json()["error"]["message"]
|
||
|
||
|
||
def test_login_v3_returns_code(monkeypatch):
|
||
"""v3/login returns an opaque code, not the JWT directly."""
|
||
mock_prisma_client = MagicMock()
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
AsyncMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.create_ui_token_object",
|
||
MagicMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"control_plane_url": "https://cp.example.com"},
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
mock_config = MagicMock()
|
||
mock_config.worker_registry = []
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
|
||
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
|
||
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v3/login",
|
||
json={"username": "alice", "password": "secret"},
|
||
)
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert "code" in data
|
||
assert data["expires_in"] == 60
|
||
assert "token" not in data
|
||
|
||
|
||
def test_login_v3_exchange_happy_path(monkeypatch):
|
||
"""Full flow: v3/login returns code, v3/login/exchange redeems it for JWT."""
|
||
mock_prisma_client = MagicMock()
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
AsyncMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.create_ui_token_object",
|
||
MagicMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"control_plane_url": "https://cp.example.com"},
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
mock_config = MagicMock()
|
||
mock_config.worker_registry = []
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
|
||
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
|
||
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
|
||
|
||
client = TestClient(app)
|
||
|
||
# Step 1: login — get code
|
||
login_response = client.post(
|
||
"/v3/login",
|
||
json={"username": "alice", "password": "secret"},
|
||
)
|
||
assert login_response.status_code == 200
|
||
code = login_response.json()["code"]
|
||
|
||
# Step 2: exchange — get JWT
|
||
exchange_response = client.post(
|
||
"/v3/login/exchange",
|
||
json={"code": code},
|
||
)
|
||
assert exchange_response.status_code == 200
|
||
exchange_data = exchange_response.json()
|
||
assert exchange_data["token"] == "signed-token"
|
||
assert "redirect_url" in exchange_data
|
||
assert exchange_response.cookies.get("token") == "signed-token"
|
||
|
||
|
||
def test_login_v3_exchange_sets_secure_cookie_behind_trusted_tls_terminating_proxy(monkeypatch):
|
||
"""Regression: /v3/login/exchange's token cookie must be Secure behind a trusted
|
||
TLS-terminating reverse proxy even though litellm only sees a plain-HTTP hop."""
|
||
mock_prisma_client = MagicMock()
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
AsyncMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.create_ui_token_object",
|
||
MagicMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{
|
||
"control_plane_url": "https://cp.example.com",
|
||
"use_x_forwarded_for": True,
|
||
"mcp_trusted_proxy_ranges": ["10.0.0.0/8"],
|
||
},
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
mock_config = MagicMock()
|
||
mock_config.worker_registry = []
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
|
||
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
|
||
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
|
||
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
|
||
|
||
client = TestClient(app, base_url="http://testserver", client=("10.0.0.5", 50000))
|
||
|
||
login_response = client.post("/v3/login", json={"username": "alice", "password": "secret"})
|
||
code = login_response.json()["code"]
|
||
|
||
exchange_response = client.post(
|
||
"/v3/login/exchange",
|
||
json={"code": code},
|
||
headers={"X-Forwarded-Proto": "https"},
|
||
)
|
||
assert exchange_response.status_code == 200
|
||
assert "Secure" in exchange_response.headers.get("set-cookie")
|
||
|
||
|
||
def test_login_v3_exchange_single_use(monkeypatch):
|
||
"""Code can only be redeemed once."""
|
||
mock_prisma_client = MagicMock()
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
AsyncMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.create_ui_token_object",
|
||
MagicMock(return_value={"user_id": "test-user"}),
|
||
)
|
||
monkeypatch.setattr("jwt.encode", MagicMock(return_value="signed-token"))
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"control_plane_url": "https://cp.example.com"},
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
mock_config = MagicMock()
|
||
mock_config.worker_registry = []
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_config", mock_config)
|
||
monkeypatch.setattr("litellm.proxy.utils.get_server_root_path", lambda: "")
|
||
monkeypatch.setattr("litellm.proxy.utils.get_proxy_base_url", lambda: None)
|
||
|
||
client = TestClient(app)
|
||
|
||
login_response = client.post(
|
||
"/v3/login",
|
||
json={"username": "alice", "password": "secret"},
|
||
)
|
||
code = login_response.json()["code"]
|
||
|
||
# First exchange succeeds
|
||
first = client.post("/v3/login/exchange", json={"code": code})
|
||
assert first.status_code == 200
|
||
|
||
# Second exchange fails
|
||
second = client.post("/v3/login/exchange", json={"code": code})
|
||
assert second.status_code == 401
|
||
|
||
|
||
def test_login_v3_exchange_invalid_code(monkeypatch):
|
||
"""Random code returns 401."""
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"control_plane_url": "https://cp.example.com"},
|
||
)
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v3/login/exchange",
|
||
json={"code": "nonexistent-code"},
|
||
)
|
||
assert response.status_code == 401
|
||
|
||
|
||
def test_login_v3_exchange_rejected_without_control_plane_url(monkeypatch):
|
||
"""v3/login/exchange returns 404 when control_plane_url is not configured."""
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v3/login/exchange",
|
||
json={"code": "some-code"},
|
||
)
|
||
|
||
assert response.status_code == 404
|
||
assert "control_plane_url" in response.json()["error"]["message"]
|
||
|
||
|
||
def test_login_v3_returns_json_on_proxy_exception(monkeypatch):
|
||
"""Test that /v3/login returns JSON error when ProxyException is raised"""
|
||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_authenticate_user = AsyncMock(
|
||
side_effect=ProxyException(
|
||
message="Invalid credentials",
|
||
type=ProxyErrorTypes.auth_error,
|
||
param="password",
|
||
code=401,
|
||
)
|
||
)
|
||
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.auth.login_utils.authenticate_user",
|
||
mock_authenticate_user,
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.master_key", "test-master-key")
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"control_plane_url": "https://cp.example.com"},
|
||
)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
|
||
client = TestClient(app)
|
||
response = client.post(
|
||
"/v3/login",
|
||
json={"username": "alice", "password": "wrong"},
|
||
)
|
||
|
||
assert response.status_code == 401
|
||
assert response.headers["content-type"] == "application/json"
|
||
data = response.json()
|
||
assert "error" in data
|
||
assert data["error"]["message"] == "Invalid credentials"
|
||
assert data["error"]["type"] == "auth_error"
|
||
|
||
|
||
def test_fallback_login_has_no_deprecation_banner(client_no_auth):
|
||
response = client_no_auth.get("/fallback/login")
|
||
|
||
assert response.status_code == 200
|
||
html = response.text
|
||
assert '<div class="deprecation-banner">' not in html
|
||
assert "Deprecated:" not in html
|
||
assert "<form" in html
|
||
|
||
|
||
@pytest.mark.parametrize(
|
||
"ui_logo_path",
|
||
[
|
||
"/etc/litellm/secret-config.json",
|
||
"/var/secrets/admin.key",
|
||
"/proc/self/environ",
|
||
"relative/path/logo.png",
|
||
],
|
||
)
|
||
def test_get_logo_url_does_not_disclose_local_paths(client_no_auth, monkeypatch, ui_logo_path):
|
||
# ``/get_logo_url`` is unauthenticated. Returning a local filesystem
|
||
# path verbatim discloses admin-only config to any caller. Only
|
||
# browser-loadable HTTP(S) URLs should be returned; for local paths
|
||
# the dashboard falls back to ``/get_image``.
|
||
monkeypatch.setenv("UI_LOGO_PATH", ui_logo_path)
|
||
|
||
response = client_no_auth.get("/get_logo_url")
|
||
|
||
assert response.status_code == 200
|
||
assert response.json() == {"logo_url": ""}
|
||
|
||
|
||
def test_get_logo_url_returns_https_url(client_no_auth, monkeypatch):
|
||
monkeypatch.setenv("UI_LOGO_PATH", "https://cdn.public.example/logo.png")
|
||
|
||
response = client_no_auth.get("/get_logo_url")
|
||
|
||
assert response.status_code == 200
|
||
assert response.json() == {"logo_url": "https://cdn.public.example/logo.png"}
|
||
|
||
|
||
def test_get_logo_url_returns_http_url(client_no_auth, monkeypatch):
|
||
# HTTP URLs (typically internal CDN) are still returned — those are
|
||
# intended to be loaded directly by the browser.
|
||
monkeypatch.setenv("UI_LOGO_PATH", "http://internal-cdn.corp:8080/logo.png")
|
||
|
||
response = client_no_auth.get("/get_logo_url")
|
||
|
||
assert response.status_code == 200
|
||
assert response.json() == {"logo_url": "http://internal-cdn.corp:8080/logo.png"}
|
||
|
||
|
||
def test_get_logo_url_returns_empty_when_unset(client_no_auth, monkeypatch):
|
||
monkeypatch.delenv("UI_LOGO_PATH", raising=False)
|
||
|
||
response = client_no_auth.get("/get_logo_url")
|
||
|
||
assert response.status_code == 200
|
||
assert response.json() == {"logo_url": ""}
|
||
|
||
|
||
def test_sso_key_generate_shows_deprecation_banner(client_no_auth, monkeypatch):
|
||
# Ensure the route returns the HTML form instead of redirecting
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.management_endpoints.ui_sso.show_missing_vars_in_env",
|
||
lambda: None,
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.get_redirect_url_for_sso",
|
||
lambda *args, **kwargs: "http://test/redirect",
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler._get_cli_state",
|
||
lambda *args, **kwargs: None,
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.should_use_sso_handler",
|
||
lambda *args, **kwargs: False,
|
||
)
|
||
# Mock premium_user to bypass enterprise check (prevents 403 Forbidden)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.premium_user",
|
||
True,
|
||
)
|
||
monkeypatch.setenv("UI_USERNAME", "admin")
|
||
|
||
response = client_no_auth.get("/sso/key/generate")
|
||
|
||
assert response.status_code == 200
|
||
html = response.text
|
||
assert '<div class="deprecation-banner">' in html
|
||
assert "Deprecated:" in html
|
||
|
||
|
||
def test_restructure_ui_html_files_handles_nested_routes(tmp_path):
|
||
"""
|
||
Test that _restructure_ui_html_files correctly restructures HTML files.
|
||
Note: This function is always called now, both in development and non-root Docker environments.
|
||
"""
|
||
from litellm.proxy import proxy_server
|
||
|
||
ui_root = tmp_path / "ui"
|
||
ui_root.mkdir()
|
||
|
||
def write_file(path: Path, content: str) -> None:
|
||
path.parent.mkdir(parents=True, exist_ok=True)
|
||
path.write_text(content)
|
||
|
||
write_file(ui_root / "home.html", "home")
|
||
write_file(ui_root / "mcp" / "oauth" / "callback.html", "callback")
|
||
write_file(ui_root / "existing" / "index.html", "keep")
|
||
write_file(ui_root / "_next" / "ignore.html", "asset")
|
||
write_file(ui_root / "litellm-asset-prefix" / "ignore.html", "asset")
|
||
|
||
proxy_server._restructure_ui_html_files(str(ui_root))
|
||
|
||
assert not (ui_root / "home.html").exists()
|
||
assert (ui_root / "home" / "index.html").read_text() == "home"
|
||
assert not (ui_root / "mcp" / "oauth" / "callback.html").exists()
|
||
assert (ui_root / "mcp" / "oauth" / "callback" / "index.html").read_text() == "callback"
|
||
assert (ui_root / "existing" / "index.html").read_text() == "keep"
|
||
assert (ui_root / "_next" / "ignore.html").read_text() == "asset"
|
||
assert (ui_root / "litellm-asset-prefix" / "ignore.html").read_text() == "asset"
|
||
|
||
|
||
def test_ui_extensionless_route_requires_restructure(tmp_path):
|
||
"""
|
||
Regression for non-root fallback: /ui/login expects login/index.html.
|
||
Note: Restructuring always happens now, both in development and non-root Docker environments.
|
||
"""
|
||
|
||
from litellm.proxy import proxy_server
|
||
|
||
ui_root = tmp_path / "ui"
|
||
ui_root.mkdir()
|
||
(ui_root / "index.html").write_text("index")
|
||
(ui_root / "login.html").write_text("login")
|
||
|
||
fastapi_app = FastAPI()
|
||
fastapi_app.mount("/ui", StaticFiles(directory=str(ui_root), html=True), name="ui")
|
||
client = TestClient(fastapi_app)
|
||
|
||
assert client.get("/ui/login.html").status_code == 200
|
||
assert client.get("/ui/login").status_code == 404
|
||
|
||
proxy_server._restructure_ui_html_files(str(ui_root))
|
||
|
||
response = client.get("/ui/login")
|
||
assert response.status_code == 200
|
||
assert "login" in response.text
|
||
|
||
|
||
def test_admin_ui_export_serves_nested_extensionless_routes():
|
||
out_dir = Path(litellm.__file__).parent / "proxy" / "_experimental" / "out"
|
||
assert out_dir.is_dir(), f"missing UI export at {out_dir}"
|
||
|
||
nested_html_offenders = [
|
||
path.relative_to(out_dir).as_posix()
|
||
for path in out_dir.rglob("*.html")
|
||
if path.parent != out_dir
|
||
and path.name != "index.html"
|
||
and "_next" not in path.parts
|
||
and "litellm-asset-prefix" not in path.parts
|
||
]
|
||
assert not nested_html_offenders, f"Nested routes must be named index.html. Offenders: {nested_html_offenders}"
|
||
|
||
callback_index = out_dir / "mcp" / "oauth" / "callback" / "index.html"
|
||
assert callback_index.is_file(), (
|
||
f"MCP OAuth callback page must exist at {callback_index}; "
|
||
"without it /ui/mcp/oauth/callback 404s after Linear redirects back."
|
||
)
|
||
|
||
fastapi_app = FastAPI()
|
||
fastapi_app.mount("/ui", StaticFiles(directory=str(out_dir), html=True), name="ui")
|
||
client = TestClient(fastapi_app)
|
||
|
||
redirect = client.get(
|
||
"/ui/mcp/oauth/callback?code=abc&state=xyz",
|
||
follow_redirects=False,
|
||
)
|
||
assert redirect.status_code == 307
|
||
assert redirect.headers["location"].endswith("/ui/mcp/oauth/callback/?code=abc&state=xyz")
|
||
|
||
landed = client.get("/ui/mcp/oauth/callback?code=abc&state=xyz")
|
||
assert landed.status_code == 200
|
||
assert "<html" in landed.text.lower()
|
||
|
||
|
||
def test_restructure_always_happens(monkeypatch):
|
||
"""
|
||
Test that restructuring logic always executes regardless of LITELLM_NON_ROOT setting.
|
||
In development (is_non_root=False), restructuring happens directly in _experimental/out.
|
||
In non-root Docker (is_non_root=True), restructuring happens in /var/lib/litellm/ui.
|
||
"""
|
||
# Test Case 1: is_non_root is True - restructuring happens in /var/lib/litellm/ui
|
||
monkeypatch.setenv("LITELLM_NON_ROOT", "true")
|
||
|
||
runtime_ui_path = "/var/lib/litellm/ui"
|
||
packaged_ui_path = "/some/packaged/ui/path"
|
||
|
||
# Simulate the logic from proxy_server.py
|
||
is_non_root = os.getenv("LITELLM_NON_ROOT", "").lower() == "true"
|
||
if is_non_root:
|
||
ui_path = runtime_ui_path
|
||
else:
|
||
ui_path = packaged_ui_path
|
||
|
||
# Restructuring always happens now, regardless of ui_path vs packaged_ui_path
|
||
should_restructure = True
|
||
|
||
assert is_non_root is True
|
||
assert should_restructure is True
|
||
assert ui_path == runtime_ui_path
|
||
|
||
# Test Case 2: is_non_root is False - restructuring happens directly in packaged_ui_path
|
||
monkeypatch.delenv("LITELLM_NON_ROOT", raising=False)
|
||
|
||
# Simulate the logic from proxy_server.py
|
||
is_non_root = os.getenv("LITELLM_NON_ROOT", "").lower() == "true"
|
||
if is_non_root:
|
||
ui_path = runtime_ui_path
|
||
else:
|
||
ui_path = packaged_ui_path
|
||
|
||
# Restructuring always happens now, even when ui_path == packaged_ui_path
|
||
should_restructure = True
|
||
|
||
assert is_non_root is False
|
||
assert should_restructure is True
|
||
assert ui_path == packaged_ui_path
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_initialize_scheduled_jobs_credentials(monkeypatch):
|
||
"""
|
||
Test that get_credentials is only called when store_model_in_db is True
|
||
"""
|
||
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
# Mock dependencies
|
||
mock_prisma_client = MagicMock()
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
): # set store_model_in_db to False
|
||
# Test when store_model_in_db is False
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
# Verify get_credentials was not called
|
||
mock_proxy_config.get_credentials.assert_not_called()
|
||
|
||
# Now test with store_model_in_db = True
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
# Verify get_credentials was called both directly and scheduled
|
||
assert mock_proxy_config.get_credentials.call_count == 1 # Direct call
|
||
|
||
# Verify a scheduled job was added for get_credentials
|
||
mock_scheduler_calls = [call[0] for call in mock_proxy_config.get_credentials.mock_calls]
|
||
assert len(mock_scheduler_calls) > 0
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_periodic_reload_job_scheduled_without_store_model_in_db(monkeypatch):
|
||
"""
|
||
Regression (LIT-4882): reload schedules configured from the Admin UI live in the DB and
|
||
must fire even without store_model_in_db, which used to gate the job that ran them
|
||
"""
|
||
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from apscheduler.schedulers.asyncio import AsyncIOScheduler
|
||
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
scheduler = AsyncIOScheduler()
|
||
|
||
try:
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False),
|
||
patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=scheduler),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
assert scheduler.get_job("periodic_reload_job") is not None
|
||
assert scheduler.get_job("add_deployment_job") is None
|
||
finally:
|
||
scheduler.shutdown(wait=False)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_initialize_scheduled_jobs_uses_configured_config_reload_interval(monkeypatch):
|
||
"""
|
||
The DB config-reload jobs (add_deployment, get_credentials) that keep multi-pod
|
||
deployments in sync must be scheduled at the configured
|
||
proxy_config_reload_interval_seconds, not a hardcoded value.
|
||
"""
|
||
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
mock_scheduler = MagicMock()
|
||
|
||
configured_interval = 47
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||
patch(
|
||
"litellm.proxy.proxy_server.proxy_config_reload_interval_seconds",
|
||
configured_interval,
|
||
),
|
||
patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=mock_scheduler),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
scheduled_seconds = {
|
||
job_call.kwargs["id"]: job_call.kwargs.get("seconds")
|
||
for job_call in mock_scheduler.add_job.call_args_list
|
||
if "id" in job_call.kwargs
|
||
}
|
||
assert scheduled_seconds["add_deployment_job"] == configured_interval
|
||
assert scheduled_seconds["get_credentials_job"] == configured_interval
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_initialize_scheduled_jobs_rejects_non_positive_config_reload_interval(monkeypatch):
|
||
"""
|
||
A non-positive proxy_config_reload_interval_seconds (misconfig via env/config/DB) would
|
||
make APScheduler reject the job and crash startup, so the scheduler must fall back to the
|
||
30s default instead of forwarding the bad value.
|
||
"""
|
||
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
mock_scheduler = MagicMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||
patch("litellm.proxy.proxy_server.proxy_config_reload_interval_seconds", 0),
|
||
patch("litellm.proxy.proxy_server.AsyncIOScheduler", return_value=mock_scheduler),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
scheduled_seconds = {
|
||
job_call.kwargs["id"]: job_call.kwargs.get("seconds")
|
||
for job_call in mock_scheduler.add_job.call_args_list
|
||
if "id" in job_call.kwargs
|
||
}
|
||
assert scheduled_seconds["add_deployment_job"] == 30
|
||
assert scheduled_seconds["get_credentials_job"] == 30
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_initialize_scheduled_jobs_hydrates_mcp_when_store_model_in_db_false(monkeypatch):
|
||
"""
|
||
Regression (LIT-4128): MCP servers created via the UI are persisted to the DB
|
||
regardless of store_model_in_db, but the in-memory registry that GET
|
||
/v1/mcp/server reads is hydrated from the DB only by the store_model_in_db
|
||
model-sync loop (add_deployment). On a DB-backed proxy with store_model_in_db
|
||
unset the registry must still be hydrated on startup so previously-added
|
||
servers survive a restart instead of showing an empty list until a write.
|
||
"""
|
||
monkeypatch.delenv("DISABLE_PRISMA_SCHEMA_UPDATE", raising=False)
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
mock_proxy_config.add_deployment.assert_not_called()
|
||
mock_proxy_config.init_mcp_servers_from_db.assert_awaited_once()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_mcp_servers_from_db_respects_supported_db_objects(monkeypatch):
|
||
"""
|
||
init_mcp_servers_from_db hydrates MCP from the DB by default but skips it when
|
||
an explicit supported_db_objects allowlist omits "mcp".
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
config = ProxyConfig()
|
||
with patch.object(config, "_init_mcp_servers_in_db", new=AsyncMock()) as mock_init:
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", {})
|
||
await config.init_mcp_servers_from_db()
|
||
mock_init.assert_awaited_once()
|
||
|
||
mock_init.reset_mock()
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"supported_db_objects": ["models"]},
|
||
)
|
||
await config.init_mcp_servers_from_db()
|
||
mock_init.assert_not_awaited()
|
||
|
||
|
||
def test_update_config_fields_deep_merge_db_wins():
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
current_config = {
|
||
"router_settings": {
|
||
"routing_mode": "cost_optimized",
|
||
"model_group_alias": {
|
||
# Existing alias with older model + different hidden flag
|
||
"claude-sonnet-4": {
|
||
"model": "claude-sonnet-4-20240219",
|
||
"hidden": True,
|
||
},
|
||
# An extra alias that should remain untouched unless DB overrides it
|
||
"legacy-sonnet": {
|
||
"model": "claude-2.1",
|
||
"hidden": True,
|
||
},
|
||
},
|
||
}
|
||
}
|
||
|
||
db_param_value = {
|
||
"model_group_alias": {
|
||
# Conflict: DB should win (both 'model' and 'hidden')
|
||
"claude-sonnet-4": {
|
||
"model": "claude-sonnet-4-20250514",
|
||
"hidden": False,
|
||
},
|
||
# New alias to be added by the merge
|
||
"claude-sonnet-latest": {
|
||
"model": "claude-sonnet-4-20250514",
|
||
"hidden": True,
|
||
},
|
||
# Demonstrate that None values from DB are skipped (preserve existing)
|
||
"legacy-sonnet": {"hidden": None}, # should not clobber current True
|
||
}
|
||
}
|
||
|
||
updated = proxy_config._update_config_fields(
|
||
current_config=current_config,
|
||
param_name="router_settings",
|
||
db_param_value=db_param_value,
|
||
)
|
||
|
||
rs = updated["router_settings"]
|
||
aliases = rs["model_group_alias"]
|
||
|
||
# DB wins on conflicts (deep) for existing alias
|
||
assert aliases["claude-sonnet-4"]["model"] == "claude-sonnet-4-20250514"
|
||
assert aliases["claude-sonnet-4"]["hidden"] is False
|
||
|
||
# New alias introduced by DB is present with its values
|
||
assert "claude-sonnet-latest" in aliases
|
||
assert aliases["claude-sonnet-latest"]["model"] == "claude-sonnet-4-20250514"
|
||
assert aliases["claude-sonnet-latest"]["hidden"] is True
|
||
|
||
# None in DB does not overwrite existing values
|
||
assert aliases["legacy-sonnet"]["model"] == "claude-2.1"
|
||
assert aliases["legacy-sonnet"]["hidden"] is True
|
||
|
||
# Unrelated router_settings keys are preserved
|
||
assert rs["routing_mode"] == "cost_optimized"
|
||
|
||
|
||
def test_get_config_custom_callback_api_env_vars(monkeypatch):
|
||
"""
|
||
Ensure /get/config/callbacks returns custom callback env vars when both custom values are provided.
|
||
"""
|
||
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
|
||
|
||
# Mock config with custom_callback_api enabled and generic logger env vars present
|
||
config_data = {
|
||
"litellm_settings": {"success_callback": ["custom_callback_api"]},
|
||
"general_settings": {},
|
||
"environment_variables": {
|
||
"GENERIC_LOGGER_ENDPOINT": "https://callback.example.com",
|
||
"GENERIC_LOGGER_HEADERS": "Auth: token",
|
||
},
|
||
}
|
||
|
||
# Mock proxy_config.get_config and router settings
|
||
mock_router = MagicMock()
|
||
mock_router.get_settings.return_value = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
|
||
monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
|
||
|
||
# Bypass auth dependency
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"
|
||
)
|
||
|
||
client = TestClient(app)
|
||
try:
|
||
response = client.get("/get/config/callbacks")
|
||
finally:
|
||
app.dependency_overrides = original_overrides
|
||
|
||
assert response.status_code == 200
|
||
callbacks = response.json()["callbacks"]
|
||
custom_cb = next((cb for cb in callbacks if cb["name"] == "custom_callback_api"), None)
|
||
|
||
assert custom_cb is not None
|
||
assert custom_cb["variables"] == {
|
||
"GENERIC_LOGGER_ENDPOINT": "https://callback.example.com",
|
||
"GENERIC_LOGGER_HEADERS": "Auth: token",
|
||
}
|
||
|
||
|
||
@patch(
|
||
"litellm.proxy.common_utils.callback_utils.CustomLogger.get_callback_env_vars",
|
||
return_value=["LANGFUSE_PUBLIC_KEY", "LANGFUSE_SECRET_KEY", "LANGFUSE_HOST"],
|
||
)
|
||
def test_get_config_callbacks_fall_back_to_process_env(mock_env_vars, monkeypatch):
|
||
"""A callback configured purely via process env vars is surfaced.
|
||
|
||
An IaC deployment sets LANGFUSE_* on the gateway and never touches the UI,
|
||
so nothing is stored in the config environment_variables overlay. The read
|
||
endpoint must still report the live values instead of blanks.
|
||
"""
|
||
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
|
||
|
||
monkeypatch.setenv("LANGFUSE_PUBLIC_KEY", "pk-env-only")
|
||
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-env-only")
|
||
monkeypatch.setenv("LANGFUSE_HOST", "https://cloud.langfuse.com")
|
||
|
||
config_data = {
|
||
"litellm_settings": {"success_callback": ["langfuse"]},
|
||
"general_settings": {},
|
||
"environment_variables": {},
|
||
}
|
||
mock_router = MagicMock()
|
||
mock_router.get_settings.return_value = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
|
||
monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
|
||
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"
|
||
)
|
||
|
||
client = TestClient(app)
|
||
try:
|
||
response = client.get("/get/config/callbacks")
|
||
finally:
|
||
app.dependency_overrides = original_overrides
|
||
|
||
assert response.status_code == 200
|
||
langfuse_cb = next((cb for cb in response.json()["callbacks"] if cb["name"] == "langfuse"), None)
|
||
assert langfuse_cb is not None
|
||
assert langfuse_cb["variables"] == {
|
||
"LANGFUSE_PUBLIC_KEY": "pk-env-only",
|
||
"LANGFUSE_SECRET_KEY": "sk-env-only",
|
||
"LANGFUSE_HOST": "https://cloud.langfuse.com",
|
||
}
|
||
|
||
|
||
@patch(
|
||
"litellm.proxy.common_utils.callback_utils.CustomLogger.get_callback_env_vars",
|
||
return_value=["LANGFUSE_SECRET_KEY", "LANGFUSE_HOST"],
|
||
)
|
||
def test_get_config_callback_env_secrets_redacted_for_non_admin(mock_env_vars, monkeypatch):
|
||
"""Surfacing env vars must not widen who can read secret values.
|
||
|
||
The callback role gate redacts sensitive keys for anyone below full admin,
|
||
and that must hold whether the value came from the stored config or the
|
||
process env. A non-secret var (LANGFUSE_HOST) still resolves for context.
|
||
"""
|
||
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
|
||
|
||
monkeypatch.setenv("LANGFUSE_SECRET_KEY", "sk-env-only-secret")
|
||
monkeypatch.setenv("LANGFUSE_HOST", "https://cloud.langfuse.com")
|
||
|
||
config_data = {
|
||
"litellm_settings": {"success_callback": ["langfuse"]},
|
||
"general_settings": {},
|
||
"environment_variables": {},
|
||
}
|
||
mock_router = MagicMock()
|
||
mock_router.get_settings.return_value = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
|
||
monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
|
||
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-user"
|
||
)
|
||
|
||
client = TestClient(app)
|
||
try:
|
||
response = client.get("/get/config/callbacks")
|
||
finally:
|
||
app.dependency_overrides = original_overrides
|
||
|
||
assert response.status_code == 200
|
||
langfuse_cb = next((cb for cb in response.json()["callbacks"] if cb["name"] == "langfuse"), None)
|
||
assert langfuse_cb is not None
|
||
assert langfuse_cb["variables"]["LANGFUSE_SECRET_KEY"] == "REDACTED"
|
||
assert langfuse_cb["variables"]["LANGFUSE_HOST"] == "https://cloud.langfuse.com"
|
||
|
||
|
||
def test_get_config_returns_email_settings(monkeypatch):
|
||
"""
|
||
Regression for https://github.com/BerriAI/litellm/issues/19221
|
||
|
||
proxy_config.get_config() already returns environment_variables decrypted
|
||
(the DB-overlay path decrypts them, and YAML values are plaintext). The
|
||
/get/config/callbacks email block must therefore surface those values as-is
|
||
instead of decrypting a second time. The old code ran decrypt_value_helper()
|
||
on the already-plaintext value, which failed and returned None, so every
|
||
SMTP_* field came back blank on UI refresh.
|
||
"""
|
||
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
|
||
|
||
smtp_password = "super-secret-app-password"
|
||
config_data = {
|
||
"litellm_settings": {},
|
||
"general_settings": {"alerting": ["email"]},
|
||
"environment_variables": {
|
||
"SMTP_HOST": "smtp.resend.com",
|
||
"SMTP_PORT": "587",
|
||
"SMTP_USERNAME": "resend",
|
||
"SMTP_PASSWORD": smtp_password,
|
||
"SMTP_SENDER_EMAIL": "alerts@example.com",
|
||
"TEST_EMAIL_ADDRESS": "admin@example.com",
|
||
},
|
||
}
|
||
|
||
mock_router = MagicMock()
|
||
mock_router.get_settings.return_value = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
|
||
monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
|
||
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"
|
||
)
|
||
|
||
client = TestClient(app)
|
||
try:
|
||
response = client.get("/get/config/callbacks")
|
||
finally:
|
||
app.dependency_overrides = original_overrides
|
||
|
||
assert response.status_code == 200
|
||
email_alert = next((a for a in response.json()["alerts"] if a["name"] == "email"), None)
|
||
assert email_alert is not None
|
||
variables = email_alert["variables"]
|
||
|
||
# Non-sensitive fields round-trip verbatim (None before the fix).
|
||
assert variables["SMTP_HOST"] == "smtp.resend.com"
|
||
assert variables["SMTP_PORT"] == "587"
|
||
assert variables["SMTP_USERNAME"] == "resend"
|
||
assert variables["SMTP_SENDER_EMAIL"] == "alerts@example.com"
|
||
assert variables["TEST_EMAIL_ADDRESS"] == "admin@example.com"
|
||
|
||
# Password is present but masked: never None, never the raw secret.
|
||
assert variables["SMTP_PASSWORD"] is not None
|
||
assert variables["SMTP_PASSWORD"] != smtp_password
|
||
assert "*" in variables["SMTP_PASSWORD"]
|
||
|
||
|
||
def _get_email_alert_variables(monkeypatch, config_data):
|
||
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
|
||
|
||
mock_router = MagicMock()
|
||
mock_router.get_settings.return_value = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
|
||
monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
|
||
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"
|
||
)
|
||
|
||
client = TestClient(app)
|
||
try:
|
||
response = client.get("/get/config/callbacks")
|
||
finally:
|
||
app.dependency_overrides = original_overrides
|
||
|
||
assert response.status_code == 200
|
||
email_alert = next((a for a in response.json()["alerts"] if a["name"] == "email"), None)
|
||
assert email_alert is not None
|
||
return email_alert["variables"]
|
||
|
||
|
||
def test_get_config_returns_email_settings_set_only_in_process_env(monkeypatch):
|
||
"""
|
||
Regression for LIT-4165.
|
||
|
||
SMTP supplied purely as process env vars (helm/terraform, no UI writes) is
|
||
live at runtime because litellm/proxy/utils.py::send_email resolves every
|
||
field from os.getenv. The /get/config/callbacks email block only read the
|
||
config/DB environment_variables overlay though, so those deployments saw an
|
||
empty Email Server Settings page and could not tell SMTP was configured.
|
||
The slack block one branch above already fell back to os.getenv.
|
||
"""
|
||
smtp_password = "env-only-app-password"
|
||
monkeypatch.setenv("SMTP_HOST", "smtp.env-host.com")
|
||
monkeypatch.setenv("SMTP_PORT", "2525")
|
||
monkeypatch.setenv("SMTP_TLS", "False")
|
||
monkeypatch.setenv("SMTP_USERNAME", "env-user")
|
||
monkeypatch.setenv("SMTP_PASSWORD", smtp_password)
|
||
monkeypatch.setenv("SMTP_SENDER_EMAIL", "alerts@env-host.com")
|
||
monkeypatch.setenv("TEST_EMAIL_ADDRESS", "admin@env-host.com")
|
||
|
||
variables = _get_email_alert_variables(
|
||
monkeypatch,
|
||
{
|
||
"litellm_settings": {},
|
||
"general_settings": {"alerting": ["email"]},
|
||
"environment_variables": {},
|
||
},
|
||
)
|
||
|
||
# Every one of these was None before the fix, despite SMTP working.
|
||
assert variables["SMTP_HOST"] == "smtp.env-host.com"
|
||
assert variables["SMTP_PORT"] == "2525"
|
||
assert variables["SMTP_TLS"] == "False"
|
||
assert variables["SMTP_USERNAME"] == "env-user"
|
||
assert variables["SMTP_SENDER_EMAIL"] == "alerts@env-host.com"
|
||
assert variables["TEST_EMAIL_ADDRESS"] == "admin@env-host.com"
|
||
|
||
# An env-sourced secret is masked exactly like a stored one.
|
||
assert variables["SMTP_PASSWORD"] not in (None, smtp_password)
|
||
assert "*" in variables["SMTP_PASSWORD"]
|
||
|
||
|
||
def test_get_config_email_settings_prefer_stored_over_process_env(monkeypatch):
|
||
"""
|
||
Stored environment_variables win over the process environment, matching the
|
||
load order in ProxyConfig.get_config, which pushes stored values into
|
||
os.environ. Only a field with no stored entry falls back to os.getenv.
|
||
"""
|
||
monkeypatch.setenv("SMTP_HOST", "smtp.env-host.com")
|
||
monkeypatch.setenv("SMTP_SENDER_EMAIL", "alerts@env-host.com")
|
||
|
||
variables = _get_email_alert_variables(
|
||
monkeypatch,
|
||
{
|
||
"litellm_settings": {},
|
||
"general_settings": {"alerting": ["email"]},
|
||
"environment_variables": {"SMTP_HOST": "smtp.stored-host.com"},
|
||
},
|
||
)
|
||
|
||
assert variables["SMTP_HOST"] == "smtp.stored-host.com"
|
||
assert variables["SMTP_SENDER_EMAIL"] == "alerts@env-host.com"
|
||
|
||
|
||
def test_get_config_email_settings_absent_everywhere_stay_none(monkeypatch):
|
||
"""A field set in neither source is reported unset rather than invented."""
|
||
for var in ("SMTP_HOST", "SMTP_PORT", "SMTP_TLS", "SMTP_USERNAME", "SMTP_PASSWORD", "SMTP_SENDER_EMAIL"):
|
||
monkeypatch.delenv(var, raising=False)
|
||
|
||
variables = _get_email_alert_variables(
|
||
monkeypatch,
|
||
{
|
||
"litellm_settings": {},
|
||
"general_settings": {"alerting": ["email"]},
|
||
"environment_variables": {},
|
||
},
|
||
)
|
||
|
||
assert variables["SMTP_HOST"] is None
|
||
assert variables["SMTP_PASSWORD"] is None
|
||
|
||
|
||
def test_get_config_returns_slack_webhook(monkeypatch):
|
||
"""
|
||
Same double-decryption regression as the email block (issue #19221): the
|
||
slack alerting block must surface the already-decrypted SLACK_WEBHOOK_URL
|
||
rather than decrypting it again into None.
|
||
"""
|
||
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
|
||
|
||
webhook_url = "https://hooks.slack.com/services/T00000/B00000/abcdefghijklmnop"
|
||
config_data = {
|
||
"litellm_settings": {},
|
||
"general_settings": {"alerting": ["slack"]},
|
||
"environment_variables": {"SLACK_WEBHOOK_URL": webhook_url},
|
||
}
|
||
|
||
mock_router = MagicMock()
|
||
mock_router.get_settings.return_value = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
|
||
|
||
mock_logging = MagicMock()
|
||
mock_logging.slack_alerting_instance.alert_types = ["budget_alerts"]
|
||
mock_logging.slack_alerting_instance._all_possible_alert_types.return_value = ["budget_alerts"]
|
||
mock_logging.slack_alerting_instance.alert_to_webhook_url = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_logging)
|
||
monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
|
||
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"
|
||
)
|
||
|
||
client = TestClient(app)
|
||
try:
|
||
response = client.get("/get/config/callbacks")
|
||
finally:
|
||
app.dependency_overrides = original_overrides
|
||
|
||
assert response.status_code == 200
|
||
slack_alert = next((a for a in response.json()["alerts"] if a["name"] == "slack"), None)
|
||
assert slack_alert is not None
|
||
masked_url = slack_alert["variables"]["SLACK_WEBHOOK_URL"]
|
||
|
||
# Masked, but derived from the real URL (None before the fix).
|
||
assert masked_url is not None
|
||
assert masked_url != webhook_url
|
||
assert masked_url.startswith("http")
|
||
assert "*" in masked_url
|
||
|
||
|
||
def test_get_config_cleared_slack_webhook_not_overridden_by_os_env(monkeypatch):
|
||
"""
|
||
A webhook the admin cleared is stored as "" in environment_variables. The
|
||
slack block must surface that empty value, not silently fall back to a
|
||
SLACK_WEBHOOK_URL still present in the OS environment (which truthiness-based
|
||
`or` would do). Only a truly absent key should trigger the os.getenv lookup.
|
||
"""
|
||
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
|
||
|
||
monkeypatch.setenv("SLACK_WEBHOOK_URL", "https://hooks.slack.com/services/STALE/OS/ENVVALUE")
|
||
config_data = {
|
||
"litellm_settings": {},
|
||
"general_settings": {"alerting": ["slack"]},
|
||
"environment_variables": {"SLACK_WEBHOOK_URL": ""},
|
||
}
|
||
|
||
mock_router = MagicMock()
|
||
mock_router.get_settings.return_value = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
|
||
|
||
mock_logging = MagicMock()
|
||
mock_logging.slack_alerting_instance.alert_types = ["budget_alerts"]
|
||
mock_logging.slack_alerting_instance._all_possible_alert_types.return_value = ["budget_alerts"]
|
||
mock_logging.slack_alerting_instance.alert_to_webhook_url = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_logging)
|
||
monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
|
||
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"
|
||
)
|
||
|
||
client = TestClient(app)
|
||
try:
|
||
response = client.get("/get/config/callbacks")
|
||
finally:
|
||
app.dependency_overrides = original_overrides
|
||
|
||
assert response.status_code == 200
|
||
slack_alert = next((a for a in response.json()["alerts"] if a["name"] == "slack"), None)
|
||
assert slack_alert is not None
|
||
assert slack_alert["variables"]["SLACK_WEBHOOK_URL"] == ""
|
||
|
||
|
||
# Mock Prisma
|
||
class MockPrisma:
|
||
def __init__(self, database_url=None, proxy_logging_obj=None, http_client=None):
|
||
self.database_url = database_url
|
||
self.proxy_logging_obj = proxy_logging_obj
|
||
self.http_client = http_client
|
||
|
||
async def connect(self):
|
||
pass
|
||
|
||
async def disconnect(self):
|
||
pass
|
||
|
||
|
||
mock_prisma = MockPrisma()
|
||
|
||
|
||
@patch(
|
||
"litellm.proxy.proxy_server.ProxyStartupEvent._setup_prisma_client",
|
||
return_value=mock_prisma,
|
||
)
|
||
@pytest.mark.asyncio
|
||
async def test_aaaproxy_startup_master_key(mock_prisma, monkeypatch, tmp_path):
|
||
"""
|
||
Test that master_key is correctly loaded from either config.yaml or environment variables
|
||
"""
|
||
import yaml
|
||
from fastapi import FastAPI
|
||
|
||
# Import happens here - this is when the module probably reads the config path
|
||
from litellm.proxy.proxy_server import proxy_startup_event
|
||
|
||
# Mock the Prisma import
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.PrismaClient", MockPrisma)
|
||
|
||
# Create test app
|
||
app = FastAPI()
|
||
|
||
# Test Case 1: Master key from config.yaml
|
||
test_master_key = "sk-12345"
|
||
test_config = {"general_settings": {"master_key": test_master_key}}
|
||
|
||
# Create a temporary config file
|
||
config_path = tmp_path / "config.yaml"
|
||
with open(config_path, "w") as f:
|
||
yaml.dump(test_config, f)
|
||
|
||
print(f"SET ENV VARIABLE - CONFIG_FILE_PATH, str(config_path): {str(config_path)}")
|
||
# Second setting of CONFIG_FILE_PATH to a different value
|
||
monkeypatch.setenv("CONFIG_FILE_PATH", str(config_path))
|
||
print(f"config_path: {config_path}")
|
||
print(f"os.getenv('CONFIG_FILE_PATH'): {os.getenv('CONFIG_FILE_PATH')}")
|
||
async with proxy_startup_event(app):
|
||
from litellm.proxy.proxy_server import master_key
|
||
|
||
assert master_key == test_master_key
|
||
|
||
# Test Case 2: Master key from environment variable
|
||
test_env_master_key = "sk-test-67890"
|
||
|
||
# Create empty config
|
||
empty_config = {"general_settings": {}}
|
||
with open(config_path, "w") as f:
|
||
yaml.dump(empty_config, f)
|
||
|
||
monkeypatch.setenv("LITELLM_MASTER_KEY", test_env_master_key)
|
||
print("test_env_master_key: {}".format(test_env_master_key))
|
||
async with proxy_startup_event(app):
|
||
from litellm.proxy.proxy_server import master_key
|
||
|
||
assert master_key == test_env_master_key
|
||
|
||
# Test Case 3: Master key with os.environ prefix
|
||
test_resolved_key = "sk-resolved-key"
|
||
test_config_with_prefix = {"general_settings": {"master_key": "os.environ/CUSTOM_MASTER_KEY"}}
|
||
|
||
# Create config with os.environ prefix
|
||
with open(config_path, "w") as f:
|
||
yaml.dump(test_config_with_prefix, f)
|
||
|
||
monkeypatch.setenv("CUSTOM_MASTER_KEY", test_resolved_key)
|
||
async with proxy_startup_event(app):
|
||
from litellm.proxy.proxy_server import master_key
|
||
|
||
assert master_key == test_resolved_key
|
||
|
||
|
||
def test_team_info_masking():
|
||
"""
|
||
Test that sensitive team information is properly masked
|
||
|
||
Ref: https://huntr.com/bounties/661b388a-44d8-4ad5-862b-4dc5b80be30a
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
# Test team object with sensitive data
|
||
team1_info = {
|
||
"success_callback": "['langfuse', 's3']",
|
||
"langfuse_secret": "secret-test-key",
|
||
"langfuse_public_key": "public-test-key",
|
||
}
|
||
|
||
with pytest.raises(Exception, match="secr\\*\\*\\*\\*\\*\\*\\*-key', 'langfuse_public_key':") as exc_info:
|
||
proxy_config._get_team_config(
|
||
team_id="test_dev",
|
||
all_teams_config=[team1_info],
|
||
)
|
||
|
||
print("Got exception: {}".format(exc_info.value))
|
||
assert "secret-test-key" not in str(exc_info.value)
|
||
assert "public-test-key" not in str(exc_info.value)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_all_team_models():
|
||
"""
|
||
Test get_all_team_models function with both "*" and specific team IDs
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import get_all_team_models
|
||
|
||
# Mock team data
|
||
mock_team1 = MagicMock()
|
||
mock_team1.model_dump.return_value = {
|
||
"team_id": "team1",
|
||
"models": ["gpt-4", "gpt-3.5-turbo"],
|
||
"team_alias": "Team 1",
|
||
}
|
||
|
||
mock_team2 = MagicMock()
|
||
mock_team2.model_dump.return_value = {
|
||
"team_id": "team2",
|
||
"models": ["claude-3", "gpt-4"],
|
||
"team_alias": "Team 2",
|
||
}
|
||
|
||
# Mock model data returned by router
|
||
mock_models_gpt4 = [
|
||
{"model_info": {"id": "gpt-4-model-1"}},
|
||
{"model_info": {"id": "gpt-4-model-2"}},
|
||
]
|
||
mock_models_gpt35 = [
|
||
{"model_info": {"id": "gpt-3.5-turbo-model-1"}},
|
||
]
|
||
mock_models_claude = [
|
||
{"model_info": {"id": "claude-3-model-1"}},
|
||
]
|
||
|
||
# Mock prisma client
|
||
mock_prisma_client = MagicMock()
|
||
mock_db = MagicMock()
|
||
mock_litellm_teamtable = MagicMock()
|
||
|
||
mock_prisma_client.db = mock_db
|
||
mock_db.litellm_teamtable = mock_litellm_teamtable
|
||
|
||
# Make find_many async
|
||
mock_litellm_teamtable.find_many = AsyncMock()
|
||
|
||
# Mock router
|
||
mock_router = MagicMock()
|
||
|
||
def mock_get_model_list(model_name, team_id=None):
|
||
if model_name == "gpt-4":
|
||
return mock_models_gpt4
|
||
elif model_name == "gpt-3.5-turbo":
|
||
return mock_models_gpt35
|
||
elif model_name == "claude-3":
|
||
return mock_models_claude
|
||
return None
|
||
|
||
mock_router.get_model_list.side_effect = mock_get_model_list
|
||
|
||
# Test Case 1: user_teams = "*" (all teams)
|
||
mock_litellm_teamtable.find_many.return_value = [mock_team1, mock_team2]
|
||
|
||
with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class:
|
||
# Configure the mock class to return proper instances
|
||
def mock_team_table_constructor(data):
|
||
mock_instance = MagicMock()
|
||
mock_instance.team_id = data["team_id"]
|
||
mock_instance.models = data["models"]
|
||
mock_instance.access_group_ids = data.get("access_group_ids")
|
||
return mock_instance
|
||
|
||
mock_team_table_class.model_validate.side_effect = mock_team_table_constructor
|
||
|
||
result = await get_all_team_models(
|
||
user_teams="*",
|
||
prisma_client=mock_prisma_client,
|
||
llm_router=mock_router,
|
||
)
|
||
|
||
# Verify find_many was called without where clause for "*"
|
||
mock_litellm_teamtable.find_many.assert_called_with()
|
||
|
||
# Verify router.get_model_list was called for each model
|
||
expected_calls = [
|
||
mock.call(model_name="gpt-4", team_id="team1"),
|
||
mock.call(model_name="gpt-3.5-turbo", team_id="team1"),
|
||
mock.call(model_name="claude-3", team_id="team2"),
|
||
mock.call(model_name="gpt-4", team_id="team2"),
|
||
]
|
||
mock_router.get_model_list.assert_has_calls(expected_calls, any_order=True)
|
||
|
||
# Test Case 2: user_teams = specific list
|
||
mock_litellm_teamtable.reset_mock()
|
||
mock_router.reset_mock()
|
||
mock_router.get_model_list.side_effect = mock_get_model_list
|
||
|
||
# Only return team1 for specific team query
|
||
mock_litellm_teamtable.find_many.return_value = [mock_team1]
|
||
|
||
with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class:
|
||
mock_team_table_class.model_validate.side_effect = mock_team_table_constructor
|
||
|
||
result = await get_all_team_models(
|
||
user_teams=["team1"],
|
||
prisma_client=mock_prisma_client,
|
||
llm_router=mock_router,
|
||
)
|
||
|
||
# Verify find_many was called with where clause for specific teams
|
||
mock_litellm_teamtable.find_many.assert_called_with(where={"team_id": {"in": ["team1"]}})
|
||
|
||
# Verify router.get_model_list was called only for team1 models
|
||
expected_calls = [
|
||
mock.call(model_name="gpt-4", team_id="team1"),
|
||
mock.call(model_name="gpt-3.5-turbo", team_id="team1"),
|
||
]
|
||
mock_router.get_model_list.assert_has_calls(expected_calls, any_order=True)
|
||
|
||
# Test Case 3: Empty teams list
|
||
mock_litellm_teamtable.reset_mock()
|
||
mock_router.reset_mock()
|
||
mock_litellm_teamtable.find_many.return_value = []
|
||
|
||
result = await get_all_team_models(
|
||
user_teams=[],
|
||
prisma_client=mock_prisma_client,
|
||
llm_router=mock_router,
|
||
)
|
||
|
||
# Verify find_many was called with empty list
|
||
mock_litellm_teamtable.find_many.assert_called_with(where={"team_id": {"in": []}})
|
||
|
||
# Should return empty list when no teams
|
||
assert result == {}
|
||
|
||
# Test Case 4: Router returns None for some models
|
||
mock_litellm_teamtable.reset_mock()
|
||
mock_router.reset_mock()
|
||
mock_litellm_teamtable.find_many.return_value = [mock_team1]
|
||
|
||
def mock_get_model_list_with_none(model_name, team_id=None):
|
||
if model_name == "gpt-4":
|
||
return mock_models_gpt4
|
||
# Return None for gpt-3.5-turbo to test None handling
|
||
return None
|
||
|
||
mock_router.get_model_list.side_effect = mock_get_model_list_with_none
|
||
|
||
with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_team_table_class:
|
||
mock_team_table_class.model_validate.side_effect = mock_team_table_constructor
|
||
|
||
result = await get_all_team_models(
|
||
user_teams=["team1"],
|
||
prisma_client=mock_prisma_client,
|
||
llm_router=mock_router,
|
||
)
|
||
|
||
# Should handle None return gracefully
|
||
assert isinstance(result, dict)
|
||
print("result: ", result)
|
||
assert result == {"gpt-4-model-1": ["team1"], "gpt-4-model-2": ["team1"]}
|
||
|
||
|
||
def test_add_team_models_to_all_models():
|
||
"""
|
||
Test add_team_models_to_all_models function
|
||
"""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_team_models_to_all_models
|
||
|
||
team_db_objects_typed = MagicMock(spec=LiteLLM_TeamTable)
|
||
team_db_objects_typed.team_id = "team1"
|
||
team_db_objects_typed.models = ["all-proxy-models"]
|
||
|
||
llm_router = MagicMock()
|
||
llm_router.get_model_list.return_value = [
|
||
{"model_info": {"id": "gpt-4-model-1", "team_id": "team2"}},
|
||
{"model_info": {"id": "gpt-4-model-2"}},
|
||
]
|
||
|
||
result = _add_team_models_to_all_models(
|
||
team_db_objects_typed=[team_db_objects_typed],
|
||
llm_router=llm_router,
|
||
)
|
||
assert result == {"gpt-4-model-2": {"team1"}}
|
||
|
||
|
||
def _make_router_with_access_groups(model_names, model_access_groups, deployments):
|
||
llm_router = MagicMock()
|
||
llm_router.get_model_names.return_value = model_names
|
||
llm_router.get_model_access_groups.return_value = model_access_groups
|
||
|
||
def get_model_list(model_name=None, team_id=None):
|
||
matched = [
|
||
deployment
|
||
for deployment in deployments
|
||
if deployment["model_name"] == model_name
|
||
and (
|
||
team_id is None
|
||
or deployment.get("model_info", {}).get("team_id") is None
|
||
or deployment.get("model_info", {}).get("team_id") == team_id
|
||
)
|
||
]
|
||
return matched or None
|
||
|
||
llm_router.get_model_list.side_effect = get_model_list
|
||
return llm_router
|
||
|
||
|
||
def test_add_team_models_to_all_models_resolves_config_access_group():
|
||
"""
|
||
LIT-4433: a CONFIG-defined access group (model_info.access_groups) named in
|
||
team.models must resolve to its member deployments' ids. The pre-fix code
|
||
passed the group name straight to get_model_list, which never matched, so the
|
||
team's /v2/model/info?include_team_models=true result was empty.
|
||
"""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_team_models_to_all_models
|
||
|
||
team = MagicMock(spec=LiteLLM_TeamTable)
|
||
team.team_id = "team-a"
|
||
team.models = ["test-access-group"]
|
||
|
||
llm_router = _make_router_with_access_groups(
|
||
model_names=["team-allowed-model-a"],
|
||
model_access_groups={"test-access-group": ["team-allowed-model-a"]},
|
||
deployments=[{"model_name": "team-allowed-model-a", "model_info": {"id": "model-a-id"}}],
|
||
)
|
||
|
||
result = _add_team_models_to_all_models(team_db_objects_typed=[team], llm_router=llm_router)
|
||
assert result == {"model-a-id": {"team-a"}}
|
||
|
||
|
||
def test_add_team_models_to_all_models_resolves_mixed_literal_and_access_group():
|
||
"""A team.models list mixing a literal model name and a config access-group
|
||
name must resolve both to their deployment ids."""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_team_models_to_all_models
|
||
|
||
team = MagicMock(spec=LiteLLM_TeamTable)
|
||
team.team_id = "team-a"
|
||
team.models = ["team-allowed-model-b", "test-access-group"]
|
||
|
||
llm_router = _make_router_with_access_groups(
|
||
model_names=["team-allowed-model-a", "team-allowed-model-b"],
|
||
model_access_groups={"test-access-group": ["team-allowed-model-a"]},
|
||
deployments=[
|
||
{"model_name": "team-allowed-model-a", "model_info": {"id": "model-a-id"}},
|
||
{"model_name": "team-allowed-model-b", "model_info": {"id": "model-b-id"}},
|
||
],
|
||
)
|
||
|
||
result = _add_team_models_to_all_models(team_db_objects_typed=[team], llm_router=llm_router)
|
||
assert result == {"model-a-id": {"team-a"}, "model-b-id": {"team-a"}}
|
||
|
||
|
||
def test_add_team_models_to_all_models_keeps_literal_model_colliding_with_group_name():
|
||
"""A team.models entry that names BOTH a deployed model and an access group
|
||
grants both at runtime, so the /v2 team map must contain the literal
|
||
deployment's id alongside the group members' ids."""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_team_models_to_all_models
|
||
|
||
team = MagicMock(spec=LiteLLM_TeamTable)
|
||
team.team_id = "team-a"
|
||
team.models = ["beta-models"]
|
||
|
||
llm_router = _make_router_with_access_groups(
|
||
model_names=["beta-models", "member-a"],
|
||
model_access_groups={"beta-models": ["member-a"]},
|
||
deployments=[
|
||
{"model_name": "beta-models", "model_info": {"id": "collision-id"}},
|
||
{"model_name": "member-a", "model_info": {"id": "member-a-id"}},
|
||
],
|
||
)
|
||
|
||
result = _add_team_models_to_all_models(team_db_objects_typed=[team], llm_router=llm_router)
|
||
assert result == {"collision-id": {"team-a"}, "member-a-id": {"team-a"}}
|
||
|
||
|
||
def test_add_team_models_to_all_models_excludes_other_access_group():
|
||
"""Only the access group named in team.models is expanded; deployments that
|
||
belong solely to a different access group must not leak into the team map."""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_team_models_to_all_models
|
||
|
||
team = MagicMock(spec=LiteLLM_TeamTable)
|
||
team.team_id = "team-a"
|
||
team.models = ["test-access-group"]
|
||
|
||
llm_router = _make_router_with_access_groups(
|
||
model_names=["team-allowed-model-a", "forbidden-model"],
|
||
model_access_groups={
|
||
"test-access-group": ["team-allowed-model-a"],
|
||
"other-access-group": ["forbidden-model"],
|
||
},
|
||
deployments=[
|
||
{"model_name": "team-allowed-model-a", "model_info": {"id": "model-a-id"}},
|
||
{"model_name": "forbidden-model", "model_info": {"id": "forbidden-id"}},
|
||
],
|
||
)
|
||
|
||
result = _add_team_models_to_all_models(team_db_objects_typed=[team], llm_router=llm_router)
|
||
assert result == {"model-a-id": {"team-a"}}
|
||
|
||
|
||
def test_add_team_models_to_all_models_excludes_other_teams_byok_with_shared_name():
|
||
"""A BYOK deployment owned by a DIFFERENT team but sharing the resolved model
|
||
name must not be added for this team. Guards the team_id filter passed to
|
||
get_model_list: dropping it would leak the other team's private deployment."""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_team_models_to_all_models
|
||
|
||
team = MagicMock(spec=LiteLLM_TeamTable)
|
||
team.team_id = "team-a"
|
||
team.models = ["test-access-group"]
|
||
|
||
llm_router = _make_router_with_access_groups(
|
||
model_names=["team-allowed-model-a"],
|
||
model_access_groups={"test-access-group": ["team-allowed-model-a"]},
|
||
deployments=[
|
||
{"model_name": "team-allowed-model-a", "model_info": {"id": "model-a-id", "team_id": "team-a"}},
|
||
{"model_name": "team-allowed-model-a", "model_info": {"id": "other-team-byok-id", "team_id": "team-b"}},
|
||
],
|
||
)
|
||
|
||
result = _add_team_models_to_all_models(team_db_objects_typed=[team], llm_router=llm_router)
|
||
assert result == {"model-a-id": {"team-a"}}
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_non_admin_all_models_returns_user_models_when_user_row_missing():
|
||
"""
|
||
Regression test: /key/generate mints keys without a LiteLLM_UserTable row, so
|
||
find_unique returns None for such a user. That miss must neither raise (a 400
|
||
here, or the AttributeError on `user_row.teams` that used to surface as a 500)
|
||
nor leak team models: the user belongs to no team, so only the models they
|
||
added themselves come back.
|
||
"""
|
||
from litellm.proxy.proxy_server import non_admin_all_models
|
||
|
||
user_added_model = {"model_name": "my-model", "model_info": {"id": "user-model-1"}}
|
||
prisma_client = MagicMock()
|
||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
||
prisma_client.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=MagicMock(created_by="ghost-user"))
|
||
|
||
llm_router = MagicMock()
|
||
llm_router.get_model_list.return_value = [
|
||
user_added_model,
|
||
{"model_name": "team-model", "model_info": {"id": "team-model-1", "team_id": "team-a"}},
|
||
]
|
||
|
||
result = await non_admin_all_models(
|
||
all_models=[user_added_model],
|
||
llm_router=llm_router,
|
||
user_api_key_dict=UserAPIKeyAuth(api_key="sk-test", user_id="ghost-user"),
|
||
prisma_client=prisma_client,
|
||
)
|
||
|
||
assert result == [user_added_model]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_apply_search_filter_matches_team_public_model_name():
|
||
"""
|
||
Regression test: team BYOK models persist an internal model_name
|
||
(e.g. `model_name_{team_id}_{uuid}`) and surface the user-facing name
|
||
via `model_info.team_public_model_name`. The /v2/model/info search
|
||
filter must match that public name so BYOK rows appear in results.
|
||
"""
|
||
from litellm.proxy.proxy_server import _apply_search_filter_to_models
|
||
|
||
byok_model = {
|
||
"model_name": "model_name_team-abc-123_4a6b8",
|
||
"litellm_params": {"model": "claude-sonnet-4-5"},
|
||
"model_info": {
|
||
"id": "byok-id-1",
|
||
"team_id": "team-abc-123",
|
||
"team_public_model_name": "team-claude-sonnet",
|
||
"db_model": True,
|
||
},
|
||
}
|
||
unrelated_model = {
|
||
"model_name": "gpt-4",
|
||
"litellm_params": {"model": "gpt-4"},
|
||
"model_info": {"id": "normal-id-1", "db_model": False},
|
||
}
|
||
|
||
# Search matching only team_public_model_name should still include BYOK
|
||
filtered, _ = await _apply_search_filter_to_models(
|
||
all_models=[byok_model, unrelated_model],
|
||
search="claude",
|
||
prisma_client=None,
|
||
proxy_config=MagicMock(),
|
||
)
|
||
filtered_ids = {m["model_info"]["id"] for m in filtered}
|
||
assert "byok-id-1" in filtered_ids
|
||
assert "normal-id-1" not in filtered_ids
|
||
|
||
# Search by internal model_name still matches as before
|
||
filtered, _ = await _apply_search_filter_to_models(
|
||
all_models=[byok_model, unrelated_model],
|
||
search="model_name_team-abc-123",
|
||
prisma_client=None,
|
||
proxy_config=MagicMock(),
|
||
)
|
||
assert [m["model_info"]["id"] for m in filtered] == ["byok-id-1"]
|
||
|
||
# Non-matching search returns nothing
|
||
filtered, _ = await _apply_search_filter_to_models(
|
||
all_models=[byok_model, unrelated_model],
|
||
search="gemini",
|
||
prisma_client=None,
|
||
proxy_config=MagicMock(),
|
||
)
|
||
assert filtered == []
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_apply_search_filter_scopes_byok_to_caller_teams():
|
||
"""
|
||
Regression test: `/v2/model/info?search=...` must not leak BYOK rows
|
||
from teams the caller is not a member of. Even with a bounded
|
||
`model_name`-contains DB query, a non-admin caller could otherwise
|
||
see other teams' BYOK rows that happen to match by internal name.
|
||
The post-fetch team scope drops those.
|
||
"""
|
||
from litellm.proxy.proxy_server import _apply_search_filter_to_models
|
||
|
||
# In-router BYOK rows: one in the caller's team, one in someone else's.
|
||
caller_team_byok = {
|
||
"model_name": "model_name_team-mine_internal",
|
||
"litellm_params": {"model": "claude-sonnet"},
|
||
"model_info": {
|
||
"id": "byok-mine",
|
||
"team_id": "team-mine",
|
||
"team_public_model_name": "claude-sonnet-prod",
|
||
"db_model": True,
|
||
},
|
||
}
|
||
other_team_byok = {
|
||
"model_name": "model_name_team-other_internal",
|
||
"litellm_params": {"model": "claude-sonnet"},
|
||
"model_info": {
|
||
"id": "byok-other",
|
||
"team_id": "team-other",
|
||
"team_public_model_name": "claude-sonnet-staging",
|
||
"db_model": True,
|
||
},
|
||
}
|
||
# Non-team row stays in the router-side result regardless of teams.
|
||
public_model = {
|
||
"model_name": "claude-public",
|
||
"litellm_params": {"model": "claude-sonnet"},
|
||
"model_info": {"id": "public-id", "db_model": False},
|
||
}
|
||
|
||
# DB-only BYOK rows fetched by the over-broad JSON branch.
|
||
db_caller_row = MagicMock()
|
||
db_caller_row.model_id = "byok-db-mine"
|
||
db_caller_row.model_name = "model_name_team-mine_db"
|
||
db_caller_row.model_info = {
|
||
"id": "byok-db-mine",
|
||
"team_id": "team-mine",
|
||
"team_public_model_name": "Claude DB Mine",
|
||
"db_model": True,
|
||
}
|
||
db_other_row = MagicMock()
|
||
db_other_row.model_id = "byok-db-other"
|
||
db_other_row.model_name = "model_name_team-other_db"
|
||
db_other_row.model_info = {
|
||
"id": "byok-db-other",
|
||
"team_id": "team-other",
|
||
"team_public_model_name": "Claude DB Other",
|
||
"db_model": True,
|
||
}
|
||
|
||
prisma_client = MagicMock()
|
||
prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=2)
|
||
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[db_caller_row, db_other_row])
|
||
caller_user_row = MagicMock()
|
||
caller_user_row.teams = ["team-mine"]
|
||
prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=caller_user_row)
|
||
|
||
proxy_config = MagicMock()
|
||
proxy_config.decrypt_model_list_from_db = lambda rows: [
|
||
{
|
||
"model_name": r.model_name,
|
||
"model_info": r.model_info,
|
||
"litellm_params": {"model": "claude-sonnet"},
|
||
}
|
||
for r in rows
|
||
]
|
||
|
||
non_admin = MagicMock(spec=UserAPIKeyAuth)
|
||
non_admin.user_role = LitellmUserRoles.INTERNAL_USER
|
||
non_admin.user_id = "user-mine"
|
||
non_admin.team_id = None
|
||
|
||
filtered, total_count = await _apply_search_filter_to_models(
|
||
all_models=[caller_team_byok, other_team_byok, public_model],
|
||
search="claude",
|
||
prisma_client=prisma_client,
|
||
proxy_config=proxy_config,
|
||
user_api_key_dict=non_admin,
|
||
)
|
||
|
||
filtered_ids = {m["model_info"]["id"] for m in filtered}
|
||
assert "byok-mine" in filtered_ids
|
||
assert "byok-db-mine" in filtered_ids
|
||
assert "public-id" in filtered_ids
|
||
assert "byok-other" not in filtered_ids, (
|
||
"router-side BYOK from another team must be dropped from search when caller doesn't belong to that team"
|
||
)
|
||
assert "byok-db-other" not in filtered_ids, (
|
||
"DB-only BYOK from another team must be dropped from search when caller doesn't belong to that team"
|
||
)
|
||
# total_count is router_models_count (2: caller_team_byok + public_model,
|
||
# other_team_byok dropped router-side) + DB count (2 from the mocked
|
||
# `count()`). The DB count is the *unscoped* match count; non-admin
|
||
# team scoping applies only to the returned page so the count can be
|
||
# over-reported, but it must never under-report (callers can paginate
|
||
# within the bound).
|
||
assert total_count == 4
|
||
|
||
# Admins keep the un-scoped view across teams.
|
||
admin = MagicMock(spec=UserAPIKeyAuth)
|
||
admin.user_role = LitellmUserRoles.PROXY_ADMIN
|
||
admin.user_id = "admin-1"
|
||
admin.team_id = None
|
||
|
||
filtered_admin, _ = await _apply_search_filter_to_models(
|
||
all_models=[caller_team_byok, other_team_byok, public_model],
|
||
search="claude",
|
||
prisma_client=prisma_client,
|
||
proxy_config=proxy_config,
|
||
user_api_key_dict=admin,
|
||
)
|
||
admin_ids = {m["model_info"]["id"] for m in filtered_admin}
|
||
assert "byok-other" in admin_ids
|
||
assert "byok-db-other" in admin_ids
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_apply_search_filter_bounds_db_fetch_by_page_and_cap():
|
||
"""
|
||
Regression test: a broad search term must not force a full BYOK-table
|
||
read + decrypt on each request.
|
||
|
||
* Unsorted searches: `find_many(take=N)` where N is just enough to
|
||
fill the current page after counting router-side matches.
|
||
* Sorted searches: `find_many(take=cap)` falls back to
|
||
`_SORTED_SEARCH_DB_FETCH_CAP` so ordering still works across a
|
||
large match set without scanning the whole table.
|
||
"""
|
||
from litellm.proxy.proxy_server import (
|
||
_SORTED_SEARCH_DB_FETCH_CAP,
|
||
_apply_search_filter_to_models,
|
||
)
|
||
|
||
prisma_client = MagicMock()
|
||
prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=10_000)
|
||
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||
|
||
proxy_config = MagicMock()
|
||
proxy_config.decrypt_model_list_from_db = lambda rows: []
|
||
|
||
# Unsorted: page=1, size=50, no router-side matches -> take must be 50.
|
||
await _apply_search_filter_to_models(
|
||
all_models=[],
|
||
search="model",
|
||
prisma_client=prisma_client,
|
||
proxy_config=proxy_config,
|
||
page=1,
|
||
size=50,
|
||
sort_by=None,
|
||
)
|
||
take = prisma_client.db.litellm_proxymodeltable.find_many.call_args.kwargs["take"]
|
||
assert take == 50, "unsorted search must take just one page's worth of rows"
|
||
|
||
# Sorted: still bounded, but by the hard cap rather than the page.
|
||
prisma_client.db.litellm_proxymodeltable.find_many.reset_mock()
|
||
await _apply_search_filter_to_models(
|
||
all_models=[],
|
||
search="model",
|
||
prisma_client=prisma_client,
|
||
proxy_config=proxy_config,
|
||
page=1,
|
||
size=50,
|
||
sort_by="model_name",
|
||
)
|
||
take = prisma_client.db.litellm_proxymodeltable.find_many.call_args.kwargs["take"]
|
||
assert take == _SORTED_SEARCH_DB_FETCH_CAP
|
||
assert take < 10_000, "sorted search must cap below the full match set"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_apply_search_filter_honours_exact_model_name_in_db_query():
|
||
"""
|
||
`/v2/model/info?model=<group>&search=<term>`: the router list is already
|
||
narrowed to the exact group, so the DB count and fetch must be too, or
|
||
other groups' rows leak into the page and inflate total_count.
|
||
"""
|
||
from litellm.proxy.proxy_server import _apply_search_filter_to_models
|
||
|
||
prisma_client = MagicMock()
|
||
prisma_client.db.litellm_proxymodeltable.count = AsyncMock(return_value=0)
|
||
prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||
proxy_config = MagicMock()
|
||
proxy_config.decrypt_model_list_from_db = lambda rows: []
|
||
|
||
await _apply_search_filter_to_models(
|
||
all_models=[],
|
||
search="sonnet",
|
||
prisma_client=prisma_client,
|
||
proxy_config=proxy_config,
|
||
model_name="anthropic-sonnet-5",
|
||
)
|
||
where = prisma_client.db.litellm_proxymodeltable.count.call_args.kwargs["where"]
|
||
assert where["model_name"] == "anthropic-sonnet-5"
|
||
assert prisma_client.db.litellm_proxymodeltable.find_many.call_args.kwargs["where"] == where
|
||
|
||
prisma_client.db.litellm_proxymodeltable.count.reset_mock()
|
||
_, total_count = await _apply_search_filter_to_models(
|
||
all_models=[],
|
||
search="opus",
|
||
prisma_client=prisma_client,
|
||
proxy_config=proxy_config,
|
||
model_name="anthropic-sonnet-5",
|
||
)
|
||
prisma_client.db.litellm_proxymodeltable.count.assert_not_called()
|
||
assert total_count == 0
|
||
|
||
await _apply_search_filter_to_models(
|
||
all_models=[],
|
||
search="sonnet",
|
||
prisma_client=prisma_client,
|
||
proxy_config=proxy_config,
|
||
)
|
||
where = prisma_client.db.litellm_proxymodeltable.count.call_args.kwargs["where"]
|
||
assert where["model_name"] == {"contains": "sonnet", "mode": "insensitive"}
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_filter_models_by_team_id_excludes_viewer_direct_access():
|
||
"""
|
||
Regression test: when the UI picks a specific team in the Current Team
|
||
selector, the model list must show only that team's BYOK rows + the
|
||
models assigned to the team. The admin viewer's `direct_access` flag
|
||
(set on every non-team model upstream) must NOT widen the team's
|
||
visible set, or selecting team-111 still shows every public model.
|
||
"""
|
||
from litellm.proxy.proxy_server import _filter_models_by_team_id
|
||
|
||
public_model = {
|
||
"model_name": "gpt-4",
|
||
"litellm_params": {"model": "gpt-4"},
|
||
"model_info": {
|
||
"id": "public-id",
|
||
# admin viewer has direct_access on this public model
|
||
"direct_access": True,
|
||
# team-111 is NOT in access_via_team_ids -> shouldn't show for team-111
|
||
"access_via_team_ids": ["team-222"],
|
||
},
|
||
}
|
||
team111_byok = {
|
||
"model_name": "model_name_team-111_uuid",
|
||
"litellm_params": {"model": "claude-sonnet"},
|
||
"model_info": {
|
||
"id": "byok-team-111",
|
||
"team_id": "team-111",
|
||
"team_public_model_name": "team-claude",
|
||
"access_via_team_ids": ["team-111"],
|
||
},
|
||
}
|
||
team222_byok = {
|
||
"model_name": "model_name_team-222_uuid",
|
||
"litellm_params": {"model": "claude-haiku"},
|
||
"model_info": {
|
||
"id": "byok-team-222",
|
||
"team_id": "team-222",
|
||
"team_public_model_name": "team-haiku",
|
||
"access_via_team_ids": ["team-222"],
|
||
},
|
||
}
|
||
|
||
prisma = MagicMock()
|
||
team_db = MagicMock()
|
||
team_db.model_dump.return_value = {
|
||
"team_id": "team-111",
|
||
"team_alias": "Team 111",
|
||
# specific models list that doesn't include the BYOK's internal name
|
||
"models": ["some-other-model"],
|
||
"access_group_ids": None,
|
||
}
|
||
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db)
|
||
prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||
|
||
router = MagicMock()
|
||
router.get_model_access_groups = MagicMock(return_value={})
|
||
# team-111 only resolves "some-other-model", which has no deployments
|
||
router.get_model_list = MagicMock(return_value=[])
|
||
|
||
filtered = await _filter_models_by_team_id(
|
||
all_models=[public_model, team111_byok, team222_byok],
|
||
team_id="team-111",
|
||
prisma_client=prisma,
|
||
llm_router=router,
|
||
)
|
||
visible_ids = sorted(m["model_info"]["id"] for m in filtered)
|
||
|
||
assert "byok-team-111" in visible_ids, "team-111's own BYOK must always be visible"
|
||
assert "byok-team-222" not in visible_ids, "must not leak other teams' BYOK"
|
||
assert "public-id" not in visible_ids, "viewer's direct_access must not widen the team's visible set"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_filter_models_by_team_id_rejects_non_member():
|
||
"""
|
||
Regression test: /v2/model/info?teamId=X includes BYOK rows solely on
|
||
`model_info.team_id == X`. Without an auth check, any authenticated user
|
||
could enumerate another team's BYOK metadata by guessing its id. Callers
|
||
that are neither proxy admins nor members of `team_id` must get 403.
|
||
"""
|
||
from fastapi import HTTPException
|
||
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import _filter_models_by_team_id
|
||
|
||
byok = {
|
||
"model_name": "model_name_team-111_uuid",
|
||
"litellm_params": {"model": "claude"},
|
||
"model_info": {"id": "byok-team-111", "team_id": "team-111"},
|
||
}
|
||
|
||
prisma = MagicMock()
|
||
# Caller is in team-222 only
|
||
user_row = MagicMock()
|
||
user_row.teams = ["team-222"]
|
||
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
|
||
|
||
caller = UserAPIKeyAuth(
|
||
user_id="alice",
|
||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||
api_key="sk-test",
|
||
)
|
||
|
||
with pytest.raises(HTTPException) as excinfo:
|
||
await _filter_models_by_team_id(
|
||
all_models=[byok],
|
||
team_id="team-111",
|
||
prisma_client=prisma,
|
||
llm_router=MagicMock(),
|
||
user_api_key_dict=caller,
|
||
)
|
||
assert excinfo.value.status_code == 403
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_filter_models_by_team_id_allows_team_member():
|
||
"""
|
||
A caller who IS a member of `team_id` must be allowed to filter, and
|
||
should see that team's BYOK rows.
|
||
"""
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import _filter_models_by_team_id
|
||
|
||
byok = {
|
||
"model_name": "model_name_team-111_uuid",
|
||
"litellm_params": {"model": "claude"},
|
||
"model_info": {"id": "byok-team-111", "team_id": "team-111"},
|
||
}
|
||
|
||
prisma = MagicMock()
|
||
user_row = MagicMock()
|
||
user_row.teams = ["team-111", "team-999"]
|
||
prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
|
||
team_db = MagicMock()
|
||
team_db.model_dump.return_value = {
|
||
"team_id": "team-111",
|
||
"team_alias": "Team 111",
|
||
"models": [],
|
||
"access_group_ids": None,
|
||
}
|
||
prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_db)
|
||
prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||
|
||
router = MagicMock()
|
||
router.get_model_access_groups = MagicMock(return_value={})
|
||
router.get_model_list = MagicMock(return_value=[byok])
|
||
|
||
caller = UserAPIKeyAuth(
|
||
user_id="bob",
|
||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||
api_key="sk-test",
|
||
)
|
||
|
||
result = await _filter_models_by_team_id(
|
||
all_models=[byok],
|
||
team_id="team-111",
|
||
prisma_client=prisma,
|
||
llm_router=router,
|
||
user_api_key_dict=caller,
|
||
)
|
||
assert [m["model_info"]["id"] for m in result] == ["byok-team-111"]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_caller_byok_team_scope_treats_view_only_admin_as_unscoped():
|
||
"""
|
||
Regression test: `PROXY_ADMIN_VIEW_ONLY` is an admin role
|
||
("can login, view all own keys, view all spend"). Search results for
|
||
this role must show BYOK rows across all teams, not be silently scoped
|
||
to the user-id's `teams` field — that path narrows results to whatever
|
||
teams the admin happens to be a member of, regressing pre-PR behavior.
|
||
"""
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import _get_caller_byok_team_scope
|
||
|
||
caller = UserAPIKeyAuth(
|
||
user_id="view-admin",
|
||
user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY,
|
||
api_key="sk-test",
|
||
)
|
||
scope = await _get_caller_byok_team_scope(
|
||
user_api_key_dict=caller,
|
||
prisma_client=MagicMock(),
|
||
)
|
||
assert scope is None, "PROXY_ADMIN_VIEW_ONLY must be unscoped, like PROXY_ADMIN"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_access_group_models_to_team_models():
|
||
"""
|
||
Test that models reachable via team access groups are included in team_models.
|
||
|
||
Scenario: A team has models=["gpt-4"] and access_group_ids=["premium"].
|
||
The "premium" access group contains ["claude-3", "gemini"].
|
||
After resolution, the team should see gpt-4 (direct) + claude-3/gemini (via access group).
|
||
"""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_access_group_models_to_team_models
|
||
|
||
# Team with specific models AND access groups
|
||
team_with_access_groups = MagicMock(spec=LiteLLM_TeamTable)
|
||
team_with_access_groups.team_id = "team1"
|
||
team_with_access_groups.models = ["gpt-4"] # non-empty = specific models
|
||
team_with_access_groups.access_group_ids = ["premium"]
|
||
|
||
# Team with no access groups — should be skipped
|
||
team_without_access_groups = MagicMock(spec=LiteLLM_TeamTable)
|
||
team_without_access_groups.team_id = "team2"
|
||
team_without_access_groups.models = ["gpt-4"]
|
||
team_without_access_groups.access_group_ids = None
|
||
|
||
# Team with empty access_group_ids list — should be skipped
|
||
team_empty_access_groups = MagicMock(spec=LiteLLM_TeamTable)
|
||
team_empty_access_groups.team_id = "team2b"
|
||
team_empty_access_groups.models = ["gpt-4"]
|
||
team_empty_access_groups.access_group_ids = []
|
||
|
||
# Team with empty models (all access) — should be skipped
|
||
team_all_access = MagicMock(spec=LiteLLM_TeamTable)
|
||
team_all_access.team_id = "team3"
|
||
team_all_access.models = []
|
||
team_all_access.access_group_ids = ["premium"]
|
||
|
||
# Team with all-proxy-models sentinel (all access) — should be skipped
|
||
team_all_proxy = MagicMock(spec=LiteLLM_TeamTable)
|
||
team_all_proxy.team_id = "team4"
|
||
team_all_proxy.models = ["all-proxy-models"]
|
||
team_all_proxy.access_group_ids = ["premium"]
|
||
|
||
# Mock router
|
||
mock_router = MagicMock()
|
||
|
||
def mock_get_model_list(model_name, team_id=None):
|
||
if model_name == "claude-3":
|
||
return [{"model_info": {"id": "claude-3-id"}}]
|
||
elif model_name == "gemini":
|
||
return [{"model_info": {"id": "gemini-id"}}]
|
||
return None
|
||
|
||
mock_router.get_model_list.side_effect = mock_get_model_list
|
||
|
||
# Pre-existing team_models (e.g., from _add_team_models_to_all_models)
|
||
existing_team_models = {
|
||
"gpt-4-id": {"team1"},
|
||
}
|
||
|
||
# Mock prisma client with batch find_many returning access group rows
|
||
mock_ag_row = MagicMock()
|
||
mock_ag_row.access_group_id = "premium"
|
||
mock_ag_row.access_model_names = ["claude-3", "gemini"]
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_ag_row])
|
||
|
||
result = await _add_access_group_models_to_team_models(
|
||
team_db_objects_typed=[
|
||
team_with_access_groups,
|
||
team_without_access_groups,
|
||
team_empty_access_groups,
|
||
team_all_access,
|
||
team_all_proxy,
|
||
],
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
team_models=existing_team_models,
|
||
)
|
||
|
||
# Single batch query with only the eligible team's access group IDs
|
||
mock_prisma_client.db.litellm_accessgrouptable.find_many.assert_called_once()
|
||
call_args = mock_prisma_client.db.litellm_accessgrouptable.find_many.call_args
|
||
queried_ids = call_args[1]["where"]["access_group_id"]["in"]
|
||
assert set(queried_ids) == {"premium"}
|
||
|
||
# Original model still present
|
||
assert "gpt-4-id" in result
|
||
assert "team1" in result["gpt-4-id"]
|
||
|
||
# Access group models added for team1
|
||
assert "claude-3-id" in result
|
||
assert "team1" in result["claude-3-id"]
|
||
assert "gemini-id" in result
|
||
assert "team1" in result["gemini-id"]
|
||
|
||
# Skipped teams should NOT have added these models
|
||
for skipped_team in ["team2", "team2b", "team3", "team4"]:
|
||
assert skipped_team not in result.get("claude-3-id", set())
|
||
assert skipped_team not in result.get("gemini-id", set())
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_access_group_models_multiple_teams_shared_group():
|
||
"""
|
||
Test that multiple teams sharing the same access group each get the models,
|
||
and only one batch DB query is made.
|
||
"""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_access_group_models_to_team_models
|
||
|
||
team_a = MagicMock(spec=LiteLLM_TeamTable)
|
||
team_a.team_id = "team-a"
|
||
team_a.models = ["gpt-4"]
|
||
team_a.access_group_ids = ["shared-group"]
|
||
|
||
team_b = MagicMock(spec=LiteLLM_TeamTable)
|
||
team_b.team_id = "team-b"
|
||
team_b.models = ["gpt-3.5"]
|
||
team_b.access_group_ids = ["shared-group", "extra-group"]
|
||
|
||
mock_router = MagicMock()
|
||
|
||
def mock_get_model_list(model_name, team_id=None):
|
||
if model_name == "claude-3":
|
||
return [{"model_info": {"id": "claude-3-id"}}]
|
||
elif model_name == "gemini":
|
||
return [{"model_info": {"id": "gemini-id"}}]
|
||
return None
|
||
|
||
mock_router.get_model_list.side_effect = mock_get_model_list
|
||
|
||
mock_shared_row = MagicMock()
|
||
mock_shared_row.access_group_id = "shared-group"
|
||
mock_shared_row.access_model_names = ["claude-3"]
|
||
|
||
mock_extra_row = MagicMock()
|
||
mock_extra_row.access_group_id = "extra-group"
|
||
mock_extra_row.access_model_names = ["gemini"]
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_shared_row, mock_extra_row])
|
||
|
||
result = await _add_access_group_models_to_team_models(
|
||
team_db_objects_typed=[team_a, team_b],
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
team_models={},
|
||
)
|
||
|
||
# Single batch query for both groups
|
||
mock_prisma_client.db.litellm_accessgrouptable.find_many.assert_called_once()
|
||
call_args = mock_prisma_client.db.litellm_accessgrouptable.find_many.call_args
|
||
queried_ids = set(call_args[1]["where"]["access_group_id"]["in"])
|
||
assert queried_ids == {"shared-group", "extra-group"}
|
||
|
||
# Both teams get claude-3 from the shared group
|
||
assert "claude-3-id" in result
|
||
assert "team-a" in result["claude-3-id"]
|
||
assert "team-b" in result["claude-3-id"]
|
||
|
||
# Only team-b gets gemini (from extra-group)
|
||
assert "gemini-id" in result
|
||
assert "team-b" in result["gemini-id"]
|
||
assert "team-a" not in result["gemini-id"]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_access_group_models_no_eligible_teams():
|
||
"""
|
||
When no teams have access groups, find_many should not be called at all.
|
||
"""
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
from litellm.proxy.proxy_server import _add_access_group_models_to_team_models
|
||
|
||
team = MagicMock(spec=LiteLLM_TeamTable)
|
||
team.team_id = "team1"
|
||
team.models = ["gpt-4"]
|
||
team.access_group_ids = None
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_accessgrouptable.find_many = AsyncMock()
|
||
|
||
result = await _add_access_group_models_to_team_models(
|
||
team_db_objects_typed=[team],
|
||
llm_router=MagicMock(),
|
||
prisma_client=mock_prisma_client,
|
||
team_models={"existing-id": {"team1"}},
|
||
)
|
||
|
||
# No DB call made
|
||
mock_prisma_client.db.litellm_accessgrouptable.find_many.assert_not_called()
|
||
|
||
# Original data unchanged
|
||
assert result == {"existing-id": {"team1"}}
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_all_team_models_with_access_groups():
|
||
"""
|
||
End-to-end test: get_all_team_models includes models from access groups.
|
||
|
||
Scenario: User is on team1 which has models=["gpt-4"] and
|
||
access_group_ids=["premium"]. The "premium" group has ["claude-3"].
|
||
The result should include both gpt-4 and claude-3 deployments for team1.
|
||
"""
|
||
from litellm.proxy.proxy_server import get_all_team_models
|
||
|
||
mock_team1 = MagicMock()
|
||
mock_team1.model_dump.return_value = {
|
||
"team_id": "team1",
|
||
"models": ["gpt-4"],
|
||
"team_alias": "Team 1",
|
||
"access_group_ids": ["premium"],
|
||
}
|
||
|
||
# Mock access group row returned by batch find_many
|
||
mock_ag_row = MagicMock()
|
||
mock_ag_row.access_group_id = "premium"
|
||
mock_ag_row.access_model_names = ["claude-3"]
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_db = MagicMock()
|
||
mock_litellm_teamtable = MagicMock()
|
||
mock_prisma_client.db = mock_db
|
||
mock_db.litellm_teamtable = mock_litellm_teamtable
|
||
mock_litellm_teamtable.find_many = AsyncMock(return_value=[mock_team1])
|
||
mock_db.litellm_accessgrouptable = MagicMock()
|
||
mock_db.litellm_accessgrouptable.find_many = AsyncMock(return_value=[mock_ag_row])
|
||
|
||
mock_router = MagicMock()
|
||
|
||
def mock_get_model_list(model_name, team_id=None):
|
||
if model_name == "gpt-4":
|
||
return [{"model_info": {"id": "gpt-4-deploy-1"}}]
|
||
elif model_name == "claude-3":
|
||
return [{"model_info": {"id": "claude-3-deploy-1"}}]
|
||
return None
|
||
|
||
mock_router.get_model_list.side_effect = mock_get_model_list
|
||
|
||
with patch("litellm.proxy.proxy_server.LiteLLM_TeamTable") as mock_tt_class:
|
||
|
||
def mock_team_table_constructor(data):
|
||
mock_instance = MagicMock()
|
||
mock_instance.team_id = data["team_id"]
|
||
mock_instance.models = data["models"]
|
||
mock_instance.access_group_ids = data.get("access_group_ids")
|
||
return mock_instance
|
||
|
||
mock_tt_class.model_validate.side_effect = mock_team_table_constructor
|
||
|
||
result = await get_all_team_models(
|
||
user_teams=["team1"],
|
||
prisma_client=mock_prisma_client,
|
||
llm_router=mock_router,
|
||
)
|
||
|
||
# gpt-4 from direct team.models
|
||
assert "gpt-4-deploy-1" in result
|
||
assert "team1" in result["gpt-4-deploy-1"]
|
||
|
||
# claude-3 from access group
|
||
assert "claude-3-deploy-1" in result
|
||
assert "team1" in result["claude-3-deploy-1"]
|
||
|
||
# Return type is Dict[str, List[str]]
|
||
assert isinstance(result["gpt-4-deploy-1"], list)
|
||
assert isinstance(result["claude-3-deploy-1"], list)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_delete_deployment_type_mismatch():
|
||
"""
|
||
Test that the _delete_deployment function handles type mismatches correctly.
|
||
Specifically test that models 12345678 and 12345679 are NOT deleted when
|
||
they exist in both combined_id_list (as integers) and router_model_ids (as strings).
|
||
|
||
This test reproduces the bug where type mismatch causes valid models to be deleted.
|
||
"""
|
||
from unittest.mock import MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
# Create mock ProxyConfig instance
|
||
pc = ProxyConfig()
|
||
|
||
# Mock llm_router with string IDs (this is the source of the type mismatch)
|
||
mock_llm_router = MagicMock()
|
||
mock_llm_router.get_model_ids.return_value = [
|
||
"a96e12e76b36a57cfae57a41288eb41567629cac89b4828c6f7074afc3534695",
|
||
"a40186dd0fdb9b7282380277d7f57044d29de95bfbfcd7f4322b3493702d5cd3",
|
||
"12345678", # String ID
|
||
"12345679", # String ID
|
||
]
|
||
|
||
# Track which deployments were deleted
|
||
deleted_ids = []
|
||
|
||
def mock_delete_deployment(id):
|
||
deleted_ids.append(id)
|
||
return True # Simulate successful deletion
|
||
|
||
mock_llm_router.delete_deployment = MagicMock(side_effect=mock_delete_deployment)
|
||
|
||
async def mock_get_config(config_file_path):
|
||
return {
|
||
"model_list": [
|
||
{
|
||
"model_name": "openai-gpt-4o",
|
||
"litellm_params": {"model": "gpt-4o"},
|
||
"model_info": {"id": 12345678},
|
||
},
|
||
{
|
||
"model_name": "openai-gpt-4o",
|
||
"litellm_params": {"model": "gpt-4o"},
|
||
"model_info": {"id": 12345679},
|
||
},
|
||
]
|
||
}
|
||
|
||
pc.get_config = AsyncMock(side_effect=mock_get_config)
|
||
|
||
# Patch the global llm_router
|
||
with (
|
||
patch("litellm.proxy.proxy_server.llm_router", mock_llm_router),
|
||
patch("litellm.proxy.proxy_server.user_config_file_path", "test_config.yaml"),
|
||
):
|
||
# Call the function under test
|
||
still_desired = await pc._delete_deployment(db_models=[])
|
||
|
||
# The two SHA-hash models have no corresponding entry in combined_id_list
|
||
# and must be evicted.
|
||
assert len(deleted_ids) == 2, f"Expected 2 deletions (SHA-hash models), got {deleted_ids}"
|
||
assert "a96e12e76b36a57cfae57a41288eb41567629cac89b4828c6f7074afc3534695" in deleted_ids
|
||
assert "a40186dd0fdb9b7282380277d7f57044d29de95bfbfcd7f4322b3493702d5cd3" in deleted_ids
|
||
|
||
# Models 12345678 and 12345679 exist in the config (as integers); str()
|
||
# conversion in _delete_deployment makes them match the router's string IDs,
|
||
# so they must NOT be evicted.
|
||
assert "12345678" not in deleted_ids, f"Model 12345678 should NOT be deleted. Deleted IDs: {deleted_ids}"
|
||
assert "12345679" not in deleted_ids, f"Model 12345679 should NOT be deleted. Deleted IDs: {deleted_ids}"
|
||
|
||
assert still_desired is not None
|
||
assert {"12345678", "12345679"} <= still_desired, (
|
||
"the int-keyed config models must come back as strings in the desired set, so a "
|
||
f"caller judging its own reload reads them as wanted rather than evicted; got {still_desired}"
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_config_from_file(tmp_path, monkeypatch):
|
||
"""
|
||
Test the _get_config_from_file method of ProxyConfig class.
|
||
Tests various scenarios: valid file, non-existent file, no file path, None config.
|
||
"""
|
||
import yaml
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
# Create a ProxyConfig instance
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Test Case 1: Valid YAML config file exists
|
||
test_config = {
|
||
"model_list": [{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4"}}],
|
||
"general_settings": {"master_key": "sk-test"},
|
||
"router_settings": {"enable_pre_call_checks": True},
|
||
"litellm_settings": {"drop_params": True},
|
||
}
|
||
|
||
config_file = tmp_path / "test_config.yaml"
|
||
with open(config_file, "w") as f:
|
||
yaml.dump(test_config, f)
|
||
|
||
# Clear global user_config_file_path for this test
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", None)
|
||
|
||
result = await proxy_config._get_config_from_file(str(config_file))
|
||
assert result == test_config
|
||
|
||
# Verify that user_config_file_path was set
|
||
from litellm.proxy.proxy_server import user_config_file_path
|
||
|
||
assert user_config_file_path == str(config_file)
|
||
|
||
# Test Case 2: File path provided but file doesn't exist
|
||
non_existent_file = tmp_path / "non_existent.yaml"
|
||
|
||
with pytest.raises(Exception, match=f"Config file not found: {non_existent_file}"):
|
||
await proxy_config._get_config_from_file(str(non_existent_file))
|
||
|
||
# Test Case 3: No file path provided (should return default config)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", None)
|
||
|
||
expected_default = {
|
||
"model_list": [],
|
||
"general_settings": {},
|
||
"router_settings": {},
|
||
"litellm_settings": {},
|
||
}
|
||
|
||
result = await proxy_config._get_config_from_file(None)
|
||
assert result == expected_default
|
||
|
||
# Test Case 4: Empty YAML file (should raise exception for None config)
|
||
empty_file = tmp_path / "empty_config.yaml"
|
||
with open(empty_file, "w") as f:
|
||
f.write("") # Write empty content which will result in None when loaded
|
||
|
||
with pytest.raises(Exception, match=re.escape("Config cannot be None or Empty.")):
|
||
await proxy_config._get_config_from_file(str(empty_file))
|
||
|
||
# Test Case 5: Using global user_config_file_path when no config_file_path provided
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", str(config_file))
|
||
|
||
result = await proxy_config._get_config_from_file(None)
|
||
assert result == test_config
|
||
|
||
|
||
def test_normalize_datetime_for_sorting():
|
||
"""
|
||
Test the _normalize_datetime_for_sorting function.
|
||
Tests various scenarios: None values, ISO format strings, datetime objects (naive and aware).
|
||
"""
|
||
from litellm.proxy.proxy_server import _normalize_datetime_for_sorting
|
||
|
||
# Test Case 1: None value
|
||
assert _normalize_datetime_for_sorting(None) is None
|
||
|
||
# Test Case 2: ISO format string with 'Z' suffix
|
||
dt_str_z = "2024-01-15T10:30:00Z"
|
||
result = _normalize_datetime_for_sorting(dt_str_z)
|
||
assert result is not None
|
||
assert isinstance(result, datetime)
|
||
assert result.tzinfo == timezone.utc
|
||
assert result.year == 2024
|
||
assert result.month == 1
|
||
assert result.day == 15
|
||
assert result.hour == 10
|
||
assert result.minute == 30
|
||
|
||
# Test Case 3: ISO format string without 'Z' suffix (naive)
|
||
dt_str_naive = "2024-01-15T10:30:00"
|
||
result = _normalize_datetime_for_sorting(dt_str_naive)
|
||
assert result is not None
|
||
assert isinstance(result, datetime)
|
||
assert result.tzinfo == timezone.utc
|
||
|
||
# Test Case 4: ISO format string with timezone offset
|
||
dt_str_tz = "2024-01-15T10:30:00+05:00"
|
||
result = _normalize_datetime_for_sorting(dt_str_tz)
|
||
assert result is not None
|
||
assert isinstance(result, datetime)
|
||
assert result.tzinfo == timezone.utc
|
||
# Should convert from +05:00 to UTC (subtract 5 hours)
|
||
assert result.hour == 5 # 10:30 - 5 hours = 5:30 UTC
|
||
|
||
# Test Case 5: Naive datetime object
|
||
naive_dt = datetime(2024, 1, 15, 10, 30, 0)
|
||
result = _normalize_datetime_for_sorting(naive_dt)
|
||
assert result is not None
|
||
assert isinstance(result, datetime)
|
||
assert result.tzinfo == timezone.utc
|
||
assert result.year == 2024
|
||
assert result.month == 1
|
||
assert result.day == 15
|
||
|
||
# Test Case 6: Timezone-aware datetime object (non-UTC)
|
||
from datetime import timedelta
|
||
|
||
aware_dt = datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone(timedelta(hours=5)))
|
||
result = _normalize_datetime_for_sorting(aware_dt)
|
||
assert result is not None
|
||
assert isinstance(result, datetime)
|
||
assert result.tzinfo == timezone.utc
|
||
# Should convert from +05:00 to UTC
|
||
assert result.hour == 5
|
||
|
||
# Test Case 7: UTC-aware datetime object
|
||
utc_dt = datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone.utc)
|
||
result = _normalize_datetime_for_sorting(utc_dt)
|
||
assert result is not None
|
||
assert isinstance(result, datetime)
|
||
assert result.tzinfo == timezone.utc
|
||
assert result == utc_dt
|
||
|
||
# Test Case 8: Invalid string format
|
||
invalid_str = "not-a-date"
|
||
result = _normalize_datetime_for_sorting(invalid_str)
|
||
assert result is None
|
||
|
||
# Test Case 9: Invalid type (should return None)
|
||
result = _normalize_datetime_for_sorting(12345)
|
||
assert result is None
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_proxy_budget_to_db_only_creates_user_no_keys():
|
||
"""
|
||
Test that _add_proxy_budget_to_db only creates a user and no keys are added.
|
||
|
||
This validates that generate_key_helper_fn is called with table_name="user"
|
||
which should prevent key creation in LiteLLM_VerificationToken table.
|
||
|
||
Also guards the row identity: the budget must land on the proxy-wide
|
||
aggregate row "litellm-proxy-budget" (the one the spend writer increments
|
||
per request), not the admin user's own row ("default_user_id"). Budgeting
|
||
the admin row leaves the global budget without a resettable counter.
|
||
"""
|
||
from unittest.mock import AsyncMock, patch
|
||
|
||
import litellm
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
|
||
# Set up required litellm settings
|
||
litellm.budget_duration = "30d"
|
||
litellm.max_budget = 100.0
|
||
|
||
litellm_proxy_budget_name = "litellm-proxy-budget"
|
||
|
||
# Mock generate_key_helper_fn to capture its call arguments
|
||
mock_generate_key_helper = AsyncMock(
|
||
return_value={
|
||
"user_id": litellm_proxy_budget_name,
|
||
"max_budget": 100.0,
|
||
"budget_duration": "30d",
|
||
"spend": 0,
|
||
"models": [],
|
||
}
|
||
)
|
||
|
||
# Patch generate_key_helper_fn in proxy_server where it's being called from
|
||
with patch("litellm.proxy.proxy_server.generate_key_helper_fn", mock_generate_key_helper):
|
||
# Call the function under test
|
||
ProxyStartupEvent._add_proxy_budget_to_db()
|
||
|
||
# Allow async task to complete
|
||
import asyncio
|
||
|
||
await asyncio.sleep(0.1)
|
||
|
||
# Verify that generate_key_helper_fn was called
|
||
mock_generate_key_helper.assert_called_once()
|
||
call_args = mock_generate_key_helper.call_args
|
||
|
||
# Verify critical parameters that prevent key creation
|
||
assert call_args.kwargs["request_type"] == "user"
|
||
assert call_args.kwargs["table_name"] == "user"
|
||
assert call_args.kwargs["user_id"] == litellm_proxy_budget_name
|
||
assert call_args.kwargs["max_budget"] == 100.0
|
||
assert call_args.kwargs["budget_duration"] == "30d"
|
||
assert call_args.kwargs["query_type"] == "update_data"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_proxy_budget_to_db_backfills_budget_reset_at():
|
||
"""
|
||
Test that _upsert_proxy_budget_with_reset_at_backfill issues a conditional
|
||
update_many with `WHERE budget_reset_at IS NULL` to backfill the column on
|
||
rows that pre-existed without a reset schedule. Without this, the proxy
|
||
budget row stays at NULL and reset_budget_for_litellm_users never matches
|
||
it (NULL < now() is unknown in SQL), so the global proxy budget never
|
||
resets.
|
||
|
||
The same conditional update must zero spend: a row that was never on a
|
||
reset schedule holds lifetime accrual, which must not gate the first
|
||
duration window.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
import litellm
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
|
||
litellm.budget_duration = "30d"
|
||
litellm.max_budget = 100.0
|
||
litellm_proxy_budget_name = "litellm-proxy-budget"
|
||
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_usertable.update_many = AsyncMock(return_value={"count": 1})
|
||
|
||
mock_generate_key_helper = AsyncMock(
|
||
return_value={
|
||
"user_id": litellm_proxy_budget_name,
|
||
"max_budget": 100.0,
|
||
"budget_duration": "30d",
|
||
"spend": 0,
|
||
"models": [],
|
||
}
|
||
)
|
||
|
||
with (
|
||
patch(
|
||
"litellm.proxy.proxy_server.generate_key_helper_fn",
|
||
mock_generate_key_helper,
|
||
),
|
||
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma),
|
||
):
|
||
await ProxyStartupEvent._upsert_proxy_budget_with_reset_at_backfill()
|
||
|
||
# Upsert ran with the configured budget
|
||
mock_generate_key_helper.assert_called_once()
|
||
|
||
# Backfill update_many ran with the conditional WHERE
|
||
mock_prisma.db.litellm_usertable.update_many.assert_called_once()
|
||
backfill_call = mock_prisma.db.litellm_usertable.update_many.call_args
|
||
assert backfill_call.kwargs["where"]["user_id"] == litellm_proxy_budget_name
|
||
assert backfill_call.kwargs["where"]["budget_reset_at"] is None
|
||
|
||
# The backfilled value must be a real future datetime — anything else and
|
||
# reset_budget_for_litellm_users would still skip the row.
|
||
from datetime import datetime, timezone
|
||
|
||
backfilled_reset_at = backfill_call.kwargs["data"]["budget_reset_at"]
|
||
assert isinstance(backfilled_reset_at, datetime)
|
||
assert backfilled_reset_at > datetime.now(timezone.utc)
|
||
|
||
assert backfill_call.kwargs["data"]["spend"] == 0
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_custom_ui_sso_sign_in_handler_config_loading():
|
||
"""
|
||
Test that custom_ui_sso_sign_in_handler from config gets properly loaded into the global variable
|
||
"""
|
||
import tempfile
|
||
from unittest.mock import MagicMock, patch
|
||
|
||
import yaml
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
# Create a test config with custom_ui_sso_sign_in_handler
|
||
test_config = {
|
||
"general_settings": {
|
||
"custom_ui_sso_sign_in_handler": "custom_hooks.custom_ui_sso_hook.custom_ui_sso_sign_in_handler"
|
||
},
|
||
"model_list": [],
|
||
"router_settings": {},
|
||
"litellm_settings": {},
|
||
}
|
||
|
||
# Create temporary config file
|
||
with tempfile.NamedTemporaryFile(mode="w", suffix=".yaml", delete=False) as f:
|
||
yaml.dump(test_config, f)
|
||
config_file_path = f.name
|
||
|
||
# Mock the get_instance_fn to return a mock handler
|
||
mock_custom_handler = MagicMock()
|
||
|
||
try:
|
||
with patch(
|
||
"litellm.proxy.proxy_server.get_instance_fn",
|
||
return_value=mock_custom_handler,
|
||
) as mock_get_instance:
|
||
# Create ProxyConfig instance and load config
|
||
proxy_config = ProxyConfig()
|
||
# Create a mock router since load_config requires it
|
||
mock_router = MagicMock()
|
||
await proxy_config.load_config(router=mock_router, config_file_path=config_file_path)
|
||
|
||
# Verify get_instance_fn was called with correct parameters
|
||
mock_get_instance.assert_called_with(
|
||
value="custom_hooks.custom_ui_sso_hook.custom_ui_sso_sign_in_handler",
|
||
config_file_path=config_file_path,
|
||
)
|
||
|
||
# Verify the global variable was set
|
||
from litellm.proxy.proxy_server import user_custom_ui_sso_sign_in_handler
|
||
|
||
assert user_custom_ui_sso_sign_in_handler == mock_custom_handler
|
||
|
||
finally:
|
||
# Clean up temporary file
|
||
import os
|
||
|
||
os.unlink(config_file_path)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_config_max_budget_env_var_coerced_to_float(tmp_path, monkeypatch):
|
||
"""
|
||
max_budget configured as os.environ/MAX_BUDGET resolves to a string;
|
||
load_config must coerce it to float so the startup check
|
||
`litellm.max_budget > 0` doesn't raise TypeError.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setenv("MAX_BUDGET", "10")
|
||
test_config = {
|
||
"model_list": [],
|
||
"litellm_settings": {"max_budget": "os.environ/MAX_BUDGET"},
|
||
}
|
||
config_file = tmp_path / "config.yaml"
|
||
config_file.write_text(yaml.dump(test_config))
|
||
|
||
original_max_budget = litellm.max_budget
|
||
try:
|
||
proxy_config = ProxyConfig()
|
||
await proxy_config.load_config(router=MagicMock(), config_file_path=str(config_file))
|
||
assert isinstance(litellm.max_budget, float)
|
||
assert litellm.max_budget == 10.0
|
||
assert litellm.max_budget > 0
|
||
finally:
|
||
litellm.max_budget = original_max_budget
|
||
|
||
|
||
def test_max_ui_session_budget_default_is_one_dollar():
|
||
"""LIT-4662: the dashboard session budget default is a product decision; the
|
||
old 0.25 default locked admins out of auto router Test Connection and the
|
||
playground mid-session with an error that looked like a hardcoded cap."""
|
||
assert litellm.max_ui_session_budget == 1.0
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_config_max_ui_session_budget_applied_and_coerced(tmp_path, monkeypatch):
|
||
"""
|
||
max_ui_session_budget configured via os.environ resolves to a string;
|
||
load_config must coerce it to float so every dashboard session key is
|
||
minted with a numeric max_budget.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setenv("UI_SESSION_BUDGET", "2.5")
|
||
test_config = {
|
||
"model_list": [],
|
||
"litellm_settings": {"max_ui_session_budget": "os.environ/UI_SESSION_BUDGET"},
|
||
}
|
||
config_file = tmp_path / "config.yaml"
|
||
config_file.write_text(yaml.dump(test_config))
|
||
|
||
original_budget = litellm.max_ui_session_budget
|
||
try:
|
||
proxy_config = ProxyConfig()
|
||
await proxy_config.load_config(router=MagicMock(), config_file_path=str(config_file))
|
||
assert isinstance(litellm.max_ui_session_budget, float)
|
||
assert litellm.max_ui_session_budget == 2.5
|
||
finally:
|
||
litellm.max_ui_session_budget = original_budget
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_config_max_ui_session_budget_none_disables_cap(tmp_path):
|
||
"""
|
||
max_ui_session_budget: null in config disables the dashboard session cap
|
||
entirely (session keys minted with no max_budget); load_config must pass
|
||
None through instead of raising on float(None).
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
test_config = {
|
||
"model_list": [],
|
||
"litellm_settings": {"max_ui_session_budget": None},
|
||
}
|
||
config_file = tmp_path / "config.yaml"
|
||
config_file.write_text(yaml.dump(test_config))
|
||
|
||
original_budget = litellm.max_ui_session_budget
|
||
try:
|
||
proxy_config = ProxyConfig()
|
||
await proxy_config.load_config(router=MagicMock(), config_file_path=str(config_file))
|
||
assert litellm.max_ui_session_budget is None
|
||
finally:
|
||
litellm.max_ui_session_budget = original_budget
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_config_default_internal_user_params_max_budget_scientific_notation(tmp_path):
|
||
"""
|
||
Helm's toYaml renders large floats in scientific notation without a
|
||
decimal mantissa (e.g. 1e+09), which PyYAML parses as a string.
|
||
load_config must coerce default_internal_user_params.max_budget to
|
||
float, otherwise every consumer of the raw dict (/user/new, SSO,
|
||
SCIM user creation) passes the string to Prisma, which rejects it
|
||
since max_budget must be Float or Null. Keys outside the coercion
|
||
(including ones not on DefaultInternalUserParams, like
|
||
auto_create_key) must pass through unchanged.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
config_file = tmp_path / "config.yaml"
|
||
config_file.write_text(
|
||
"model_list: []\n"
|
||
"litellm_settings:\n"
|
||
" default_internal_user_params:\n"
|
||
" user_role: internal_user\n"
|
||
" max_budget: 1e+09\n"
|
||
" budget_duration: 30d\n"
|
||
" auto_create_key: false\n"
|
||
)
|
||
|
||
original_params = litellm.default_internal_user_params
|
||
try:
|
||
await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(config_file))
|
||
assert litellm.default_internal_user_params == {
|
||
"user_role": "internal_user",
|
||
"max_budget": 1000000000.0,
|
||
"budget_duration": "30d",
|
||
"auto_create_key": False,
|
||
}
|
||
assert isinstance(litellm.default_internal_user_params["max_budget"], float)
|
||
finally:
|
||
litellm.default_internal_user_params = original_params
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_config_default_internal_user_params_without_max_budget(tmp_path):
|
||
"""
|
||
default_internal_user_params without max_budget (or with an explicit
|
||
null) must be stored as-is and not gain a max_budget key.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
absent_config_file = tmp_path / "absent_config.yaml"
|
||
absent_config_file.write_text(
|
||
"model_list: []\nlitellm_settings:\n default_internal_user_params:\n user_role: internal_user\n"
|
||
)
|
||
|
||
null_config_file = tmp_path / "null_config.yaml"
|
||
null_config_file.write_text(
|
||
"model_list: []\n"
|
||
"litellm_settings:\n"
|
||
" default_internal_user_params:\n"
|
||
" user_role: internal_user\n"
|
||
" max_budget: null\n"
|
||
)
|
||
|
||
original_params = litellm.default_internal_user_params
|
||
try:
|
||
await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(absent_config_file))
|
||
assert litellm.default_internal_user_params == {"user_role": "internal_user"}
|
||
|
||
await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(null_config_file))
|
||
assert litellm.default_internal_user_params == {
|
||
"user_role": "internal_user",
|
||
"max_budget": None,
|
||
}
|
||
finally:
|
||
litellm.default_internal_user_params = original_params
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_config_user_url_validation_handles_null_and_string_false(tmp_path, monkeypatch):
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setattr(litellm, "user_url_validation", True)
|
||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", ["internal.example"])
|
||
monkeypatch.setattr(litellm, "provider_url_destination_allowed_hosts", ["provider.example"])
|
||
null_config_file = tmp_path / "null_config.yaml"
|
||
null_config_file.write_text(
|
||
yaml.dump(
|
||
{
|
||
"model_list": [],
|
||
"general_settings": {
|
||
"user_url_allowed_hosts": None,
|
||
"user_url_validation": None,
|
||
"provider_url_destination_allowed_hosts": None,
|
||
},
|
||
}
|
||
)
|
||
)
|
||
|
||
await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(null_config_file))
|
||
assert litellm.user_url_validation is True
|
||
assert litellm.user_url_allowed_hosts is None
|
||
assert litellm.provider_url_destination_allowed_hosts is None
|
||
|
||
false_config_file = tmp_path / "false_config.yaml"
|
||
false_config_file.write_text(
|
||
yaml.dump(
|
||
{
|
||
"model_list": [],
|
||
"general_settings": {"user_url_validation": "false"},
|
||
}
|
||
)
|
||
)
|
||
|
||
await ProxyConfig().load_config(router=MagicMock(), config_file_path=str(false_config_file))
|
||
assert litellm.user_url_validation is False
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_environment_variables_direct_and_os_environ():
|
||
"""
|
||
Test _load_environment_variables method with direct values and os.environ/ prefixed values
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Test config with both direct values and os.environ/ prefixed values
|
||
test_config = {
|
||
"environment_variables": {
|
||
"DIRECT_VAR": "direct_value",
|
||
"NUMERIC_VAR": 12345,
|
||
"BOOL_VAR": True,
|
||
"SECRET_VAR": "os.environ/ACTUAL_SECRET_VAR",
|
||
}
|
||
}
|
||
|
||
# Mock get_secret_str to return a resolved value
|
||
mock_secret_value = "resolved_secret_value"
|
||
|
||
with patch("litellm.proxy.proxy_server.get_secret_str", return_value=mock_secret_value) as mock_get_secret:
|
||
with patch.dict(os.environ, {}, clear=False): # Don't clear existing env vars, just track changes
|
||
# Call the method under test
|
||
proxy_config._load_environment_variables(test_config)
|
||
|
||
# Verify direct environment variables were set correctly
|
||
assert os.environ["DIRECT_VAR"] == "direct_value"
|
||
assert os.environ["NUMERIC_VAR"] == "12345" # Should be converted to string
|
||
assert os.environ["BOOL_VAR"] == "True" # Should be converted to string
|
||
|
||
# Verify os.environ/ prefixed variable was resolved and set
|
||
assert os.environ["SECRET_VAR"] == mock_secret_value
|
||
|
||
# Verify get_secret_str was called with the correct value
|
||
mock_get_secret.assert_called_once_with(secret_name="os.environ/ACTUAL_SECRET_VAR")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_environment_variables_litellm_license_and_edge_cases():
|
||
"""
|
||
Test _load_environment_variables method with LITELLM_LICENSE special handling and edge cases
|
||
"""
|
||
from unittest.mock import MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Test Case 1: LITELLM_LICENSE in environment_variables
|
||
test_config_with_license = {
|
||
"environment_variables": {
|
||
"LITELLM_LICENSE": "test_license_key",
|
||
"OTHER_VAR": "other_value",
|
||
}
|
||
}
|
||
|
||
# Mock _license_check
|
||
mock_license_check = MagicMock()
|
||
mock_license_check.is_premium.return_value = True
|
||
|
||
with patch("litellm.proxy.proxy_server._license_check", mock_license_check):
|
||
with patch.dict(os.environ, {}, clear=False):
|
||
# Call the method under test
|
||
proxy_config._load_environment_variables(test_config_with_license)
|
||
|
||
# Verify LITELLM_LICENSE was set in environment
|
||
assert os.environ["LITELLM_LICENSE"] == "test_license_key"
|
||
|
||
# Verify license check was updated
|
||
assert mock_license_check.license_str == "test_license_key"
|
||
mock_license_check.is_premium.assert_called_once()
|
||
|
||
# Test Case 2: No environment_variables in config
|
||
test_config_no_env_vars = {}
|
||
|
||
# This should not raise any errors and should return without doing anything
|
||
result = proxy_config._load_environment_variables(test_config_no_env_vars)
|
||
assert result is None # Method returns None
|
||
|
||
# Test Case 3: environment_variables is None
|
||
test_config_none_env_vars = {"environment_variables": None}
|
||
|
||
# This should not raise any errors and should return without doing anything
|
||
result = proxy_config._load_environment_variables(test_config_none_env_vars)
|
||
assert result is None # Method returns None
|
||
|
||
# Test Case 4: os.environ/ prefix but get_secret_str returns None
|
||
test_config_secret_none = {"environment_variables": {"FAILED_SECRET": "os.environ/NONEXISTENT_SECRET"}}
|
||
|
||
with patch("litellm.proxy.proxy_server.get_secret_str", return_value=None):
|
||
with patch.dict(os.environ, {}, clear=False):
|
||
# Call the method under test
|
||
proxy_config._load_environment_variables(test_config_secret_none)
|
||
|
||
# Verify that the environment variable was not set when secret resolution fails
|
||
assert "FAILED_SECRET" not in os.environ
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_environment_variables_blocks_dangerous_keys():
|
||
"""
|
||
Test that _load_environment_variables rejects dangerous env var keys
|
||
like PATH, LD_PRELOAD, PYTHONPATH, etc.
|
||
"""
|
||
import logging
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
original_path = os.environ.get("PATH", "")
|
||
|
||
test_config = {
|
||
"environment_variables": {
|
||
"PATH": "/tmp/evil",
|
||
"LD_PRELOAD": "/tmp/evil.so",
|
||
"PYTHONPATH": "/tmp/evil",
|
||
"SAFE_CUSTOM_VAR": "safe_value",
|
||
}
|
||
}
|
||
|
||
with patch.dict(os.environ, {}, clear=False):
|
||
proxy_config._load_environment_variables(test_config)
|
||
|
||
# Blocked keys should not be set to the attacker value
|
||
assert os.environ.get("PATH") != "/tmp/evil"
|
||
assert "LD_PRELOAD" not in os.environ or os.environ["LD_PRELOAD"] != "/tmp/evil.so"
|
||
assert os.environ.get("PYTHONPATH") != "/tmp/evil"
|
||
|
||
# Safe keys should still be set
|
||
assert os.environ["SAFE_CUSTOM_VAR"] == "safe_value"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_environment_variables_allows_proxy_keys():
|
||
"""
|
||
Test that HTTP_PROXY/HTTPS_PROXY are allowed since they are commonly used
|
||
in corporate environments to route outbound API calls.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
test_config = {
|
||
"environment_variables": {
|
||
"HTTP_PROXY": "http://corp-proxy:8080",
|
||
"HTTPS_PROXY": "http://corp-proxy:8080",
|
||
}
|
||
}
|
||
|
||
with patch.dict(os.environ, {}, clear=False):
|
||
proxy_config._load_environment_variables(test_config)
|
||
|
||
assert os.environ["HTTP_PROXY"] == "http://corp-proxy:8080"
|
||
assert os.environ["HTTPS_PROXY"] == "http://corp-proxy:8080"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_environment_variables_blocks_no_proxy():
|
||
"""
|
||
Test that NO_PROXY/no_proxy are blocked to prevent bypassing proxy-based
|
||
network monitoring.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
test_config = {
|
||
"environment_variables": {
|
||
"NO_PROXY": "internal-service",
|
||
"no_proxy": "internal-service",
|
||
}
|
||
}
|
||
|
||
with patch.dict(os.environ, {}, clear=False):
|
||
proxy_config._load_environment_variables(test_config)
|
||
|
||
assert os.environ.get("NO_PROXY") != "internal-service"
|
||
assert os.environ.get("no_proxy") != "internal-service"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_write_config_to_file(monkeypatch):
|
||
"""
|
||
Do not write config to file if store_model_in_db is True
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
# Set store_model_in_db to True
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
|
||
|
||
# Mock prisma_client to not be None (so DB path is taken)
|
||
mock_prisma_client = AsyncMock()
|
||
mock_prisma_client.insert_data = AsyncMock()
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
|
||
|
||
# Mock general_settings
|
||
mock_general_settings = {"store_model_in_db": True}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", mock_general_settings)
|
||
|
||
# Mock user_config_file_path
|
||
test_config_path = "/tmp/test_config.yaml"
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", test_config_path)
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Mock the open function to track if file writing is attempted
|
||
mock_file_open = mock_open()
|
||
|
||
with patch("builtins.open", mock_file_open), patch("yaml.dump") as mock_yaml_dump:
|
||
# Call save_config with test data
|
||
test_config = {"key": "value", "model_list": ["model1", "model2"]}
|
||
await proxy_config.save_config(new_config=test_config)
|
||
|
||
# Verify that file was NOT opened for writing (since store_model_in_db=True)
|
||
mock_file_open.assert_not_called()
|
||
mock_yaml_dump.assert_not_called()
|
||
|
||
# Verify that database insert was called instead
|
||
mock_prisma_client.insert_data.assert_called_once()
|
||
|
||
# Verify the config passed to DB has model_list removed
|
||
call_args = mock_prisma_client.insert_data.call_args
|
||
assert call_args.kwargs["data"] == {"key": "value"} # model_list should be popped
|
||
assert call_args.kwargs["table_name"] == "config"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_write_config_to_file_when_store_model_in_db_false(monkeypatch):
|
||
"""
|
||
Test that config IS written to file when store_model_in_db is False
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
# Set store_model_in_db to False
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False)
|
||
|
||
# Mock prisma_client to be None (so file path is taken)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None)
|
||
|
||
# Mock general_settings
|
||
mock_general_settings = {"store_model_in_db": False}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", mock_general_settings)
|
||
|
||
# Mock user_config_file_path
|
||
test_config_path = "/tmp/test_config.yaml"
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.user_config_file_path", test_config_path)
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Mock the open function and yaml.dump
|
||
mock_file_open = mock_open()
|
||
|
||
with patch("builtins.open", mock_file_open), patch("yaml.dump") as mock_yaml_dump:
|
||
# Call save_config with test data
|
||
test_config = {"key": "value", "other_key": "other_value"}
|
||
await proxy_config.save_config(new_config=test_config)
|
||
|
||
# Verify that file WAS opened for writing (since store_model_in_db=False)
|
||
mock_file_open.assert_called_once_with(f"{test_config_path}", "w")
|
||
|
||
# Verify yaml.dump was called with the config
|
||
mock_yaml_dump.assert_called_once_with(
|
||
test_config,
|
||
mock_file_open.return_value.__enter__.return_value,
|
||
default_flow_style=False,
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_midstream_error():
|
||
"""
|
||
Test async_data_generator handles midstream error from async_post_call_streaming_hook
|
||
Specifically testing the case where Azure Content Safety Guardrail returns an error
|
||
"""
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
# Create mock objects
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gpt-3.5-turbo",
|
||
"messages": [{"role": "user", "content": "test"}],
|
||
}
|
||
|
||
# Mock response chunks - simulating normal streaming that gets interrupted
|
||
mock_chunks = [
|
||
{"choices": [{"delta": {"content": "Hello"}}]},
|
||
{"choices": [{"delta": {"content": " world"}}]},
|
||
{"choices": [{"delta": {"content": " this"}}]},
|
||
]
|
||
|
||
# Mock the proxy_logging_obj
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
|
||
# Mock async_post_call_streaming_iterator_hook to yield chunks
|
||
async def mock_streaming_iterator(*args, **kwargs):
|
||
for chunk in mock_chunks:
|
||
yield chunk
|
||
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = mock_streaming_iterator
|
||
|
||
# Mock async_post_call_streaming_hook to return error on third chunk
|
||
def mock_streaming_hook(*args, **kwargs):
|
||
chunk = kwargs.get("response")
|
||
# Return error message for the third chunk (simulating guardrail trigger)
|
||
if chunk == mock_chunks[2]:
|
||
return (
|
||
'data: {"error": {"error": "Azure Content Safety Guardrail: Hate crossed severity 2, Got severity: 2"}}'
|
||
)
|
||
# Return normal chunks for first two
|
||
return chunk
|
||
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(side_effect=mock_streaming_hook)
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
# Mock the global proxy_logging_obj
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
# Create a mock response object
|
||
mock_response = MagicMock()
|
||
|
||
# Collect all yielded data from the generator
|
||
yielded_data = []
|
||
try:
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
except Exception as e:
|
||
# If there's an exception, that's also part of what we want to test
|
||
pass
|
||
|
||
# Verify the results
|
||
assert len(yielded_data) >= 3, f"Expected at least 3 chunks, got {len(yielded_data)}: {yielded_data}"
|
||
|
||
# First two chunks should be normal data
|
||
assert yielded_data[0].startswith("data: "), f"First chunk should start with 'data: ', got: {yielded_data[0]}"
|
||
assert yielded_data[1].startswith("data: "), f"Second chunk should start with 'data: ', got: {yielded_data[1]}"
|
||
|
||
# The error message should be yielded
|
||
error_found = False
|
||
done_found = False
|
||
|
||
for data in yielded_data:
|
||
if "Azure Content Safety Guardrail: Hate crossed severity 2" in data:
|
||
error_found = True
|
||
if "data: [DONE]" in data:
|
||
done_found = True
|
||
|
||
assert error_found, f"Error message should be found in yielded data. Got: {yielded_data}"
|
||
assert done_found, f"[DONE] message should be found at the end. Got: {yielded_data}"
|
||
|
||
# Verify that the streaming hook was called for each chunk
|
||
assert mock_proxy_logging_obj.async_post_call_streaming_hook.call_count == len(mock_chunks)
|
||
|
||
# Verify that post_call_failure_hook was NOT called (since this is not an exception case)
|
||
mock_proxy_logging_obj.post_call_failure_hook.assert_not_called()
|
||
|
||
|
||
def _has_nested_none_values(obj, path="root"):
|
||
"""
|
||
Recursively check if an object contains nested None values.
|
||
|
||
Args:
|
||
obj: The object to check
|
||
path: Current path in the object tree (for debugging)
|
||
|
||
Returns:
|
||
List of paths where None values were found
|
||
"""
|
||
none_paths = []
|
||
|
||
if obj is None:
|
||
none_paths.append(path)
|
||
elif isinstance(obj, dict):
|
||
for key, value in obj.items():
|
||
none_paths.extend(_has_nested_none_values(value, f"{path}.{key}"))
|
||
elif isinstance(obj, (list, tuple)):
|
||
for i, item in enumerate(obj):
|
||
none_paths.extend(_has_nested_none_values(item, f"{path}[{i}]"))
|
||
elif hasattr(obj, "__dict__"):
|
||
# Handle object attributes
|
||
for key, value in obj.__dict__.items():
|
||
if not key.startswith("_"): # Skip private attributes
|
||
none_paths.extend(_has_nested_none_values(value, f"{path}.{key}"))
|
||
|
||
return none_paths
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_chat_completion_result_no_nested_none_values():
|
||
"""
|
||
Test that chat_completion result doesn't have nested None values when using exclude_none=True
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from fastapi import Request, Response
|
||
from pydantic import BaseModel
|
||
|
||
import litellm
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import chat_completion
|
||
|
||
# Create a mock ModelResponse with nested None values
|
||
mock_model_response = litellm.ModelResponse()
|
||
mock_model_response.id = "test-id"
|
||
mock_model_response.model = "gpt-3.5-turbo"
|
||
mock_model_response.object = "chat.completion"
|
||
mock_model_response.created = 1234567890
|
||
|
||
# Create message with None values that should be excluded
|
||
mock_message = litellm.Message(
|
||
content="Hello, world!",
|
||
role="assistant",
|
||
function_call=None, # This should be excluded
|
||
tool_calls=None, # This should be excluded
|
||
audio=None, # This should be excluded
|
||
reasoning_content=None, # This should be excluded
|
||
thinking_blocks=None, # This should be excluded
|
||
annotations=None, # This should be excluded
|
||
)
|
||
|
||
# Create choice with potential None values
|
||
mock_choice = litellm.Choices(
|
||
finish_reason="stop",
|
||
index=0,
|
||
message=mock_message,
|
||
logprobs=None, # This should be excluded when exclude_none=True
|
||
)
|
||
|
||
mock_model_response.choices = [mock_choice]
|
||
setattr(
|
||
mock_model_response,
|
||
"usage",
|
||
litellm.Usage(prompt_tokens=10, completion_tokens=5, total_tokens=15),
|
||
)
|
||
|
||
# Verify the mock has None values before serialization
|
||
raw_dict = mock_model_response.model_dump()
|
||
none_paths_before = _has_nested_none_values(raw_dict)
|
||
assert len(none_paths_before) > 0, "Mock should have None values before exclude_none=True"
|
||
|
||
# Mock the request processing to return our mock response
|
||
mock_base_processor = MagicMock()
|
||
mock_base_processor.base_process_llm_request = AsyncMock(return_value=mock_model_response)
|
||
|
||
# Mock other dependencies
|
||
mock_request = MagicMock(spec=Request)
|
||
mock_response = MagicMock(spec=Response)
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
|
||
with (
|
||
patch(
|
||
"litellm.proxy.proxy_server._read_request_body",
|
||
return_value={"model": "gpt-3.5-turbo", "messages": []},
|
||
),
|
||
patch(
|
||
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing",
|
||
return_value=mock_base_processor,
|
||
),
|
||
):
|
||
# Call the chat_completion function
|
||
result = await chat_completion(
|
||
request=mock_request,
|
||
fastapi_response=mock_response,
|
||
user_api_key_dict=mock_user_api_key_dict,
|
||
)
|
||
|
||
# Verify the result is a dict (since isinstance(result, BaseModel) was True)
|
||
assert isinstance(result, dict), f"Expected dict result, got {type(result)}"
|
||
|
||
# Check that there are no nested None values in the result
|
||
none_paths_after = _has_nested_none_values(result)
|
||
assert len(none_paths_after) == 0, (
|
||
f"Result should not contain nested None values. Found None at: {none_paths_after}"
|
||
)
|
||
|
||
# Verify essential fields are present
|
||
assert "id" in result
|
||
assert "model" in result
|
||
assert "object" in result
|
||
assert "created" in result
|
||
assert "choices" in result
|
||
assert "usage" in result
|
||
|
||
# Verify that the choices contain the expected message content
|
||
assert len(result["choices"]) == 1
|
||
assert result["choices"][0]["message"]["content"] == "Hello, world!"
|
||
assert result["choices"][0]["message"]["role"] == "assistant"
|
||
|
||
# Verify that None fields were excluded (should not be present in the dict)
|
||
message = result["choices"][0]["message"]
|
||
excluded_fields = [
|
||
"function_call",
|
||
"tool_calls",
|
||
"audio",
|
||
"reasoning_content",
|
||
"thinking_blocks",
|
||
"annotations",
|
||
]
|
||
for field in excluded_fields:
|
||
assert field not in message, f"Field '{field}' should be excluded when it's None"
|
||
|
||
|
||
# ============================================================================
|
||
# Price Data Reload Tests
|
||
# ============================================================================
|
||
|
||
|
||
def _reload_schedule_row(
|
||
param_value: dict,
|
||
*,
|
||
reload_revision: int = 0,
|
||
last_run_at: datetime | None = None,
|
||
) -> types.SimpleNamespace:
|
||
"""LiteLLM_Config row shape: admin-owned interval in param_value, run state in dedicated columns"""
|
||
return types.SimpleNamespace(
|
||
param_value=param_value,
|
||
reload_revision=reload_revision,
|
||
last_run_at=last_run_at,
|
||
)
|
||
|
||
|
||
class TestPriceDataReloadAPI:
|
||
"""Test cases for price data reload API endpoints"""
|
||
|
||
@pytest.fixture
|
||
def client_with_auth(self):
|
||
"""Create a test client with authentication"""
|
||
from litellm.proxy._types import LitellmUserRoles
|
||
from litellm.proxy.proxy_server import cleanup_router_config_variables
|
||
|
||
cleanup_router_config_variables()
|
||
filepath = os.path.dirname(os.path.abspath(__file__))
|
||
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
|
||
asyncio.run(initialize(config=config_fp, debug=True))
|
||
|
||
# Mock admin user authentication
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_role = LitellmUserRoles.PROXY_ADMIN
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
|
||
return TestClient(app)
|
||
|
||
def test_reload_model_cost_map_admin_access(self, client_with_auth):
|
||
"""Test that admin users can access the reload endpoint"""
|
||
# Save the original model_cost so the endpoint's direct assignment
|
||
# (litellm.model_cost = new_model_cost_map) does not contaminate
|
||
# subsequent tests running in the same worker process.
|
||
from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map",
|
||
new=AsyncMock(
|
||
return_value=ModelCostMapReloaded(model_cost_map={"gpt-3.5-turbo": {"input_cost_per_token": 0.001}})
|
||
),
|
||
):
|
||
# Mock the database connection
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(
|
||
return_value=_reload_schedule_row({}, reload_revision=1)
|
||
)
|
||
|
||
response = client_with_auth.post("/reload/model_cost_map")
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert data["status"] == "success"
|
||
assert "message" in data
|
||
assert "timestamp" in data
|
||
assert "models_count" in data
|
||
# The new implementation immediately reloads and returns the count
|
||
assert "Price data reloaded successfully! 1 models updated." in data["message"]
|
||
assert data["models_count"] == 1
|
||
finally:
|
||
# Restore the full model cost map so subsequent tests are not affected
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_reload_model_cost_map_non_admin_access(self, client_with_auth):
|
||
"""Test that non-admin users cannot access the reload endpoint"""
|
||
# Mock non-admin user
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_role = "user" # Non-admin role
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
|
||
response = client_with_auth.post("/reload/model_cost_map")
|
||
|
||
assert response.status_code == 403
|
||
data = response.json()
|
||
assert "Access denied" in data["detail"]
|
||
assert "Admin role required" in data["detail"]
|
||
|
||
def test_get_model_cost_map_public_access(self, client_no_auth):
|
||
"""Test that the model cost map endpoint is publicly accessible"""
|
||
with patch("litellm.model_cost", {"gpt-3.5-turbo": {"input_cost_per_token": 0.001}}):
|
||
response = client_no_auth.get("/public/litellm_model_cost_map")
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert "gpt-3.5-turbo" in data
|
||
|
||
def test_reload_model_cost_map_error_handling(self, client_with_auth):
|
||
"""Test error handling in the reload endpoint"""
|
||
with patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map",
|
||
new=AsyncMock(side_effect=Exception("Network error")),
|
||
):
|
||
# Mock the database connection
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(
|
||
return_value=_reload_schedule_row({}, reload_revision=1)
|
||
)
|
||
|
||
response = client_with_auth.post("/reload/model_cost_map")
|
||
|
||
assert response.status_code == 500 # An unexpected exception still maps to 500
|
||
data = response.json()
|
||
assert "Failed to reload model cost map" in data["detail"]
|
||
|
||
def test_schedule_model_cost_map_reload_admin_access(self, client_with_auth):
|
||
"""Admin schedule write owns param_value only, so it can't clobber the job-owned run columns"""
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
# Mock database upsert
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
|
||
|
||
response = client_with_auth.post("/schedule/model_cost_map_reload?hours=6")
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert data["status"] == "success"
|
||
assert data["interval_hours"] == 6
|
||
assert "message" in data
|
||
assert "timestamp" in data
|
||
|
||
call_args = mock_prisma.db.litellm_config.upsert.call_args
|
||
assert call_args[1]["where"] == {"param_name": "model_cost_map_reload_config"}
|
||
update_payload = call_args[1]["data"]["update"]
|
||
assert set(update_payload.keys()) == {"param_value"}
|
||
assert json.loads(update_payload["param_value"]) == {"interval_hours": 6}
|
||
create_payload = call_args[1]["data"]["create"]
|
||
assert set(create_payload.keys()) == {"param_name", "param_value"}
|
||
assert json.loads(create_payload["param_value"]) == {"interval_hours": 6}
|
||
|
||
def test_schedule_model_cost_map_reload_non_admin_access(self, client_with_auth):
|
||
"""Test that non-admin users cannot schedule periodic reload"""
|
||
# Mock non-admin user
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_role = "user" # Non-admin role
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
|
||
response = client_with_auth.post("/schedule/model_cost_map_reload?hours=6")
|
||
|
||
assert response.status_code == 403
|
||
data = response.json()
|
||
assert "Access denied" in data["detail"]
|
||
assert "Admin role required" in data["detail"]
|
||
|
||
def test_schedule_model_cost_map_reload_invalid_hours(self, client_with_auth):
|
||
"""Test that invalid hours parameter is rejected"""
|
||
response = client_with_auth.post("/schedule/model_cost_map_reload?hours=0")
|
||
|
||
assert response.status_code == 400
|
||
data = response.json()
|
||
assert "Hours must be greater than 0" in data["detail"]
|
||
|
||
def test_cancel_model_cost_map_reload_admin_access(self, client_with_auth):
|
||
"""Test that admin users can cancel periodic reload"""
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=1)
|
||
mock_prisma.db.litellm_config.delete = AsyncMock(return_value=None)
|
||
|
||
response = client_with_auth.delete("/schedule/model_cost_map_reload")
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert data["status"] == "success"
|
||
assert "message" in data
|
||
assert "timestamp" in data
|
||
assert json.loads(mock_prisma.db.litellm_config.update_many.await_args.kwargs["data"]["param_value"]) == {
|
||
"interval_hours": None
|
||
}
|
||
mock_prisma.db.litellm_config.delete.assert_not_called()
|
||
|
||
def test_cancel_model_cost_map_reload_non_admin_access(self, client_with_auth):
|
||
"""Test that non-admin users cannot cancel periodic reload"""
|
||
# Mock non-admin user
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_role = "user" # Non-admin role
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
|
||
response = client_with_auth.delete("/schedule/model_cost_map_reload")
|
||
|
||
assert response.status_code == 403
|
||
data = response.json()
|
||
assert "Access denied" in data["detail"]
|
||
assert "Admin role required" in data["detail"]
|
||
|
||
def test_get_model_cost_map_reload_status_admin_access(self, client_with_auth):
|
||
"""
|
||
Regression (LIT-4882): status is served purely from the DB row, so a restarted pod
|
||
(whose in-memory clock only knows its own boot) still reports the real last/next run
|
||
"""
|
||
proxy_server_module.proxy_config.model_cost_map_loaded_at = datetime(2030, 6, 1, tzinfo=timezone.utc)
|
||
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row(
|
||
{"interval_hours": 6},
|
||
last_run_at=datetime(2024, 1, 1, 6, 0, tzinfo=timezone.utc),
|
||
)
|
||
)
|
||
|
||
response = client_with_auth.get("/schedule/model_cost_map_reload/status")
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert data["scheduled"] is True
|
||
assert data["interval_hours"] == 6
|
||
assert data["last_run"] == "2024-01-01T06:00:00+00:00"
|
||
assert data["next_run"] == "2024-01-01T12:00:00+00:00"
|
||
|
||
def test_get_model_cost_map_reload_status_non_admin_access(self, client_with_auth):
|
||
"""Test that non-admin users cannot get reload status"""
|
||
# Mock non-admin user
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_role = "user" # Non-admin role
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
|
||
response = client_with_auth.get("/schedule/model_cost_map_reload/status")
|
||
|
||
assert response.status_code == 403
|
||
data = response.json()
|
||
assert "Access denied" in data["detail"]
|
||
assert "Admin role required" in data["detail"]
|
||
|
||
def test_get_model_cost_map_reload_status_no_config(self, client_with_auth):
|
||
"""Test that status returns not scheduled when no config exists"""
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||
|
||
response = client_with_auth.get("/schedule/model_cost_map_reload/status")
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert data["scheduled"] == False
|
||
assert data["interval_hours"] == None
|
||
assert data["last_run"] == None
|
||
assert data["next_run"] == None
|
||
|
||
def test_get_model_cost_map_reload_status_no_interval(self, client_with_auth):
|
||
"""A row left behind by a manual reload (no interval) must not read as scheduled"""
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row(
|
||
{"interval_hours": None},
|
||
reload_revision=3,
|
||
)
|
||
)
|
||
|
||
response = client_with_auth.get("/schedule/model_cost_map_reload/status")
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert data["scheduled"] is False
|
||
assert data["interval_hours"] is None
|
||
assert data["last_run"] is None
|
||
assert data["next_run"] is None
|
||
|
||
def test_get_model_cost_map_reload_status_before_first_run(self, client_with_auth):
|
||
"""Scheduled but never executed: no last_run_at means no next_run can be computed"""
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row({"interval_hours": 6})
|
||
)
|
||
|
||
response = client_with_auth.get("/schedule/model_cost_map_reload/status")
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert data["scheduled"] is True
|
||
assert data["interval_hours"] == 6
|
||
assert data["last_run"] is None
|
||
assert data["next_run"] is None
|
||
|
||
|
||
class TestPriceDataReloadIntegration:
|
||
"""Integration tests for the complete price data reload feature"""
|
||
|
||
@pytest.fixture
|
||
def client_with_auth(self):
|
||
"""Create a test client with authentication"""
|
||
from litellm.proxy._types import LitellmUserRoles
|
||
from litellm.proxy.proxy_server import cleanup_router_config_variables
|
||
|
||
cleanup_router_config_variables()
|
||
filepath = os.path.dirname(os.path.abspath(__file__))
|
||
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
|
||
asyncio.run(initialize(config=config_fp, debug=True))
|
||
|
||
# Mock admin user authentication
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_role = LitellmUserRoles.PROXY_ADMIN
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
|
||
return TestClient(app)
|
||
|
||
def test_complete_reload_flow(self, client_with_auth):
|
||
"""Test the complete reload flow from API to model cost update"""
|
||
# Mock the model cost map
|
||
mock_cost_map = {
|
||
"gpt-3.5-turbo": {
|
||
"input_cost_per_token": 0.001,
|
||
"output_cost_per_token": 0.002,
|
||
},
|
||
"gpt-4": {"input_cost_per_token": 0.03, "output_cost_per_token": 0.06},
|
||
}
|
||
|
||
from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map",
|
||
new=AsyncMock(return_value=ModelCostMapReloaded(model_cost_map=mock_cost_map)),
|
||
):
|
||
# Mock the database connection
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(
|
||
return_value=_reload_schedule_row({}, reload_revision=1)
|
||
)
|
||
|
||
# Test reload endpoint
|
||
response = client_with_auth.post("/reload/model_cost_map")
|
||
assert response.status_code == 200
|
||
|
||
# Test get endpoint
|
||
response = client_with_auth.get("/public/litellm_model_cost_map")
|
||
assert response.status_code == 200
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_pod_data_clock_seeded_from_actual_cost_map_load(self):
|
||
"""Regression: seeding from ProxyConfig construction time instead of the real
|
||
import-time fetch let a manual request stamped during startup be skipped"""
|
||
from datetime import datetime, timezone
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
fetch_time = datetime(2024, 1, 1, 6, 0, tzinfo=timezone.utc)
|
||
with patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map_loaded_at",
|
||
return_value=fetch_time,
|
||
):
|
||
assert ProxyConfig().model_cost_map_loaded_at == fetch_time
|
||
|
||
def test_distributed_reload_check_function(self):
|
||
"""
|
||
A revision this pod has not applied takes effect here even one minute into a 6h
|
||
interval; a missing row is a no-op
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=None)
|
||
|
||
boot_loaded_at = proxy_config.model_cost_map_loaded_at
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
mock_prisma.db.litellm_config.update_many.assert_not_called()
|
||
assert proxy_config.model_cost_map_loaded_at == boot_loaded_at
|
||
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row(
|
||
{"interval_hours": 6},
|
||
reload_revision=4,
|
||
last_run_at=datetime(2024, 1, 1, 6, 59, 30, tzinfo=timezone.utc),
|
||
)
|
||
)
|
||
proxy_config.model_cost_map_loaded_at = frozen_now - timedelta(minutes=1)
|
||
proxy_config.model_cost_map_applied_revision = 3
|
||
|
||
from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock
|
||
) as mock_get_map,
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
mock_get_map.return_value = ModelCostMapReloaded(
|
||
model_cost_map={"gpt-3.5-turbo": {"input_cost_per_token": 0.001}}
|
||
)
|
||
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
assert litellm.model_cost["gpt-3.5-turbo"] == {"input_cost_per_token": 0.001}
|
||
assert proxy_config.model_cost_map_loaded_at == frozen_now
|
||
assert mock_prisma.db.litellm_config.update_many.call_args[1] == {
|
||
"data": {"last_run_at": frozen_now},
|
||
"where": {"param_name": "model_cost_map_reload_config"},
|
||
}
|
||
mock_prisma.db.litellm_config.upsert.assert_not_called()
|
||
assert proxy_config.model_cost_map_applied_revision == 4
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_distributed_reload_ignores_already_applied_request(self):
|
||
"""
|
||
A revision this pod already applied must not re-trigger on every job tick for the
|
||
rest of the interval
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row(
|
||
{"interval_hours": 6},
|
||
reload_revision=4,
|
||
last_run_at=datetime(2024, 1, 1, 6, 0, tzinfo=timezone.utc),
|
||
)
|
||
)
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
pod_data_loaded_at = frozen_now - timedelta(minutes=1)
|
||
proxy_config.model_cost_map_loaded_at = pod_data_loaded_at
|
||
proxy_config.model_cost_map_applied_revision = 4
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock
|
||
) as mock_get_map,
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
mock_get_map.assert_not_called()
|
||
mock_prisma.db.litellm_config.update_many.assert_not_called()
|
||
assert proxy_config.model_cost_map_loaded_at == pod_data_loaded_at
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_periodic_reload_uses_pod_local_data_age(self):
|
||
"""
|
||
Each pod decides from the age of its own data, so a pod holding a stale copy
|
||
refreshes even when the shared row was just stamped by another pod, and stays
|
||
put while its copy is inside the interval
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row(
|
||
{"interval_hours": 6},
|
||
last_run_at=datetime(2024, 1, 1, 6, 59, tzinfo=timezone.utc),
|
||
)
|
||
)
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
proxy_config.model_cost_map_loaded_at = datetime(2024, 1, 1, 0, 0, tzinfo=timezone.utc)
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock
|
||
) as mock_get_map,
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
mock_get_map.return_value = ModelCostMapReloaded(
|
||
model_cost_map={"gpt-4-test": {"input_cost_per_token": 0.5}}
|
||
)
|
||
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
assert litellm.model_cost["gpt-4-test"] == {"input_cost_per_token": 0.5}
|
||
assert proxy_config.model_cost_map_loaded_at == frozen_now
|
||
assert mock_prisma.db.litellm_config.update_many.call_args[1]["data"] == {"last_run_at": frozen_now}
|
||
|
||
mock_get_map.reset_mock()
|
||
mock_prisma.db.litellm_config.update_many.reset_mock()
|
||
proxy_config.model_cost_map_loaded_at = frozen_now - timedelta(hours=1)
|
||
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
mock_get_map.assert_not_called()
|
||
mock_prisma.db.litellm_config.update_many.assert_not_called()
|
||
assert proxy_config.model_cost_map_loaded_at == frozen_now - timedelta(hours=1)
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_every_pod_applies_a_manual_revision_exactly_once(self):
|
||
"""The fleet property: no pod clears the revision, so each one reloads on the tick
|
||
after it is published and then stops, whatever order the pods poll in"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
pods = [ProxyConfig(), ProxyConfig(), ProxyConfig()]
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
|
||
for pod in pods:
|
||
pod.model_cost_map_applied_revision = 0
|
||
pod.model_cost_map_loaded_at = frozen_now
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock
|
||
) as mock_get_map,
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
mock_get_map.return_value = ModelCostMapReloaded(
|
||
model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}
|
||
)
|
||
|
||
for _ in range(3):
|
||
for pod in pods:
|
||
asyncio.run(pod._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
assert mock_get_map.call_count == len(pods)
|
||
assert all(p.model_cost_map_applied_revision == 1 for p in pods)
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
@pytest.mark.parametrize(
|
||
"published_revision, expect_reload",
|
||
[(4, True), (0, False)],
|
||
)
|
||
def test_booting_pod_serves_an_outstanding_request_once(self, published_revision, expect_reload):
|
||
"""
|
||
Regression: a manual reload published while this pod was starting must still be
|
||
served. The pod cannot prove its import-time fetch already covers that request, so
|
||
it applies it on the first poll and adopts the revision, leaving later polls quiet.
|
||
A row nobody has ever reloaded (revision 0) costs the pod nothing
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
proxy_config.model_cost_map_loaded_at = frozen_now
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row({}, reload_revision=published_revision)
|
||
)
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock
|
||
) as mock_get_map,
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
mock_get_map.return_value = ModelCostMapReloaded(
|
||
model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}
|
||
)
|
||
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
assert mock_get_map.call_count == (1 if expect_reload else 0)
|
||
assert proxy_config.model_cost_map_applied_revision == published_revision
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_distributed_reload_stamps_last_run_without_creating_row(self):
|
||
"""
|
||
Regression: the job's write carries neither param_value (which would clobber the
|
||
admin-configured interval) nor a create branch (which would resurrect a schedule
|
||
a concurrent cancel just deleted)
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=_reload_schedule_row({"interval_hours": 24}))
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock
|
||
) as mock_get_map,
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
mock_get_map.return_value = ModelCostMapReloaded(
|
||
model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}
|
||
)
|
||
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
assert mock_prisma.db.litellm_config.update_many.call_args[1] == {
|
||
"data": {"last_run_at": frozen_now},
|
||
"where": {"param_name": "model_cost_map_reload_config"},
|
||
}
|
||
mock_prisma.db.litellm_config.upsert.assert_not_called()
|
||
mock_prisma.db.litellm_config.create.assert_not_called()
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_distributed_reload_leaves_request_unserved_when_status_write_fails(self):
|
||
"""
|
||
A run that never reached the row must not be recorded as served. Adopting the
|
||
revision here would leave the card reporting the previous run until someone clicks
|
||
again, because a manual request is published once and never republished
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
proxy_config.model_cost_map_loaded_at = frozen_now - timedelta(hours=9)
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7)
|
||
)
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(side_effect=Exception("connection reset"))
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map",
|
||
new_callable=AsyncMock,
|
||
) as mock_get_map,
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
mock_get_map.return_value = ModelCostMapReloaded(
|
||
model_cost_map={"gpt-4": {"input_cost_per_token": 0.1}}
|
||
)
|
||
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
assert proxy_config.model_cost_map_applied_revision == 0
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_distributed_reload_keeps_current_map_when_fetch_fails(self):
|
||
"""Fetch failure during a periodic reload must not downgrade the pod or count the
|
||
request as served.
|
||
|
||
Regression: a 429/network failure used to silently replace litellm.model_cost with
|
||
the stale packaged backup and stamp last_run. Adopting the revision here would be
|
||
the same bug one level up: a manual request is published once and never republished,
|
||
so a pod that records it as applied without the data stays mispriced until someone
|
||
clicks again
|
||
"""
|
||
from litellm.litellm_core_utils.get_model_cost_map import (
|
||
ModelCostMapReloadUnavailable,
|
||
)
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
pod_data_loaded_at = frozen_now - timedelta(hours=9)
|
||
proxy_config.model_cost_map_loaded_at = pod_data_loaded_at
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7)
|
||
)
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None)
|
||
|
||
original_model_cost = litellm.model_cost
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map",
|
||
new=AsyncMock(return_value=ModelCostMapReloadUnavailable(reason="HTTP 429 from upstream")),
|
||
),
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
assert litellm.model_cost is original_model_cost, (
|
||
"a failed reload must keep the currently loaded cost map, not swap in the packaged backup"
|
||
)
|
||
assert proxy_config.model_cost_map_loaded_at == pod_data_loaded_at, (
|
||
"a failed reload must not stamp the pod's data age, otherwise the retry waits a full interval"
|
||
)
|
||
assert proxy_config.model_cost_map_applied_revision == 0, (
|
||
"a failed reload must leave the revision unapplied so the next poll retries it"
|
||
)
|
||
mock_prisma.db.litellm_config.update_many.assert_not_called()
|
||
mock_prisma.db.litellm_config.upsert.assert_not_called()
|
||
|
||
def test_scheduled_reload_replays_runtime_registrations(self):
|
||
"""The scheduled reload is the trigger a pod hits on its own, so it must
|
||
both preserve runtime-registered model metadata and run to completion.
|
||
The swap happens early in the handler, so a failure in the bookkeeping
|
||
after it is swallowed by the surrounding except and would otherwise
|
||
leave the metadata correct while the path is quietly broken"""
|
||
from litellm import utils as litellm_utils
|
||
from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
proxy_config.model_cost_map_loaded_at = frozen_now - timedelta(hours=9)
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(
|
||
return_value=_reload_schedule_row({"interval_hours": 6}, reload_revision=7)
|
||
)
|
||
mock_prisma.db.litellm_config.update_many = AsyncMock(return_value=None)
|
||
|
||
original_model_cost = litellm.model_cost
|
||
original_registry = dict(litellm_utils._runtime_registered_model_cost)
|
||
try:
|
||
litellm.register_model(
|
||
model_cost={"custom/deployment-model": {"litellm_provider": "custom", "max_input_tokens": 4321}}
|
||
)
|
||
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map",
|
||
new=AsyncMock(
|
||
return_value=ModelCostMapReloaded(
|
||
model_cost_map={"gpt-4o": {"litellm_provider": "openai", "mode": "chat"}}
|
||
)
|
||
),
|
||
),
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
patch("litellm.proxy.proxy_server.verbose_proxy_logger") as mock_logger,
|
||
):
|
||
asyncio.run(proxy_config._check_and_reload_model_cost_map(mock_prisma))
|
||
|
||
mock_logger.exception.assert_not_called()
|
||
assert litellm.model_cost["custom/deployment-model"]["max_input_tokens"] == 4321
|
||
assert "gpt-4o" in litellm.model_cost
|
||
assert proxy_config.model_cost_map_applied_revision == 7
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
litellm_utils._runtime_registered_model_cost.clear()
|
||
litellm_utils._runtime_registered_model_cost.update(original_registry)
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_swap_in_model_cost_map_counts_the_fetched_catalog_only(self):
|
||
"""The count the reload endpoints report describes the price data, so it
|
||
is taken before the runtime registrations are written back into the same
|
||
dict. Counting after would inflate it by however many deployments and
|
||
overrides this pod happens to be carrying"""
|
||
from litellm import utils as litellm_utils
|
||
from litellm.proxy.proxy_server import _swap_in_model_cost_map
|
||
|
||
original_model_cost = litellm.model_cost
|
||
original_registry = dict(litellm_utils._runtime_registered_model_cost)
|
||
try:
|
||
litellm.register_model(
|
||
model_cost={"custom/deployment-model": {"litellm_provider": "custom", "max_input_tokens": 4321}}
|
||
)
|
||
|
||
models_count = _swap_in_model_cost_map({"gpt-4o": {"litellm_provider": "openai", "mode": "chat"}})
|
||
|
||
assert models_count == 1
|
||
assert litellm.model_cost["custom/deployment-model"]["max_input_tokens"] == 4321
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
litellm_utils._runtime_registered_model_cost.clear()
|
||
litellm_utils._runtime_registered_model_cost.update(original_registry)
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_manual_reload_preserves_interval_hours(self):
|
||
"""
|
||
Regression: manual reload owns only the run columns, so it never reads or rewrites
|
||
param_value and cannot destroy an existing schedule
|
||
"""
|
||
from litellm.proxy._types import LitellmUserRoles
|
||
from litellm.proxy.proxy_server import cleanup_router_config_variables
|
||
|
||
cleanup_router_config_variables()
|
||
filepath = os.path.dirname(os.path.abspath(__file__))
|
||
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
|
||
asyncio.run(initialize(config=config_fp, debug=True))
|
||
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_role = LitellmUserRoles.PROXY_ADMIN
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
client = TestClient(app)
|
||
frozen_now = datetime(2024, 1, 1, 7, 0, tzinfo=timezone.utc)
|
||
|
||
from litellm.litellm_core_utils.get_model_cost_map import ModelCostMapReloaded
|
||
|
||
original_model_cost = litellm.model_cost.copy()
|
||
try:
|
||
with (
|
||
patch(
|
||
"litellm.litellm_core_utils.get_model_cost_map.refetch_model_cost_map", new_callable=AsyncMock
|
||
) as mock_get_map,
|
||
patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma,
|
||
patch("litellm.proxy.proxy_server.utc_now", return_value=frozen_now),
|
||
):
|
||
mock_get_map.return_value = ModelCostMapReloaded(
|
||
model_cost_map={"gpt-4": {"input_cost_per_token": 0.001}}
|
||
)
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(
|
||
return_value=_reload_schedule_row({}, reload_revision=9)
|
||
)
|
||
|
||
response = client.post("/reload/model_cost_map")
|
||
assert response.status_code == 200
|
||
|
||
mock_prisma.db.litellm_config.find_unique.assert_not_called()
|
||
call_args = mock_prisma.db.litellm_config.upsert.call_args
|
||
assert call_args[1]["data"]["update"] == {
|
||
"last_run_at": frozen_now,
|
||
"reload_revision": {"increment": 1},
|
||
}
|
||
assert call_args[1]["data"]["create"] == {
|
||
"param_name": "model_cost_map_reload_config",
|
||
"last_run_at": frozen_now,
|
||
"reload_revision": 1,
|
||
}
|
||
assert proxy_server_module.proxy_config.model_cost_map_loaded_at == frozen_now
|
||
assert proxy_server_module.proxy_config.model_cost_map_applied_revision == 9, (
|
||
"the serving pod must adopt the revision it published, not reload again"
|
||
)
|
||
finally:
|
||
litellm.model_cost = original_model_cost
|
||
_invalidate_model_cost_lowercase_map()
|
||
|
||
def test_anthropic_beta_headers_reload_preserves_interval_hours(self):
|
||
"""Test that _check_and_reload_anthropic_beta_headers preserves interval_hours after reload.
|
||
|
||
Regression test: the update branch of the upsert was dropping interval_hours,
|
||
identical to the model cost map bug.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
from litellm.proxy.utils import litellm_config_cache
|
||
|
||
litellm_config_cache.flush_cache()
|
||
proxy_config = ProxyConfig()
|
||
mock_prisma = MagicMock()
|
||
|
||
# Set up config with interval_hours=12 and force_reload=True to trigger reload
|
||
mock_config = MagicMock()
|
||
mock_config.param_value = {"interval_hours": 12, "force_reload": True}
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_config)
|
||
# _check_and_reload_anthropic_beta_headers now reads through get_generic_data.
|
||
mock_prisma.get_generic_data = AsyncMock(return_value=mock_config)
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=_reload_schedule_row({}, reload_revision=1))
|
||
|
||
with patch("litellm.anthropic_beta_headers_manager.reload_beta_headers_config") as mock_reload:
|
||
mock_reload.return_value = {"anthropic": {"beta_header": "test-value"}}
|
||
|
||
asyncio.run(proxy_config._check_and_reload_anthropic_beta_headers(mock_prisma))
|
||
|
||
# Verify the upsert update branch preserves interval_hours
|
||
mock_prisma.db.litellm_config.upsert.assert_called()
|
||
call_args = mock_prisma.db.litellm_config.upsert.call_args
|
||
param_value_json = call_args[1]["data"]["update"]["param_value"]
|
||
param_value_dict = json.loads(param_value_json)
|
||
assert param_value_dict["force_reload"] == False
|
||
assert param_value_dict["interval_hours"] == 12, (
|
||
"interval_hours must be preserved in the update branch; "
|
||
"dropping it causes the schedule to self-destruct"
|
||
)
|
||
|
||
def test_anthropic_beta_headers_manual_reload_preserves_interval_hours(self):
|
||
"""Test that manual reload via /reload/anthropic_beta_headers preserves existing interval_hours.
|
||
|
||
Regression test: the manual reload endpoint was overwriting param_value with
|
||
only force_reload=True, dropping any existing interval_hours schedule.
|
||
"""
|
||
from litellm.proxy._types import LitellmUserRoles
|
||
from litellm.proxy.proxy_server import cleanup_router_config_variables
|
||
|
||
cleanup_router_config_variables()
|
||
filepath = os.path.dirname(os.path.abspath(__file__))
|
||
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
|
||
asyncio.run(initialize(config=config_fp, debug=True))
|
||
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_role = LitellmUserRoles.PROXY_ADMIN
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
client = TestClient(app)
|
||
|
||
with patch("litellm.anthropic_beta_headers_manager.reload_beta_headers_config") as mock_reload:
|
||
mock_reload.return_value = {"anthropic": {"beta_header": "test-value"}}
|
||
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
# Simulate existing config with a schedule
|
||
mock_existing = MagicMock()
|
||
mock_existing.param_value = {"interval_hours": 8, "force_reload": False}
|
||
mock_prisma.db.litellm_config.find_unique = AsyncMock(return_value=mock_existing)
|
||
mock_prisma.db.litellm_config.upsert = AsyncMock(
|
||
return_value=_reload_schedule_row({}, reload_revision=1)
|
||
)
|
||
|
||
response = client.post("/reload/anthropic_beta_headers")
|
||
assert response.status_code == 200
|
||
|
||
# Verify interval_hours was preserved in the upsert
|
||
mock_prisma.db.litellm_config.upsert.assert_called()
|
||
call_args = mock_prisma.db.litellm_config.upsert.call_args
|
||
param_value_json = call_args[1]["data"]["update"]["param_value"]
|
||
param_value_dict = json.loads(param_value_json)
|
||
assert param_value_dict["force_reload"] == True
|
||
assert param_value_dict["interval_hours"] == 8, (
|
||
"interval_hours must be preserved when manual reload sets force_reload; "
|
||
"dropping it destroys any existing schedule"
|
||
)
|
||
|
||
def test_config_file_parsing(self):
|
||
"""Test parsing of config file with reload settings"""
|
||
config_content = """
|
||
general_settings:
|
||
master_key: sk-1234
|
||
model_cost_map_reload_interval: 21600
|
||
|
||
model_list:
|
||
- model_name: gpt-3.5-turbo
|
||
litellm_params:
|
||
model: gpt-3.5-turbo
|
||
- model_name: gpt-4
|
||
litellm_params:
|
||
model: gpt-4
|
||
"""
|
||
|
||
# Parse the config
|
||
config = yaml.safe_load(config_content)
|
||
|
||
# Verify the reload setting is present
|
||
assert "general_settings" in config
|
||
assert "model_cost_map_reload_interval" in config["general_settings"]
|
||
assert config["general_settings"]["model_cost_map_reload_interval"] == 21600
|
||
|
||
# Verify models are present
|
||
assert "model_list" in config
|
||
assert len(config["model_list"]) == 2
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_router_settings_from_db_config_merge_logic():
|
||
"""
|
||
Test the _add_router_settings_from_db_config method's merge logic.
|
||
|
||
This tests how router settings from config file and database are combined,
|
||
including scenarios where nested dictionaries should be properly merged.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
# Create ProxyConfig instance
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Mock router
|
||
mock_router = MagicMock()
|
||
mock_router.update_settings = MagicMock()
|
||
|
||
# Test Case 1: Both config and DB settings exist - should merge them
|
||
config_data = {
|
||
"router_settings": {
|
||
"routing_strategy": "usage-based-routing",
|
||
"model_group_alias": {"gpt-4": "openai-gpt-4"},
|
||
"enable_pre_call_checks": True,
|
||
"timeout": 30,
|
||
"nested_config": {"setting1": "config_value1", "setting2": "config_value2"},
|
||
}
|
||
}
|
||
|
||
# Mock database config record
|
||
mock_db_config = MagicMock()
|
||
mock_db_config.param_value = {
|
||
"routing_strategy": "least-busy", # This should override config value
|
||
"retry_delay": 2, # This is new, should be added
|
||
"nested_config": {
|
||
"setting2": "db_value2", # This should override config value
|
||
"setting3": "db_value3", # This is new, should be added
|
||
},
|
||
}
|
||
|
||
# Mock prisma client
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
|
||
|
||
# Call the method under test
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data=config_data,
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
)
|
||
|
||
# Verify find_first was called with correct parameters
|
||
mock_prisma_client.db.litellm_config.find_first.assert_called_once_with(where={"param_name": "router_settings"})
|
||
|
||
# Verify update_settings was called
|
||
mock_router.update_settings.assert_called_once()
|
||
|
||
# Get the actual settings passed to update_settings
|
||
call_args = mock_router.update_settings.call_args
|
||
combined_settings = call_args[1] # kwargs
|
||
|
||
# Verify the merge results
|
||
# DB values should override config values
|
||
assert combined_settings["routing_strategy"] == "least-busy"
|
||
|
||
# Config-only values should be preserved
|
||
assert combined_settings["model_group_alias"] == {"gpt-4": "openai-gpt-4"}
|
||
assert combined_settings["enable_pre_call_checks"] == True
|
||
assert combined_settings["timeout"] == 30
|
||
|
||
# DB-only values should be added
|
||
assert combined_settings["retry_delay"] == 2
|
||
|
||
# Nested dictionaries should be merged (but this is shallow merge)
|
||
expected_nested = {
|
||
"setting1": "config_value1",
|
||
"setting2": "db_value2",
|
||
"setting3": "db_value3",
|
||
}
|
||
assert combined_settings["nested_config"] == expected_nested
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_router_settings_from_db_config_empty_db_lists_do_not_clobber_config_fallbacks():
|
||
"""
|
||
Regression test for DB router_settings rows carrying explicit empty lists
|
||
(e.g. {"fallbacks": []} written by the dashboard's delete-last-fallback flow):
|
||
empty lists are "no value" and must not clobber config.yaml fallbacks,
|
||
matching _deep_merge_dicts semantics. Non-empty DB lists still win.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_router = MagicMock()
|
||
mock_router.update_settings = MagicMock()
|
||
|
||
config_data = {
|
||
"router_settings": {
|
||
"fallbacks": [{"gpt-oss-120b": ["granite-4-h-small"]}],
|
||
"context_window_fallbacks": [{"gpt-oss-120b": ["granite-4-h-small"]}],
|
||
"content_policy_fallbacks": [{"gpt-oss-120b": ["granite-4-h-small"]}],
|
||
}
|
||
}
|
||
|
||
mock_db_config = MagicMock()
|
||
mock_db_config.param_value = {
|
||
"fallbacks": [],
|
||
"context_window_fallbacks": [],
|
||
"content_policy_fallbacks": [{"gpt-oss-120b": ["other-model"]}],
|
||
"model_group_alias": {},
|
||
"num_retries": 3,
|
||
}
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
|
||
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data=config_data,
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
)
|
||
|
||
combined_settings = mock_router.update_settings.call_args.kwargs
|
||
assert combined_settings["fallbacks"] == [{"gpt-oss-120b": ["granite-4-h-small"]}]
|
||
assert combined_settings["context_window_fallbacks"] == [{"gpt-oss-120b": ["granite-4-h-small"]}]
|
||
assert combined_settings["content_policy_fallbacks"] == [{"gpt-oss-120b": ["other-model"]}]
|
||
assert combined_settings["num_retries"] == 3
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_router_settings_from_db_config_empty_db_list_still_clears_unconfigured_key():
|
||
"""
|
||
An empty DB list only yields to config.yaml where the yaml configures that key.
|
||
When the yaml router_settings has no fallbacks, a DB {"fallbacks": []} (the
|
||
dashboard's delete-last-fallback write) must still reach the router so the
|
||
running pods drop the deleted fallback without a restart.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_router = MagicMock()
|
||
mock_router.update_settings = MagicMock()
|
||
|
||
config_data = {"router_settings": {"num_retries": 1}}
|
||
|
||
mock_db_config = MagicMock()
|
||
mock_db_config.param_value = {"fallbacks": [], "model_group_alias": {}}
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
|
||
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data=config_data,
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
)
|
||
|
||
combined_settings = mock_router.update_settings.call_args.kwargs
|
||
assert combined_settings["fallbacks"] == []
|
||
assert combined_settings["num_retries"] == 1
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_router_settings_from_db_config_edge_cases():
|
||
"""
|
||
Test edge cases for _add_router_settings_from_db_config method.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_router = MagicMock()
|
||
mock_router.update_settings = MagicMock()
|
||
|
||
# Test Case 1: No router provided
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data={"router_settings": {"test": "value"}},
|
||
llm_router=None,
|
||
prisma_client=MagicMock(),
|
||
)
|
||
# Should not call anything when router is None
|
||
mock_router.update_settings.assert_not_called()
|
||
|
||
# Test Case 2: No prisma client provided
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data={"router_settings": {"test": "value"}},
|
||
llm_router=mock_router,
|
||
prisma_client=None,
|
||
)
|
||
# Should not call anything when prisma_client is None
|
||
mock_router.update_settings.assert_not_called()
|
||
|
||
# Test Case 3: DB returns None (no router_settings in DB)
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||
|
||
config_data = {"router_settings": {"routing_strategy": "usage-based"}}
|
||
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data=config_data,
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
)
|
||
|
||
# Should use only config settings
|
||
mock_router.update_settings.assert_called_once_with(routing_strategy="usage-based")
|
||
mock_router.reset_mock()
|
||
|
||
# Test Case 4: Config has no router_settings
|
||
mock_db_config = MagicMock()
|
||
mock_db_config.param_value = {"db_setting": "db_value"}
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
|
||
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data={}, # No router_settings in config
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
)
|
||
|
||
# Should use only DB settings
|
||
mock_router.update_settings.assert_called_once_with(db_setting="db_value")
|
||
mock_router.reset_mock()
|
||
|
||
# Test Case 5: Both config and DB router_settings are None/empty
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data={}, llm_router=mock_router, prisma_client=mock_prisma_client
|
||
)
|
||
|
||
# Should not call update_settings when no settings exist
|
||
mock_router.update_settings.assert_not_called()
|
||
|
||
# Test Case 6: DB config exists but param_value is not a dict
|
||
mock_db_config_invalid = MagicMock()
|
||
mock_db_config_invalid.param_value = "not_a_dict"
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config_invalid)
|
||
|
||
config_data = {"router_settings": {"config_setting": "config_value"}}
|
||
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data=config_data,
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
)
|
||
|
||
# Should use only config settings when DB param_value is invalid
|
||
mock_router.update_settings.assert_called_once_with(config_setting="config_value")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_add_router_settings_shallow_merge_behavior():
|
||
"""
|
||
Test that the merge behavior is shallow (nested dicts get replaced, not merged).
|
||
This documents the current behavior using _update_dictionary.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_router = MagicMock()
|
||
mock_router.update_settings = MagicMock()
|
||
|
||
# Config with nested dictionary
|
||
config_data = {
|
||
"router_settings": {
|
||
"nested_setting": {
|
||
"key1": "config_value1",
|
||
"key2": "config_value2",
|
||
"key3": "config_value3",
|
||
},
|
||
"top_level": "config_top",
|
||
}
|
||
}
|
||
|
||
# DB config that partially overlaps the nested dictionary
|
||
mock_db_config = MagicMock()
|
||
mock_db_config.param_value = {
|
||
"nested_setting": {
|
||
"key2": "db_value2", # Override existing key
|
||
"key4": "db_value4", # Add new key
|
||
# Note: key1 and key3 from config will be lost due to shallow merge
|
||
},
|
||
"top_level": "db_top", # Override top level
|
||
}
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_config)
|
||
|
||
await proxy_config._add_router_settings_from_db_config(
|
||
config_data=config_data,
|
||
llm_router=mock_router,
|
||
prisma_client=mock_prisma_client,
|
||
)
|
||
|
||
# Get the merged settings
|
||
call_args = mock_router.update_settings.call_args
|
||
merged_settings = call_args[1]
|
||
|
||
# Verify shallow merge behavior:
|
||
# The entire nested_setting dict from config is replaced by the DB version
|
||
expected_nested = {
|
||
"key1": "config_value1",
|
||
"key3": "config_value3",
|
||
"key2": "db_value2",
|
||
"key4": "db_value4",
|
||
}
|
||
|
||
assert merged_settings["nested_setting"] == expected_nested
|
||
assert merged_settings["top_level"] == "db_top"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_model_info_v1_oci_secrets_not_leaked():
|
||
"""
|
||
Test that model_info_v1 endpoint properly masks OCI sensitive parameters and does not leak secrets.
|
||
"""
|
||
from unittest.mock import MagicMock, patch
|
||
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import model_info_v1
|
||
|
||
# Mock user authentication
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_user_api_key_dict.user_id = "test-user"
|
||
mock_user_api_key_dict.api_key = "test-key"
|
||
mock_user_api_key_dict.team_models = []
|
||
mock_user_api_key_dict.models = ["oci-grok-test"]
|
||
|
||
# Mock model data with OCI sensitive information
|
||
mock_model_data = {
|
||
"model_name": "oci-grok-test",
|
||
"litellm_params": {
|
||
"model": "oci/xai.grok-4",
|
||
"oci_key": "ocid1.api_key.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk",
|
||
"oci_region": "us-phoenix-1",
|
||
"oci_user": "ocid1.user.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk",
|
||
"oci_fingerprint": "aa:bb:cc:dd:ee:ff:11:22:33:44:55:66:77:88:99:00",
|
||
"oci_tenancy": "ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk",
|
||
"oci_key_file": "/path/to/oci_api_key.pem",
|
||
"oci_compartment_id": "ocid1.compartment.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk",
|
||
"drop_params": True,
|
||
},
|
||
"model_info": {"mode": "completion", "id": "test-model-id"},
|
||
}
|
||
|
||
# Mock the llm_router to return our test data
|
||
mock_router = MagicMock()
|
||
mock_router.model_list = [mock_model_data]
|
||
mock_router.get_model_names.return_value = ["oci-grok-test"]
|
||
mock_router.get_model_access_groups.return_value = {}
|
||
|
||
# Mock global variables
|
||
with (
|
||
patch("litellm.proxy.proxy_server.llm_router", mock_router),
|
||
patch("litellm.proxy.proxy_server.llm_model_list", [mock_model_data]),
|
||
patch("litellm.proxy.proxy_server.prisma_client", None),
|
||
patch(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"infer_model_from_keys": False},
|
||
),
|
||
patch("litellm.proxy.proxy_server.user_model", None),
|
||
):
|
||
# Call the model_info_v1 endpoint
|
||
result = await model_info_v1(user_api_key_dict=mock_user_api_key_dict, litellm_model_id=None)
|
||
|
||
# Verify the result structure
|
||
assert "data" in result
|
||
assert len(result["data"]) == 1
|
||
|
||
model_info = result["data"][0]
|
||
litellm_params = model_info["litellm_params"]
|
||
|
||
# Verify that sensitive OCI fields are masked
|
||
assert "****" in litellm_params["oci_key"], "oci_key should be masked"
|
||
assert "****" in litellm_params["oci_fingerprint"], "oci_fingerprint should be masked"
|
||
assert "****" in litellm_params["oci_tenancy"], "oci_tenancy should be masked"
|
||
assert "****" in litellm_params["oci_key_file"], "oci_key_file should be masked"
|
||
|
||
# Verify that non-sensitive fields are NOT masked
|
||
assert litellm_params["model"] == "oci/xai.grok-4", "model field should not be masked"
|
||
assert litellm_params["oci_region"] == "us-phoenix-1", "oci_region should not be masked"
|
||
assert litellm_params["drop_params"] is True, "drop_params should not be masked"
|
||
|
||
# Verify the model field specifically is not masked (this was the original issue)
|
||
assert "****" not in litellm_params["model"], "model field should never be masked"
|
||
assert litellm_params["model"].startswith("oci/"), "model should retain its full value"
|
||
|
||
# Verify that actual secret values are not present in the response
|
||
result_str = str(result)
|
||
assert "ocid1.api_key.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk" not in result_str
|
||
assert "aa:bb:cc:dd:ee:ff:11:22:33:44:55:66:77:88:99:00" not in result_str
|
||
assert "ocid1.tenancy.oc1..aaaaaaaa7kbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbkbk" not in result_str
|
||
assert "/path/to/oci_api_key.pem" not in result_str
|
||
|
||
|
||
def test_add_callback_from_db_to_in_memory_litellm_callbacks():
|
||
"""
|
||
Test that _add_callback_from_db_to_in_memory_litellm_callbacks correctly adds callbacks
|
||
for success, failure, and combined event types.
|
||
"""
|
||
from unittest.mock import MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Mock the callback manager
|
||
mock_callback_manager = MagicMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.litellm") as mock_litellm:
|
||
# Set up mock litellm attributes
|
||
mock_litellm._known_custom_logger_compatible_callbacks = []
|
||
mock_litellm.logging_callback_manager = mock_callback_manager
|
||
|
||
# Test Case 1: Add success callback
|
||
mock_success_callbacks = []
|
||
proxy_config._add_callback_from_db_to_in_memory_litellm_callbacks(
|
||
callback="prometheus",
|
||
event_types=["success"],
|
||
existing_callbacks=mock_success_callbacks,
|
||
)
|
||
mock_callback_manager.add_litellm_success_callback.assert_called_once_with("prometheus")
|
||
mock_callback_manager.reset_mock()
|
||
|
||
# Test Case 2: Add failure callback
|
||
mock_failure_callbacks = []
|
||
proxy_config._add_callback_from_db_to_in_memory_litellm_callbacks(
|
||
callback="langfuse",
|
||
event_types=["failure"],
|
||
existing_callbacks=mock_failure_callbacks,
|
||
)
|
||
mock_callback_manager.add_litellm_failure_callback.assert_called_once_with("langfuse")
|
||
mock_callback_manager.reset_mock()
|
||
|
||
# Test Case 3: Add callback for both success and failure
|
||
mock_callbacks = []
|
||
proxy_config._add_callback_from_db_to_in_memory_litellm_callbacks(
|
||
callback="s3",
|
||
event_types=["success", "failure"],
|
||
existing_callbacks=mock_callbacks,
|
||
)
|
||
mock_callback_manager.add_litellm_callback.assert_called_once_with("s3")
|
||
mock_callback_manager.reset_mock()
|
||
|
||
# Test Case 4: Don't add callback if it already exists
|
||
existing_callbacks_with_item = ["prometheus"]
|
||
proxy_config._add_callback_from_db_to_in_memory_litellm_callbacks(
|
||
callback="prometheus",
|
||
event_types=["success"],
|
||
existing_callbacks=existing_callbacks_with_item,
|
||
)
|
||
mock_callback_manager.add_litellm_success_callback.assert_not_called()
|
||
|
||
|
||
def test_should_load_db_object_with_supported_db_objects():
|
||
"""
|
||
Test _should_load_db_object method with supported_db_objects configuration.
|
||
|
||
Verifies that when supported_db_objects is set, only specified object types
|
||
are loaded from the database.
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Test Case 1: supported_db_objects not set - all objects should be loaded
|
||
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||
|
||
# Test Case 2: supported_db_objects set to only load MCP
|
||
with patch(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"supported_db_objects": ["mcp"]},
|
||
):
|
||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||
assert proxy_config._should_load_db_object(object_type="guardrails") is False
|
||
assert proxy_config._should_load_db_object(object_type="vector_stores") is False
|
||
assert proxy_config._should_load_db_object(object_type="prompts") is False
|
||
|
||
# Test Case 3: supported_db_objects set to load multiple types
|
||
with patch(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"supported_db_objects": ["mcp", "guardrails", "vector_stores"]},
|
||
):
|
||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||
assert proxy_config._should_load_db_object(object_type="prompts") is False
|
||
|
||
# Test Case 4: supported_db_objects is not a list (should default to loading all)
|
||
with patch(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"supported_db_objects": "invalid_type"},
|
||
):
|
||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||
|
||
# Test Case 5: supported_db_objects is an empty list (nothing should be loaded)
|
||
with patch(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"supported_db_objects": []},
|
||
):
|
||
assert proxy_config._should_load_db_object(object_type="models") is False
|
||
assert proxy_config._should_load_db_object(object_type="mcp") is False
|
||
assert proxy_config._should_load_db_object(object_type="guardrails") is False
|
||
|
||
# Test Case 6: Test all available object types
|
||
with patch(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{
|
||
"supported_db_objects": [
|
||
"models",
|
||
"mcp",
|
||
"guardrails",
|
||
"vector_stores",
|
||
"pass_through_endpoints",
|
||
"prompts",
|
||
"model_cost_map",
|
||
]
|
||
},
|
||
):
|
||
assert proxy_config._should_load_db_object(object_type="models") is True
|
||
assert proxy_config._should_load_db_object(object_type="mcp") is True
|
||
assert proxy_config._should_load_db_object(object_type="guardrails") is True
|
||
assert proxy_config._should_load_db_object(object_type="vector_stores") is True
|
||
assert proxy_config._should_load_db_object(object_type="pass_through_endpoints") is True
|
||
assert proxy_config._should_load_db_object(object_type="prompts") is True
|
||
assert proxy_config._should_load_db_object(object_type="model_cost_map") is True
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_tag_cache_update_called():
|
||
"""
|
||
Test that update_cache updates tag cache when tags are provided.
|
||
"""
|
||
from litellm.caching.caching import DualCache
|
||
from litellm.proxy.proxy_server import user_api_key_cache
|
||
|
||
cache = DualCache()
|
||
|
||
setattr(
|
||
litellm.proxy.proxy_server,
|
||
"user_api_key_cache",
|
||
cache,
|
||
)
|
||
|
||
mock_tag_obj = {
|
||
"tag_name": "test-tag",
|
||
"spend": 10.0,
|
||
}
|
||
|
||
with patch.object(cache, "async_get_cache", new=AsyncMock(return_value=mock_tag_obj)) as mock_get_cache:
|
||
with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache:
|
||
await litellm.proxy.proxy_server.update_cache(
|
||
token=None,
|
||
user_id=None,
|
||
end_user_id=None,
|
||
team_id=None,
|
||
response_cost=5.0,
|
||
parent_otel_span=None,
|
||
tags=["test-tag"],
|
||
)
|
||
|
||
await asyncio.sleep(0.1)
|
||
|
||
mock_get_cache.assert_awaited_once_with(key="tag:test-tag")
|
||
mock_set_cache.assert_awaited_once()
|
||
|
||
call_args = mock_set_cache.call_args
|
||
cache_list = call_args.kwargs["cache_list"]
|
||
|
||
assert len(cache_list) == 1
|
||
cache_key, cache_value = cache_list[0]
|
||
assert cache_key == "tag:test-tag"
|
||
assert cache_value["spend"] == 15.0
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_tag_cache_update_multiple_tags():
|
||
"""
|
||
Test that multiple tags are updated in cache.
|
||
"""
|
||
from litellm.caching.caching import DualCache
|
||
from litellm.proxy.proxy_server import user_api_key_cache
|
||
|
||
cache = DualCache()
|
||
|
||
setattr(
|
||
litellm.proxy.proxy_server,
|
||
"user_api_key_cache",
|
||
cache,
|
||
)
|
||
|
||
mock_tag1_obj = {"tag_name": "tag1", "spend": 10.0}
|
||
mock_tag2_obj = {"tag_name": "tag2", "spend": 20.0}
|
||
|
||
async def mock_get_cache_side_effect(key):
|
||
if key == "tag:tag1":
|
||
return mock_tag1_obj
|
||
elif key == "tag:tag2":
|
||
return mock_tag2_obj
|
||
return None
|
||
|
||
with patch.object(
|
||
cache, "async_get_cache", new=AsyncMock(side_effect=mock_get_cache_side_effect)
|
||
) as mock_get_cache:
|
||
with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache:
|
||
await litellm.proxy.proxy_server.update_cache(
|
||
token=None,
|
||
user_id=None,
|
||
end_user_id=None,
|
||
team_id=None,
|
||
response_cost=5.0,
|
||
parent_otel_span=None,
|
||
tags=["tag1", "tag2"],
|
||
)
|
||
|
||
await asyncio.sleep(0.1)
|
||
|
||
assert mock_get_cache.call_count == 2
|
||
mock_set_cache.assert_awaited_once()
|
||
|
||
call_args = mock_set_cache.call_args
|
||
cache_list = call_args.kwargs["cache_list"]
|
||
|
||
assert len(cache_list) == 2
|
||
|
||
tag_updates = {cache_key: cache_value for cache_key, cache_value in cache_list}
|
||
assert "tag:tag1" in tag_updates
|
||
assert "tag:tag2" in tag_updates
|
||
assert tag_updates["tag:tag1"]["spend"] == 15.0
|
||
assert tag_updates["tag:tag2"]["spend"] == 25.0
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_cache_pipeline_honors_user_api_key_cache_ttl():
|
||
"""
|
||
Regression for LIT-3338: the spend-update writeback must honor
|
||
``user_api_key_cache_ttl`` (configured as ``default_in_memory_ttl``) instead of
|
||
a hardcoded 60s, otherwise every priced request resets an active key's cache
|
||
entry back to 60s and the configured TTL is never observed.
|
||
"""
|
||
from litellm.caching.caching import DualCache
|
||
|
||
original_cache = litellm.proxy.proxy_server.user_api_key_cache
|
||
cache = DualCache(default_in_memory_ttl=300)
|
||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache)
|
||
try:
|
||
with patch.object(
|
||
cache,
|
||
"async_get_cache",
|
||
new=AsyncMock(return_value={"tag_name": "active-tag", "spend": 1.0}),
|
||
):
|
||
with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache:
|
||
await litellm.proxy.proxy_server.update_cache(
|
||
token=None,
|
||
user_id=None,
|
||
end_user_id=None,
|
||
team_id=None,
|
||
response_cost=5.0,
|
||
parent_otel_span=None,
|
||
tags=["active-tag"],
|
||
)
|
||
|
||
await asyncio.sleep(0.1)
|
||
|
||
mock_set_cache.assert_awaited_once()
|
||
assert mock_set_cache.call_args.kwargs["ttl"] == 300
|
||
finally:
|
||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_spend_tracking_never_writes_the_auth_object_back():
|
||
"""Spend tracking must never write the auth object back into the cache.
|
||
|
||
Writing the mutated auth object back after every priced request let a
|
||
stale copy be re-published with a fresh TTL: to shared Redis it defeated
|
||
/key/update and /key/delete across replicas, and even a local-only write
|
||
could race an invalidation and resurrect a revoked key on this worker.
|
||
Spend is tracked through the spend:key:* counters, so the auth object is
|
||
only ever written by the DB-load paths.
|
||
"""
|
||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||
|
||
original_cache = litellm.proxy.proxy_server.user_api_key_cache
|
||
cache = UserApiKeyCache()
|
||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache)
|
||
try:
|
||
hashed_token = "spend-tracking-no-writeback-token"
|
||
await cache.async_set_cache(
|
||
key=hashed_token,
|
||
value=UserAPIKeyAuth(token=hashed_token, spend=1.0),
|
||
model_type=UserAPIKeyAuth,
|
||
)
|
||
with (
|
||
patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_pipeline,
|
||
patch.object(cache, "async_set_cache", new=AsyncMock()) as mock_set,
|
||
):
|
||
await litellm.proxy.proxy_server.update_cache(
|
||
token=hashed_token,
|
||
user_id=None,
|
||
end_user_id=None,
|
||
team_id=None,
|
||
response_cost=5.0,
|
||
parent_otel_span=None,
|
||
)
|
||
pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
|
||
if pending:
|
||
await asyncio.wait(pending, timeout=5)
|
||
|
||
key_pipeline_writes = [
|
||
call
|
||
for call in mock_pipeline.call_args_list
|
||
if any(k == hashed_token for k, _ in call.kwargs["cache_list"])
|
||
]
|
||
assert key_pipeline_writes == []
|
||
mock_set.assert_not_called()
|
||
finally:
|
||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_cache_global_proxy_spend_scalar_stays_shared():
|
||
"""
|
||
The proxy-wide spend estimate must keep flowing to Redis when the spend
|
||
writeback goes per-pod: the global max_budget check reads the
|
||
``{litellm_proxy_admin_name}:spend`` cache entry between authoritative DB
|
||
reloads, so keeping it pod-local would let traffic spread across replicas
|
||
exceed the proxy budget by roughly a factor of the replica count within a
|
||
cache TTL. Sharing this scalar is safe because it carries no limits or
|
||
permissions, so it cannot resurrect an invalidated auth blob.
|
||
"""
|
||
from litellm.caching.caching import DualCache
|
||
|
||
admin_name = litellm.proxy.proxy_server.litellm_proxy_admin_name
|
||
global_key = "{}:spend".format(admin_name)
|
||
|
||
async def fake_get(key, **kwargs):
|
||
if key == "user-lit":
|
||
return {"user_id": "user-lit", "spend": 1.0}
|
||
if key == global_key:
|
||
return 10.0
|
||
return None
|
||
|
||
original_cache = litellm.proxy.proxy_server.user_api_key_cache
|
||
cache = DualCache(default_in_memory_ttl=300)
|
||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", cache)
|
||
try:
|
||
with patch.object(cache, "async_get_cache", new=AsyncMock(side_effect=fake_get)):
|
||
with patch.object(cache, "async_set_cache_pipeline", new=AsyncMock()) as mock_set_cache:
|
||
await litellm.proxy.proxy_server.update_cache(
|
||
token=None,
|
||
user_id="user-lit",
|
||
end_user_id=None,
|
||
team_id=None,
|
||
response_cost=5.0,
|
||
parent_otel_span=None,
|
||
)
|
||
|
||
pending = [t for t in asyncio.all_tasks() if t is not asyncio.current_task()]
|
||
if pending:
|
||
await asyncio.wait(pending, timeout=5)
|
||
|
||
calls = mock_set_cache.await_args_list
|
||
local_keys = [k for c in calls if c.kwargs.get("local_only") is True for k, _ in c.kwargs["cache_list"]]
|
||
shared_keys = [
|
||
k for c in calls if c.kwargs.get("local_only") is not True for k, _ in c.kwargs["cache_list"]
|
||
]
|
||
assert "user-lit" in local_keys
|
||
assert global_key not in local_keys
|
||
assert shared_keys == [global_key]
|
||
finally:
|
||
setattr(litellm.proxy.proxy_server, "user_api_key_cache", original_cache)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_sso_settings_in_db():
|
||
"""
|
||
Test that _init_sso_settings_in_db properly loads SSO settings from database,
|
||
uppercases keys, and calls _decrypt_and_set_db_env_variables.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Test Case 1: SSO settings exist in database
|
||
mock_sso_config = MagicMock()
|
||
mock_sso_config.sso_settings = {
|
||
"google_client_id": "test-client-id",
|
||
"google_client_secret": "test-client-secret",
|
||
"microsoft_client_id": "ms-client-id",
|
||
"microsoft_client_secret": "ms-client-secret",
|
||
}
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config)
|
||
|
||
# Mock _decrypt_and_set_db_env_variables
|
||
with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set:
|
||
await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client)
|
||
|
||
# Verify find_unique was called with correct parameters
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with(where={"id": "sso_config"})
|
||
|
||
# Verify _decrypt_and_set_db_env_variables was called with uppercased keys
|
||
mock_decrypt_and_set.assert_called_once()
|
||
call_args = mock_decrypt_and_set.call_args
|
||
uppercased_settings = call_args.kwargs["environment_variables"]
|
||
|
||
# Verify all keys are uppercased
|
||
assert "GOOGLE_CLIENT_ID" in uppercased_settings
|
||
assert "GOOGLE_CLIENT_SECRET" in uppercased_settings
|
||
assert "MICROSOFT_CLIENT_ID" in uppercased_settings
|
||
assert "MICROSOFT_CLIENT_SECRET" in uppercased_settings
|
||
|
||
# Verify values are preserved
|
||
assert uppercased_settings["GOOGLE_CLIENT_ID"] == "test-client-id"
|
||
assert uppercased_settings["GOOGLE_CLIENT_SECRET"] == "test-client-secret"
|
||
assert uppercased_settings["MICROSOFT_CLIENT_ID"] == "ms-client-id"
|
||
assert uppercased_settings["MICROSOFT_CLIENT_SECRET"] == "ms-client-secret"
|
||
|
||
# Verify original lowercase keys are not present
|
||
assert "google_client_id" not in uppercased_settings
|
||
assert "microsoft_client_id" not in uppercased_settings
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_sso_settings_in_db_no_settings():
|
||
"""
|
||
Test that _init_sso_settings_in_db handles the case when no SSO settings exist in database.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Mock prisma client to return None (no SSO settings)
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None)
|
||
|
||
# Mock _decrypt_and_set_db_env_variables
|
||
with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set:
|
||
await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client)
|
||
|
||
# Verify find_unique was called
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with(where={"id": "sso_config"})
|
||
|
||
# Verify _decrypt_and_set_db_env_variables was NOT called when no settings exist
|
||
mock_decrypt_and_set.assert_not_called()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_sso_settings_in_db_error_handling():
|
||
"""
|
||
Test that _init_sso_settings_in_db handles database errors gracefully.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Mock prisma client to raise an exception
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=Exception("Database connection error"))
|
||
|
||
# The method should not raise an exception, it should log it instead
|
||
try:
|
||
await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client)
|
||
# If we get here, the exception was handled properly
|
||
assert True
|
||
except Exception as e:
|
||
# The exception should be caught and logged, not propagated
|
||
pytest.fail(f"Exception should have been caught and logged, but was raised: {e}")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_sso_settings_in_db_empty_settings():
|
||
"""
|
||
Test that _init_sso_settings_in_db handles empty SSO settings dictionary.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Mock SSO config with empty settings dictionary
|
||
mock_sso_config = MagicMock()
|
||
mock_sso_config.sso_settings = {}
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_sso_config)
|
||
|
||
# Mock _decrypt_and_set_db_env_variables
|
||
with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt_and_set:
|
||
await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client)
|
||
|
||
# Verify find_unique was called
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique.assert_awaited_once_with(where={"id": "sso_config"})
|
||
|
||
# Verify _decrypt_and_set_db_env_variables was called with empty dict
|
||
mock_decrypt_and_set.assert_called_once()
|
||
call_args = mock_decrypt_and_set.call_args
|
||
uppercased_settings = call_args.kwargs["environment_variables"]
|
||
|
||
# Verify empty dictionary
|
||
assert uppercased_settings == {}
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_sso_settings_in_db_retries_on_transport_error():
|
||
"""`_init_sso_settings_in_db` self-heals across one ClientNotConnectedError
|
||
via call_with_db_reconnect_retry — mirrors the auth-path behavior so
|
||
startup/reload bursts don't spam the log."""
|
||
import prisma
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_sso_config = MagicMock()
|
||
mock_sso_config.sso_settings = {"GOOGLE_CLIENT_ID": "xxx"}
|
||
|
||
invocations: list = []
|
||
|
||
async def _flaky_find_unique(**kwargs):
|
||
invocations.append(None)
|
||
if len(invocations) == 1:
|
||
raise prisma.errors.ClientNotConnectedError()
|
||
return mock_sso_config
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=_flaky_find_unique)
|
||
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True)
|
||
mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0
|
||
mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1
|
||
|
||
with patch.object(proxy_config, "_decrypt_and_set_db_env_variables") as mock_decrypt:
|
||
await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client)
|
||
|
||
assert len(invocations) == 2
|
||
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
|
||
reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs
|
||
assert reconnect_kwargs["reason"] == "init_sso_settings_in_db_lookup_failure"
|
||
mock_decrypt.assert_called_once()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_sso_settings_in_db_propagates_when_reconnect_fails():
|
||
"""When reconnect returns False (cooldown / lock contention), the original
|
||
ClientNotConnectedError is caught by the function's `except Exception` and
|
||
logged — no retry storm, no crash."""
|
||
import prisma
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_ssoconfig.find_unique = AsyncMock(side_effect=prisma.errors.ClientNotConnectedError())
|
||
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=False)
|
||
mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0
|
||
mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1
|
||
|
||
# Should NOT raise — the function's own try/except swallows the propagated error.
|
||
await proxy_config._init_sso_settings_in_db(prisma_client=mock_prisma_client)
|
||
|
||
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_hashicorp_vault_config_override_retries_on_transport_error():
|
||
"""`_init_hashicorp_vault_config_override` self-heals across one
|
||
ClientNotConnectedError via call_with_db_reconnect_retry."""
|
||
import prisma
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
proxy_config._last_hashicorp_vault_config = None
|
||
|
||
invocations: list = []
|
||
|
||
async def _flaky_find_unique(**kwargs):
|
||
invocations.append(None)
|
||
if len(invocations) == 1:
|
||
raise prisma.errors.ClientNotConnectedError()
|
||
return None # No config in DB → function returns early after retry.
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_configoverrides.find_unique = AsyncMock(side_effect=_flaky_find_unique)
|
||
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True)
|
||
mock_prisma_client._db_auth_reconnect_timeout_seconds = 2.0
|
||
mock_prisma_client._db_auth_reconnect_lock_timeout_seconds = 0.1
|
||
|
||
await proxy_config._init_hashicorp_vault_config_override(prisma_client=mock_prisma_client)
|
||
|
||
assert len(invocations) == 2
|
||
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
|
||
reconnect_kwargs = mock_prisma_client.attempt_db_reconnect.await_args.kwargs
|
||
assert reconnect_kwargs["reason"] == "init_hashicorp_vault_config_override_lookup_failure"
|
||
|
||
|
||
def test_update_config_fields_uppercases_env_vars(monkeypatch):
|
||
"""
|
||
Ensure environment variables pulled from DB are uppercased when applied so
|
||
integrations like Datadog that expect uppercase env keys can read them.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
for key in ["DD_API_KEY", "DD_SITE", "dd_api_key", "dd_site"]:
|
||
monkeypatch.delenv(key, raising=False)
|
||
|
||
proxy_config = ProxyConfig()
|
||
updated_config = proxy_config._update_config_fields(
|
||
current_config={},
|
||
param_name="environment_variables",
|
||
db_param_value={"dd_api_key": "test-api-key", "dd_site": "us5.datadoghq.com"},
|
||
)
|
||
|
||
env_vars = updated_config.get("environment_variables", {})
|
||
assert env_vars["DD_API_KEY"] == "test-api-key"
|
||
assert env_vars["DD_SITE"] == "us5.datadoghq.com"
|
||
assert os.environ.get("DD_API_KEY") == "test-api-key"
|
||
assert os.environ.get("DD_SITE") == "us5.datadoghq.com"
|
||
|
||
|
||
def test_encrypt_env_variables_for_db_is_idempotent(monkeypatch):
|
||
"""
|
||
Regression: /config/update and save_config must not stack a second
|
||
encryption layer when a caller re-submits a value that is already
|
||
ciphertext (the Admin UI reads config back from /get/config/callbacks —
|
||
which returns the stored, still-encrypted value — and re-POSTs it on the
|
||
next save). _encrypt_env_variables_for_db must yield a value that decrypts
|
||
to the original plaintext in exactly ONE layer, no matter how many times
|
||
its own output is fed back in. It must also not mutate os.environ (write
|
||
path — loading into the process env is the read path's job).
|
||
"""
|
||
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
|
||
decrypt_value_helper,
|
||
)
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key")
|
||
monkeypatch.delenv("LANGFUSE_PUBLIC_KEY", raising=False)
|
||
|
||
proxy_config = ProxyConfig()
|
||
plaintext = "pk-langfuse-secret-value"
|
||
|
||
# First write: plaintext in -> single-encrypted out.
|
||
enc1 = proxy_config._encrypt_env_variables_for_db({"LANGFUSE_PUBLIC_KEY": plaintext})
|
||
assert enc1["LANGFUSE_PUBLIC_KEY"] != plaintext
|
||
assert decrypt_value_helper(value=enc1["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY") == plaintext
|
||
|
||
# UI round-trip: feed the ciphertext back in. Must NOT double-encrypt.
|
||
enc2 = proxy_config._encrypt_env_variables_for_db(enc1)
|
||
assert decrypt_value_helper(value=enc2["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY") == plaintext
|
||
|
||
# And again, ×3 total ciphertext re-feeds — still exactly one layer,
|
||
# never stacked, no matter how many times the UI re-saves.
|
||
enc3 = proxy_config._encrypt_env_variables_for_db(enc2)
|
||
enc4 = proxy_config._encrypt_env_variables_for_db(enc3)
|
||
for stacked in (enc3, enc4):
|
||
assert decrypt_value_helper(value=stacked["LANGFUSE_PUBLIC_KEY"], key="LANGFUSE_PUBLIC_KEY") == plaintext
|
||
|
||
# Write path must not leak the value into the process environment.
|
||
assert os.environ.get("LANGFUSE_PUBLIC_KEY") is None
|
||
|
||
|
||
def test_get_prompt_spec_for_db_prompt_with_versions():
|
||
"""
|
||
Test that _get_prompt_spec_for_db_prompt correctly converts database prompts
|
||
to PromptSpec with versioned naming convention.
|
||
"""
|
||
from unittest.mock import MagicMock
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Mock database prompt version 1
|
||
mock_prompt_v1 = MagicMock()
|
||
mock_prompt_v1.model_dump.return_value = {
|
||
"id": "uuid-1",
|
||
"prompt_id": "chat_prompt",
|
||
"version": 1,
|
||
"litellm_params": '{"prompt_id": "chat_prompt", "prompt_integration": "dotprompt", "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "v1 content"}]}',
|
||
"prompt_info": '{"prompt_type": "db"}',
|
||
"created_at": "2024-01-01T00:00:00",
|
||
"updated_at": "2024-01-01T00:00:00",
|
||
}
|
||
|
||
# Mock database prompt version 2
|
||
mock_prompt_v2 = MagicMock()
|
||
mock_prompt_v2.model_dump.return_value = {
|
||
"id": "uuid-2",
|
||
"prompt_id": "chat_prompt",
|
||
"version": 2,
|
||
"litellm_params": '{"prompt_id": "chat_prompt", "prompt_integration": "dotprompt", "model": "gpt-4", "messages": [{"role": "user", "content": "v2 content"}]}',
|
||
"prompt_info": '{"prompt_type": "db"}',
|
||
"created_at": "2024-01-02T00:00:00",
|
||
"updated_at": "2024-01-02T00:00:00",
|
||
}
|
||
|
||
# Test version 1
|
||
prompt_spec_v1 = proxy_config._get_prompt_spec_for_db_prompt(db_prompt=mock_prompt_v1)
|
||
assert prompt_spec_v1.prompt_id == "chat_prompt.v1"
|
||
|
||
# Test version 2
|
||
prompt_spec_v2 = proxy_config._get_prompt_spec_for_db_prompt(db_prompt=mock_prompt_v2)
|
||
assert prompt_spec_v2.prompt_id == "chat_prompt.v2"
|
||
|
||
|
||
def test_root_redirect_when_docs_url_not_root_and_redirect_url_set(monkeypatch):
|
||
from fastapi.responses import RedirectResponse
|
||
|
||
from litellm.proxy.proxy_server import cleanup_router_config_variables
|
||
from litellm.proxy.utils import _get_docs_url
|
||
|
||
cleanup_router_config_variables()
|
||
filepath = os.path.dirname(os.path.abspath(__file__))
|
||
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
|
||
# Ensure docs are mounted on a non-root path to trigger redirect logic
|
||
monkeypatch.setenv("DOCS_URL", "/docs")
|
||
|
||
test_redirect_url = "/ui"
|
||
monkeypatch.setenv("ROOT_REDIRECT_URL", test_redirect_url)
|
||
|
||
asyncio.run(initialize(config=config_fp, debug=True))
|
||
|
||
docs_url = _get_docs_url()
|
||
root_redirect_url = os.getenv("ROOT_REDIRECT_URL")
|
||
|
||
# Remove any existing "/" route that might interfere
|
||
routes_to_remove = []
|
||
for route in app.routes:
|
||
if hasattr(route, "path") and route.path == "/":
|
||
if hasattr(route, "methods") and "GET" in route.methods:
|
||
routes_to_remove.append(route)
|
||
elif not hasattr(route, "methods"): # Catch-all routes
|
||
routes_to_remove.append(route)
|
||
|
||
for route in routes_to_remove:
|
||
app.routes.remove(route)
|
||
|
||
# Add the redirect route if conditions are met (matching the actual implementation)
|
||
if docs_url != "/" and root_redirect_url:
|
||
|
||
@app.get("/", include_in_schema=False)
|
||
async def root_redirect():
|
||
return RedirectResponse(url=root_redirect_url)
|
||
|
||
client = TestClient(app)
|
||
response = client.get("/", follow_redirects=False)
|
||
assert response.status_code == 307
|
||
assert response.headers["location"] == test_redirect_url
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_image_non_root_uses_var_lib_assets_dir(monkeypatch):
|
||
"""
|
||
Test that get_image uses /var/lib/litellm/assets when LITELLM_NON_ROOT is true.
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
from litellm.proxy.proxy_server import get_image
|
||
|
||
# Set LITELLM_NON_ROOT to true
|
||
monkeypatch.setenv("LITELLM_NON_ROOT", "true")
|
||
monkeypatch.delenv("UI_LOGO_PATH", raising=False)
|
||
|
||
# Mock os.path operations - exists=False for assets_dir so makedirs gets called
|
||
def exists_side_effect(path):
|
||
return False if path == "/var/lib/litellm/assets" else True
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs,
|
||
patch("litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect),
|
||
patch("litellm.proxy.proxy_server.os.access", return_value=True),
|
||
patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv,
|
||
patch("litellm.proxy.proxy_server.FileResponse") as mock_file_response,
|
||
):
|
||
# Setup mock_getenv to return empty string for UI_LOGO_PATH
|
||
def getenv_side_effect(key, default=""):
|
||
if key == "UI_LOGO_PATH":
|
||
return ""
|
||
elif key == "LITELLM_NON_ROOT":
|
||
return "true"
|
||
return default
|
||
|
||
mock_getenv.side_effect = getenv_side_effect
|
||
|
||
# Call the function
|
||
await get_image()
|
||
|
||
# Verify makedirs was called with /var/lib/litellm/assets
|
||
mock_makedirs.assert_called_once_with("/var/lib/litellm/assets", exist_ok=True)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_image_non_root_fallback_to_default_logo(monkeypatch):
|
||
"""
|
||
Test that get_image falls back to default_site_logo when logo doesn't exist
|
||
in /var/lib/litellm/assets for non-root case.
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
from litellm.proxy.proxy_server import get_image
|
||
|
||
# Set LITELLM_NON_ROOT to true
|
||
monkeypatch.setenv("LITELLM_NON_ROOT", "true")
|
||
monkeypatch.delenv("UI_LOGO_PATH", raising=False)
|
||
|
||
# Track path.exists calls to verify it checks /var/lib/litellm/assets/logo.jpg
|
||
exists_calls = []
|
||
|
||
def exists_side_effect(path):
|
||
exists_calls.append(path)
|
||
# Return False for /var/lib/litellm/assets* so: makedirs is called, logo fallback
|
||
# triggers, and we don't return early with cached file
|
||
if "/var/lib/litellm/assets" in path:
|
||
return False
|
||
return True
|
||
|
||
# Mock os.path operations
|
||
with (
|
||
patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs,
|
||
patch("litellm.proxy.proxy_server.os.path.exists", side_effect=exists_side_effect),
|
||
patch("litellm.proxy.proxy_server.os.access", return_value=True),
|
||
patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv,
|
||
patch("litellm.proxy.proxy_server.FileResponse") as mock_file_response,
|
||
):
|
||
# Setup mock_getenv
|
||
def getenv_side_effect(key, default=""):
|
||
if key == "UI_LOGO_PATH":
|
||
return ""
|
||
elif key == "LITELLM_NON_ROOT":
|
||
return "true"
|
||
return default
|
||
|
||
mock_getenv.side_effect = getenv_side_effect
|
||
|
||
# Call the function
|
||
await get_image()
|
||
|
||
# Verify makedirs was called with /var/lib/litellm/assets
|
||
mock_makedirs.assert_called_once_with("/var/lib/litellm/assets", exist_ok=True)
|
||
|
||
# Verify that exists was called to check /var/lib/litellm/assets/logo.jpg
|
||
assets_logo_path = "/var/lib/litellm/assets/logo.jpg"
|
||
assert any(assets_logo_path in str(call) for call in exists_calls), f"Should check if {assets_logo_path} exists"
|
||
|
||
# Verify FileResponse was called (with fallback logo)
|
||
assert mock_file_response.called, "FileResponse should be called"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_image_root_case_uses_current_dir(monkeypatch):
|
||
"""
|
||
Test that get_image uses current_dir when LITELLM_NON_ROOT is not true.
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
from litellm.proxy.proxy_server import get_image
|
||
|
||
# Don't set LITELLM_NON_ROOT (or set it to false)
|
||
monkeypatch.delenv("LITELLM_NON_ROOT", raising=False)
|
||
monkeypatch.delenv("UI_LOGO_PATH", raising=False)
|
||
|
||
# Mock os.path operations
|
||
with (
|
||
patch("litellm.proxy.proxy_server.os.makedirs") as mock_makedirs,
|
||
patch("litellm.proxy.proxy_server.os.path.exists", return_value=True),
|
||
patch("litellm.proxy.proxy_server.os.getenv") as mock_getenv,
|
||
patch("litellm.proxy.proxy_server.FileResponse") as mock_file_response,
|
||
):
|
||
# Setup mock_getenv
|
||
def getenv_side_effect(key, default=""):
|
||
if key == "UI_LOGO_PATH":
|
||
return ""
|
||
elif key == "LITELLM_NON_ROOT":
|
||
return "" # Not set or empty
|
||
return default
|
||
|
||
mock_getenv.side_effect = getenv_side_effect
|
||
|
||
# Call the function
|
||
await get_image()
|
||
|
||
# Verify makedirs was NOT called with /var/lib/litellm/assets (should not create it for root case)
|
||
var_lib_assets_calls = [call for call in mock_makedirs.call_args_list if "/var/lib/litellm/assets" in str(call)]
|
||
assert len(var_lib_assets_calls) == 0, "Should not create /var/lib/litellm/assets for root case"
|
||
|
||
# Verify FileResponse was called
|
||
assert mock_file_response.called, "FileResponse should be called"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_image_custom_local_logo_bypasses_cache(monkeypatch, tmp_path):
|
||
"""
|
||
Test that when UI_LOGO_PATH is set to a local file, get_image serves it
|
||
directly and does not return a stale cached_logo.jpg.
|
||
|
||
Regression test: previously the cache check ran before reading UI_LOGO_PATH,
|
||
so a pre-existing cached_logo.jpg (e.g. from the base Docker image) would
|
||
always be returned, ignoring the user's custom logo.
|
||
"""
|
||
from litellm.proxy.proxy_server import get_image
|
||
|
||
custom_logo = tmp_path / "custom_logo.jpg"
|
||
custom_logo.write_bytes(b"\xff\xd8\xff custom logo")
|
||
monkeypatch.setenv("UI_LOGO_PATH", str(custom_logo))
|
||
monkeypatch.delenv("LITELLM_NON_ROOT", raising=False)
|
||
monkeypatch.delenv("LITELLM_ASSETS_PATH", raising=False)
|
||
|
||
calls_to_file_response = []
|
||
|
||
def fake_file_response(path, **kwargs):
|
||
calls_to_file_response.append(path)
|
||
return MagicMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response),
|
||
):
|
||
await get_image()
|
||
|
||
assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once"
|
||
assert calls_to_file_response[0] == str(custom_logo.resolve()), (
|
||
f"Expected custom logo path, got {calls_to_file_response[0]}. "
|
||
"A stale cached_logo.jpg may have been returned instead."
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_image_default_logo_ignores_stale_cache(monkeypatch, tmp_path):
|
||
"""
|
||
Test that when UI_LOGO_PATH is NOT set, stale pre-fix cached_logo.jpg
|
||
files are ignored and the default logo is served.
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
from litellm.proxy.proxy_server import get_image
|
||
|
||
cache_path = tmp_path / "cached_logo.jpg"
|
||
cache_path.write_bytes(b"\xff\xd8\xff cached logo")
|
||
monkeypatch.delenv("UI_LOGO_PATH", raising=False)
|
||
monkeypatch.delenv("LITELLM_NON_ROOT", raising=False)
|
||
monkeypatch.setenv("LITELLM_ASSETS_PATH", str(tmp_path))
|
||
|
||
calls_to_file_response = []
|
||
|
||
def fake_file_response(path, **kwargs):
|
||
calls_to_file_response.append(path)
|
||
return MagicMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response),
|
||
):
|
||
await get_image()
|
||
|
||
assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once"
|
||
served_path = calls_to_file_response[0]
|
||
assert served_path != str(cache_path.resolve())
|
||
assert served_path.endswith("logo.jpg")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_image_custom_logo_missing_falls_through_to_default(monkeypatch, tmp_path):
|
||
"""
|
||
Test that when UI_LOGO_PATH points to a non-existent local file,
|
||
get_image falls through to the default logo instead of failing.
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
from litellm.proxy.proxy_server import get_image
|
||
|
||
custom_logo_path = tmp_path / "nonexistent_logo.jpg"
|
||
monkeypatch.setenv("UI_LOGO_PATH", str(custom_logo_path))
|
||
monkeypatch.delenv("LITELLM_NON_ROOT", raising=False)
|
||
monkeypatch.setenv("LITELLM_ASSETS_PATH", str(tmp_path))
|
||
|
||
calls_to_file_response = []
|
||
|
||
def fake_file_response(path, **kwargs):
|
||
calls_to_file_response.append(path)
|
||
return MagicMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response),
|
||
):
|
||
await get_image()
|
||
|
||
assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once"
|
||
served_path = calls_to_file_response[0]
|
||
assert served_path != str(custom_logo_path), "Should not attempt to serve a non-existent custom logo"
|
||
assert served_path.endswith("logo.jpg")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_image_custom_logo_missing_no_cache_serves_default(monkeypatch, tmp_path):
|
||
"""
|
||
Test that when UI_LOGO_PATH points to a non-existent file AND there is no
|
||
cached_logo.jpg, get_image serves the default logo instead of the non-existent
|
||
custom path.
|
||
"""
|
||
from unittest.mock import patch
|
||
|
||
from litellm.proxy.proxy_server import get_image
|
||
|
||
custom_logo_path = tmp_path / "nonexistent_logo.jpg"
|
||
monkeypatch.setenv("UI_LOGO_PATH", str(custom_logo_path))
|
||
monkeypatch.delenv("LITELLM_NON_ROOT", raising=False)
|
||
monkeypatch.setenv("LITELLM_ASSETS_PATH", str(tmp_path))
|
||
|
||
calls_to_file_response = []
|
||
|
||
def fake_file_response(path, **kwargs):
|
||
calls_to_file_response.append(path)
|
||
return MagicMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response),
|
||
):
|
||
await get_image()
|
||
|
||
assert len(calls_to_file_response) == 1, "FileResponse should be called exactly once"
|
||
served_path = calls_to_file_response[0]
|
||
assert served_path != str(custom_logo_path), "Should not attempt to serve a non-existent custom logo"
|
||
assert served_path.endswith("logo.jpg"), f"Expected fallback to default logo.jpg, got {served_path}"
|
||
|
||
|
||
def test_get_config_normalizes_string_callbacks(monkeypatch):
|
||
"""
|
||
Test that /get/config/callbacks normalizes string callbacks to lists.
|
||
"""
|
||
from litellm.proxy.proxy_server import app, proxy_config, user_api_key_auth
|
||
|
||
config_data = {
|
||
"litellm_settings": {
|
||
"success_callback": "langfuse",
|
||
"failure_callback": None,
|
||
"callbacks": ["prometheus", "datadog"],
|
||
},
|
||
"general_settings": {},
|
||
"environment_variables": {},
|
||
}
|
||
|
||
mock_router = MagicMock()
|
||
mock_router.get_settings.return_value = {}
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_router)
|
||
monkeypatch.setattr(proxy_config, "get_config", AsyncMock(return_value=config_data))
|
||
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-1234"
|
||
)
|
||
|
||
client = TestClient(app)
|
||
try:
|
||
response = client.get("/get/config/callbacks")
|
||
finally:
|
||
app.dependency_overrides = original_overrides
|
||
|
||
assert response.status_code == 200
|
||
callbacks = response.json()["callbacks"]
|
||
|
||
success_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "success"]
|
||
failure_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "failure"]
|
||
success_and_failure_callbacks = [cb["name"] for cb in callbacks if cb.get("type") == "success_and_failure"]
|
||
|
||
assert "langfuse" in success_callbacks
|
||
assert len(failure_callbacks) == 0
|
||
assert "prometheus" in success_and_failure_callbacks
|
||
assert "datadog" in success_and_failure_callbacks
|
||
|
||
|
||
def test_deep_merge_dicts_skips_none_and_empty_lists(monkeypatch):
|
||
"""
|
||
Test that _update_config_fields deep merge skips None values and empty lists.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
current_config = {
|
||
"general_settings": {
|
||
"max_parallel_requests": 10,
|
||
"allowed_models": ["gpt-3.5-turbo", "gpt-4"],
|
||
"nested": {
|
||
"key1": "value1",
|
||
"key2": "value2",
|
||
},
|
||
}
|
||
}
|
||
|
||
db_param_value = {
|
||
"max_parallel_requests": None,
|
||
"allowed_models": [],
|
||
"new_key": "new_value",
|
||
"nested": {
|
||
"key1": "updated_value1",
|
||
"key3": "value3",
|
||
},
|
||
}
|
||
|
||
result = proxy_config._update_config_fields(current_config, "general_settings", db_param_value)
|
||
|
||
assert result["general_settings"]["max_parallel_requests"] == 10
|
||
assert result["general_settings"]["allowed_models"] == ["gpt-3.5-turbo", "gpt-4"]
|
||
assert result["general_settings"]["new_key"] == "new_value"
|
||
assert result["general_settings"]["nested"]["key1"] == "updated_value1"
|
||
assert result["general_settings"]["nested"]["key2"] == "value2"
|
||
assert result["general_settings"]["nested"]["key3"] == "value3"
|
||
|
||
|
||
class TestInvitationEndpoints:
|
||
"""Tests for /invitation/new and /invitation/delete endpoints."""
|
||
|
||
@pytest.fixture
|
||
def client_with_auth(self):
|
||
"""Create a test client with admin authentication."""
|
||
from litellm.proxy._types import LitellmUserRoles
|
||
from litellm.proxy.proxy_server import cleanup_router_config_variables
|
||
|
||
cleanup_router_config_variables()
|
||
filepath = os.path.dirname(os.path.abspath(__file__))
|
||
config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml"
|
||
asyncio.run(initialize(config=config_fp, debug=True))
|
||
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_id = "admin-user-id"
|
||
mock_auth.user_role = LitellmUserRoles.PROXY_ADMIN
|
||
mock_auth.api_key = "sk-test"
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
|
||
return TestClient(app)
|
||
|
||
@pytest.mark.parametrize(
|
||
"endpoint,payload,mock_return",
|
||
[
|
||
(
|
||
"/invitation/new",
|
||
{"user_id": "target-user-123"},
|
||
{
|
||
"id": "inv-123",
|
||
"user_id": "target-user-123",
|
||
"is_accepted": False,
|
||
"accepted_at": None,
|
||
"expires_at": "2025-02-18T00:00:00",
|
||
"created_at": "2025-02-11T00:00:00",
|
||
"created_by": "admin-user-id",
|
||
"updated_at": "2025-02-11T00:00:00",
|
||
"updated_by": "admin-user-id",
|
||
},
|
||
),
|
||
(
|
||
"/invitation/delete",
|
||
{"invitation_id": "inv-456"},
|
||
{
|
||
"id": "inv-456",
|
||
"user_id": "target-user-123",
|
||
"is_accepted": False,
|
||
"accepted_at": None,
|
||
"expires_at": "2025-02-18T00:00:00",
|
||
"created_at": "2025-02-11T00:00:00",
|
||
"created_by": "admin-user-id",
|
||
"updated_at": "2025-02-11T00:00:00",
|
||
"updated_by": "admin-user-id",
|
||
},
|
||
),
|
||
],
|
||
)
|
||
def test_invitation_endpoints_proxy_admin_success(self, client_with_auth, endpoint, payload, mock_return):
|
||
"""Proxy admin can successfully create and delete invitations."""
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_invitationlink = MagicMock()
|
||
if endpoint == "/invitation/new":
|
||
mock_create = AsyncMock(return_value=mock_return)
|
||
with patch(
|
||
"litellm.proxy.management_helpers.user_invitation.create_invitation_for_user",
|
||
mock_create,
|
||
):
|
||
response = client_with_auth.post(endpoint, json=payload)
|
||
else:
|
||
mock_prisma.db.litellm_invitationlink.find_unique = AsyncMock(
|
||
return_value={**mock_return, "created_by": "admin-user-id"}
|
||
)
|
||
mock_prisma.db.litellm_invitationlink.delete = AsyncMock(return_value=mock_return)
|
||
response = client_with_auth.post(endpoint, json=payload)
|
||
|
||
assert response.status_code == 200
|
||
data = response.json()
|
||
assert data["id"] == mock_return["id"]
|
||
assert data["user_id"] == mock_return["user_id"]
|
||
|
||
@pytest.mark.parametrize(
|
||
"endpoint,payload",
|
||
[
|
||
("/invitation/new", {"user_id": "target-user-123"}),
|
||
("/invitation/delete", {"invitation_id": "inv-456"}),
|
||
],
|
||
)
|
||
def test_invitation_endpoints_non_admin_denied(self, client_with_auth, endpoint, payload):
|
||
"""Non-admin users cannot access invitation endpoints."""
|
||
from litellm.proxy._types import LitellmUserRoles
|
||
|
||
mock_auth = MagicMock()
|
||
mock_auth.user_id = "regular-user"
|
||
mock_auth.user_role = LitellmUserRoles.INTERNAL_USER
|
||
mock_auth.api_key = "sk-regular"
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
|
||
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma:
|
||
mock_prisma.db.litellm_invitationlink = MagicMock()
|
||
# Avoid triggering async DB calls in _user_has_admin_privileges
|
||
with patch(
|
||
"litellm.proxy.proxy_server._user_has_admin_privileges",
|
||
new_callable=AsyncMock,
|
||
return_value=False,
|
||
):
|
||
response = client_with_auth.post(endpoint, json=payload)
|
||
|
||
assert response.status_code == 400
|
||
body = response.json()
|
||
# ProxyException handler returns {"error": {...}}, HTTPException returns {"detail": {...}}
|
||
error_content = body.get("error", body.get("detail", body))
|
||
assert "not allowed" in str(error_content).lower()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_cleanup_on_early_exit():
|
||
"""
|
||
Test that async_data_generator calls response.aclose() in the finally block
|
||
when the generator is abandoned mid-stream (client disconnect).
|
||
"""
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gpt-3.5-turbo",
|
||
"messages": [{"role": "user", "content": "test"}],
|
||
}
|
||
|
||
mock_chunks = [
|
||
{"choices": [{"delta": {"content": "Hello"}}]},
|
||
{"choices": [{"delta": {"content": " world"}}]},
|
||
{"choices": [{"delta": {"content": " more"}}]},
|
||
]
|
||
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
|
||
async def mock_streaming_iterator(*args, **kwargs):
|
||
for chunk in mock_chunks:
|
||
yield chunk
|
||
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = mock_streaming_iterator
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(
|
||
side_effect=lambda **kwargs: kwargs.get("response")
|
||
)
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
# Create a mock response with aclose
|
||
mock_response = MagicMock()
|
||
mock_response.aclose = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
# Consume only the first chunk then abandon the generator (simulates client disconnect)
|
||
gen = async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data)
|
||
first_chunk = await gen.__anext__()
|
||
assert first_chunk.startswith("data: ")
|
||
|
||
# Close the generator early (simulates what ASGI does on client disconnect)
|
||
await gen.aclose()
|
||
|
||
# Verify aclose was called on the response to release the HTTP connection
|
||
mock_response.aclose.assert_awaited_once()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_uses_direct_stream_fast_path_without_callbacks():
|
||
"""
|
||
When there are no streaming callbacks, async_data_generator should avoid
|
||
per-chunk hook machinery and iterate the provider stream directly.
|
||
"""
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gpt-3.5-turbo",
|
||
"messages": [{"role": "user", "content": "test"}],
|
||
}
|
||
mock_chunks = [
|
||
{"choices": [{"delta": {"content": "Hello"}}]},
|
||
{"choices": [{"delta": {"content": " world"}}]},
|
||
]
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
for chunk in mock_chunks:
|
||
yield chunk
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
with patch.object(ProxyLogging, "_fire_deferred_stream_logging") as mock_deferred_logging:
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data]
|
||
assert len([chunk for chunk in yielded_text if chunk.startswith("data: {")]) == 2
|
||
assert yielded_text[-1] == "data: [DONE]\n\n"
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook.assert_not_called()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook.assert_not_awaited()
|
||
mock_deferred_logging.assert_called_once_with(mock_request_data)
|
||
mock_response.aclose.assert_awaited_once()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_preserves_non_raw_sse_like_bytes():
|
||
"""
|
||
Already formatted SSE bytes from non-raw streams keep the legacy passthrough
|
||
behavior, including appending a missing event terminator.
|
||
"""
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gemini-2.0-flash",
|
||
"messages": [{"role": "user", "content": "test"}],
|
||
}
|
||
gemini_event = b'data: {"candidates": [{"content": "hi"}]}\n\n'
|
||
gemini_event_without_terminator = b'data: {"candidates": [{"content": "there"}]}'
|
||
raw_payload = b'{"partial": true}'
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
yield gemini_event
|
||
yield gemini_event_without_terminator
|
||
yield raw_payload
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data]
|
||
assert yielded_text[0] == gemini_event.decode("utf-8")
|
||
assert yielded_text[1] == gemini_event_without_terminator.decode("utf-8") + "\n\n"
|
||
assert yielded_text[2] == f"data: {raw_payload.decode('utf-8')}\n\n"
|
||
assert "b'data:" not in "".join(yielded_text)
|
||
assert yielded_text[-1] == "data: [DONE]\n\n"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_buffers_split_google_native_sse_json_frame():
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gemini-3.5-flash",
|
||
"_litellm_skip_openai_stream_done": True,
|
||
"_litellm_raw_sse_stream": True,
|
||
}
|
||
payload = (
|
||
'data: {"candidates": [{"content": {"role": "model", "parts": '
|
||
'[{"text": "", "thoughtSignature": "abc123def456"}]}}]}\n\n'
|
||
)
|
||
raw_chunks = [
|
||
payload[:2].encode("utf-8"),
|
||
payload[2 : payload.index("thoughtSignature") + len('thoughtSignature": "abc')].encode("utf-8"),
|
||
payload[payload.index("thoughtSignature") + len('thoughtSignature": "abc') :].encode("utf-8"),
|
||
]
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
for chunk in raw_chunks:
|
||
yield chunk
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data]
|
||
|
||
assert yielded_text == [payload]
|
||
for chunk in yielded_text:
|
||
assert chunk.endswith("\n\n")
|
||
assert json.loads(chunk.removeprefix("data: ").strip())
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_flushes_raw_sse_stream_without_trailing_delimiter():
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gemini-3.5-flash",
|
||
"_litellm_skip_openai_stream_done": True,
|
||
"_litellm_raw_sse_stream": True,
|
||
}
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
yield b'data: {"candidates": [{"content": "unterminated"}]'
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj),
|
||
patch.object(ProxyLogging, "_fire_deferred_stream_logging"),
|
||
):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data]
|
||
assert len(yielded_text) == 1
|
||
assert yielded_text[0] == 'data: {"candidates": [{"content": "unterminated"}]\n\n'
|
||
assert "[DONE]" not in yielded_text[0]
|
||
mock_proxy_logging_obj.post_call_failure_hook.assert_not_awaited()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_errors_when_raw_sse_frame_exceeds_buffer_limit():
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gemini-3.5-flash",
|
||
"_litellm_skip_openai_stream_done": True,
|
||
"_litellm_raw_sse_stream": True,
|
||
}
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
yield b"data: "
|
||
yield b'{"candidates": [{"content": "unterminated"}]'
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj),
|
||
patch("litellm.proxy.proxy_server._MAX_RAW_SSE_BUFFER_CHARS", 8),
|
||
patch.object(ProxyLogging, "_fire_deferred_stream_logging"),
|
||
):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data]
|
||
assert len(yielded_text) == 1
|
||
assert "maximum buffered size" in yielded_text[0]
|
||
assert "[DONE]" not in yielded_text[0]
|
||
mock_proxy_logging_obj.post_call_failure_hook.assert_awaited_once()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.parametrize("as_bytes", [True, False])
|
||
async def test_async_data_generator_checks_raw_sse_buffer_limit_after_complete_frames(
|
||
as_bytes,
|
||
):
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
complete_frame = 'data: {"candidates": [{"content": "ok"}]}\n\n'
|
||
partial_frame = "data: "
|
||
raw_chunk = complete_frame + partial_frame
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gemini-3.5-flash",
|
||
"_litellm_skip_openai_stream_done": True,
|
||
"_litellm_raw_sse_stream": True,
|
||
}
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
yield raw_chunk.encode("utf-8") if as_bytes else raw_chunk
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj),
|
||
patch("litellm.proxy.proxy_server._MAX_RAW_SSE_BUFFER_CHARS", 8),
|
||
patch.object(ProxyLogging, "_fire_deferred_stream_logging"),
|
||
):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data]
|
||
assert yielded_text[0] == complete_frame
|
||
assert yielded_text[1] == partial_frame + "\n\n"
|
||
assert "[DONE]" not in "".join(yielded_text)
|
||
mock_proxy_logging_obj.post_call_failure_hook.assert_not_awaited()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_google_genai_stream_omits_openai_done():
|
||
"""
|
||
google-genai SDK streamGenerateContent?alt=sse must not receive data: [DONE].
|
||
"""
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gemini-2.0-flash",
|
||
"_litellm_skip_openai_stream_done": True,
|
||
}
|
||
gemini_event = b'data: {"candidates": [{"content": {"parts": [{"text": "Hi"}]}}]}\n\n'
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
yield gemini_event
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data]
|
||
assert yielded_text == [gemini_event.decode("utf-8")]
|
||
assert "[DONE]" not in "".join(yielded_text)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_does_not_mark_completed_stream_as_disconnect():
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {"model": "gpt-4o", "metadata": {}}
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
yield {"choices": [{"delta": {"content": "done"}}]}
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_request = MagicMock()
|
||
mock_request.is_disconnected = AsyncMock(return_value=True)
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
|
||
yielded_data = []
|
||
async for data in async_data_generator(
|
||
mock_response,
|
||
mock_user_api_key_dict,
|
||
mock_request_data,
|
||
request=mock_request,
|
||
):
|
||
yielded_data.append(data)
|
||
|
||
assert yielded_data[-1] == "data: [DONE]\n\n"
|
||
mock_request.is_disconnected.assert_not_awaited()
|
||
assert "client_disconnected" not in mock_request_data["metadata"]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_google_genai_stream_forwards_error_without_done():
|
||
"""Stream errors must still reach the client when OpenAI [DONE] is skipped."""
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
error_sse = 'data: {"error": {"message": "stream failed"}}\n\n'
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gemini-2.0-flash",
|
||
"_litellm_skip_openai_stream_done": True,
|
||
}
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
yield error_sse
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_response = MockStream()
|
||
mock_response.aclose = AsyncMock()
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.has_streaming_callbacks.return_value = False
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = MagicMock()
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock()
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
with patch.object(ProxyLogging, "_fire_deferred_stream_logging"):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
yielded_text = [chunk.decode("utf-8") if isinstance(chunk, bytes) else chunk for chunk in yielded_data]
|
||
assert yielded_text == [error_sse]
|
||
assert "[DONE]" not in "".join(yielded_text)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_cleanup_on_normal_completion():
|
||
"""
|
||
Test that async_data_generator calls response.aclose() even on normal completion.
|
||
"""
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gpt-3.5-turbo",
|
||
"messages": [{"role": "user", "content": "test"}],
|
||
}
|
||
|
||
mock_chunks = [
|
||
{"choices": [{"delta": {"content": "Hello"}}]},
|
||
]
|
||
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
|
||
async def mock_streaming_iterator(*args, **kwargs):
|
||
for chunk in mock_chunks:
|
||
yield chunk
|
||
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = mock_streaming_iterator
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(
|
||
side_effect=lambda **kwargs: kwargs.get("response")
|
||
)
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
mock_response = MagicMock()
|
||
mock_response.aclose = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
# Should have completed normally with [DONE]
|
||
assert any("[DONE]" in d for d in yielded_data)
|
||
# aclose should still be called via finally block
|
||
mock_response.aclose.assert_awaited_once()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_cleanup_on_midstream_error():
|
||
"""
|
||
Test that async_data_generator calls response.aclose() via finally block
|
||
even when an exception occurs mid-stream.
|
||
"""
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_user_api_key_dict = MagicMock(spec=UserAPIKeyAuth)
|
||
mock_request_data = {
|
||
"model": "gpt-3.5-turbo",
|
||
"messages": [{"role": "user", "content": "test"}],
|
||
}
|
||
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
|
||
async def mock_streaming_iterator_with_error(*args, **kwargs):
|
||
yield {"choices": [{"delta": {"content": "Hello"}}]}
|
||
raise RuntimeError("upstream connection reset")
|
||
|
||
mock_proxy_logging_obj.async_post_call_streaming_iterator_hook = mock_streaming_iterator_with_error
|
||
mock_proxy_logging_obj.async_post_call_streaming_hook = AsyncMock(
|
||
side_effect=lambda **kwargs: kwargs.get("response")
|
||
)
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
mock_response = MagicMock()
|
||
mock_response.aclose = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
yielded_data = []
|
||
async for data in async_data_generator(mock_response, mock_user_api_key_dict, mock_request_data):
|
||
yielded_data.append(data)
|
||
|
||
# Should have yielded data chunk and then an error chunk
|
||
assert len(yielded_data) >= 2
|
||
assert any("error" in d for d in yielded_data)
|
||
# aclose must still be called via finally block despite the error
|
||
mock_response.aclose.assert_awaited_once()
|
||
|
||
|
||
# ============================================================================
|
||
# store_model_in_db DB Config Override Tests
|
||
# ============================================================================
|
||
|
||
|
||
def test_store_model_in_db_in_config_general_settings():
|
||
"""
|
||
Verify store_model_in_db is a valid field in ConfigGeneralSettings
|
||
and validates correctly for True/False values.
|
||
"""
|
||
from litellm.proxy._types import ConfigGeneralSettings
|
||
|
||
assert "store_model_in_db" in ConfigGeneralSettings.model_fields
|
||
|
||
# Should validate with True
|
||
config = ConfigGeneralSettings(store_model_in_db=True)
|
||
assert config.store_model_in_db is True
|
||
|
||
# Should validate with False
|
||
config = ConfigGeneralSettings(store_model_in_db=False)
|
||
assert config.store_model_in_db is False
|
||
|
||
# Should validate with None (default)
|
||
config = ConfigGeneralSettings(store_model_in_db=None)
|
||
assert config.store_model_in_db is None
|
||
|
||
# Should validate with no value
|
||
config = ConfigGeneralSettings()
|
||
assert config.store_model_in_db is None
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_store_model_in_db_true():
|
||
"""
|
||
Verify _update_general_settings sets global store_model_in_db to True
|
||
when DB general_settings has store_model_in_db=True.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False) as mock_store,
|
||
patch("litellm.proxy.proxy_server.general_settings", {}) as mock_gs,
|
||
):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": True})
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.store_model_in_db is True
|
||
assert ps.general_settings["store_model_in_db"] is True
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_store_model_in_db_false():
|
||
"""
|
||
Verify _update_general_settings sets global store_model_in_db to False
|
||
when DB general_settings has store_model_in_db=False.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||
):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": False})
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.store_model_in_db is False
|
||
assert ps.general_settings["store_model_in_db"] is False
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_propagates_apply_user_budget_to_team_keys():
|
||
"""The Admin UI toggle writes to the DB config, so the flag has to be in the
|
||
runtime propagation allowlist. The reverted skip_user_budget_on_team_key was
|
||
exposed in /config/list but never propagated, so its toggle did nothing."""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
||
await proxy_config._update_general_settings(db_general_settings={"apply_user_budget_to_team_keys": "true"})
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.general_settings["apply_user_budget_to_team_keys"] is True
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_propagates_spend_log_cleanup_bounds():
|
||
"""The dashboard writes the cleanup bounds straight to the DB config, so
|
||
without runtime propagation the scheduled job never sees them and the knobs
|
||
do nothing until the process restarts."""
|
||
from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import (
|
||
SPEND_LOG_CLEANUP_BOUND_SETTINGS,
|
||
)
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
db_settings = {
|
||
"maximum_spend_logs_cleanup_batch_size": 2000,
|
||
"maximum_spend_logs_cleanup_max_batches": 250,
|
||
"maximum_spend_logs_cleanup_run_budget": "90s",
|
||
"maximum_spend_logs_cleanup_batch_timeout": "10s",
|
||
}
|
||
assert set(db_settings) == set(SPEND_LOG_CLEANUP_BOUND_SETTINGS)
|
||
|
||
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
||
await proxy_config._update_general_settings(db_general_settings=db_settings)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert {key: ps.general_settings.get(key) for key in db_settings} == db_settings
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_clears_a_spend_log_cleanup_bound_dropped_from_the_db():
|
||
"""Blanking the field in the dashboard deletes the key outright, so leaving
|
||
the last value in memory would keep a bound the operator just removed."""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
with patch(
|
||
"litellm.proxy.proxy_server.general_settings",
|
||
{"maximum_spend_logs_cleanup_run_budget": "90s", "maximum_spend_logs_cleanup_batch_timeout": "10s"},
|
||
):
|
||
await proxy_config._update_general_settings(
|
||
db_general_settings={"maximum_spend_logs_cleanup_batch_timeout": "10s"}
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.general_settings["maximum_spend_logs_cleanup_run_budget"] is None
|
||
assert ps.general_settings["maximum_spend_logs_cleanup_batch_timeout"] == "10s"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_keeps_a_yaml_set_spend_log_cleanup_bound():
|
||
"""A YAML-set bound never appears in the DB object, so treating its absence
|
||
as a dashboard clear would discard the deployed config on every reload."""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
proxy_config._yaml_spend_log_cleanup_bounds = {"maximum_spend_logs_cleanup_run_budget": "90s"}
|
||
|
||
with patch("litellm.proxy.proxy_server.general_settings", {"maximum_spend_logs_cleanup_run_budget": "90s"}):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": True})
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.general_settings["maximum_spend_logs_cleanup_run_budget"] == "90s"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_clearing_a_db_override_falls_back_to_the_yaml_bound():
|
||
"""Clearing a dashboard override of a YAML-declared bound must restore the
|
||
YAML value. Leaving the deleted override in memory would keep enforcing the
|
||
bound the operator just removed, until the process restarted."""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
proxy_config._yaml_spend_log_cleanup_bounds = {"maximum_spend_logs_cleanup_run_budget": "90s"}
|
||
|
||
# Memory currently holds the dashboard override, and the DB no longer carries it.
|
||
with patch("litellm.proxy.proxy_server.general_settings", {"maximum_spend_logs_cleanup_run_budget": "30s"}):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": True})
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.general_settings["maximum_spend_logs_cleanup_run_budget"] == "90s"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_apply_user_budget_to_team_keys_yaml_wins():
|
||
"""A DB value must not silently override an explicit YAML setting on reload."""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
proxy_config._yaml_general_settings_keys = {"apply_user_budget_to_team_keys"}
|
||
|
||
with patch("litellm.proxy.proxy_server.general_settings", {"apply_user_budget_to_team_keys": True}):
|
||
await proxy_config._update_general_settings(db_general_settings={"apply_user_budget_to_team_keys": False})
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.general_settings["apply_user_budget_to_team_keys"] is True
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.parametrize(
|
||
"db_value,expected",
|
||
[(True, True), (False, False), ("true", True), ("false", False), (None, None)],
|
||
)
|
||
async def test_update_general_settings_disable_auto_add_proxy_admin_to_teams(db_value, expected):
|
||
"""
|
||
Verify _update_general_settings propagates disable_auto_add_proxy_admin_to_teams
|
||
from the DB config into the live general_settings dict, so a UI toggle via
|
||
/config/field/update takes effect on the next config poll instead of
|
||
requiring a proxy restart.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
with patch("litellm.proxy.proxy_server.general_settings", {}):
|
||
await proxy_config._update_general_settings(
|
||
db_general_settings={"disable_auto_add_proxy_admin_to_teams": db_value}
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.general_settings["disable_auto_add_proxy_admin_to_teams"] is expected
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_store_model_in_db_string_normalization():
|
||
"""
|
||
Verify _update_general_settings normalizes string values for store_model_in_db.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# Test "true" string
|
||
with (
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||
):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": "true"})
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.store_model_in_db is True
|
||
|
||
# Test "True" string
|
||
with (
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||
):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": "True"})
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.store_model_in_db is True
|
||
|
||
# Test "false" string
|
||
with (
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||
):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": "false"})
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.store_model_in_db is False
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_general_settings_store_model_in_db_none_keeps_current():
|
||
"""
|
||
Verify _update_general_settings does not change store_model_in_db
|
||
when DB value is None.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
|
||
# When current is True and DB sends None, should stay True
|
||
with (
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||
):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": None})
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.store_model_in_db is True
|
||
|
||
# When current is False and DB sends None, should stay False
|
||
with (
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
patch("litellm.proxy.proxy_server.general_settings", {}),
|
||
):
|
||
await proxy_config._update_general_settings(db_general_settings={"store_model_in_db": None})
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.store_model_in_db is False
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_batch_cost_poller_is_confirmed_before_serving(monkeypatch):
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.openai_files_endpoints.common_utils import batch_cost_poller_is_active
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||
mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=None)
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", AsyncMock()),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
patch("litellm.proxy.proxy_server.llm_router", MagicMock()),
|
||
patch("litellm.proxy.proxy_server.PROXY_BATCH_POLLING_ENABLED", True),
|
||
patch("litellm.constants.PROXY_BATCH_POLLING_ENABLED", True),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
poller = proxy_server_module.scheduler.get_job("check_batch_cost_job").func.__self__
|
||
assert poller.batch_processed_support_confirmed is True
|
||
assert batch_cost_poller_is_active() is True
|
||
probe_where = mock_prisma_client.db.litellm_managedobjecttable.find_first.call_args[1]["where"]
|
||
assert probe_where["batch_processed"] is False
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_store_model_in_db_db_override_when_config_false():
|
||
"""
|
||
Verify the early DB check in initialize_scheduled_background_jobs
|
||
overrides store_model_in_db=False when DB has True.
|
||
"""
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
|
||
# Mock DB returning store_model_in_db=True in general_settings
|
||
mock_db_record = MagicMock()
|
||
mock_db_record.param_value = {"store_model_in_db": True}
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=mock_db_record)
|
||
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
# store_model_in_db should now be True (overridden by DB)
|
||
assert ps.store_model_in_db is True
|
||
|
||
# add_deployment and get_credentials should have been called
|
||
# since store_model_in_db is now True
|
||
assert mock_proxy_config.add_deployment.call_count == 1
|
||
assert mock_proxy_config.get_credentials.call_count == 1
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_store_model_in_db_db_check_skipped_when_already_true(monkeypatch):
|
||
"""
|
||
Verify the early DB check is skipped when store_model_in_db is already True.
|
||
The DB query for the early check should not be called.
|
||
"""
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
# The early DB check uses find_first with param_name="general_settings".
|
||
# When store_model_in_db is already True, the early check should be skipped.
|
||
# However, add_deployment may also call find_first.
|
||
# We just verify that store_model_in_db stays True and jobs are scheduled.
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.store_model_in_db is True
|
||
assert mock_proxy_config.add_deployment.call_count == 1
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_store_model_in_db_db_failure_graceful(monkeypatch):
|
||
"""
|
||
Verify the early DB check handles DB failures gracefully
|
||
without crashing and keeps store_model_in_db as False.
|
||
"""
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
# Simulate DB failure
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(side_effect=Exception("DB connection error"))
|
||
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", False),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=False),
|
||
):
|
||
# Should not raise an exception
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
# store_model_in_db should remain False
|
||
assert ps.store_model_in_db is False
|
||
|
||
# add_deployment should NOT have been called since store_model_in_db is False
|
||
mock_proxy_config.add_deployment.assert_not_called()
|
||
|
||
|
||
# =====================================================================
|
||
# Spend counter tests (v2 — Redis-backed spend counters)
|
||
# =====================================================================
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_current_spend_reads_redis_first():
|
||
"""get_current_spend should prefer Redis over in-memory."""
|
||
from litellm.caching.dual_cache import DualCache
|
||
|
||
counter_cache = DualCache()
|
||
|
||
# In-memory has stale value
|
||
counter_cache.in_memory_cache.set_cache(key="spend:key:test", value=0.30)
|
||
|
||
# Mock Redis with cross-pod authoritative value
|
||
mock_redis = AsyncMock()
|
||
mock_redis.async_get_cache = AsyncMock(return_value=0.90)
|
||
counter_cache.redis_cache = mock_redis
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
original = ps.spend_counter_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
|
||
try:
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
result = await get_current_spend(
|
||
counter_key="spend:key:test",
|
||
fallback_spend=0.0,
|
||
)
|
||
# Should return Redis value (0.90), not in-memory (0.30)
|
||
assert result == 0.90
|
||
mock_redis.async_get_cache.assert_called_once_with(key="spend:key:test")
|
||
finally:
|
||
ps.spend_counter_cache = original
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_current_spend_fallback_to_in_memory():
|
||
"""When Redis is not configured, get_current_spend uses in-memory."""
|
||
from litellm.caching.dual_cache import DualCache
|
||
|
||
counter_cache = DualCache() # no redis_cache
|
||
counter_cache.in_memory_cache.set_cache(key="spend:key:test", value=0.50)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
original = ps.spend_counter_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
|
||
try:
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
result = await get_current_spend(
|
||
counter_key="spend:key:test",
|
||
fallback_spend=0.0,
|
||
)
|
||
assert result == 0.50
|
||
finally:
|
||
ps.spend_counter_cache = original
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_increment_spend_counters_initializes_and_increments():
|
||
"""Counter should initialize from cached object spend, then increment.
|
||
|
||
Uses a pre-hashed token to match production: metadata["user_api_key"]
|
||
is always hashed by the auth flow before reaching the cost callback.
|
||
"""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy._types import LiteLLM_VerificationTokenView, hash_token
|
||
|
||
key_cache = DualCache()
|
||
counter_cache = DualCache()
|
||
|
||
# In production, the auth flow hashes the raw key before it reaches
|
||
# the cost callback. Simulate that by passing the hashed token.
|
||
hashed_token = hash_token("sk-test-token-for-counter")
|
||
|
||
# Simulate a cached key object with existing spend from DB
|
||
cached_key = LiteLLM_VerificationTokenView(
|
||
token=hashed_token,
|
||
spend=5.0,
|
||
max_budget=10.0,
|
||
)
|
||
key_cache.in_memory_cache.set_cache(key=hashed_token, value=cached_key)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
original_key_cache = ps.user_api_key_cache
|
||
original_counter_cache = ps.spend_counter_cache
|
||
ps.user_api_key_cache = key_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
|
||
try:
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
# Pass pre-hashed token (as the cost callback would in production)
|
||
await increment_spend_counters(
|
||
token=hashed_token,
|
||
team_id=None,
|
||
user_id=None,
|
||
response_cost=0.50,
|
||
)
|
||
|
||
# Counter should be: base(5.0) + increment(0.50) = 5.50
|
||
counter = counter_cache.in_memory_cache.get_cache(key=f"spend:key:{hashed_token}")
|
||
assert counter == 5.50
|
||
|
||
# Second increment — counter already exists, just increment
|
||
await increment_spend_counters(
|
||
token=hashed_token,
|
||
team_id=None,
|
||
user_id=None,
|
||
response_cost=0.25,
|
||
)
|
||
|
||
counter = counter_cache.in_memory_cache.get_cache(key=f"spend:key:{hashed_token}")
|
||
assert counter == 5.75
|
||
finally:
|
||
ps.user_api_key_cache = original_key_cache
|
||
ps.spend_counter_cache = original_counter_cache
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_increment_spend_counters_team_and_member():
|
||
"""Counter should track team and team member spend separately."""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy._types import LiteLLM_TeamTable
|
||
|
||
key_cache = DualCache()
|
||
counter_cache = DualCache()
|
||
|
||
# Cached team object
|
||
team_obj = LiteLLM_TeamTable(team_id="team-1", spend=2.0)
|
||
key_cache.in_memory_cache.set_cache(key="team_id:team-1", value=team_obj)
|
||
|
||
# Cached team membership
|
||
key_cache.in_memory_cache.set_cache(
|
||
key="team_membership:user-1:team-1",
|
||
value={"user_id": "user-1", "team_id": "team-1", "spend": 1.0},
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
original_key_cache = ps.user_api_key_cache
|
||
original_counter_cache = ps.spend_counter_cache
|
||
ps.user_api_key_cache = key_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
|
||
try:
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
await increment_spend_counters(
|
||
token=None,
|
||
team_id="team-1",
|
||
user_id="user-1",
|
||
response_cost=0.30,
|
||
)
|
||
|
||
team_counter = counter_cache.in_memory_cache.get_cache(key="spend:team:team-1")
|
||
assert team_counter == 2.30
|
||
|
||
member_counter = counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1")
|
||
assert member_counter == 1.30
|
||
finally:
|
||
ps.user_api_key_cache = original_key_cache
|
||
ps.spend_counter_cache = original_counter_cache
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_and_increment_spend_counter_reseeds_from_db_on_counter_miss():
|
||
"""When the Redis counter is missing, the reseed path reads the
|
||
authoritative spend from the DB (not a stale cache), so the next
|
||
increment continues from the correct base value."""
|
||
from litellm.caching.dual_cache import DualCache
|
||
|
||
counter_cache = DualCache()
|
||
recorded_increments: list = []
|
||
|
||
async def record_increment(key, value, ttl=None, **kwargs):
|
||
recorded_increments.append({"key": key, "value": value, "ttl": ttl})
|
||
return value
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_increment = AsyncMock(side_effect=record_increment)
|
||
fake_redis.async_get_cache = AsyncMock(return_value=None) # counter missing
|
||
fake_redis.async_set_cache = AsyncMock(return_value=True) # SET NX wins
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
# Prisma returns spend=42.0 (authoritative) while the stale cached
|
||
# value (would be read only if prisma is None) is 10.0. The counter
|
||
# must seed from 42, not 10.
|
||
db_row = MagicMock()
|
||
db_row.spend = 42.0
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row)
|
||
|
||
stale_cache = DualCache()
|
||
stale_team = MagicMock()
|
||
stale_team.spend = 10.0
|
||
stale_cache.in_memory_cache.set_cache(key="team_id:team-9", value=stale_team)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy.proxy_server import _init_and_increment_spend_counter
|
||
|
||
orig_user, orig_counter, orig_prisma = (
|
||
ps.user_api_key_cache,
|
||
ps.spend_counter_cache,
|
||
ps.prisma_client,
|
||
)
|
||
ps.user_api_key_cache = stale_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
await _init_and_increment_spend_counter(
|
||
counter_key="spend:team:team-9",
|
||
source_cache_key="team_id:team-9",
|
||
increment=1.5,
|
||
)
|
||
|
||
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-9"})
|
||
# Seed uses SET NX with db_spend (42) — cross-pod safe, no INCR of 42.
|
||
# Only the per-request delta (1.5) goes through INCRBYFLOAT.
|
||
fake_redis.async_set_cache.assert_awaited_once_with(key="spend:team:team-9", value=42.0, nx=True)
|
||
writes = [(c["key"], c["value"]) for c in recorded_increments]
|
||
assert writes == [("spend:team:team-9", 1.5)]
|
||
finally:
|
||
ps.user_api_key_cache = orig_user
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_primary_spend_counter_redis_concurrent_seed_does_not_double_seed():
|
||
"""Two pods both observing a missing Redis counter must not both
|
||
INCRBYFLOAT the full DB spend. SpendCounterReseed.coalesced uses SET NX
|
||
so the loser reads the winner's value; final Redis = db_spend, not
|
||
2 * db_spend.
|
||
|
||
The per-counter asyncio.Lock is per-process, so it does NOT coordinate
|
||
across pods. We simulate two pods by patching _get_lock to return a
|
||
fresh lock per call (each "pod" has its own lock registry in real life).
|
||
"""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
|
||
|
||
counter_key = "spend:team:team-concurrent-seed"
|
||
redis_store: dict = {}
|
||
db_read_count = 0
|
||
set_results: list = []
|
||
get_after_set_count = 0
|
||
set_completed_count = 0
|
||
|
||
async def redis_set_cache(key, value, nx=False, **_):
|
||
# Yield BEFORE the membership check so two concurrent callers
|
||
# interleave the way real atomic Redis SET NX does: the first
|
||
# to resume runs check + write atomically and wins; the second
|
||
# resumes after the key exists and loses. Yielding *after* the
|
||
# check would let both callers pass the empty-store check before
|
||
# either writes, so neither would ever lose.
|
||
await asyncio.sleep(0)
|
||
if nx and key in redis_store:
|
||
set_results.append(False)
|
||
return False
|
||
redis_store[key] = float(value)
|
||
set_results.append(True)
|
||
nonlocal set_completed_count
|
||
set_completed_count += 1
|
||
return True
|
||
|
||
async def redis_get_cache(key):
|
||
# Track reads that happen after at least one SET NX has completed
|
||
# — those are the loser-path fallback reads we want to verify.
|
||
if set_completed_count > 0:
|
||
nonlocal get_after_set_count
|
||
get_after_set_count += 1
|
||
return redis_store.get(key)
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache)
|
||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||
|
||
async def slow_find_unique(**_):
|
||
nonlocal db_read_count
|
||
db_read_count += 1
|
||
# Both pods read DB before either's SET NX lands.
|
||
await asyncio.sleep(0)
|
||
row = MagicMock()
|
||
row.spend = 506.0
|
||
return row
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(side_effect=slow_find_unique)
|
||
|
||
pod_a = DualCache()
|
||
pod_a.redis_cache = fake_redis
|
||
pod_b = DualCache()
|
||
pod_b.redis_cache = fake_redis
|
||
|
||
# Each "pod" has its own per-process lock registry. Patch _get_lock to
|
||
# always return a fresh lock so the two coalesced calls do not serialize
|
||
# via one in-process lock (which is what would happen across pods).
|
||
async def fresh_lock(_counter_key):
|
||
return asyncio.Lock()
|
||
|
||
with patch.object(SpendCounterReseed, "_get_lock", side_effect=fresh_lock):
|
||
results = await asyncio.gather(
|
||
SpendCounterReseed.coalesced(
|
||
prisma_client=fake_prisma,
|
||
spend_counter_cache=pod_a,
|
||
counter_key=counter_key,
|
||
),
|
||
SpendCounterReseed.coalesced(
|
||
prisma_client=fake_prisma,
|
||
spend_counter_cache=pod_b,
|
||
counter_key=counter_key,
|
||
),
|
||
)
|
||
|
||
assert all(r == 506.0 for r in results), results
|
||
assert redis_store[counter_key] == pytest.approx(506.0), redis_store
|
||
# Both pods read the DB and both attempted SET NX; exactly one wrote
|
||
# (winner) and one was rejected (loser).
|
||
assert db_read_count == 2
|
||
assert fake_redis.async_set_cache.await_count == 2
|
||
nx_writes = [call for call in fake_redis.async_set_cache.await_args_list if call.kwargs.get("nx") is True]
|
||
assert len(nx_writes) == 2
|
||
assert sorted(set_results) == [
|
||
False,
|
||
True,
|
||
], f"expected exactly one SET NX winner and one loser, got {set_results}"
|
||
# Loser path executed: after the winner's SET NX returned True, the
|
||
# losing coalesced() call falls back to async_get_cache to read the
|
||
# winner's value rather than re-seeding.
|
||
assert get_after_set_count >= 1, "loser branch (else: read back winner's value) was never exercised"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_reseed_spend_from_db_user_and_org_prefixes():
|
||
"""User and org counters reseed from their own DB tables.
|
||
|
||
End-user and tag counters use the already fetched auth objects passed as
|
||
fallback_spend, so this reseed helper must not add extra per-request DB
|
||
reads for them.
|
||
"""
|
||
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
|
||
|
||
user_row = MagicMock()
|
||
user_row.spend = 17.0
|
||
org_row = MagicMock()
|
||
org_row.spend = 305.0
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=user_row)
|
||
fake_prisma.db.litellm_endusertable.find_unique = AsyncMock()
|
||
fake_prisma.db.litellm_tagtable.find_unique = AsyncMock()
|
||
fake_prisma.db.litellm_organizationtable.find_unique = AsyncMock(return_value=org_row)
|
||
|
||
assert await SpendCounterReseed.from_db(fake_prisma, "spend:user:alice") == 17.0
|
||
fake_prisma.db.litellm_usertable.find_unique.assert_awaited_once_with(where={"user_id": "alice"})
|
||
|
||
assert (
|
||
await SpendCounterReseed.from_db(
|
||
fake_prisma,
|
||
"spend:end_user:customer-1",
|
||
)
|
||
is None
|
||
)
|
||
fake_prisma.db.litellm_endusertable.find_unique.assert_not_awaited()
|
||
|
||
assert await SpendCounterReseed.from_db(fake_prisma, "spend:tag:paid-tag") is None
|
||
fake_prisma.db.litellm_tagtable.find_unique.assert_not_awaited()
|
||
|
||
assert await SpendCounterReseed.from_db(fake_prisma, "spend:org:acme") == 305.0
|
||
fake_prisma.db.litellm_organizationtable.find_unique.assert_awaited_once_with(where={"organization_id": "acme"})
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_reseed_spend_from_db_skips_window_variant_keys():
|
||
"""Window counters (spend:*:window:{duration}) share prefixes with
|
||
primary counters but don't correspond to a DB row. The guard must
|
||
short-circuit without querying the DB."""
|
||
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_verificationtoken.find_unique = AsyncMock()
|
||
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock()
|
||
|
||
assert await SpendCounterReseed.from_db(fake_prisma, "spend:key:sk-abc:window:1h") is None
|
||
assert await SpendCounterReseed.from_db(fake_prisma, "spend:team:team-1:window:1d") is None
|
||
fake_prisma.db.litellm_verificationtoken.find_unique.assert_not_awaited()
|
||
fake_prisma.db.litellm_teamtable.find_unique.assert_not_awaited()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_window_spend_counter_reseeds_from_spend_logs_on_counter_miss():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||
|
||
counter_cache = DualCache()
|
||
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None)
|
||
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
|
||
return_value=[{"api_key": "key-window", "_sum": {"spend": 2.25}}]
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
await _init_and_increment_window_spend_counter(
|
||
counter_key="spend:key:key-window:window:1h",
|
||
entity_type="Key",
|
||
entity_id="key-window",
|
||
window_duration="1h",
|
||
window_start=window_start,
|
||
increment=0.5,
|
||
)
|
||
|
||
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
|
||
by=["api_key"],
|
||
where={"api_key": "key-window", "startTime": {"gte": window_start}},
|
||
sum={"spend": True},
|
||
)
|
||
assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-window:window:1h") == pytest.approx(2.75)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import _init_and_increment_spend_counter
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:team:team-stale-local"
|
||
counter_cache.in_memory_cache.set_cache(key=counter_key, value=10.0)
|
||
|
||
redis_store: dict = {}
|
||
|
||
async def redis_increment(key, value, **_):
|
||
redis_store[key] = (redis_store.get(key) or 0.0) + value
|
||
return redis_store[key]
|
||
|
||
async def redis_set_cache(key, value, nx=False, **_):
|
||
if nx and key in redis_store:
|
||
return False
|
||
redis_store[key] = float(value)
|
||
return True
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
db_row = MagicMock()
|
||
db_row.spend = 42.0
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=db_row)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma, orig_user = (
|
||
ps.spend_counter_cache,
|
||
ps.prisma_client,
|
||
ps.user_api_key_cache,
|
||
)
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
ps.user_api_key_cache = DualCache()
|
||
try:
|
||
await _init_and_increment_spend_counter(
|
||
counter_key=counter_key,
|
||
source_cache_key="team_id:team-stale-local",
|
||
increment=1.5,
|
||
)
|
||
|
||
fake_prisma.db.litellm_teamtable.find_unique.assert_awaited_once_with(where={"team_id": "team-stale-local"})
|
||
# Seed via SET NX (42) + delta via INCRBYFLOAT (1.5) = 43.5.
|
||
assert redis_store[counter_key] == pytest.approx(43.5)
|
||
assert counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(43.5)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
ps.user_api_key_cache = orig_user
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_window_spend_counter_redis_clean_miss_skips_stale_in_memory():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:key:key-window-stale-local:window:1h"
|
||
counter_cache.in_memory_cache.set_cache(key=counter_key, value=100.0)
|
||
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
|
||
|
||
redis_store: dict = {}
|
||
|
||
async def redis_increment(key, value, **_):
|
||
redis_store[key] = (redis_store.get(key) or 0.0) + value
|
||
return redis_store[key]
|
||
|
||
async def redis_set_cache(key, value, **_):
|
||
if key in redis_store:
|
||
return False
|
||
redis_store[key] = value
|
||
return True
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None)
|
||
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
|
||
return_value=[{"api_key": "key-window-stale-local", "_sum": {"spend": 2.25}}]
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
await _init_and_increment_window_spend_counter(
|
||
counter_key=counter_key,
|
||
entity_type="Key",
|
||
entity_id="key-window-stale-local",
|
||
window_duration="1h",
|
||
window_start=window_start,
|
||
increment=0.5,
|
||
)
|
||
|
||
fake_prisma.db.litellm_spendlogs.group_by.assert_awaited_once_with(
|
||
by=["api_key"],
|
||
where={
|
||
"api_key": "key-window-stale-local",
|
||
"startTime": {"gte": window_start},
|
||
},
|
||
sum={"spend": True},
|
||
)
|
||
assert redis_store[counter_key] == pytest.approx(2.75)
|
||
assert counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(2.75)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_window_spend_counter_redis_concurrent_seed_does_not_double_seed():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:key:key-window-concurrent-seed:window:1h"
|
||
window_start = datetime.now(timezone.utc) - timedelta(hours=1)
|
||
redis_store = {counter_key: 2.75}
|
||
redis_reads = 0
|
||
|
||
async def redis_get_cache(key):
|
||
nonlocal redis_reads
|
||
redis_reads += 1
|
||
if redis_reads <= 2:
|
||
return None
|
||
return redis_store.get(key)
|
||
|
||
async def redis_increment(key, value, **_):
|
||
redis_store[key] = (redis_store.get(key) or 0.0) + value
|
||
return redis_store[key]
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get_cache)
|
||
fake_redis.async_set_cache = AsyncMock(return_value=False)
|
||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_budgetwindowspend.find_unique = AsyncMock(return_value=None)
|
||
fake_prisma.db.litellm_spendlogs.group_by = AsyncMock(
|
||
return_value=[{"api_key": "key-window-concurrent-seed", "_sum": {"spend": 2.25}}]
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
await _init_and_increment_window_spend_counter(
|
||
counter_key=counter_key,
|
||
entity_type="Key",
|
||
entity_id="key-window-concurrent-seed",
|
||
window_duration="1h",
|
||
window_start=window_start,
|
||
increment=0.5,
|
||
)
|
||
|
||
fake_redis.async_set_cache.assert_awaited_once_with(
|
||
key=counter_key,
|
||
value=2.25,
|
||
nx=True,
|
||
)
|
||
assert redis_store[counter_key] == pytest.approx(3.25)
|
||
assert counter_cache.in_memory_cache.get_cache(key=counter_key) == pytest.approx(3.25)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_window_spend_counter_skips_invalid_window_start():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import _init_and_increment_window_spend_counter
|
||
|
||
counter_cache = DualCache()
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter = ps.spend_counter_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
try:
|
||
await _init_and_increment_window_spend_counter(
|
||
counter_key="spend:key:key-invalid-window:window:not-a-duration",
|
||
entity_type="Key",
|
||
entity_id="key-invalid-window",
|
||
window_duration="not-a-duration",
|
||
window_start=None,
|
||
increment=0.5,
|
||
)
|
||
|
||
assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-invalid-window:window:not-a-duration") is None
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_window_spend_counter_does_not_seed_zero_when_db_unavailable():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import _ensure_window_spend_counter_initialized
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:key:key-window-db-unavailable:window:1h"
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = None
|
||
try:
|
||
initialized = await _ensure_window_spend_counter_initialized(
|
||
counter_key=counter_key,
|
||
entity_type="Key",
|
||
entity_id="key-window-db-unavailable",
|
||
window_duration="1h",
|
||
window_start=datetime.now(timezone.utc) - timedelta(hours=1),
|
||
)
|
||
|
||
assert initialized is False
|
||
assert counter_cache.in_memory_cache.get_cache(key=counter_key) is None
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_increment_spend_counters_finalizes_after_unreserved_increments():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
counter_cache = DualCache()
|
||
counter_cache.in_memory_cache.set_cache(
|
||
key="spend:key:key-finalize-after-increments",
|
||
value=0.5,
|
||
)
|
||
budget_reservation = {
|
||
"reserved_cost": 0.5,
|
||
"entries": [
|
||
{
|
||
"counter_key": "spend:key:key-finalize-after-increments",
|
||
"entity_type": "Key",
|
||
"entity_id": "key-finalize-after-increments",
|
||
"reserved_cost": 0.5,
|
||
"applied_adjustment": 0.0,
|
||
}
|
||
],
|
||
"finalized": False,
|
||
}
|
||
incremented_counters = []
|
||
|
||
async def assert_reservation_not_finalized_yet(**kwargs):
|
||
assert budget_reservation["finalized"] is False
|
||
incremented_counters.append(kwargs["counter_key"])
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_user = ps.spend_counter_cache, ps.user_api_key_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.user_api_key_cache = DualCache()
|
||
try:
|
||
with patch(
|
||
"litellm.proxy.proxy_server._init_and_increment_spend_counter",
|
||
new=AsyncMock(side_effect=assert_reservation_not_finalized_yet),
|
||
):
|
||
await increment_spend_counters(
|
||
token="key-finalize-after-increments",
|
||
team_id="team-finalize-after-increments",
|
||
user_id=None,
|
||
response_cost=0.25,
|
||
budget_reservation=budget_reservation,
|
||
)
|
||
|
||
assert incremented_counters == ["spend:team:team-finalize-after-increments"]
|
||
assert budget_reservation["finalized"] is True
|
||
assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-finalize-after-increments") == pytest.approx(
|
||
0.25
|
||
)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.user_api_key_cache = orig_user
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_increment_spend_counters_finalizes_none_cost_reservation():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
counter_cache = DualCache()
|
||
counter_cache.in_memory_cache.set_cache(
|
||
key="spend:key:key-finalize-none-cost",
|
||
value=0.5,
|
||
)
|
||
budget_reservation = {
|
||
"reserved_cost": 0.5,
|
||
"entries": [
|
||
{
|
||
"counter_key": "spend:key:key-finalize-none-cost",
|
||
"entity_type": "Key",
|
||
"entity_id": "key-finalize-none-cost",
|
||
"reserved_cost": 0.5,
|
||
"applied_adjustment": 0.0,
|
||
}
|
||
],
|
||
"finalized": False,
|
||
}
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter = ps.spend_counter_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
try:
|
||
await increment_spend_counters(
|
||
token="key-finalize-none-cost",
|
||
team_id=None,
|
||
user_id=None,
|
||
response_cost=None,
|
||
budget_reservation=budget_reservation,
|
||
)
|
||
|
||
assert budget_reservation["finalized"] is True
|
||
assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-finalize-none-cost") == pytest.approx(0.0)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_increment_spend_counters_reseeds_from_db_on_bad_reserved_counter():
|
||
"""When the reservation reconcile finds the counter in an inconsistent state
|
||
(here: missing), it must NOT delete the counter and fail open (the old
|
||
behavior, which left the counter unenforced after a Redis reload). It reseeds
|
||
from the authoritative DB so the counter reflects the recorded total and
|
||
budget gating continues."""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
|
||
|
||
counter_cache = DualCache()
|
||
budget_reservation = {
|
||
"reserved_cost": 0.5,
|
||
"entries": [
|
||
{
|
||
"counter_key": "spend:key:key-bad-reserved-counter",
|
||
"entity_type": "Key",
|
||
"entity_id": "key-bad-reserved-counter",
|
||
"reserved_cost": 0.5,
|
||
"applied_adjustment": 0.0,
|
||
}
|
||
],
|
||
"finalized": False,
|
||
}
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter = ps.spend_counter_cache
|
||
orig_prisma = ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = MagicMock() # truthy so reseed reaches from_db
|
||
try:
|
||
with patch.object(SpendCounterReseed, "from_db", AsyncMock(return_value=0.6)):
|
||
await increment_spend_counters(
|
||
token="key-bad-reserved-counter",
|
||
team_id=None,
|
||
user_id=None,
|
||
response_cost=0.25,
|
||
budget_reservation=budget_reservation,
|
||
)
|
||
|
||
assert budget_reservation["finalized"] is True
|
||
# counter reseeded to the authoritative DB value, not deleted/left None
|
||
# and not double-counted via a direct increment
|
||
assert counter_cache.in_memory_cache.get_cache(key="spend:key:key-bad-reserved-counter") == pytest.approx(0.6)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_increment_spend_counter_invalidates_stale_cache_on_redis_failure():
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import _increment_spend_counter_cache
|
||
|
||
counter_cache = DualCache()
|
||
counter_cache.in_memory_cache.set_cache(key="spend:team:redis-fail", value=4.0)
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_increment = AsyncMock(side_effect=RuntimeError("redis down"))
|
||
fake_redis.async_delete_cache = AsyncMock()
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter = ps.spend_counter_cache
|
||
ps.spend_counter_cache = counter_cache
|
||
try:
|
||
with pytest.raises(RuntimeError):
|
||
await _increment_spend_counter_cache(
|
||
counter_key="spend:team:redis-fail",
|
||
increment=0.5,
|
||
)
|
||
|
||
assert counter_cache.in_memory_cache.get_cache(key="spend:team:redis-fail") is None
|
||
fake_redis.async_delete_cache.assert_awaited_once_with(key="spend:team:redis-fail")
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_current_spend_reseeds_from_db_when_counter_missing():
|
||
"""
|
||
When both the Redis and in-memory counters are missing, the enforcement
|
||
read path must reseed from the authoritative DB, not fall back to the
|
||
caller-supplied stale value. Otherwise, every Redis TTL expiry lets a
|
||
request through against a stale in-process `team_membership.spend`.
|
||
"""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
counter_cache = DualCache()
|
||
recorded_seeds: list = []
|
||
|
||
async def record_set_cache(key, value, nx=False, **kwargs):
|
||
recorded_seeds.append({"key": key, "value": value, "nx": nx})
|
||
return True
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_set_cache = AsyncMock(side_effect=record_set_cache)
|
||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
# DB has authoritative spend=362.0; caller hands us stale fallback=30.0
|
||
# (the in-process team_membership.spend that hasn't caught up to DB).
|
||
db_row = MagicMock()
|
||
db_row.spend = 362.0
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
spend = await get_current_spend(
|
||
counter_key="spend:team_member:user-1:team-1",
|
||
fallback_spend=30.0,
|
||
)
|
||
assert spend == 362.0, (
|
||
f"expected DB reseed to return 362.0, got {spend} (fallback would have returned 30.0 and caused bypass)"
|
||
)
|
||
# Counter warmed via SET NX so subsequent reads are fast.
|
||
assert ("spend:team_member:user-1:team-1", 362.0, True) in [
|
||
(s["key"], s["value"], s["nx"]) for s in recorded_seeds
|
||
]
|
||
assert counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == pytest.approx(362.0)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_current_spend_uses_fallback_when_db_unavailable():
|
||
"""
|
||
If prisma is unavailable and both counters are missing, the read path
|
||
must degrade to the caller-supplied fallback rather than raising.
|
||
"""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
counter_cache = DualCache()
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = None # simulate prisma unavailable
|
||
try:
|
||
spend = await get_current_spend(
|
||
counter_key="spend:team_member:user-1:team-1",
|
||
fallback_spend=15.5,
|
||
)
|
||
assert spend == 15.5
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_current_spend_coalesces_concurrent_reseeds():
|
||
"""
|
||
When several concurrent calls hit a cold counter on the same pod,
|
||
only one DB query should fire. The rest should wait for the lock,
|
||
re-check the warmed counter, and return without hitting the DB.
|
||
"""
|
||
import asyncio as _asyncio
|
||
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:team_member:user-1:team-coalesce"
|
||
|
||
# Track DB query calls and inject a small delay so the concurrent
|
||
# callers actually overlap in the lock-acquire window.
|
||
db_call_count = 0
|
||
|
||
async def slow_find_unique(**kwargs):
|
||
nonlocal db_call_count
|
||
db_call_count += 1
|
||
await _asyncio.sleep(0.05)
|
||
row = MagicMock()
|
||
row.spend = 100.0
|
||
return row
|
||
|
||
fake_redis = AsyncMock()
|
||
redis_store: dict = {}
|
||
|
||
async def redis_get(key, **_):
|
||
return redis_store.get(key)
|
||
|
||
async def redis_increment(key, value, **_):
|
||
redis_store[key] = (redis_store.get(key) or 0.0) + value
|
||
return redis_store[key]
|
||
|
||
async def redis_set_cache(key, value, nx=False, **_):
|
||
if nx and key in redis_store:
|
||
return False
|
||
redis_store[key] = float(value)
|
||
return True
|
||
|
||
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get)
|
||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=slow_find_unique)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
results = await _asyncio.gather(
|
||
*[get_current_spend(counter_key=counter_key, fallback_spend=0.0) for _ in range(5)]
|
||
)
|
||
assert results == [100.0] * 5, f"all callers should see DB value, got {results}"
|
||
assert db_call_count == 1, f"expected exactly 1 DB query for 5 concurrent reseeds, got {db_call_count}"
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_current_spend_uses_db_zero_over_stale_fallback():
|
||
"""
|
||
When DB returns spend=0 (e.g. just after a budget period reset), the
|
||
authoritative DB value must win over a stale non-zero fallback. The
|
||
fallback in production is the in-process team_membership.spend, which
|
||
can still hold the pre-reset value across pods.
|
||
"""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
counter_cache = DualCache()
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
db_row = MagicMock()
|
||
db_row.spend = 0.0
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
spend = await get_current_spend(
|
||
counter_key="spend:team_member:user-1:team-after-reset",
|
||
fallback_spend=42.0,
|
||
)
|
||
assert spend == 0.0, f"DB authoritative 0 must override stale fallback 42, got {spend}"
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_concurrent_read_and_write_paths_share_one_db_query():
|
||
"""
|
||
The read path (`get_current_spend`) and the write path
|
||
(`_init_and_increment_spend_counter`) both reseed cold counters from
|
||
the DB. They must share the per-counter lock so a concurrent pre-call
|
||
enforcement read and post-call increment for the same counter collapse
|
||
to one DB query, not two.
|
||
"""
|
||
import asyncio as _asyncio
|
||
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import (
|
||
_init_and_increment_spend_counter,
|
||
get_current_spend,
|
||
)
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:team_member:user-1:team-cross-path"
|
||
|
||
db_call_count = 0
|
||
|
||
async def slow_find_unique(**kwargs):
|
||
nonlocal db_call_count
|
||
db_call_count += 1
|
||
await _asyncio.sleep(0.05)
|
||
row = MagicMock()
|
||
row.spend = 50.0
|
||
return row
|
||
|
||
redis_store: dict = {}
|
||
|
||
async def redis_get(key, **_):
|
||
return redis_store.get(key)
|
||
|
||
async def redis_increment(key, value, **_):
|
||
redis_store[key] = (redis_store.get(key) or 0.0) + value
|
||
return redis_store[key]
|
||
|
||
async def redis_set_cache(key, value, nx=False, **_):
|
||
if nx and key in redis_store:
|
||
return False
|
||
redis_store[key] = float(value)
|
||
return True
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get)
|
||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=slow_find_unique)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma, orig_user = (
|
||
ps.spend_counter_cache,
|
||
ps.prisma_client,
|
||
ps.user_api_key_cache,
|
||
)
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
ps.user_api_key_cache = DualCache()
|
||
try:
|
||
results = await _asyncio.gather(
|
||
get_current_spend(counter_key=counter_key, fallback_spend=0.0),
|
||
_init_and_increment_spend_counter(
|
||
counter_key=counter_key,
|
||
source_cache_key="ignored",
|
||
increment=1.5,
|
||
),
|
||
get_current_spend(counter_key=counter_key, fallback_spend=0.0),
|
||
)
|
||
assert db_call_count == 1, f"expected 1 DB query for concurrent read+write+read, got {db_call_count}"
|
||
# Read-path callers see the warmed counter; the write path's
|
||
# increment may or may not have landed by then, so accept either
|
||
# the seeded value or seeded+increment.
|
||
assert results[0] in (50.0, 51.5), f"got {results[0]}"
|
||
assert results[2] in (50.0, 51.5), f"got {results[2]}"
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
ps.user_api_key_cache = orig_user
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_reseed_locks_dict_is_bounded():
|
||
"""
|
||
`SpendCounterReseed._locks` is an LRU bounded at
|
||
`SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE` to prevent unbounded growth in
|
||
long-lived deployments with high counter-key churn. Inserting more
|
||
than the cap evicts the oldest entries.
|
||
"""
|
||
import litellm.constants as constants
|
||
from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed
|
||
|
||
orig_locks = SpendCounterReseed._locks.copy()
|
||
SpendCounterReseed._locks.clear()
|
||
orig_max = constants.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
|
||
constants.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = 5
|
||
# The class reads the constant via module-level import, so patch the
|
||
# module-level name on the spend_counter_reseed module too.
|
||
import litellm.proxy.db.spend_counter_reseed as scr
|
||
|
||
orig_module_max = scr.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE
|
||
scr.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = 5
|
||
try:
|
||
for i in range(7):
|
||
await SpendCounterReseed._get_lock(f"spend:key:test-key-{i}")
|
||
assert len(SpendCounterReseed._locks) == 5, f"got {len(SpendCounterReseed._locks)}"
|
||
# Oldest two evicted
|
||
assert "spend:key:test-key-0" not in SpendCounterReseed._locks
|
||
assert "spend:key:test-key-1" not in SpendCounterReseed._locks
|
||
# Most recent retained
|
||
assert "spend:key:test-key-6" in SpendCounterReseed._locks
|
||
finally:
|
||
constants.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = orig_max
|
||
scr.SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE = orig_module_max
|
||
SpendCounterReseed._locks.clear()
|
||
SpendCounterReseed._locks.update(orig_locks)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_reseed_warms_cache_even_on_zero_db_spend():
|
||
"""
|
||
When DB returns 0.0 (fresh entity / just after reset), the cache must
|
||
still be warmed so subsequent reads hit the cache instead of issuing
|
||
another DB query. Skipping the warm causes O(requests) DB load on
|
||
zero-spend entities.
|
||
"""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:team_member:user-1:team-zero-warm"
|
||
redis_store: dict = {}
|
||
|
||
async def redis_get(key, **_):
|
||
return redis_store.get(key)
|
||
|
||
async def redis_increment(key, value, **_):
|
||
redis_store[key] = (redis_store.get(key) or 0.0) + value
|
||
return redis_store[key]
|
||
|
||
async def redis_set_cache(key, value, nx=False, **_):
|
||
if nx and key in redis_store:
|
||
return False
|
||
redis_store[key] = float(value)
|
||
return True
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(side_effect=redis_get)
|
||
fake_redis.async_increment = AsyncMock(side_effect=redis_increment)
|
||
fake_redis.async_set_cache = AsyncMock(side_effect=redis_set_cache)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
db_call_count = 0
|
||
|
||
async def find_unique(**kwargs):
|
||
nonlocal db_call_count
|
||
db_call_count += 1
|
||
row = MagicMock()
|
||
row.spend = 0.0
|
||
return row
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(side_effect=find_unique)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
# First call: cold cache, hits DB, returns 0.
|
||
spend1 = await get_current_spend(counter_key=counter_key, fallback_spend=0.0)
|
||
# Second call: cache should be warmed at 0, no second DB query.
|
||
spend2 = await get_current_spend(counter_key=counter_key, fallback_spend=0.0)
|
||
assert spend1 == 0.0 and spend2 == 0.0
|
||
assert db_call_count == 1, f"second read should hit warmed cache, got {db_call_count} DB queries"
|
||
assert redis_store.get(counter_key) == 0.0, "cache must be warmed at 0"
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
# -----------------------------------------------------------------------------
|
||
# /config/update — critical paths only.
|
||
#
|
||
# These exercise the four behaviors that broke or changed in the rewrite of
|
||
# update_config (litellm/proxy/proxy_server.py): targeted per-section writes,
|
||
# the removal of the store_model_in_db gate, env var encryption, and the
|
||
# success_callback / litellm_settings merge semantics. All other branches
|
||
# (auth, missing-DB, slack auto-enable, router_settings merge) are covered
|
||
# implicitly or by upstream tests.
|
||
# -----------------------------------------------------------------------------
|
||
|
||
|
||
class _FakeRow:
|
||
def __init__(self, param_name, param_value):
|
||
self.param_name = param_name
|
||
self.param_value = param_value
|
||
|
||
|
||
class _FakeLitellmConfig:
|
||
def __init__(self, initial_rows=None):
|
||
self.rows = dict(initial_rows or {})
|
||
self.upsert_calls: list = []
|
||
self.find_first = AsyncMock(side_effect=self._find_first)
|
||
self.upsert = AsyncMock(side_effect=self._upsert)
|
||
|
||
async def _find_first(self, where=None):
|
||
if where and "param_name" in where:
|
||
name = where["param_name"]
|
||
if name in self.rows:
|
||
return _FakeRow(name, self.rows[name])
|
||
return None
|
||
|
||
async def _upsert(self, where=None, data=None):
|
||
name = where["param_name"]
|
||
raw = data["update"]["param_value"]
|
||
value = json.loads(raw) if isinstance(raw, str) else raw
|
||
self.rows[name] = value
|
||
self.upsert_calls.append((name, value))
|
||
|
||
|
||
class _FakePrismaClient:
|
||
def __init__(self, initial_rows=None):
|
||
self.db = mock.MagicMock()
|
||
self.db.litellm_config = _FakeLitellmConfig(initial_rows=initial_rows)
|
||
self.jsonify_object = lambda obj: obj
|
||
|
||
|
||
@pytest.fixture
|
||
def _update_config_setup(monkeypatch):
|
||
"""Install fakes for the /config/update endpoint and return (client, prisma)."""
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth as auth_dep
|
||
|
||
def _install(initial_rows=None, store_model_in_db=True):
|
||
prisma = _FakePrismaClient(initial_rows=initial_rows)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma)
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", store_model_in_db)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.encrypt_value_helper",
|
||
lambda value, **_: f"enc:{value}",
|
||
)
|
||
monkeypatch.setattr(
|
||
"litellm.proxy.proxy_server.invalidate_config_param",
|
||
AsyncMock(return_value=None),
|
||
)
|
||
from litellm.proxy.proxy_server import proxy_config as real_proxy_config
|
||
|
||
monkeypatch.setattr(real_proxy_config, "add_deployment", AsyncMock(return_value=None))
|
||
|
||
original_overrides = app.dependency_overrides.copy()
|
||
app.dependency_overrides[auth_dep] = lambda: UserAPIKeyAuth(
|
||
user_id="test_admin",
|
||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||
api_key="sk-1234",
|
||
)
|
||
client = TestClient(app)
|
||
|
||
def _restore():
|
||
app.dependency_overrides = original_overrides
|
||
|
||
return client, prisma, _restore
|
||
|
||
return _install
|
||
|
||
|
||
def test_update_config_writes_only_sent_section(_update_config_setup):
|
||
"""A request that only touches general_settings must not write any other
|
||
section row, and must leave previously-written rows byte-identical."""
|
||
client, prisma, restore = _update_config_setup(
|
||
initial_rows={
|
||
"litellm_settings": {"drop_params": True},
|
||
"environment_variables": {"FOO": "enc:bar"},
|
||
}
|
||
)
|
||
try:
|
||
resp = client.post(
|
||
"/config/update",
|
||
json={"general_settings": {"store_prompts_in_spend_logs": True}},
|
||
)
|
||
assert resp.status_code == 200
|
||
written = {name for name, _ in prisma.db.litellm_config.upsert_calls}
|
||
assert written == {"general_settings"}
|
||
assert prisma.db.litellm_config.rows["litellm_settings"] == {"drop_params": True}
|
||
assert prisma.db.litellm_config.rows["environment_variables"] == {"FOO": "enc:bar"}
|
||
finally:
|
||
restore()
|
||
|
||
|
||
def test_update_config_env_var_round_trip_not_double_encrypted(_update_config_setup, monkeypatch):
|
||
"""Endpoint-level regression for the /config/update double-encryption bug.
|
||
|
||
The Admin UI reads config back via /get/config/callbacks (which returns
|
||
the stored, still-encrypted value) and re-POSTs it on the next save. The
|
||
handler must NOT stack a second encryption layer on the re-submitted
|
||
ciphertext, and must leave untouched keys byte-identical.
|
||
|
||
Uses an invertible fake encrypt/decrypt pair ("enc:" prefix) so the
|
||
decrypt-then-encrypt chokepoint round-trips faithfully. On the pre-fix
|
||
code this stored "enc:enc:..."; the assertions below would fail there.
|
||
"""
|
||
|
||
def _fake_decrypt(value, key=None, exception_type="error", return_original_value=False):
|
||
if isinstance(value, str) and value.startswith("enc:"):
|
||
return value[len("enc:") :]
|
||
return value if return_original_value else None
|
||
|
||
monkeypatch.setattr("litellm.proxy.proxy_server.decrypt_value_helper", _fake_decrypt)
|
||
|
||
client, prisma, restore = _update_config_setup(
|
||
initial_rows={"environment_variables": {"PREEXISTING_KEY": "enc:keepme"}}
|
||
)
|
||
try:
|
||
# First write: plaintext in -> single-encrypted at rest.
|
||
resp = client.post(
|
||
"/config/update",
|
||
json={"environment_variables": {"LANGFUSE_SECRET_KEY": "sk-secret"}},
|
||
)
|
||
assert resp.status_code == 200
|
||
stored = prisma.db.litellm_config.rows["environment_variables"]
|
||
assert stored["LANGFUSE_SECRET_KEY"] == "enc:sk-secret"
|
||
|
||
# UI round-trip: re-POST the stored ciphertext (no field change).
|
||
resp = client.post(
|
||
"/config/update",
|
||
json={"environment_variables": {"LANGFUSE_SECRET_KEY": stored["LANGFUSE_SECRET_KEY"]}},
|
||
)
|
||
assert resp.status_code == 200
|
||
stored = prisma.db.litellm_config.rows["environment_variables"]
|
||
|
||
# The bug: this would be "enc:enc:sk-secret". The fix keeps it single.
|
||
assert stored["LANGFUSE_SECRET_KEY"] == "enc:sk-secret"
|
||
assert _fake_decrypt(stored["LANGFUSE_SECRET_KEY"], return_original_value=True) == "sk-secret"
|
||
|
||
# Untouched key preserved byte-for-byte (only sent keys rewritten).
|
||
assert stored["PREEXISTING_KEY"] == "enc:keepme"
|
||
finally:
|
||
restore()
|
||
|
||
|
||
def test_update_config_can_flip_store_model_in_db_when_currently_false(
|
||
_update_config_setup,
|
||
):
|
||
"""The endpoint used to refuse all writes when store_model_in_db was
|
||
False, blocking the very request that would flip it to True."""
|
||
client, prisma, restore = _update_config_setup(store_model_in_db=False)
|
||
try:
|
||
resp = client.post("/config/update", json={"general_settings": {"store_model_in_db": True}})
|
||
assert resp.status_code == 200
|
||
assert prisma.db.litellm_config.rows["general_settings"]["store_model_in_db"] is True
|
||
finally:
|
||
restore()
|
||
|
||
|
||
def test_update_config_environment_variables_encrypted_before_write(
|
||
_update_config_setup,
|
||
):
|
||
"""env var values must be encrypted before they hit the DB row."""
|
||
client, prisma, restore = _update_config_setup()
|
||
try:
|
||
resp = client.post(
|
||
"/config/update",
|
||
json={"environment_variables": {"OPENAI_API_KEY": "sk-secret"}},
|
||
)
|
||
assert resp.status_code == 200
|
||
stored = prisma.db.litellm_config.rows["environment_variables"]
|
||
assert stored == {"OPENAI_API_KEY": "enc:sk-secret"}
|
||
finally:
|
||
restore()
|
||
|
||
|
||
def test_update_config_litellm_settings_request_wins_for_non_callback_keys(
|
||
_update_config_setup,
|
||
):
|
||
"""Sending {"drop_params": False} when the row holds drop_params: True
|
||
must persist False (request wins). Untouched keys preserved."""
|
||
client, prisma, restore = _update_config_setup(
|
||
initial_rows={
|
||
"litellm_settings": {"drop_params": True, "set_verbose": True},
|
||
}
|
||
)
|
||
try:
|
||
resp = client.post("/config/update", json={"litellm_settings": {"drop_params": False}})
|
||
assert resp.status_code == 200
|
||
stored = prisma.db.litellm_config.rows["litellm_settings"]
|
||
assert stored["drop_params"] is False
|
||
assert stored["set_verbose"] is True
|
||
finally:
|
||
restore()
|
||
|
||
|
||
def test_update_config_success_callback_normalizes_existing_mixed_case(
|
||
_update_config_setup,
|
||
):
|
||
"""Existing mixed-case callback names (written elsewhere) must be
|
||
normalized to lowercase before union, otherwise the union dedup misses
|
||
against the lowercase incoming entry and delete_callback (lowercase
|
||
lookup) cannot find the original."""
|
||
client, prisma, restore = _update_config_setup(
|
||
initial_rows={"litellm_settings": {"success_callback": ["Langfuse", "SQS"]}}
|
||
)
|
||
try:
|
||
resp = client.post(
|
||
"/config/update",
|
||
json={"litellm_settings": {"success_callback": ["langfuse"]}},
|
||
)
|
||
assert resp.status_code == 200
|
||
stored = prisma.db.litellm_config.rows["litellm_settings"]["success_callback"]
|
||
assert set(stored) == {"langfuse", "sqs"}
|
||
finally:
|
||
restore()
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Lazy feature loading (LazyFeatureMiddleware) — verifies that optional
|
||
# routers are NOT imported at module load and ARE imported on first request
|
||
# to a matching path prefix. The same module isn't re-imported on subsequent
|
||
# requests.
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
class TestLazyFeatureRegistry:
|
||
"""Sanity checks on the registry shape — guards against accidental edits."""
|
||
|
||
def test_registry_entries_have_required_fields(self):
|
||
from litellm.proxy._lazy_features import LAZY_FEATURES, LazyFeature
|
||
|
||
assert len(LAZY_FEATURES) > 0
|
||
for feat in LAZY_FEATURES:
|
||
assert isinstance(feat, LazyFeature)
|
||
assert feat.name
|
||
assert feat.module_path
|
||
assert feat.path_prefixes
|
||
assert all(p.startswith("/") for p in feat.path_prefixes)
|
||
assert callable(feat.register_fn)
|
||
|
||
def test_registry_names_unique(self):
|
||
from litellm.proxy._lazy_features import LAZY_FEATURES
|
||
|
||
names = [f.name for f in LAZY_FEATURES]
|
||
assert len(names) == len(set(names)), "duplicate feature names"
|
||
|
||
def test_matches_covers_prefix_and_suffix(self):
|
||
"""``matches`` is the single matcher shared by the middleware (request
|
||
paths) and the warm endpoint (registered route paths), so a route that
|
||
only matches via suffix — e.g. ``/v1/a2a/{id}/message/send`` against the
|
||
``/a2a`` prefix — must still be claimed by the feature."""
|
||
from litellm.proxy._lazy_features import LazyFeature
|
||
|
||
feat = LazyFeature(
|
||
name="a2a",
|
||
module_path="json",
|
||
path_prefixes=("/a2a",),
|
||
path_suffixes=("/message/send",),
|
||
)
|
||
assert feat.matches("/a2a/abc/message/send")
|
||
assert feat.matches("/v1/a2a/abc/message/send")
|
||
assert feat.matches("/a2a/abc/.well-known/agent-card.json")
|
||
assert not feat.matches("/v1/a2a/discover")
|
||
assert not feat.matches("/unrelated")
|
||
|
||
|
||
class TestLazyFeaturesNotImportedAtStartup:
|
||
"""
|
||
The whole point of the refactor: gated feature modules must NOT be
|
||
present in `sys.modules` immediately after `proxy_server` imports.
|
||
"""
|
||
|
||
def test_heavy_modules_absent_at_startup(self):
|
||
# Static scan of proxy_server.py source — catches any top-level
|
||
# `from <lazy_module> import` that would defeat lazy loading.
|
||
# Importing proxy_server in a subprocess and diffing sys.modules
|
||
# would also work, but takes 60-120 s and flakes on slow CI runners.
|
||
import re
|
||
from pathlib import Path
|
||
|
||
from litellm.proxy._lazy_features import LAZY_FEATURES
|
||
|
||
proxy_server_src = (Path(__file__).resolve().parents[3] / "litellm/proxy/proxy_server.py").read_text()
|
||
|
||
leaks = []
|
||
for feat in LAZY_FEATURES:
|
||
# Anchor at column 0 — indented imports inside function bodies
|
||
# are fine (deferred until the function runs).
|
||
pattern = (
|
||
rf"^(from\s+{re.escape(feat.module_path)}\s+import|"
|
||
rf"import\s+{re.escape(feat.module_path)})"
|
||
)
|
||
if re.search(pattern, proxy_server_src, re.MULTILINE):
|
||
leaks.append(feat.module_path)
|
||
|
||
assert not leaks, (
|
||
"proxy_server.py top-level imports a lazy feature module — these "
|
||
f"should be loaded via LazyFeatureMiddleware: {leaks}"
|
||
)
|
||
|
||
|
||
class TestLazyFeatureMiddleware:
|
||
"""Behavior of the middleware itself, exercised in isolation."""
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_first_request_triggers_load_subsequent_does_not(self):
|
||
from fastapi import FastAPI
|
||
|
||
from litellm.proxy._lazy_features import (
|
||
LazyFeature,
|
||
LazyFeatureMiddleware,
|
||
)
|
||
|
||
loads = []
|
||
|
||
def fake_register(app, module):
|
||
loads.append(getattr(module, "__name__", "?"))
|
||
|
||
feat = LazyFeature(
|
||
name="dummy",
|
||
module_path="json", # any always-importable stdlib module
|
||
path_prefixes=("/dummy",),
|
||
register_fn=fake_register,
|
||
)
|
||
|
||
# Build a minimal ASGI receiver to satisfy the middleware contract
|
||
async def downstream(scope, receive, send):
|
||
# echo back; no-op handler
|
||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||
await send({"type": "http.response.body", "body": b""})
|
||
|
||
target_app = FastAPI()
|
||
mw = LazyFeatureMiddleware(downstream, fastapi_app=target_app, features=(feat,))
|
||
|
||
async def receive():
|
||
return {"type": "http.request", "body": b"", "more_body": False}
|
||
|
||
sent: list = []
|
||
|
||
async def send(message):
|
||
sent.append(message)
|
||
|
||
# First request matching the prefix triggers register
|
||
await mw(
|
||
{"type": "http", "path": "/dummy/x", "method": "GET", "headers": []},
|
||
receive,
|
||
send,
|
||
)
|
||
assert loads == ["json"]
|
||
|
||
# Second matching request must NOT re-register
|
||
sent.clear()
|
||
await mw(
|
||
{"type": "http", "path": "/dummy/y", "method": "GET", "headers": []},
|
||
receive,
|
||
send,
|
||
)
|
||
assert loads == ["json"], "register_fn called twice for the same feature"
|
||
|
||
# Non-matching path must not trigger anything
|
||
await mw(
|
||
{"type": "http", "path": "/unrelated", "method": "GET", "headers": []},
|
||
receive,
|
||
send,
|
||
)
|
||
assert loads == ["json"]
|
||
|
||
@pytest.mark.asyncio
|
||
@pytest.mark.parametrize(
|
||
"server_root_path,request_path,should_load,case",
|
||
[
|
||
# SERVER_ROOT_PATH set: incoming path includes prefix → strip and match.
|
||
("/api/v1", "/api/v1/dummy/x", True, "root_path strip + match"),
|
||
# Trailing-slash env var must be normalized.
|
||
("/api/v1/", "/api/v1/dummy/x", True, "trailing-slash env normalization"),
|
||
# Reverse proxy already stripped the prefix → original path still matches.
|
||
("/api/v1", "/dummy/x", True, "pre-stripped path still loads"),
|
||
# No SERVER_ROOT_PATH set → unchanged behavior.
|
||
("", "/dummy/x", True, "no root path"),
|
||
# SERVER_ROOT_PATH=/ must be a no-op (not strip every leading slash).
|
||
("/", "/dummy/x", True, "root_path='/' is no-op"),
|
||
# Boundary check: /apiv2 must not match root /api.
|
||
("/api", "/apiv2/foo", False, "boundary check prevents false match"),
|
||
# Genuine non-match under root_path.
|
||
("/api/v1", "/api/v1/unrelated", False, "unrelated path under root"),
|
||
],
|
||
)
|
||
async def test_root_path_handling(self, monkeypatch, server_root_path, request_path, should_load, case):
|
||
"""
|
||
The middleware must strip SERVER_ROOT_PATH before prefix-matching so
|
||
lazy features load under deployments that set a server root path,
|
||
while handling boundary, trailing-slash, and reverse-proxy edge cases
|
||
correctly.
|
||
"""
|
||
from fastapi import FastAPI
|
||
|
||
from litellm.proxy._lazy_features import (
|
||
LazyFeature,
|
||
LazyFeatureMiddleware,
|
||
)
|
||
|
||
monkeypatch.setenv("SERVER_ROOT_PATH", server_root_path)
|
||
|
||
loads = []
|
||
|
||
def fake_register(app, module):
|
||
loads.append(getattr(module, "__name__", "?"))
|
||
|
||
feat = LazyFeature(
|
||
name=f"dummy_{case}",
|
||
module_path="json",
|
||
path_prefixes=("/dummy",),
|
||
register_fn=fake_register,
|
||
)
|
||
|
||
async def downstream(scope, receive, send):
|
||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||
await send({"type": "http.response.body", "body": b""})
|
||
|
||
target_app = FastAPI()
|
||
mw = LazyFeatureMiddleware(downstream, fastapi_app=target_app, features=(feat,))
|
||
|
||
async def receive():
|
||
return {"type": "http.request", "body": b"", "more_body": False}
|
||
|
||
async def send(message):
|
||
pass
|
||
|
||
await mw(
|
||
{
|
||
"type": "http",
|
||
"path": request_path,
|
||
"method": "GET",
|
||
"headers": [],
|
||
},
|
||
receive,
|
||
send,
|
||
)
|
||
if should_load:
|
||
assert loads == ["json"], f"{case}: expected feature to load"
|
||
else:
|
||
assert loads == [], f"{case}: feature must not load"
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_concurrent_first_requests_only_register_once(self):
|
||
"""
|
||
Two requests to the same prefix arriving in parallel must result in
|
||
exactly one `register_fn` invocation — the lock prevents the import +
|
||
register from racing with itself.
|
||
"""
|
||
from fastapi import FastAPI
|
||
|
||
from litellm.proxy._lazy_features import (
|
||
LazyFeature,
|
||
LazyFeatureMiddleware,
|
||
)
|
||
|
||
loads = []
|
||
|
||
def slow_register(app, module):
|
||
loads.append(getattr(module, "__name__", "?"))
|
||
|
||
feat = LazyFeature(
|
||
name="dummy_concurrent",
|
||
module_path="json",
|
||
path_prefixes=("/dummy_c",),
|
||
register_fn=slow_register,
|
||
)
|
||
|
||
async def downstream(scope, receive, send):
|
||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||
await send({"type": "http.response.body", "body": b""})
|
||
|
||
target_app = FastAPI()
|
||
mw = LazyFeatureMiddleware(downstream, fastapi_app=target_app, features=(feat,))
|
||
|
||
async def receive():
|
||
return {"type": "http.request", "body": b"", "more_body": False}
|
||
|
||
sent: list = []
|
||
|
||
async def send(message):
|
||
sent.append(message)
|
||
|
||
async def hit():
|
||
await mw(
|
||
{
|
||
"type": "http",
|
||
"path": "/dummy_c/x",
|
||
"method": "GET",
|
||
"headers": [],
|
||
},
|
||
receive,
|
||
send,
|
||
)
|
||
|
||
await asyncio.gather(hit(), hit(), hit(), hit(), hit())
|
||
assert loads == ["json"], f"expected one registration despite concurrent first hits, got {loads}"
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_failing_import_does_not_loop(self):
|
||
"""
|
||
If a feature's module can't be imported, the middleware should mark it
|
||
loaded anyway so subsequent requests don't repeatedly retry the failing
|
||
import (which would amplify the cost on every request).
|
||
"""
|
||
from fastapi import FastAPI
|
||
|
||
from litellm.proxy._lazy_features import (
|
||
LazyFeature,
|
||
LazyFeatureMiddleware,
|
||
)
|
||
|
||
attempts = []
|
||
|
||
def fail_register(app, module):
|
||
attempts.append("called")
|
||
raise RuntimeError("boom")
|
||
|
||
feat = LazyFeature(
|
||
name="failing",
|
||
module_path="json",
|
||
path_prefixes=("/fail",),
|
||
register_fn=fail_register,
|
||
)
|
||
|
||
async def downstream(scope, receive, send):
|
||
await send({"type": "http.response.start", "status": 200, "headers": []})
|
||
await send({"type": "http.response.body", "body": b""})
|
||
|
||
target_app = FastAPI()
|
||
mw = LazyFeatureMiddleware(downstream, fastapi_app=target_app, features=(feat,))
|
||
|
||
async def receive():
|
||
return {"type": "http.request", "body": b"", "more_body": False}
|
||
|
||
sent: list = []
|
||
|
||
async def send(message):
|
||
sent.append(message)
|
||
|
||
for _ in range(3):
|
||
await mw(
|
||
{"type": "http", "path": "/fail/x", "method": "GET", "headers": []},
|
||
receive,
|
||
send,
|
||
)
|
||
assert attempts == ["called"], (
|
||
f"failing register_fn should be invoked once, not on every request; got {attempts}"
|
||
)
|
||
|
||
|
||
class TestInjectLazyStubs:
|
||
"""Stub injection keys off the app-tracked loaded set, never sys.modules:
|
||
proxy boot imports several feature modules (mcp_management, cloudzero,
|
||
vantage, config_overrides) without mounting their routers, and their
|
||
/openapi.json entries must survive that (LIT-6275)."""
|
||
|
||
def test_imported_but_unregistered_module_still_gets_stub(self):
|
||
import sys
|
||
|
||
from litellm.proxy._lazy_features import LazyFeature, inject_lazy_stubs
|
||
|
||
feat = LazyFeature(
|
||
name="dummy_lazy_test",
|
||
module_path="json",
|
||
path_prefixes=("/dummy-lazy-test",),
|
||
)
|
||
assert feat.module_path in sys.modules
|
||
|
||
schema = inject_lazy_stubs({"paths": {}}, loaded_modules=frozenset(), features=(feat,))
|
||
assert "/dummy-lazy-test" in schema["paths"]
|
||
|
||
def test_registered_module_gets_no_stub(self):
|
||
from litellm.proxy._lazy_features import LazyFeature, inject_lazy_stubs
|
||
|
||
feat = LazyFeature(
|
||
name="dummy_lazy_test",
|
||
module_path="json",
|
||
path_prefixes=("/dummy-lazy-test",),
|
||
)
|
||
schema = inject_lazy_stubs({"paths": {}}, loaded_modules=frozenset({"json"}), features=(feat,))
|
||
assert "/dummy-lazy-test" not in schema["paths"]
|
||
|
||
def test_snapshot_fragments_injected_for_boot_imported_features(self):
|
||
from litellm.proxy._lazy_features import LAZY_FEATURES, inject_lazy_stubs
|
||
from litellm.proxy._lazy_openapi_snapshot import load_snapshot
|
||
|
||
snapshot = load_snapshot()
|
||
assert snapshot
|
||
boot_imported = tuple(
|
||
f for f in LAZY_FEATURES if f.name in ("mcp_management", "cloudzero", "vantage", "config_overrides")
|
||
)
|
||
assert len(boot_imported) == 4
|
||
|
||
schema = inject_lazy_stubs({"paths": {}}, loaded_modules=frozenset(), features=boot_imported)
|
||
for feat in boot_imported:
|
||
missing = [p for p in snapshot[feat.name]["paths"] if p not in schema["paths"]]
|
||
assert not missing, f"{feat.name} snapshot paths missing from /openapi.json: {missing}"
|
||
|
||
def test_persistent_stub_survives_load(self):
|
||
from litellm.proxy._lazy_features import LazyFeature, inject_lazy_stubs
|
||
|
||
feat = LazyFeature(
|
||
name="dummy_lazy_test",
|
||
module_path="json",
|
||
path_prefixes=("/dummy-lazy-test",),
|
||
persistent_swagger_stub=True,
|
||
)
|
||
schema = inject_lazy_stubs({"paths": {}}, loaded_modules=frozenset({"json"}), features=(feat,))
|
||
assert "/dummy-lazy-test" in schema["paths"]
|
||
|
||
def test_loaded_lazy_modules_reads_app_state(self):
|
||
from fastapi import FastAPI
|
||
|
||
from litellm.proxy._lazy_features import loaded_lazy_modules
|
||
|
||
app = FastAPI()
|
||
assert loaded_lazy_modules(app) == frozenset()
|
||
|
||
app.state.lazy_loaded = {"litellm.proxy.spend_tracking.cloudzero_endpoints"}
|
||
assert loaded_lazy_modules(app) == frozenset({"litellm.proxy.spend_tracking.cloudzero_endpoints"})
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_current_spend_redis_clean_miss_skips_stale_in_memory():
|
||
"""When Redis is reachable and cleanly returns None (TTL expired,
|
||
counter genuinely absent), the read must reseed from DB - NOT fall
|
||
through to per-pod in-memory which only contains this pod's writes.
|
||
|
||
Pre-fix in multi-pod deployments, in-memory contained a stale local
|
||
subset (e.g. $30) while DB had the true cross-pod total ($500). The
|
||
fall-through returned $30, enforcement passed, bypass.
|
||
"""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:team_member:user-1:team-1"
|
||
|
||
# Per-pod stale in-memory: only this pod's writes, not cross-pod truth.
|
||
counter_cache.in_memory_cache.set_cache(key=counter_key, value=30.0)
|
||
|
||
# Redis cleanly returns None (key expired or never written on this pod).
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
||
fake_redis.async_increment = AsyncMock(return_value=500.0)
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
# DB has the authoritative cross-pod spend.
|
||
db_row = MagicMock()
|
||
db_row.spend = 500.0
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=db_row)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
spend = await get_current_spend(counter_key=counter_key, fallback_spend=0.0)
|
||
assert spend == 500.0, (
|
||
f"expected DB-authoritative 500.0 on clean Redis miss, got {spend} "
|
||
f"(stale per-pod in-memory $30 would have caused multi-pod bypass)"
|
||
)
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_current_spend_redis_error_falls_back_to_in_memory():
|
||
"""When Redis raises, the read should still degrade to in-memory rather
|
||
than going straight to DB - in-memory is at least same-pod-fresh and
|
||
cheaper than a DB query during a Redis outage."""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.proxy_server import get_current_spend
|
||
|
||
counter_cache = DualCache()
|
||
counter_key = "spend:team_member:user-1:team-1"
|
||
|
||
counter_cache.in_memory_cache.set_cache(key=counter_key, value=42.0)
|
||
|
||
fake_redis = AsyncMock()
|
||
fake_redis.async_get_cache = AsyncMock(side_effect=ConnectionError("redis down"))
|
||
counter_cache.redis_cache = fake_redis
|
||
|
||
fake_prisma = MagicMock()
|
||
fake_prisma.db.litellm_teammembership.find_unique = AsyncMock(return_value=MagicMock(spend=999.0))
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
orig_counter, orig_prisma = ps.spend_counter_cache, ps.prisma_client
|
||
ps.spend_counter_cache = counter_cache
|
||
ps.prisma_client = fake_prisma
|
||
try:
|
||
spend = await get_current_spend(counter_key=counter_key, fallback_spend=0.0)
|
||
assert spend == 42.0, (
|
||
f"expected in-memory fallback 42.0 on Redis error, got {spend} (should not have hit DB when Redis errored)"
|
||
)
|
||
# DB query should NOT have fired - in-memory short-circuits.
|
||
fake_prisma.db.litellm_teammembership.find_unique.assert_not_awaited()
|
||
finally:
|
||
ps.spend_counter_cache = orig_counter
|
||
ps.prisma_client = orig_prisma
|
||
|
||
|
||
def test_realtime_websocket_route_aliases_registered():
|
||
"""Realtime sessions reach the proxy via three path aliases stacked on
|
||
`realtime_websocket_endpoint`. Dropping any of them silently 405s
|
||
WebSocket upgrades because the catch-all `/openai/{endpoint:path}`
|
||
HTTP passthrough only declares HTTP methods. The aliases must also be
|
||
in `LiteLLMRoutes.openai_routes` (so non-admin / team / key-scoped
|
||
auth allows them) and in `API_ROUTE_TO_CALL_TYPES` (so call-type-aware
|
||
logic such as guardrails can resolve the realtime call type)."""
|
||
from starlette.routing import WebSocketRoute
|
||
|
||
from litellm.proxy._types import LiteLLMRoutes
|
||
from litellm.proxy.proxy_server import app
|
||
from litellm.types.utils import API_ROUTE_TO_CALL_TYPES, CallTypes
|
||
|
||
websocket_paths = {route.path for route in app.routes if isinstance(route, WebSocketRoute)}
|
||
openai_routes = LiteLLMRoutes.openai_routes.value
|
||
|
||
for expected in ("/openai/v1/realtime", "/v1/realtime", "/realtime"):
|
||
assert expected in websocket_paths, (
|
||
f"{expected!r} missing from registered WebSocket routes; the "
|
||
f"realtime endpoint will 405 for clients hitting this path."
|
||
)
|
||
assert expected in openai_routes, (
|
||
f"{expected!r} missing from LiteLLMRoutes.openai_routes; "
|
||
f"non-admin / team / key-scoped users will get 403 on this path."
|
||
)
|
||
assert tuple(API_ROUTE_TO_CALL_TYPES.get(expected) or ()) == (CallTypes.arealtime,), (
|
||
f"{expected!r} missing from API_ROUTE_TO_CALL_TYPES; call-type "
|
||
f"resolution will return None and break call-type-aware features."
|
||
)
|
||
|
||
|
||
class TestTransformRequestBannedParams:
|
||
"""
|
||
/utils/transform_request applies the same banned-param check as LLM endpoints.
|
||
|
||
Without this check, any authenticated user could supply aws_sts_endpoint,
|
||
api_base, etc. and have the server forward its credentials to an
|
||
attacker-controlled endpoint during SDK credential resolution.
|
||
"""
|
||
|
||
@pytest.fixture
|
||
def client(self):
|
||
mock_auth = UserAPIKeyAuth(
|
||
user_id="test-internal",
|
||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||
)
|
||
original = app.dependency_overrides.copy()
|
||
app.dependency_overrides[user_api_key_auth] = lambda: mock_auth
|
||
try:
|
||
yield TestClient(app)
|
||
finally:
|
||
app.dependency_overrides = original
|
||
|
||
@pytest.mark.parametrize(
|
||
"banned",
|
||
[
|
||
"aws_sts_endpoint",
|
||
"api_base",
|
||
"aws_web_identity_token",
|
||
"vertex_credentials",
|
||
],
|
||
)
|
||
def test_banned_params_rejected_for_all_users(self, client, banned):
|
||
"""Banned params must be blocked for any authenticated user."""
|
||
response = client.post(
|
||
"/utils/transform_request",
|
||
json={
|
||
"call_type": "completion",
|
||
"request_body": {
|
||
"model": "gpt-3.5-turbo",
|
||
banned: "https://attacker.example",
|
||
},
|
||
},
|
||
)
|
||
assert response.status_code == 400, (
|
||
f"Expected 400 for banned param '{banned}', got {response.status_code}: {response.json()}"
|
||
)
|
||
|
||
|
||
class TestSortModelsByDisplayName:
|
||
"""Regression: team BYOK rows persist an internal `model_name` like
|
||
`model_name_{team_id}_{uuid}` and expose the user-facing name via
|
||
`model_info.team_public_model_name`. Sorting must use the displayed
|
||
name so BYOK rows interleave with non-BYOK rows alphabetically —
|
||
otherwise they clump at the end on their opaque IDs even though the
|
||
UI shows them under a normal-looking name.
|
||
"""
|
||
|
||
def test_byok_models_sort_by_team_public_model_name(self):
|
||
from litellm.proxy.proxy_server import _sort_models
|
||
|
||
models = [
|
||
{"model_name": "claude-haiku-4-5", "model_info": {}},
|
||
{
|
||
# Opaque internal name; UI displays team_public_model_name.
|
||
"model_name": "model_name_team-1_abc123",
|
||
"model_info": {"team_public_model_name": "anthropic/claude"},
|
||
},
|
||
{"model_name": "gpt-4o", "model_info": {}},
|
||
]
|
||
|
||
sorted_models = _sort_models(all_models=models, sort_by="model_name", sort_order="asc")
|
||
displayed_order = [m["model_info"].get("team_public_model_name") or m["model_name"] for m in sorted_models]
|
||
assert displayed_order == [
|
||
"anthropic/claude",
|
||
"claude-haiku-4-5",
|
||
"gpt-4o",
|
||
]
|
||
|
||
def test_byok_models_sort_descending_by_display_name(self):
|
||
from litellm.proxy.proxy_server import _sort_models
|
||
|
||
models = [
|
||
{"model_name": "claude-haiku-4-5", "model_info": {}},
|
||
{
|
||
"model_name": "model_name_team-1_zzz",
|
||
"model_info": {"team_public_model_name": "zeta/model"},
|
||
},
|
||
{"model_name": "gpt-4o", "model_info": {}},
|
||
]
|
||
|
||
sorted_models = _sort_models(all_models=models, sort_by="model_name", sort_order="desc")
|
||
displayed_order = [m["model_info"].get("team_public_model_name") or m["model_name"] for m in sorted_models]
|
||
assert displayed_order == [
|
||
"zeta/model",
|
||
"gpt-4o",
|
||
"claude-haiku-4-5",
|
||
]
|
||
|
||
def test_empty_team_public_model_name_falls_back_to_model_name(self):
|
||
# Empty string for team_public_model_name (not None) must still
|
||
# fall back to model_name — otherwise BYOK rows with a blank
|
||
# display name would sort to the top.
|
||
from litellm.proxy.proxy_server import _sort_models
|
||
|
||
models = [
|
||
{"model_name": "alpha", "model_info": {"team_public_model_name": ""}},
|
||
{"model_name": "beta", "model_info": {}},
|
||
]
|
||
|
||
sorted_models = _sort_models(all_models=models, sort_by="model_name", sort_order="asc")
|
||
assert [m["model_name"] for m in sorted_models] == ["alpha", "beta"]
|
||
|
||
|
||
class TestDeleteDeploymentSync:
|
||
@pytest.mark.asyncio
|
||
async def test_delete_deployment_evicts_model_when_all_db_models_deleted(self):
|
||
"""
|
||
Regression test for #28443.
|
||
When all DB models are deleted, _delete_deployment must evict them from
|
||
the router. The old code returned 0 early when db_models was empty.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_router = MagicMock()
|
||
mock_router.get_model_ids.return_value = ["model-id-to-evict"]
|
||
mock_router.delete_deployment.return_value = MagicMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.llm_router", mock_router):
|
||
with patch.object(proxy_config, "get_config", AsyncMock(return_value={"model_list": []})):
|
||
still_desired = await proxy_config._delete_deployment(db_models=[])
|
||
|
||
mock_router.delete_deployment.assert_called_once_with(id="model-id-to-evict")
|
||
assert still_desired == frozenset(), (
|
||
"an empty db and an empty config want nothing, which must stay distinct from "
|
||
f"the None returned when no reconcile ran at all; got {still_desired}"
|
||
)
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_llm_router_skips_update_on_db_fetch_failure(self):
|
||
"""
|
||
When _get_models_from_db returns None (transient DB failure), _update_llm_router
|
||
must return early without touching the router.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock, patch
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_router = MagicMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.llm_router", mock_router):
|
||
with patch.object(proxy_config, "get_config", AsyncMock(return_value={})):
|
||
await proxy_config._update_llm_router(new_models=None, proxy_logging_obj=MagicMock())
|
||
|
||
mock_router.delete_deployment.assert_not_called()
|
||
mock_router.upsert_deployment.assert_not_called()
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_models_from_db_returns_none_on_exception(self):
|
||
"""
|
||
_get_models_from_db must return None (not []) when the DB raises an exception,
|
||
so callers can distinguish a transient failure from a genuinely empty DB.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
proxy_config = ProxyConfig()
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(side_effect=Exception("DB connection lost"))
|
||
|
||
result = await proxy_config._get_models_from_db(prisma_client=mock_prisma)
|
||
|
||
assert result is None, f"Expected None on DB failure to signal fetch error, got {result!r}"
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_models_from_db_reads_from_writer_not_replica(self):
|
||
"""
|
||
Regression for #38556: with DATABASE_URL_READ_REPLICA configured, the model
|
||
reconcile after /model/new used to read via the replica, so a lagging replica
|
||
made the reload miss the just-committed row and fail the request with a 500.
|
||
The reconcile read must be pinned to the writer.
|
||
"""
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
writer_inner = MagicMock(name="writer_prisma")
|
||
reader_inner = MagicMock(name="reader_prisma")
|
||
committed_row = MagicMock(name="just_committed_model_row")
|
||
writer_inner.litellm_proxymodeltable.find_many = AsyncMock(return_value=[committed_row])
|
||
reader_inner.litellm_proxymodeltable.find_many = AsyncMock(return_value=[])
|
||
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db = RoutingPrismaWrapper(
|
||
writer=PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False),
|
||
reader=PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False),
|
||
)
|
||
|
||
result = await ProxyConfig()._get_models_from_db(prisma_client=mock_prisma)
|
||
|
||
assert result == [committed_row], f"Expected the writer's just-committed row, got {result!r}"
|
||
reader_inner.litellm_proxymodeltable.find_many.assert_not_awaited()
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_get_models_from_db_falls_back_to_replica_when_writer_down(self):
|
||
"""
|
||
The writer pin must not break reader-only degraded mode: a proxy that
|
||
starts during a primary outage (writer connect failed, replica healthy)
|
||
must still load DB-backed models through the replica instead of sending
|
||
the reconcile read to the unavailable writer.
|
||
"""
|
||
from types import SimpleNamespace
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||
from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
writer_inner = MagicMock(name="writer_prisma")
|
||
reader_inner = MagicMock(name="reader_prisma")
|
||
replica_row = MagicMock(name="replica_model_row")
|
||
writer_inner.litellm_proxymodeltable = SimpleNamespace(
|
||
find_many=AsyncMock(side_effect=RuntimeError("writer unreachable")),
|
||
create=MagicMock(name="writer_create"),
|
||
)
|
||
reader_inner.litellm_proxymodeltable = SimpleNamespace(
|
||
find_many=AsyncMock(return_value=[replica_row]),
|
||
create=MagicMock(name="reader_create"),
|
||
)
|
||
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db = RoutingPrismaWrapper(
|
||
writer=PrismaWrapper(original_prisma=writer_inner, iam_token_db_auth=False),
|
||
reader=PrismaWrapper(original_prisma=reader_inner, iam_token_db_auth=False),
|
||
)
|
||
mock_prisma.db._writer_unavailable = True
|
||
|
||
result = await ProxyConfig()._get_models_from_db(prisma_client=mock_prisma)
|
||
|
||
assert result == [replica_row], f"Expected the replica's rows in degraded mode, got {result!r}"
|
||
writer_inner.litellm_proxymodeltable.find_many.assert_not_awaited()
|
||
|
||
|
||
def test_get_config_list_includes_cancel_on_disconnect(monkeypatch):
|
||
"""Follow-up to #30223: the flag must be discoverable via /config/list,
|
||
which requires both the ConfigGeneralSettings field and the allowed_args
|
||
entry in get_config_list; missing either silently hides it from the UI."""
|
||
import types
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from fastapi.testclient import TestClient
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import app
|
||
|
||
mock_prisma = MagicMock()
|
||
mock_config_table = MagicMock()
|
||
mock_config_table.find_first = AsyncMock(return_value=None)
|
||
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
|
||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||
)
|
||
try:
|
||
client = TestClient(app)
|
||
resp = client.get("/config/list", params={"config_type": "general_settings"})
|
||
assert resp.status_code == 200, resp.text
|
||
fields = {item["field_name"]: item for item in resp.json()}
|
||
assert "cancel_on_disconnect" in fields
|
||
assert fields["cancel_on_disconnect"]["field_type"] == "Boolean"
|
||
finally:
|
||
app.dependency_overrides.clear()
|
||
|
||
|
||
def test_get_config_list_includes_apply_user_budget_to_team_keys(monkeypatch):
|
||
"""Related to #12905: the opt-in must be discoverable via /config/list so it
|
||
renders as a Boolean toggle on the Admin UI General Settings table. This needs
|
||
both the ConfigGeneralSettings field and the allowed_args entry."""
|
||
import types
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from fastapi.testclient import TestClient
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import app
|
||
|
||
mock_prisma = MagicMock()
|
||
mock_config_table = MagicMock()
|
||
mock_config_table.find_first = AsyncMock(return_value=None)
|
||
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
|
||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||
)
|
||
try:
|
||
client = TestClient(app)
|
||
resp = client.get("/config/list", params={"config_type": "general_settings"})
|
||
assert resp.status_code == 200, resp.text
|
||
fields = {item["field_name"]: item for item in resp.json()}
|
||
assert "apply_user_budget_to_team_keys" in fields
|
||
assert fields["apply_user_budget_to_team_keys"]["field_type"] == "Boolean"
|
||
finally:
|
||
app.dependency_overrides.clear()
|
||
|
||
|
||
def test_get_config_list_includes_budget_exceeded_throttle_percentage(monkeypatch):
|
||
"""The throttle fraction is a litellm_settings scalar surfaced on the General
|
||
Settings table as a Float field so it sits with the other global limits; it
|
||
must appear in /config/list reading its live litellm.<attr> value."""
|
||
import types
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from fastapi.testclient import TestClient
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import app
|
||
|
||
mock_prisma = MagicMock()
|
||
mock_config_table = MagicMock()
|
||
mock_config_table.find_first = AsyncMock(return_value=None)
|
||
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
|
||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.15)
|
||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||
)
|
||
try:
|
||
client = TestClient(app)
|
||
resp = client.get("/config/list", params={"config_type": "general_settings"})
|
||
assert resp.status_code == 200, resp.text
|
||
fields = {item["field_name"]: item for item in resp.json()}
|
||
assert "budget_exceeded_throttle_percentage" in fields
|
||
assert fields["budget_exceeded_throttle_percentage"]["field_type"] == "Float"
|
||
assert fields["budget_exceeded_throttle_percentage"]["field_value"] == 0.15
|
||
finally:
|
||
app.dependency_overrides.clear()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_field_throttle_persists_to_litellm_settings(monkeypatch):
|
||
"""Editing the throttle Float row on the General Settings table routes to
|
||
litellm_settings (not general_settings): it sets litellm.<attr> live and
|
||
persists under litellm_settings so the runtime read is unchanged."""
|
||
from unittest.mock import MagicMock
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import (
|
||
ConfigFieldUpdate,
|
||
LitellmUserRoles,
|
||
UserAPIKeyAuth,
|
||
)
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
saved: dict = {}
|
||
|
||
async def fake_get_config():
|
||
return {"litellm_settings": {}}
|
||
|
||
async def fake_save_config(new_config=None):
|
||
saved.update(new_config or {})
|
||
|
||
monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
|
||
monkeypatch.setattr(ps.proxy_config, "save_config", fake_save_config)
|
||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", None)
|
||
|
||
admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="budget_exceeded_throttle_percentage",
|
||
field_value=0.1,
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
|
||
assert litellm.budget_exceeded_throttle_percentage == 0.1
|
||
assert saved["litellm_settings"]["budget_exceeded_throttle_percentage"] == 0.1
|
||
|
||
|
||
def test_get_config_list_includes_anthropic_prompt_caching_fields(monkeypatch):
|
||
"""The auto prompt caching flag and its ttl are litellm_settings globals surfaced on the
|
||
General Settings table, so an admin can turn caching on without hand-writing config. The
|
||
ttl is a Select and must ship its allowed values, or the table renders no editor for it."""
|
||
import types
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from fastapi.testclient import TestClient
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import app
|
||
|
||
mock_prisma = MagicMock()
|
||
mock_config_table = MagicMock()
|
||
mock_config_table.find_first = AsyncMock(return_value=None)
|
||
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
|
||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True)
|
||
monkeypatch.setattr(litellm, "anthropic_prompt_caching_ttl", "1h")
|
||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||
)
|
||
try:
|
||
client = TestClient(app)
|
||
resp = client.get("/config/list", params={"config_type": "general_settings"})
|
||
assert resp.status_code == 200, resp.text
|
||
fields = {item["field_name"]: item for item in resp.json()}
|
||
|
||
assert fields["enable_anthropic_prompt_caching"]["field_type"] == "Boolean"
|
||
assert fields["enable_anthropic_prompt_caching"]["field_value"] is True
|
||
|
||
assert fields["anthropic_prompt_caching_ttl"]["field_type"] == "Select"
|
||
assert fields["anthropic_prompt_caching_ttl"]["field_value"] == "1h"
|
||
assert fields["anthropic_prompt_caching_ttl"]["field_options"] == ["5m", "1h"]
|
||
|
||
# Both caching fields carry their sub-tab so the Admin UI can render them on a
|
||
# dedicated Prompt Caching tab, while ungrouped fields stay on General.
|
||
assert fields["enable_anthropic_prompt_caching"]["field_tab"] == "prompt_caching"
|
||
assert fields["anthropic_prompt_caching_ttl"]["field_tab"] == "prompt_caching"
|
||
assert fields["budget_exceeded_throttle_percentage"]["field_tab"] is None
|
||
finally:
|
||
app.dependency_overrides.clear()
|
||
|
||
|
||
def test_general_settings_ui_fields_are_db_overridable():
|
||
"""Every field the Admin UI can edit is a `litellm.<attr>` set via setattr on the handling
|
||
worker (`_persist_general_settings_ui_litellm_field`). Unless it is also in
|
||
LITELLM_SETTINGS_SAFE_DB_OVERRIDES, a config reload on a peer worker merges the DB value but
|
||
never applies it to the live attribute, so peer workers stay on their startup value.
|
||
|
||
This invariant is the guard against the two registries drifting: adding a UI-editable field
|
||
without enrolling it in the DB-override allowlist silently breaks cross-worker propagation.
|
||
"""
|
||
from litellm.constants import LITELLM_SETTINGS_SAFE_DB_OVERRIDES
|
||
from litellm.proxy.proxy_server import _GENERAL_SETTINGS_UI_LITELLM_FIELDS
|
||
|
||
missing = set(_GENERAL_SETTINGS_UI_LITELLM_FIELDS) - set(LITELLM_SETTINGS_SAFE_DB_OVERRIDES)
|
||
assert not missing, (
|
||
f"UI-editable litellm_settings fields missing from LITELLM_SETTINGS_SAFE_DB_OVERRIDES: {sorted(missing)}. "
|
||
"Add them, or they will not propagate to other workers when changed from the UI."
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_field_max_ui_session_budget_sets_live_value(monkeypatch):
|
||
"""LIT-4662: the dashboard session budget is editable from the Admin UI General tab.
|
||
A Dollar field must accept values above 1 (the old Float type capped at 1, which cannot
|
||
express a dollar budget), apply live via setattr, and persist under litellm_settings."""
|
||
from unittest.mock import MagicMock
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import (
|
||
ConfigFieldUpdate,
|
||
LitellmUserRoles,
|
||
UserAPIKeyAuth,
|
||
)
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
saved: dict = {}
|
||
|
||
async def fake_get_config():
|
||
return {"litellm_settings": {}}
|
||
|
||
async def fake_save_config(new_config=None):
|
||
saved.update(new_config or {})
|
||
|
||
monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
|
||
monkeypatch.setattr(ps.proxy_config, "save_config", fake_save_config)
|
||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||
monkeypatch.setattr(litellm, "max_ui_session_budget", 1.0)
|
||
|
||
admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="max_ui_session_budget",
|
||
field_value=25.0,
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
|
||
assert litellm.max_ui_session_budget == 25.0
|
||
assert saved["litellm_settings"]["max_ui_session_budget"] == 25.0
|
||
|
||
|
||
@pytest.mark.parametrize("bad_value", [True, "abc", -1, 0, [2.5]])
|
||
def test_validate_max_ui_session_budget_rejects_malformed(bad_value):
|
||
"""A Dollar field accepts only positive numbers; zero would block every dashboard
|
||
LLM call at mint and non-numerics would break session key generation."""
|
||
from fastapi import HTTPException
|
||
|
||
from litellm.proxy.proxy_server import _validate_general_settings_ui_litellm_value
|
||
|
||
with pytest.raises(HTTPException) as exc_info:
|
||
_validate_general_settings_ui_litellm_value("max_ui_session_budget", bad_value)
|
||
assert exc_info.value.status_code == 400
|
||
|
||
|
||
@pytest.mark.parametrize("empty_value", [None, ""])
|
||
def test_validate_max_ui_session_budget_empty_restores_default(empty_value):
|
||
"""Clearing the field in the UI restores the shipped $1 default rather than None;
|
||
None would silently remove the session spend guardrail (unlimited budget), which
|
||
must stay a deliberate config.yaml act (max_ui_session_budget: null)."""
|
||
from litellm.proxy.proxy_server import _validate_general_settings_ui_litellm_value
|
||
|
||
assert _validate_general_settings_ui_litellm_value("max_ui_session_budget", empty_value) == 1.0
|
||
|
||
|
||
def test_general_settings_ui_defaults_unchanged_for_existing_fields():
|
||
"""The spec-default mechanism added for max_ui_session_budget must not change what
|
||
clearing the pre-existing fields restores (None for Float/Select, False for Boolean)."""
|
||
from litellm.proxy.proxy_server import (
|
||
_GENERAL_SETTINGS_UI_LITELLM_FIELDS,
|
||
_general_settings_ui_litellm_default,
|
||
)
|
||
|
||
assert (
|
||
_general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["budget_exceeded_throttle_percentage"])
|
||
is None
|
||
)
|
||
assert (
|
||
_general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["enable_anthropic_prompt_caching"])
|
||
is False
|
||
)
|
||
assert (
|
||
_general_settings_ui_litellm_default(_GENERAL_SETTINGS_UI_LITELLM_FIELDS["anthropic_prompt_caching_ttl"])
|
||
is None
|
||
)
|
||
|
||
|
||
@pytest.mark.parametrize(
|
||
"field_name, db_value",
|
||
[
|
||
("enable_anthropic_prompt_caching", True),
|
||
("anthropic_prompt_caching_ttl", "1h"),
|
||
],
|
||
)
|
||
def test_prompt_caching_settings_propagate_on_config_reload(monkeypatch, field_name, db_value):
|
||
"""A UI toggle on one worker persists to the DB; a peer worker picks it up only when the
|
||
config reload applies the safe-override allowlist. Regression for the fields being absent
|
||
from that allowlist, which left peer workers stale."""
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
# peer worker booted with the opposite/absent value
|
||
monkeypatch.setattr(litellm, field_name, False if isinstance(db_value, bool) else None)
|
||
|
||
pc = ps.ProxyConfig()
|
||
pc._update_config_fields(
|
||
current_config={"litellm_settings": {}},
|
||
param_name="litellm_settings",
|
||
db_param_value={field_name: db_value},
|
||
)
|
||
|
||
assert getattr(litellm, field_name) == db_value
|
||
|
||
|
||
def test_get_config_list_marks_untouched_prompt_caching_flag_as_not_set(monkeypatch):
|
||
"""The flag defaults to False rather than None, so a plain 'is not None' check would
|
||
report the default as 'In Config' and imply an admin had set it."""
|
||
import types
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from fastapi.testclient import TestClient
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import app
|
||
|
||
mock_prisma = MagicMock()
|
||
mock_config_table = MagicMock()
|
||
mock_config_table.find_first = AsyncMock(return_value=None)
|
||
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
|
||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||
monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False)
|
||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||
user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN
|
||
)
|
||
try:
|
||
client = TestClient(app)
|
||
resp = client.get("/config/list", params={"config_type": "general_settings"})
|
||
fields = {item["field_name"]: item for item in resp.json()}
|
||
assert fields["enable_anthropic_prompt_caching"]["stored_in_db"] is None
|
||
finally:
|
||
app.dependency_overrides.clear()
|
||
|
||
|
||
@pytest.mark.parametrize(
|
||
"field_name, field_value",
|
||
[
|
||
("enable_anthropic_prompt_caching", True),
|
||
("enable_anthropic_prompt_caching", False),
|
||
("anthropic_prompt_caching_ttl", "5m"),
|
||
("anthropic_prompt_caching_ttl", "1h"),
|
||
],
|
||
)
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_field_prompt_caching_persists_to_litellm_settings(monkeypatch, field_name, field_value):
|
||
"""Toggling either row must set litellm.<attr> live and persist under litellm_settings,
|
||
so the running proxy caches immediately and still does after a restart."""
|
||
from unittest.mock import MagicMock
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import (
|
||
ConfigFieldUpdate,
|
||
LitellmUserRoles,
|
||
UserAPIKeyAuth,
|
||
)
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
saved: dict = {}
|
||
|
||
async def fake_get_config():
|
||
return {"litellm_settings": {}}
|
||
|
||
async def fake_save_config(new_config=None):
|
||
saved.update(new_config or {})
|
||
|
||
monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
|
||
monkeypatch.setattr(ps.proxy_config, "save_config", fake_save_config)
|
||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||
monkeypatch.setattr(litellm, field_name, None)
|
||
|
||
admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(field_name=field_name, field_value=field_value, config_type="general_settings"),
|
||
user_api_key_dict=admin,
|
||
)
|
||
|
||
assert getattr(litellm, field_name) == field_value
|
||
assert saved["litellm_settings"][field_name] == field_value
|
||
|
||
|
||
@pytest.mark.parametrize(
|
||
"field_name, bad_value",
|
||
[
|
||
("enable_anthropic_prompt_caching", "yes"),
|
||
("enable_anthropic_prompt_caching", 1),
|
||
("anthropic_prompt_caching_ttl", "10m"),
|
||
("anthropic_prompt_caching_ttl", "1H"),
|
||
("anthropic_prompt_caching_ttl", 3600),
|
||
],
|
||
)
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_field_prompt_caching_rejects_invalid(monkeypatch, field_name, bad_value):
|
||
"""An unsupported ttl must be refused here rather than reaching Anthropic verbatim."""
|
||
from unittest.mock import MagicMock
|
||
|
||
from fastapi import HTTPException
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import (
|
||
ConfigFieldUpdate,
|
||
LitellmUserRoles,
|
||
UserAPIKeyAuth,
|
||
)
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
async def fake_get_config():
|
||
return {"litellm_settings": {}}
|
||
|
||
monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
|
||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||
monkeypatch.setattr(litellm, field_name, None)
|
||
|
||
admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||
with pytest.raises(HTTPException) as exc:
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(field_name=field_name, field_value=bad_value, config_type="general_settings"),
|
||
user_api_key_dict=admin,
|
||
)
|
||
assert exc.value.status_code == 400
|
||
assert getattr(litellm, field_name) is None
|
||
|
||
|
||
@pytest.mark.parametrize(
|
||
"field_name, expected_default",
|
||
[
|
||
("enable_anthropic_prompt_caching", False),
|
||
("anthropic_prompt_caching_ttl", None),
|
||
("budget_exceeded_throttle_percentage", None),
|
||
],
|
||
)
|
||
@pytest.mark.asyncio
|
||
async def test_reset_config_field_restores_type_default(monkeypatch, field_name, expected_default):
|
||
"""Reset must restore each field's own default. Blanket None would leave the boolean flag
|
||
set to None, which is not a bool and would read as neither on nor off."""
|
||
from unittest.mock import MagicMock
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import (
|
||
ConfigFieldDelete,
|
||
LitellmUserRoles,
|
||
UserAPIKeyAuth,
|
||
)
|
||
from litellm.proxy.proxy_server import delete_config_general_settings
|
||
|
||
saved: dict = {}
|
||
|
||
async def fake_get_config():
|
||
return {"litellm_settings": {field_name: "stale"}}
|
||
|
||
async def fake_save_config(new_config=None):
|
||
saved.update(new_config or {})
|
||
|
||
monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
|
||
monkeypatch.setattr(ps.proxy_config, "save_config", fake_save_config)
|
||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||
monkeypatch.setattr(litellm, field_name, "stale")
|
||
|
||
admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||
await delete_config_general_settings(
|
||
data=ConfigFieldDelete(field_name=field_name, config_type="general_settings"),
|
||
user_api_key_dict=admin,
|
||
)
|
||
|
||
assert getattr(litellm, field_name) is expected_default
|
||
assert field_name not in saved["litellm_settings"]
|
||
|
||
|
||
@pytest.mark.parametrize("bad_value", [0, -0.1, 1.5, True])
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_field_throttle_rejects_invalid(monkeypatch, bad_value):
|
||
from unittest.mock import MagicMock
|
||
|
||
from fastapi import HTTPException
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import (
|
||
ConfigFieldUpdate,
|
||
LitellmUserRoles,
|
||
UserAPIKeyAuth,
|
||
)
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
async def fake_get_config():
|
||
return {"litellm_settings": {}}
|
||
|
||
monkeypatch.setattr(ps.proxy_config, "get_config", fake_get_config)
|
||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", None)
|
||
|
||
admin = UserAPIKeyAuth(api_key="k", user_id="a", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||
with pytest.raises(HTTPException) as exc:
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="budget_exceeded_throttle_percentage",
|
||
field_value=bad_value,
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
assert exc.value.status_code == 400
|
||
assert litellm.budget_exceeded_throttle_percentage is None
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_field_throttle_rejected_for_non_admin(monkeypatch):
|
||
from unittest.mock import MagicMock
|
||
|
||
from fastapi import HTTPException
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import (
|
||
ConfigFieldUpdate,
|
||
LitellmUserRoles,
|
||
UserAPIKeyAuth,
|
||
)
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
monkeypatch.setattr(ps, "prisma_client", MagicMock())
|
||
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", None)
|
||
|
||
non_admin = UserAPIKeyAuth(api_key="k", user_id="u", user_role=LitellmUserRoles.INTERNAL_USER)
|
||
with pytest.raises(HTTPException):
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="budget_exceeded_throttle_percentage",
|
||
field_value=0.1,
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=non_admin,
|
||
)
|
||
assert litellm.budget_exceeded_throttle_percentage is None
|
||
|
||
|
||
def test_preserve_redacted_plugin_keys_keeps_stored_credential():
|
||
"""A redacted or blank plugin_key on update must not overwrite the real key."""
|
||
from litellm.proxy.proxy_server import _preserve_redacted_plugin_keys
|
||
|
||
existing = [{"name": "p1", "url": "https://p1", "plugin_key": "sk-real-1"}]
|
||
|
||
redacted = _preserve_redacted_plugin_keys([{"name": "p1", "url": "https://p1-new", "plugin_key": "***"}], existing)
|
||
assert redacted == [{"name": "p1", "url": "https://p1-new", "plugin_key": "sk-real-1"}]
|
||
|
||
blanked = _preserve_redacted_plugin_keys([{"name": "p1", "url": "https://p1", "plugin_key": ""}], existing)
|
||
assert blanked[0]["plugin_key"] == "sk-real-1"
|
||
|
||
|
||
def test_preserve_redacted_plugin_keys_sets_new_and_drops_orphan_placeholder():
|
||
"""A real new key replaces; a placeholder with no stored key is dropped, never persisted."""
|
||
from litellm.proxy.proxy_server import _preserve_redacted_plugin_keys
|
||
|
||
existing = [{"name": "p1", "url": "https://p1", "plugin_key": "sk-real-1"}]
|
||
|
||
rotated = _preserve_redacted_plugin_keys([{"name": "p1", "url": "https://p1", "plugin_key": "sk-new"}], existing)
|
||
assert rotated[0]["plugin_key"] == "sk-new"
|
||
|
||
new_plugin = _preserve_redacted_plugin_keys([{"name": "p2", "url": "https://p2", "plugin_key": "***"}], existing)
|
||
assert "plugin_key" not in new_plugin[0]
|
||
|
||
|
||
def _config_field_info_client(monkeypatch, user_role):
|
||
import types
|
||
from unittest.mock import AsyncMock, MagicMock
|
||
|
||
from fastapi.testclient import TestClient
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
from litellm.proxy._types import UserAPIKeyAuth
|
||
from litellm.proxy.proxy_server import app
|
||
|
||
db_record = types.SimpleNamespace(
|
||
param_value={
|
||
"master_key": "sk-super-secret-master",
|
||
"database_url": "postgresql://user:p4ssw0rd@db:5432/litellm",
|
||
"pass_through_endpoints": [
|
||
{
|
||
"path": "/upstream",
|
||
"target": "https://upstream.example.com",
|
||
"headers": {"Authorization": "Bearer sk-upstream-secret"},
|
||
}
|
||
],
|
||
"max_parallel_requests": 100,
|
||
}
|
||
)
|
||
mock_config_table = MagicMock()
|
||
mock_config_table.find_first = AsyncMock(return_value=db_record)
|
||
mock_prisma = MagicMock()
|
||
mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table)
|
||
monkeypatch.setattr(ps, "prisma_client", mock_prisma)
|
||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(user_id="u", user_role=user_role)
|
||
return TestClient(app)
|
||
|
||
|
||
def test_config_field_info_redacts_secrets_for_view_only_admin(monkeypatch):
|
||
"""/config/field/info gates on _user_has_admin_view, which also grants
|
||
PROXY_ADMIN_VIEW_ONLY. A view-only admin reading master_key/database_url verbatim is
|
||
effectively a full admin. Secret-bearing fields must come back REDACTED for anyone who
|
||
is not a FULL PROXY_ADMIN, while non-secret fields stay readable."""
|
||
from litellm.proxy._types import LitellmUserRoles
|
||
|
||
client = _config_field_info_client(monkeypatch, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)
|
||
try:
|
||
for secret_field in ("master_key", "database_url", "pass_through_endpoints"):
|
||
resp = client.get("/config/field/info", params={"field_name": secret_field})
|
||
assert resp.status_code == 200, resp.text
|
||
body = resp.json()
|
||
assert body["field_value"] == "REDACTED"
|
||
assert "secret" not in str(body["field_value"])
|
||
assert "p4ssw0rd" not in str(body["field_value"])
|
||
|
||
resp = client.get("/config/field/info", params={"field_name": "max_parallel_requests"})
|
||
assert resp.status_code == 200, resp.text
|
||
assert resp.json()["field_value"] == 100
|
||
finally:
|
||
app.dependency_overrides.clear()
|
||
|
||
|
||
def test_config_field_info_returns_raw_secrets_for_full_admin(monkeypatch):
|
||
"""the redaction must not over-apply. A FULL PROXY_ADMIN still
|
||
needs the real master_key value to populate the admin edit form."""
|
||
from litellm.proxy._types import LitellmUserRoles
|
||
|
||
client = _config_field_info_client(monkeypatch, LitellmUserRoles.PROXY_ADMIN)
|
||
try:
|
||
resp = client.get("/config/field/info", params={"field_name": "master_key"})
|
||
assert resp.status_code == 200, resp.text
|
||
assert resp.json()["field_value"] == "sk-super-secret-master"
|
||
|
||
resp = client.get("/config/field/info", params={"field_name": "pass_through_endpoints"})
|
||
assert resp.status_code == 200, resp.text
|
||
assert resp.json()["field_value"][0]["headers"]["Authorization"] == "Bearer sk-upstream-secret"
|
||
finally:
|
||
app.dependency_overrides.clear()
|
||
|
||
|
||
def _fake_prisma_with_config(existing_param_value):
|
||
"""MagicMock prisma whose litellm_config row returns existing_param_value and
|
||
whose litellm_auditlog.create records the written audit row."""
|
||
fake = MagicMock()
|
||
config_row = MagicMock()
|
||
config_row.param_value = existing_param_value
|
||
fake.db.litellm_config.find_first = AsyncMock(return_value=config_row)
|
||
fake.db.litellm_config.upsert = AsyncMock(return_value=config_row)
|
||
fake.db.litellm_auditlog.create = AsyncMock()
|
||
return fake
|
||
|
||
|
||
def test_dump_redacted_config_redacts_secret_leaves():
|
||
from litellm.proxy.proxy_server import _dump_redacted_config
|
||
|
||
assert _dump_redacted_config(None) is None
|
||
|
||
restored = json.loads(
|
||
_dump_redacted_config(
|
||
{
|
||
"api_key": "sk-leak",
|
||
"model": "gpt-4",
|
||
"nested": {"aws_secret_access_key": "abc", "region": "us-east-1"},
|
||
}
|
||
)
|
||
)
|
||
assert restored["api_key"] == "REDACTED"
|
||
assert restored["model"] == "gpt-4"
|
||
assert restored["nested"]["aws_secret_access_key"] == "REDACTED"
|
||
assert restored["nested"]["region"] == "us-east-1"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_create_config_audit_log_writes_redacted_entry(monkeypatch):
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
from litellm.proxy._types import LitellmTableNames
|
||
from litellm.proxy.proxy_server import create_config_audit_log
|
||
|
||
fake = _fake_prisma_with_config({})
|
||
monkeypatch.setattr(proxy_server_module, "prisma_client", fake)
|
||
monkeypatch.setattr(proxy_server_module, "premium_user", True)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", True)
|
||
|
||
caller = UserAPIKeyAuth(api_key="hashed-key-abc", user_id="admin-7")
|
||
await create_config_audit_log(
|
||
"router_settings",
|
||
"updated",
|
||
{"routing_strategy": "simple-shuffle", "api_key": "sk-old"},
|
||
{"routing_strategy": "latency-based", "api_key": "sk-new"},
|
||
caller,
|
||
)
|
||
|
||
fake.db.litellm_auditlog.create.assert_awaited_once()
|
||
written = fake.db.litellm_auditlog.create.call_args.kwargs["data"]
|
||
assert written["table_name"] == LitellmTableNames.CONFIG_TABLE_NAME.value
|
||
assert written["object_id"] == "router_settings"
|
||
assert written["action"] == "updated"
|
||
assert written["changed_by"] == "admin-7"
|
||
assert written["changed_by_api_key"] == "hashed-key-abc"
|
||
|
||
before = json.loads(written["before_value"])
|
||
after = json.loads(written["updated_values"])
|
||
assert before["routing_strategy"] == "simple-shuffle"
|
||
assert after["routing_strategy"] == "latency-based"
|
||
assert "sk-old" not in written["before_value"]
|
||
assert "sk-new" not in written["updated_values"]
|
||
assert before["api_key"] != "sk-old"
|
||
assert after["api_key"] != "sk-new"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_create_config_audit_log_noop_when_store_audit_logs_disabled(monkeypatch):
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
from litellm.proxy.proxy_server import create_config_audit_log
|
||
|
||
fake = _fake_prisma_with_config({})
|
||
monkeypatch.setattr(proxy_server_module, "prisma_client", fake)
|
||
monkeypatch.setattr(proxy_server_module, "premium_user", True)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||
|
||
await create_config_audit_log(
|
||
"router_settings",
|
||
"updated",
|
||
{},
|
||
{"a": 1},
|
||
UserAPIKeyAuth(api_key="k", user_id="u"),
|
||
)
|
||
fake.db.litellm_auditlog.create.assert_not_called()
|
||
|
||
|
||
def test_dump_redacted_config_serializes_non_json_native_values():
|
||
"""YAML-loaded config can contain datetime/date/custom values that plain
|
||
json.dumps refuses. Without default=str the audit write turns into a 500
|
||
after the config change has already committed; the sibling audit-log
|
||
serializers in team_endpoints.py use default=str for the same reason."""
|
||
from datetime import datetime, timezone
|
||
|
||
from litellm.proxy.proxy_server import _dump_redacted_config
|
||
|
||
out = _dump_redacted_config({"updated_at": datetime(2026, 6, 30, tzinfo=timezone.utc)})
|
||
assert out is not None
|
||
restored = json.loads(out)
|
||
assert "2026-06-30" in restored["updated_at"]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_general_settings_emits_audit_log(monkeypatch):
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
from litellm.proxy._types import ConfigFieldUpdate
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
existing = {"max_parallel_requests": 5, "some_api_key": "sk-stored-secret"}
|
||
fake = _fake_prisma_with_config(existing)
|
||
monkeypatch.setattr(proxy_server_module, "prisma_client", fake)
|
||
monkeypatch.setattr(proxy_server_module, "premium_user", True)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", True)
|
||
|
||
admin = UserAPIKeyAuth(
|
||
api_key="hashed-admin",
|
||
user_id="admin-1",
|
||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||
)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="max_parallel_requests",
|
||
field_value=42,
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
# Audit is scheduled via asyncio.create_task; yield so it runs.
|
||
await asyncio.sleep(0)
|
||
|
||
fake.db.litellm_auditlog.create.assert_awaited_once()
|
||
written = fake.db.litellm_auditlog.create.call_args.kwargs["data"]
|
||
assert written["table_name"] == "LiteLLM_Config"
|
||
assert written["object_id"] == "general_settings"
|
||
assert written["action"] == "updated"
|
||
assert written["changed_by"] == "admin-1"
|
||
|
||
before = json.loads(written["before_value"])
|
||
after = json.loads(written["updated_values"])
|
||
assert before["max_parallel_requests"] == 5
|
||
assert after["max_parallel_requests"] == 42
|
||
assert "sk-stored-secret" not in written["before_value"]
|
||
assert "sk-stored-secret" not in written["updated_values"]
|
||
assert before["some_api_key"] != "sk-stored-secret"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_field_rejects_out_of_range_alerting_args(monkeypatch):
|
||
"""Out-of-range alerting_args must be rejected at save time. If they land in the
|
||
DB, SlackAlertingArgs raises during the config reload and alerting breaks."""
|
||
from unittest.mock import MagicMock
|
||
|
||
from fastapi import HTTPException
|
||
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
from litellm.proxy._types import ConfigFieldUpdate
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
monkeypatch.setattr(proxy_server_module, "prisma_client", MagicMock())
|
||
|
||
admin = UserAPIKeyAuth(
|
||
api_key="hashed-admin",
|
||
user_id="admin-1",
|
||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||
)
|
||
with pytest.raises(HTTPException) as exc_info:
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="alerting_args",
|
||
field_value={
|
||
"daily_spend_per_user_threshold": -5.0,
|
||
"user_spend_check_interval": 20,
|
||
},
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
|
||
assert exc_info.value.status_code == 400
|
||
error_msg = exc_info.value.detail["error"]
|
||
assert "daily_spend_per_user_threshold" in error_msg
|
||
assert "user_spend_check_interval" in error_msg
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_field_accepts_valid_alerting_args(monkeypatch):
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
from litellm.proxy._types import ConfigFieldUpdate
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
fake = _fake_prisma_with_config({})
|
||
monkeypatch.setattr(proxy_server_module, "prisma_client", fake)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||
|
||
admin = UserAPIKeyAuth(
|
||
api_key="hashed-admin",
|
||
user_id="admin-1",
|
||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||
)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="alerting_args",
|
||
field_value={
|
||
"daily_spend_per_user_threshold": 5.0,
|
||
"user_spend_check_interval": 60,
|
||
},
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
|
||
written = json.loads(fake.db.litellm_config.upsert.call_args.kwargs["data"]["update"]["param_value"])
|
||
assert written["alerting_args"]["daily_spend_per_user_threshold"] == 5.0
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_update_config_general_settings_applies_ssrf_globals(monkeypatch):
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
from litellm.proxy._types import ConfigFieldUpdate
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
|
||
fake = _fake_prisma_with_config({})
|
||
monkeypatch.setattr(proxy_server_module, "prisma_client", fake)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", False)
|
||
monkeypatch.setattr(litellm, "user_url_validation", True)
|
||
monkeypatch.setattr(litellm, "user_url_allowed_hosts", [])
|
||
monkeypatch.setattr(litellm, "provider_url_destination_allowed_hosts", [])
|
||
|
||
admin = UserAPIKeyAuth(
|
||
api_key="hashed-admin",
|
||
user_id="admin-1",
|
||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||
)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="user_url_validation",
|
||
field_value="false",
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="user_url_allowed_hosts",
|
||
field_value=["internal.example"],
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="provider_url_destination_allowed_hosts",
|
||
field_value=["provider.example"],
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
await asyncio.sleep(0)
|
||
|
||
assert litellm.user_url_validation is False
|
||
assert litellm.user_url_allowed_hosts == ["internal.example"]
|
||
assert litellm.provider_url_destination_allowed_hosts == ["provider.example"]
|
||
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="user_url_allowed_hosts",
|
||
field_value=None,
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name="provider_url_destination_allowed_hosts",
|
||
field_value=None,
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
await asyncio.sleep(0)
|
||
|
||
assert litellm.user_url_allowed_hosts is None
|
||
assert litellm.provider_url_destination_allowed_hosts is None
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_delete_config_general_settings_emits_deleted_audit_log(monkeypatch):
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
from litellm.proxy._types import ConfigFieldDelete
|
||
from litellm.proxy.proxy_server import delete_config_general_settings
|
||
|
||
existing = {"max_parallel_requests": 5}
|
||
fake = _fake_prisma_with_config(existing)
|
||
monkeypatch.setattr(proxy_server_module, "prisma_client", fake)
|
||
monkeypatch.setattr(proxy_server_module, "premium_user", True)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", True)
|
||
|
||
admin = UserAPIKeyAuth(
|
||
api_key="hashed-admin",
|
||
user_id="admin-1",
|
||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||
)
|
||
await delete_config_general_settings(
|
||
data=ConfigFieldDelete(field_name="max_parallel_requests", config_type="general_settings"),
|
||
user_api_key_dict=admin,
|
||
)
|
||
# Audit is scheduled via asyncio.create_task; yield so it runs.
|
||
await asyncio.sleep(0)
|
||
|
||
fake.db.litellm_auditlog.create.assert_awaited_once()
|
||
written = fake.db.litellm_auditlog.create.call_args.kwargs["data"]
|
||
assert written["object_id"] == "general_settings"
|
||
assert written["action"] == "deleted"
|
||
before = json.loads(written["before_value"])
|
||
after = json.loads(written["updated_values"])
|
||
assert before["max_parallel_requests"] == 5
|
||
assert "max_parallel_requests" not in after
|
||
|
||
|
||
def test_update_config_audits_every_written_section(_update_config_setup, monkeypatch):
|
||
"""/config/update must emit one audit row per section it writes, so each
|
||
of the four call sites (general_settings, environment_variables,
|
||
litellm_settings, router_settings) is mutation-protected. litellm_settings
|
||
is the row that holds default_internal_user_params ("default user settings")."""
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
|
||
client, prisma, restore = _update_config_setup(initial_rows={"litellm_settings": {"drop_params": True}})
|
||
audit_create = AsyncMock()
|
||
prisma.db.litellm_auditlog.create = audit_create
|
||
monkeypatch.setattr(proxy_server_module, "premium_user", True)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", True)
|
||
try:
|
||
resp = client.post(
|
||
"/config/update",
|
||
json={
|
||
"general_settings": {"store_prompts_in_spend_logs": True},
|
||
"environment_variables": {"FOO": "bar"},
|
||
"litellm_settings": {"default_internal_user_params": {"max_budget": 10}},
|
||
"router_settings": {"routing_strategy": "latency-based-routing"},
|
||
},
|
||
)
|
||
assert resp.status_code == 200, resp.text
|
||
|
||
audited = {
|
||
call.kwargs["data"]["object_id"]: call.kwargs["data"]["action"] for call in audit_create.await_args_list
|
||
}
|
||
assert audited == {
|
||
"general_settings": "updated",
|
||
"environment_variables": "updated",
|
||
"litellm_settings": "updated",
|
||
"router_settings": "updated",
|
||
}
|
||
for call in audit_create.await_args_list:
|
||
assert call.kwargs["data"]["table_name"] == "LiteLLM_Config"
|
||
assert call.kwargs["data"]["changed_by"] == "test_admin"
|
||
|
||
ls_call = next(c for c in audit_create.await_args_list if c.kwargs["data"]["object_id"] == "litellm_settings")
|
||
after = json.loads(ls_call.kwargs["data"]["updated_values"])
|
||
assert after["default_internal_user_params"] == {"max_budget": 10}
|
||
finally:
|
||
restore()
|
||
|
||
|
||
def test_delete_callback_audits_litellm_settings_deletion(_update_config_setup, monkeypatch):
|
||
"""/config/callback/delete must emit a deleted audit row for litellm_settings
|
||
capturing the success_callback list before and after removal."""
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
|
||
client, prisma, restore = _update_config_setup()
|
||
audit_create = AsyncMock()
|
||
prisma.db.litellm_auditlog.create = audit_create
|
||
monkeypatch.setattr(proxy_server_module, "premium_user", True)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", True)
|
||
|
||
from litellm.proxy.proxy_server import proxy_config as real_proxy_config
|
||
|
||
monkeypatch.setattr(
|
||
real_proxy_config,
|
||
"get_config",
|
||
AsyncMock(return_value={"litellm_settings": {"success_callback": ["langfuse", "datadog"]}}),
|
||
)
|
||
monkeypatch.setattr(real_proxy_config, "save_config", AsyncMock(return_value=None))
|
||
try:
|
||
resp = client.post("/config/callback/delete", json={"callback_name": "datadog"})
|
||
assert resp.status_code == 200, resp.text
|
||
|
||
audit_create.assert_awaited_once()
|
||
written = audit_create.await_args.kwargs["data"]
|
||
assert written["object_id"] == "litellm_settings"
|
||
assert written["action"] == "deleted"
|
||
before = json.loads(written["before_value"])
|
||
after = json.loads(written["updated_values"])
|
||
assert before["success_callback"] == ["langfuse", "datadog"]
|
||
assert after["success_callback"] == ["langfuse"]
|
||
finally:
|
||
restore()
|
||
|
||
|
||
def test_delete_callback_audits_before_reload_failure(_update_config_setup, monkeypatch):
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
|
||
client, prisma, restore = _update_config_setup()
|
||
audit_create = AsyncMock()
|
||
prisma.db.litellm_auditlog.create = audit_create
|
||
monkeypatch.setattr(proxy_server_module, "premium_user", True)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", True)
|
||
|
||
from litellm.proxy.proxy_server import proxy_config as real_proxy_config
|
||
|
||
monkeypatch.setattr(
|
||
real_proxy_config,
|
||
"get_config",
|
||
AsyncMock(return_value={"litellm_settings": {"success_callback": ["langfuse", "datadog"]}}),
|
||
)
|
||
monkeypatch.setattr(real_proxy_config, "save_config", AsyncMock(return_value=None))
|
||
monkeypatch.setattr(
|
||
real_proxy_config,
|
||
"add_deployment",
|
||
AsyncMock(side_effect=RuntimeError("reload failed")),
|
||
)
|
||
try:
|
||
resp = client.post("/config/callback/delete", json={"callback_name": "datadog"})
|
||
assert resp.status_code == 500, resp.text
|
||
|
||
audit_create.assert_awaited_once()
|
||
written = audit_create.await_args.kwargs["data"]
|
||
assert written["object_id"] == "litellm_settings"
|
||
assert written["action"] == "deleted"
|
||
finally:
|
||
restore()
|
||
|
||
|
||
def test_update_config_redacts_all_environment_variable_values(_update_config_setup, monkeypatch):
|
||
"""environment_variables hold credentials under arbitrary uppercase keys
|
||
(DATABASE_URL) that key-name secret matching misses, so every value in the
|
||
section must be redacted before the audit row is written; a plaintext
|
||
secret must never reach LiteLLM_AuditLog."""
|
||
import litellm.proxy.proxy_server as proxy_server_module
|
||
|
||
# DATABASE_URL is the bug class: an uppercase env key that key-name secret
|
||
# matching does NOT flag, so only whole-section value redaction protects it.
|
||
client, prisma, restore = _update_config_setup(
|
||
initial_rows={"environment_variables": {"DATABASE_URL": "enc:postgresql://OLDsecret@old.host:5432/db"}}
|
||
)
|
||
audit_create = AsyncMock()
|
||
prisma.db.litellm_auditlog.create = audit_create
|
||
monkeypatch.setattr(proxy_server_module, "premium_user", True)
|
||
monkeypatch.setattr(litellm, "store_audit_logs", True)
|
||
try:
|
||
resp = client.post(
|
||
"/config/update",
|
||
json={
|
||
"environment_variables": {
|
||
"DATABASE_URL": "postgresql://u:p@db.internal:5432/litellm",
|
||
"LOG_LEVEL": "debug",
|
||
}
|
||
},
|
||
)
|
||
assert resp.status_code == 200, resp.text
|
||
|
||
env_call = next(
|
||
c for c in audit_create.await_args_list if c.kwargs["data"]["object_id"] == "environment_variables"
|
||
)
|
||
data = env_call.kwargs["data"]
|
||
|
||
# the pre-existing secret must be redacted in the before snapshot
|
||
before = json.loads(data["before_value"])
|
||
assert before == {"DATABASE_URL": "REDACTED"}
|
||
assert "OLDsecret" not in data["before_value"]
|
||
assert "old.host" not in data["before_value"]
|
||
|
||
# the newly-written values must be redacted in the after snapshot
|
||
after = json.loads(data["updated_values"])
|
||
assert after == {"DATABASE_URL": "REDACTED", "LOG_LEVEL": "REDACTED"}
|
||
assert "postgresql://" not in data["updated_values"]
|
||
assert "db.internal" not in data["updated_values"]
|
||
finally:
|
||
restore()
|
||
|
||
|
||
class _EnvBuiltRedisCache(RedisCache):
|
||
"""RedisCache stand-in that records its constructor kwargs and never
|
||
opens a network connection, so tests can assert which connection params
|
||
the proxy used to build its coordination Redis. `ping()` reports reachable
|
||
by default, matching a real Redis the env fallback should adopt."""
|
||
|
||
def __init__(self, **kwargs):
|
||
self.init_kwargs = kwargs
|
||
|
||
async def ping(self) -> bool:
|
||
return True
|
||
|
||
|
||
class _UnreachableRedisCache(_EnvBuiltRedisCache):
|
||
"""Same stand-in, but `ping()` fails like a REDIS_* env var naming a Redis
|
||
that is not actually reachable (wrong host, no service running, ...)."""
|
||
|
||
async def ping(self) -> bool:
|
||
raise ConnectionError("connection refused")
|
||
|
||
|
||
@contextlib.contextmanager
|
||
def _patched_coordination_redis_module_state(
|
||
*,
|
||
spend_cache: DualCache,
|
||
config_cache: types.SimpleNamespace,
|
||
redis_cache_class: type = _EnvBuiltRedisCache,
|
||
):
|
||
"""Stub every `litellm.proxy.proxy_server` global that
|
||
`_attach_redis_usage_cache` (and its callers) can write to, shared by the
|
||
whole coordination-Redis test family below.
|
||
|
||
Centralizing this is not just DRY: `_attach_redis_usage_cache` always sets
|
||
`cli_sso_session_cache.redis_cache` unconditionally, and a call site that
|
||
forgets to patch that one real (persistent) global leaks a throwaway
|
||
Redis stand-in into it for the rest of the pytest session, breaking
|
||
unrelated tests that run later. One patched-state helper means a new call
|
||
site cannot forget a global this family already knows to isolate.
|
||
"""
|
||
with (
|
||
patch.object(proxy_server_module, "redis_usage_cache", None),
|
||
patch.object(proxy_server_module, "spend_counter_cache", spend_cache),
|
||
patch.object(proxy_server_module, "user_api_key_cache", DualCache()),
|
||
patch.object(proxy_server_module, "cli_sso_session_cache", DualCache()),
|
||
patch.object(proxy_server_module, "llm_router", None),
|
||
patch.object(proxy_server_module, "litellm_config_cache", config_cache),
|
||
patch.object(proxy_server_module, "RedisCache", redis_cache_class),
|
||
patch.object(proxy_server_module, "RedisClusterCache", _EnvBuiltClusterCache),
|
||
):
|
||
yield
|
||
|
||
|
||
def _run_init_cache_with_backend(cache_backend, redis_env_kwargs):
|
||
"""Run ProxyConfig._init_cache with a stubbed response-cache backend and a
|
||
controlled REDIS_* environment, returning (redis_usage_cache,
|
||
spend_counter redis, config-cache redis) as observed after the call."""
|
||
mock_litellm_cache = MagicMock()
|
||
mock_litellm_cache.cache = cache_backend
|
||
fresh_spend_cache = DualCache()
|
||
fresh_config_cache = types.SimpleNamespace(redis_cache=None)
|
||
|
||
with (
|
||
_patched_coordination_redis_module_state(spend_cache=fresh_spend_cache, config_cache=fresh_config_cache),
|
||
patch(
|
||
"litellm._redis._redis_kwargs_from_environment",
|
||
return_value=redis_env_kwargs,
|
||
),
|
||
patch("litellm.Cache", return_value=mock_litellm_cache),
|
||
):
|
||
litellm.cache = None
|
||
resolved = proxy_server_module.ProxyConfig()._init_cache(cache_params={"type": "qdrant-semantic"})
|
||
return (
|
||
resolved,
|
||
fresh_spend_cache.redis_cache,
|
||
fresh_config_cache.redis_cache,
|
||
)
|
||
|
||
|
||
def test_init_cache_non_redis_backend_builds_usage_redis_from_environment():
|
||
"""A semantic (non-Redis-KV) response cache must not disable the proxy's
|
||
coordination Redis: when REDIS_* env vars provide a connection,
|
||
_init_cache builds a standalone usage cache so cross-pod rate limits,
|
||
spend tracking, and the pod lock manager stay Redis-backed."""
|
||
usage_cache, spend_redis, config_redis = _run_init_cache_with_backend(
|
||
cache_backend=object(),
|
||
redis_env_kwargs={"host": "coordination-redis", "port": "6379"},
|
||
)
|
||
|
||
assert isinstance(usage_cache, _EnvBuiltRedisCache)
|
||
assert usage_cache.init_kwargs["host"] == "coordination-redis"
|
||
assert spend_redis is usage_cache
|
||
assert config_redis is usage_cache
|
||
|
||
|
||
def test_init_cache_non_redis_backend_without_redis_env_stays_in_memory():
|
||
"""Without any REDIS_* connection info, a non-Redis response cache must
|
||
leave the coordination Redis unset instead of building a broken client."""
|
||
usage_cache, spend_redis, config_redis = _run_init_cache_with_backend(
|
||
cache_backend=object(),
|
||
redis_env_kwargs={},
|
||
)
|
||
|
||
assert usage_cache is None
|
||
assert spend_redis is None
|
||
assert config_redis is None
|
||
|
||
|
||
def test_init_cache_redis_backend_reuses_cache_backend_over_environment():
|
||
"""When the response cache itself is a plain Redis KV cache, it must be
|
||
reused as the coordination Redis; the REDIS_* environment fallback must
|
||
not construct a second client."""
|
||
redis_backend = _EnvBuiltRedisCache(host="cache-params-host")
|
||
usage_cache, spend_redis, _ = _run_init_cache_with_backend(
|
||
cache_backend=redis_backend,
|
||
redis_env_kwargs={"host": "env-host"},
|
||
)
|
||
|
||
assert usage_cache is redis_backend
|
||
assert usage_cache.init_kwargs["host"] == "cache-params-host"
|
||
assert spend_redis is redis_backend
|
||
|
||
|
||
class _EnvBuiltClusterCache(RedisClusterCache):
|
||
"""RedisClusterCache stand-in that records constructor kwargs and never
|
||
opens a network connection."""
|
||
|
||
def __init__(self, **kwargs):
|
||
self.init_kwargs = kwargs
|
||
|
||
|
||
def _run_init_coordination_redis(config, env=None):
|
||
"""Run ProxyConfig._init_coordination_redis against a stubbed module state,
|
||
returning (redis_usage_cache, spend_counter redis, config-cache redis)."""
|
||
fresh_spend_cache = DualCache()
|
||
fresh_config_cache = types.SimpleNamespace(redis_cache=None)
|
||
|
||
with (
|
||
_patched_coordination_redis_module_state(spend_cache=fresh_spend_cache, config_cache=fresh_config_cache),
|
||
mock.patch.dict(os.environ, env or {}, clear=False),
|
||
):
|
||
built = proxy_server_module.ProxyConfig()._init_coordination_redis(config=config)
|
||
return (
|
||
built,
|
||
fresh_spend_cache.redis_cache,
|
||
fresh_config_cache.redis_cache,
|
||
)
|
||
|
||
|
||
def test_init_coordination_redis_explicit_block_builds_standalone_client():
|
||
"""general_settings.coordination_redis must build the coordination Redis
|
||
even when no response cache is configured at all, and attach it to the
|
||
spend counter and config caches."""
|
||
usage_cache, spend_redis, config_redis = _run_init_coordination_redis(
|
||
config={"general_settings": {"coordination_redis": {"host": "coord-host", "port": 6380}}},
|
||
)
|
||
|
||
assert isinstance(usage_cache, _EnvBuiltRedisCache)
|
||
assert usage_cache.init_kwargs["host"] == "coord-host"
|
||
assert usage_cache.init_kwargs["port"] == 6380
|
||
assert spend_redis is usage_cache
|
||
assert config_redis is usage_cache
|
||
|
||
|
||
def test_init_coordination_redis_resolves_os_environ_references():
|
||
"""os.environ/ values inside the coordination_redis block must be resolved
|
||
the same way cache_params values are."""
|
||
usage_cache, _, _ = _run_init_coordination_redis(
|
||
config={"general_settings": {"coordination_redis": {"host": "os.environ/COORD_REDIS_HOST"}}},
|
||
env={"COORD_REDIS_HOST": "resolved-host"},
|
||
)
|
||
|
||
assert usage_cache.init_kwargs["host"] == "resolved-host"
|
||
|
||
|
||
def test_init_coordination_redis_startup_nodes_builds_cluster_client():
|
||
"""A coordination_redis block with startup_nodes must construct a cluster
|
||
client, so cluster-aware consumers (v3 rate limiter) take the cluster path."""
|
||
usage_cache, _, _ = _run_init_coordination_redis(
|
||
config={"general_settings": {"coordination_redis": {"startup_nodes": [{"host": "node-1", "port": 7000}]}}},
|
||
)
|
||
|
||
assert isinstance(usage_cache, _EnvBuiltClusterCache)
|
||
assert usage_cache.init_kwargs["startup_nodes"] == [{"host": "node-1", "port": 7000}]
|
||
|
||
|
||
def test_init_coordination_redis_without_connection_target_raises():
|
||
"""A coordination_redis block with no host, url, startup_nodes, or
|
||
sentinel_nodes is a config error and must fail startup loudly instead of
|
||
silently running without coordination."""
|
||
with pytest.raises(ValueError, match="connection target"):
|
||
_run_init_coordination_redis(
|
||
config={"general_settings": {"coordination_redis": {"ssl": True}}},
|
||
)
|
||
|
||
|
||
def test_init_coordination_redis_non_mapping_block_raises():
|
||
"""A scalar coordination_redis value is a config error."""
|
||
with pytest.raises(ValueError, match="mapping"):
|
||
_run_init_coordination_redis(
|
||
config={"general_settings": {"coordination_redis": "redis://host:6379"}},
|
||
)
|
||
|
||
|
||
def test_init_coordination_redis_absent_leaves_usage_cache_unset():
|
||
"""Without the block, nothing changes: the coordination Redis stays unset
|
||
for the downstream borrow / env fallback logic to decide."""
|
||
usage_cache, spend_redis, _ = _run_init_coordination_redis(
|
||
config={"general_settings": {}},
|
||
)
|
||
|
||
assert usage_cache is None
|
||
assert spend_redis is None
|
||
|
||
|
||
def test_explicit_coordination_redis_takes_precedence_over_cache_backend():
|
||
"""When both an explicit coordination_redis block and a plain-Redis
|
||
response cache are configured, the explicit block must win; the cache
|
||
backend must not overwrite it."""
|
||
fresh_spend_cache = DualCache()
|
||
fresh_config_cache = types.SimpleNamespace(redis_cache=None)
|
||
cache_backend = _EnvBuiltRedisCache(host="cache-backend-host")
|
||
mock_litellm_cache = MagicMock()
|
||
mock_litellm_cache.cache = cache_backend
|
||
|
||
with (
|
||
_patched_coordination_redis_module_state(spend_cache=fresh_spend_cache, config_cache=fresh_config_cache),
|
||
patch("litellm.Cache", return_value=mock_litellm_cache),
|
||
):
|
||
litellm.cache = None
|
||
proxy_config = proxy_server_module.ProxyConfig()
|
||
built = proxy_config._init_coordination_redis(
|
||
config={"general_settings": {"coordination_redis": {"host": "explicit-coord-host"}}}
|
||
)
|
||
assert built is not None
|
||
proxy_server_module.redis_usage_cache = built
|
||
usage_cache = proxy_config._init_cache(cache_params={"type": "redis"})
|
||
|
||
assert isinstance(usage_cache, _EnvBuiltRedisCache)
|
||
assert usage_cache is not cache_backend
|
||
assert usage_cache.init_kwargs["host"] == "explicit-coord-host"
|
||
assert fresh_spend_cache.redis_cache is usage_cache
|
||
|
||
|
||
async def _run_init_coordination_redis_env_fallback(
|
||
litellm_settings, redis_env_kwargs, redis_cache_class=_EnvBuiltRedisCache
|
||
):
|
||
"""Run ProxyConfig._init_coordination_redis_env_fallback against a
|
||
stubbed module state and a controlled REDIS_* environment, returning
|
||
(built, spend_counter redis)."""
|
||
fresh_spend_cache = DualCache()
|
||
fresh_config_cache = types.SimpleNamespace(redis_cache=None)
|
||
|
||
with (
|
||
_patched_coordination_redis_module_state(
|
||
spend_cache=fresh_spend_cache, config_cache=fresh_config_cache, redis_cache_class=redis_cache_class
|
||
),
|
||
patch(
|
||
"litellm._redis._redis_kwargs_from_environment",
|
||
return_value=redis_env_kwargs,
|
||
),
|
||
):
|
||
built = await proxy_server_module.ProxyConfig._init_coordination_redis_env_fallback(
|
||
litellm_settings=litellm_settings
|
||
)
|
||
return built, fresh_spend_cache.redis_cache
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_coordination_redis_env_fallback_builds_from_environment():
|
||
"""A deployment with no coordination_redis block and no litellm_settings.cache
|
||
but with bare REDIS_HOST/REDIS_PORT env vars must still get a coordination
|
||
Redis: otherwise spend counters, budget-window enforcement, and the
|
||
reset_spend cache-eviction broadcast stay per-pod local and a reset issued
|
||
on one pod never clears another pod's stale enforcement."""
|
||
built, spend_redis = await _run_init_coordination_redis_env_fallback(
|
||
litellm_settings={},
|
||
redis_env_kwargs={"host": "env-fallback-host", "port": "6390"},
|
||
)
|
||
|
||
assert isinstance(built, _EnvBuiltRedisCache)
|
||
assert built.init_kwargs["host"] == "env-fallback-host"
|
||
assert spend_redis is built
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_coordination_redis_env_fallback_without_redis_env_returns_none():
|
||
"""With no REDIS_* connection info at all, the fallback must leave the
|
||
coordination Redis unset rather than building a broken client."""
|
||
built, spend_redis = await _run_init_coordination_redis_env_fallback(
|
||
litellm_settings={},
|
||
redis_env_kwargs={},
|
||
)
|
||
|
||
assert built is None
|
||
assert spend_redis is None
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_coordination_redis_env_fallback_unreachable_stays_in_memory():
|
||
"""REDIS_* env vars can name a Redis that is not actually reachable (wrong
|
||
host, leftover from an unrelated job/service). Guessing "coordination
|
||
available" from bare env vars must not turn a previously harmless
|
||
in-memory-only proxy into one that raises on its next cache write, so an
|
||
unreachable ping must leave everything exactly as if no REDIS_* vars were
|
||
set at all."""
|
||
built, spend_redis = await _run_init_coordination_redis_env_fallback(
|
||
litellm_settings={},
|
||
redis_env_kwargs={"host": "unreachable-host", "port": "6390"},
|
||
redis_cache_class=_UnreachableRedisCache,
|
||
)
|
||
|
||
assert built is None
|
||
assert spend_redis is None
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_coordination_redis_env_fallback_malformed_cluster_nodes_stays_in_memory():
|
||
"""REDIS_CLUSTER_NODES can be set to a malformed value nothing here ever
|
||
asked to be parsed. Unlike the explicit coordination_redis block (a
|
||
deliberate opt-in, so a bad value there should fail loudly), this
|
||
inferred fallback must not abort proxy startup over it -- it has to
|
||
decline the same way it does for an absent or unreachable Redis."""
|
||
with mock.patch.dict(os.environ, {"REDIS_CLUSTER_NODES": "not-valid-json"}, clear=False):
|
||
built, spend_redis = await _run_init_coordination_redis_env_fallback(
|
||
litellm_settings={},
|
||
redis_env_kwargs={},
|
||
)
|
||
|
||
assert built is None
|
||
assert spend_redis is None
|
||
|
||
|
||
def test_env_fallback_builds_cluster_client_from_cluster_nodes_env():
|
||
"""A deployment whose only Redis env is REDIS_CLUSTER_NODES must still get
|
||
a coordination Redis from the env fallback, and it must be a cluster
|
||
client so cluster-aware consumers take the cluster path."""
|
||
nodes = '[{"host": "cnode-1", "port": 7000}]'
|
||
with (
|
||
patch.object(proxy_server_module, "RedisCache", _EnvBuiltRedisCache),
|
||
patch.object(proxy_server_module, "RedisClusterCache", _EnvBuiltClusterCache),
|
||
patch("litellm._redis._redis_kwargs_from_environment", return_value={}),
|
||
mock.patch.dict(os.environ, {"REDIS_CLUSTER_NODES": nodes}, clear=False),
|
||
):
|
||
result = proxy_server_module._build_redis_usage_cache_from_environment()
|
||
|
||
assert isinstance(result, _EnvBuiltClusterCache)
|
||
assert result.init_kwargs["startup_nodes"] == [{"host": "cnode-1", "port": 7000}]
|
||
|
||
|
||
def test_env_fallback_builds_client_from_sentinel_nodes_env():
|
||
"""A sentinel-only environment (REDIS_SENTINEL_NODES, no host or url) must
|
||
also produce a coordination Redis from the env fallback."""
|
||
with (
|
||
patch.object(proxy_server_module, "RedisCache", _EnvBuiltRedisCache),
|
||
patch.object(proxy_server_module, "RedisClusterCache", _EnvBuiltClusterCache),
|
||
patch("litellm._redis._redis_kwargs_from_environment", return_value={}),
|
||
mock.patch.dict(os.environ, {"REDIS_SENTINEL_NODES": '[["s1", 26379]]'}, clear=False),
|
||
):
|
||
result = proxy_server_module._build_redis_usage_cache_from_environment()
|
||
|
||
assert isinstance(result, _EnvBuiltRedisCache)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_startup_applies_coordination_redis_saved_in_database():
|
||
"""A coordination_redis block saved from the admin UI lives only in the
|
||
database, so startup must read it and build the coordination Redis from it.
|
||
Without this the save endpoint's "restart to apply" promise is false and the
|
||
proxy silently coordinates in per-pod memory."""
|
||
fresh_spend_cache = DualCache()
|
||
fresh_config_cache = types.SimpleNamespace(redis_cache=None)
|
||
|
||
with (
|
||
_patched_coordination_redis_module_state(spend_cache=fresh_spend_cache, config_cache=fresh_config_cache),
|
||
patch.object(
|
||
proxy_server_module,
|
||
"get_persisted_coordination_redis_settings",
|
||
AsyncMock(return_value={"host": "db-host", "port": 6381}),
|
||
),
|
||
):
|
||
result = await proxy_server_module.ProxyStartupEvent._init_coordination_redis_from_db(
|
||
litellm_settings={},
|
||
llm_router=None,
|
||
)
|
||
|
||
assert isinstance(result, _EnvBuiltRedisCache)
|
||
assert result.init_kwargs["host"] == "db-host"
|
||
assert fresh_spend_cache.redis_cache is result
|
||
assert fresh_config_cache.redis_cache is result
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_startup_ignores_database_coordination_redis_without_connection_target():
|
||
"""A persisted block with no host/url/cluster/sentinel must be ignored rather
|
||
than crashing startup or building a client that cannot connect."""
|
||
with (
|
||
patch.object(proxy_server_module, "spend_counter_cache", DualCache()),
|
||
patch.object(proxy_server_module, "litellm_config_cache", types.SimpleNamespace(redis_cache=None)),
|
||
patch.object(proxy_server_module, "RedisCache", _EnvBuiltRedisCache),
|
||
patch.object(
|
||
proxy_server_module,
|
||
"get_persisted_coordination_redis_settings",
|
||
AsyncMock(return_value={"ssl": True}),
|
||
),
|
||
):
|
||
result = await proxy_server_module.ProxyStartupEvent._init_coordination_redis_from_db(
|
||
litellm_settings={},
|
||
llm_router=None,
|
||
)
|
||
|
||
assert result is None
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_startup_survives_database_read_failure_for_coordination_redis():
|
||
"""A config-row read failure must not block proxy startup."""
|
||
with (
|
||
patch.object(
|
||
proxy_server_module,
|
||
"get_persisted_coordination_redis_settings",
|
||
AsyncMock(side_effect=RuntimeError("db unreachable")),
|
||
),
|
||
):
|
||
result = await proxy_server_module.ProxyStartupEvent._init_coordination_redis_from_db(
|
||
litellm_settings={},
|
||
llm_router=None,
|
||
)
|
||
|
||
assert result is None
|
||
|
||
|
||
def _stream_usage_test_chunks():
|
||
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Usage
|
||
|
||
content_chunk = ModelResponseStream(
|
||
model="gpt-5.4-nano",
|
||
choices=[StreamingChoices(delta=Delta(content="pong"))],
|
||
)
|
||
finish_chunk = ModelResponseStream(
|
||
model="gpt-5.4-nano",
|
||
choices=[StreamingChoices(finish_reason="stop")],
|
||
)
|
||
usage_chunk = ModelResponseStream(model="gpt-5.4-nano", choices=[])
|
||
usage_chunk.usage = Usage(prompt_tokens=50, completion_tokens=188, total_tokens=238)
|
||
return content_chunk, finish_chunk, usage_chunk
|
||
|
||
|
||
def _stream_usage_generator_chunks():
|
||
from litellm.types.utils import ModelResponseStream
|
||
|
||
content_chunk, finish_chunk, usage_chunk = _stream_usage_test_chunks()
|
||
prompt_filter_chunk = ModelResponseStream(model="gpt-5.4-nano", choices=[])
|
||
return prompt_filter_chunk, content_chunk, finish_chunk, usage_chunk
|
||
|
||
|
||
def test_is_injected_stream_usage_artifact():
|
||
from litellm.proxy.proxy_server import _is_injected_stream_usage_artifact
|
||
from litellm.types.utils import ModelResponseStream, Usage
|
||
|
||
content_chunk, finish_chunk, empty_choices_usage_chunk = _stream_usage_test_chunks()
|
||
assert _is_injected_stream_usage_artifact(empty_choices_usage_chunk) is True
|
||
|
||
synthetic_final_chunk = ModelResponseStream(model="gpt-5.4-nano")
|
||
synthetic_final_chunk.usage = Usage(prompt_tokens=50, completion_tokens=188, total_tokens=238)
|
||
assert _is_injected_stream_usage_artifact(synthetic_final_chunk) is True
|
||
|
||
azure_prompt_filter_chunk = ModelResponseStream(model="gpt-5.4-nano", choices=[])
|
||
assert _is_injected_stream_usage_artifact(azure_prompt_filter_chunk) is True
|
||
|
||
assert _is_injected_stream_usage_artifact(content_chunk) is False
|
||
assert _is_injected_stream_usage_artifact(finish_chunk) is False
|
||
|
||
content_chunk_with_usage, finish_chunk_with_usage, _ = _stream_usage_test_chunks()
|
||
content_chunk_with_usage.usage = Usage(prompt_tokens=50, completion_tokens=188, total_tokens=238)
|
||
finish_chunk_with_usage.usage = Usage(prompt_tokens=50, completion_tokens=188, total_tokens=238)
|
||
assert _is_injected_stream_usage_artifact(content_chunk_with_usage) is False
|
||
assert _is_injected_stream_usage_artifact(finish_chunk_with_usage) is False
|
||
|
||
assert _is_injected_stream_usage_artifact({"usage": {"prompt_tokens": 1}}) is False
|
||
|
||
|
||
async def _collect_async_data_generator_frames(request_data: dict) -> list:
|
||
from litellm.proxy.proxy_server import async_data_generator
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
chunks = _stream_usage_generator_chunks()
|
||
|
||
class MockStream:
|
||
def __aiter__(self):
|
||
return self._stream()
|
||
|
||
async def _stream(self):
|
||
for chunk in chunks:
|
||
yield chunk
|
||
|
||
async def aclose(self):
|
||
pass
|
||
|
||
mock_proxy_logging_obj = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging_obj.needs_iterator_wrap.return_value = False
|
||
mock_proxy_logging_obj.needs_per_chunk_streaming_hook.return_value = False
|
||
mock_proxy_logging_obj.post_call_failure_hook = AsyncMock()
|
||
|
||
with patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj):
|
||
with patch.object(proxy_server_module.ProxyLogging, "_fire_deferred_stream_logging"):
|
||
return [
|
||
frame.decode("utf-8") if isinstance(frame, bytes) else frame
|
||
async for frame in async_data_generator(MockStream(), MagicMock(spec=UserAPIKeyAuth), request_data)
|
||
]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_strips_injected_usage_chunk():
|
||
frames = await _collect_async_data_generator_frames({"model": "gpt-5.4-nano", "_litellm_strip_stream_usage": True})
|
||
|
||
data_frames = [frame for frame in frames if frame.startswith("data: {")]
|
||
assert len(data_frames) == 2
|
||
assert any("pong" in frame for frame in data_frames)
|
||
assert any("finish_reason" in frame for frame in data_frames)
|
||
assert not any('"usage"' in frame for frame in data_frames)
|
||
assert frames[-1] == "data: [DONE]\n\n"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_async_data_generator_forwards_usage_chunk_without_strip_marker():
|
||
frames = await _collect_async_data_generator_frames({"model": "gpt-5.4-nano"})
|
||
|
||
data_frames = [frame for frame in frames if frame.startswith("data: {")]
|
||
assert len(data_frames) == 4
|
||
assert any('"usage"' in frame and '"completion_tokens":188' in frame.replace(" ", "") for frame in data_frames)
|
||
assert frames[-1] == "data: [DONE]\n\n"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_config_field_update_rejects_mock_testing_flag():
|
||
"""The mock-testing opt-in is deliberately absent from
|
||
``ConfigGeneralSettings`` so that ``/config/field/update`` refuses it. If
|
||
someone later adds the field for tidiness, this test fails and tells them
|
||
they have just opened an API write path into a config-file-only setting."""
|
||
from fastapi import HTTPException
|
||
|
||
from litellm.proxy._types import ConfigFieldUpdate
|
||
from litellm.proxy.proxy_server import update_config_general_settings
|
||
from litellm.proxy.route_llm_request import MOCK_TESTING_CONFIG_KEY
|
||
|
||
admin = UserAPIKeyAuth(
|
||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||
api_key="sk-test",
|
||
)
|
||
|
||
with patch.object(proxy_server_module, "prisma_client", MagicMock()):
|
||
with pytest.raises(HTTPException) as exc_info:
|
||
await update_config_general_settings(
|
||
data=ConfigFieldUpdate(
|
||
field_name=MOCK_TESTING_CONFIG_KEY,
|
||
field_value=True,
|
||
config_type="general_settings",
|
||
),
|
||
user_api_key_dict=admin,
|
||
)
|
||
|
||
assert exc_info.value.status_code == 400
|
||
|
||
|
||
def test_config_update_body_drops_mock_testing_flag():
|
||
"""``/config/update`` parses its body as ``ConfigYAML``, whose
|
||
``general_settings`` is a ``ConfigGeneralSettings``. Undeclared keys are
|
||
dropped on parse, so the flag never reaches the DB by that route either."""
|
||
from litellm.proxy._types import ConfigYAML
|
||
from litellm.proxy.route_llm_request import MOCK_TESTING_CONFIG_KEY
|
||
|
||
parsed = ConfigYAML.model_validate({"general_settings": {MOCK_TESTING_CONFIG_KEY: True}})
|
||
|
||
assert parsed.general_settings is not None
|
||
assert MOCK_TESTING_CONFIG_KEY not in parsed.general_settings.model_dump(exclude_none=True)
|
||
|
||
|
||
def test_startup_warns_when_mock_testing_params_enabled(caplog):
|
||
"""Enabling the opt-in must announce itself, naming every param it
|
||
unlocks — the config key says ``mock_testing`` but the gate also covers
|
||
``mock_timeout`` and ``mock_delay``, so coverage cannot be inferred from
|
||
the name alone."""
|
||
import logging
|
||
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.route_llm_request import (
|
||
GATED_MOCK_PARAM_NAMES,
|
||
MOCK_TESTING_CONFIG_KEY,
|
||
)
|
||
|
||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||
ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings={MOCK_TESTING_CONFIG_KEY: True})
|
||
|
||
assert MOCK_TESTING_CONFIG_KEY in caplog.text
|
||
for param_name in GATED_MOCK_PARAM_NAMES:
|
||
assert param_name in caplog.text
|
||
|
||
|
||
def test_startup_is_silent_when_mock_testing_params_disabled(caplog):
|
||
"""A proxy that never set the opt-in must not emit the warning."""
|
||
import logging
|
||
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.route_llm_request import MOCK_TESTING_CONFIG_KEY
|
||
|
||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||
ProxyStartupEvent._warn_if_mock_testing_params_enabled(general_settings={})
|
||
|
||
assert MOCK_TESTING_CONFIG_KEY not in caplog.text
|
||
|
||
|
||
# ---------------------------------------------------------------------------
|
||
# Budget window spend row enqueue (LiteLLM_BudgetWindowSpend writer)
|
||
# ---------------------------------------------------------------------------
|
||
|
||
|
||
@contextlib.contextmanager
|
||
def _window_spend_enqueue_env(cached_objects: dict):
|
||
"""Point increment_spend_counters at throwaway caches and a real
|
||
WindowSpendUpdateQueue, and hand back the queue to inspect."""
|
||
from litellm.caching.dual_cache import DualCache
|
||
from litellm.proxy.db.db_transaction_queue.window_spend_update_queue import (
|
||
WindowSpendUpdateQueue,
|
||
)
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
user_api_key_cache = MagicMock()
|
||
user_api_key_cache.async_get_cache = AsyncMock(side_effect=lambda key, **_: cached_objects.get(key))
|
||
|
||
queue = WindowSpendUpdateQueue()
|
||
proxy_logging_obj = MagicMock()
|
||
proxy_logging_obj.db_spend_update_writer.window_spend_update_queue = queue
|
||
|
||
originals = (
|
||
ps.user_api_key_cache,
|
||
ps.spend_counter_cache,
|
||
ps.prisma_client,
|
||
ps.proxy_logging_obj,
|
||
)
|
||
ps.user_api_key_cache = user_api_key_cache
|
||
ps.spend_counter_cache = DualCache()
|
||
ps.prisma_client = None
|
||
ps.proxy_logging_obj = proxy_logging_obj
|
||
try:
|
||
yield queue
|
||
finally:
|
||
(
|
||
ps.user_api_key_cache,
|
||
ps.spend_counter_cache,
|
||
ps.prisma_client,
|
||
ps.proxy_logging_obj,
|
||
) = originals
|
||
|
||
|
||
async def _drain(queue):
|
||
return list(await queue.flush_and_get_aggregated_window_spend_transactions())
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_key_window_spend_row_is_enqueued_with_the_actual_cost():
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||
key_obj = MagicMock()
|
||
key_obj.budget_limits = [
|
||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||
]
|
||
|
||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||
await increment_spend_counters(
|
||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||
)
|
||
enqueued = await _drain(queue)
|
||
|
||
assert len(enqueued) == 1
|
||
assert enqueued[0]["entity_type"] == "key"
|
||
assert enqueued[0]["entity_id"] == "hashed-token"
|
||
assert enqueued[0]["window_duration"] == "30d"
|
||
assert enqueued[0]["spend"] == pytest.approx(0.25)
|
||
assert enqueued[0]["window_start"] == (reset_at - timedelta(days=30)).astimezone(timezone.utc).replace(
|
||
tzinfo=None
|
||
).isoformat(timespec="microseconds")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_team_window_spend_row_is_enqueued():
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
reset_at = datetime.now(timezone.utc) + timedelta(days=3)
|
||
team_obj = MagicMock()
|
||
team_obj.budget_limits = [
|
||
{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}
|
||
]
|
||
|
||
with _window_spend_enqueue_env({"team_id:team-1": team_obj}) as queue:
|
||
await increment_spend_counters(
|
||
token=None, team_id="team-1", user_id=None, response_cost=1.5
|
||
)
|
||
enqueued = await _drain(queue)
|
||
|
||
assert len(enqueued) == 1
|
||
assert enqueued[0]["entity_type"] == "team"
|
||
assert enqueued[0]["entity_id"] == "team-1"
|
||
assert enqueued[0]["window_duration"] == "7d"
|
||
assert enqueued[0]["spend"] == pytest.approx(1.5)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_window_spend_row_is_enqueued_even_when_the_counter_was_reserved():
|
||
"""A reservation only pre-charged the cache counter with an estimate; the
|
||
row still owes the actual cost, so the enqueue must not be skipped."""
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
import litellm.proxy.spend_tracking.budget_reservation as br
|
||
|
||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||
key_obj = MagicMock()
|
||
key_obj.budget_limits = [
|
||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||
]
|
||
reservation = {
|
||
"entries": [
|
||
{"counter_key": "spend:key:hashed-token", "reserved": 1.0},
|
||
{"counter_key": "spend:key:hashed-token:window:30d", "reserved": 1.0},
|
||
]
|
||
}
|
||
|
||
original_reconcile = br.reconcile_budget_reservation
|
||
br.reconcile_budget_reservation = AsyncMock(return_value=None)
|
||
try:
|
||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||
await increment_spend_counters(
|
||
token="hashed-token",
|
||
team_id=None,
|
||
user_id=None,
|
||
response_cost=0.25,
|
||
budget_reservation=reservation,
|
||
)
|
||
enqueued = await _drain(queue)
|
||
finally:
|
||
br.reconcile_budget_reservation = original_reconcile
|
||
|
||
assert len(enqueued) == 1
|
||
assert enqueued[0]["spend"] == pytest.approx(0.25)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_sliding_window_without_reset_at_is_not_enqueued():
|
||
"""Windows with no reset_at slide with wall clock, so window_start moves on
|
||
every request and no single row can represent them; the read path keeps
|
||
using its LiteLLM_SpendLogs fallback instead."""
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
key_obj = MagicMock()
|
||
key_obj.budget_limits = [{"budget_duration": "30d", "max_budget": 100.0}]
|
||
|
||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||
await increment_spend_counters(
|
||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||
)
|
||
enqueued = await _drain(queue)
|
||
|
||
assert enqueued == []
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_each_configured_window_gets_its_own_row_enqueue():
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
now = datetime.now(timezone.utc)
|
||
key_obj = MagicMock()
|
||
key_obj.budget_limits = [
|
||
{"budget_duration": "1d", "max_budget": 5.0, "reset_at": (now + timedelta(hours=5)).isoformat()},
|
||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": (now + timedelta(days=10)).isoformat()},
|
||
]
|
||
|
||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||
await increment_spend_counters(
|
||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||
)
|
||
enqueued = await _drain(queue)
|
||
|
||
assert sorted(item["window_duration"] for item in enqueued) == ["1d", "30d"]
|
||
assert all(item["spend"] == pytest.approx(0.25) for item in enqueued)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_no_window_spend_row_enqueued_without_budget_limits():
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
key_obj = MagicMock()
|
||
key_obj.budget_limits = None
|
||
|
||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||
await increment_spend_counters(
|
||
token="hashed-token", team_id=None, user_id=None, response_cost=0.25
|
||
)
|
||
enqueued = await _drain(queue)
|
||
|
||
assert enqueued == []
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_window_spend_row_carries_the_request_start_time():
|
||
"""The seed sums LiteLLM_SpendLogs only up to this point, so it must be the
|
||
same start the spend log row was written with."""
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
reset_at = datetime.now(timezone.utc) + timedelta(days=10)
|
||
key_obj = MagicMock()
|
||
key_obj.budget_limits = [
|
||
{"budget_duration": "30d", "max_budget": 100.0, "reset_at": reset_at.isoformat()}
|
||
]
|
||
|
||
with _window_spend_enqueue_env({"hashed-token": key_obj}) as queue:
|
||
await increment_spend_counters(
|
||
token="hashed-token",
|
||
team_id=None,
|
||
user_id=None,
|
||
response_cost=0.25,
|
||
request_started_at=datetime(2026, 8, 10, 12, 0, 0, 500_000, tzinfo=timezone.utc),
|
||
)
|
||
enqueued = await _drain(queue)
|
||
|
||
assert enqueued[0]["started_at"] == "2026-08-10T12:00:00.500000"
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_team_window_spend_row_carries_the_request_start_time():
|
||
from litellm.proxy.proxy_server import increment_spend_counters
|
||
|
||
reset_at = datetime.now(timezone.utc) + timedelta(days=3)
|
||
team_obj = MagicMock()
|
||
team_obj.budget_limits = [
|
||
{"budget_duration": "7d", "max_budget": 50.0, "reset_at": reset_at.isoformat()}
|
||
]
|
||
|
||
with _window_spend_enqueue_env({"team_id:team-1": team_obj}) as queue:
|
||
await increment_spend_counters(
|
||
token=None,
|
||
team_id="team-1",
|
||
user_id=None,
|
||
response_cost=1.5,
|
||
request_started_at=datetime(2026, 8, 10, 12, 0, 0, 500_000, tzinfo=timezone.utc),
|
||
)
|
||
enqueued = await _drain(queue)
|
||
|
||
assert enqueued[0]["started_at"] == "2026-08-10T12:00:00.500000"
|
||
|
||
|
||
def _mock_startup_prisma_client(health_check_error=None, connect_error=None):
|
||
client = MagicMock()
|
||
client.connect = AsyncMock(side_effect=connect_error)
|
||
client.db.start_token_refresh_task = AsyncMock()
|
||
client.check_view_exists = AsyncMock()
|
||
client._set_spend_logs_row_count_in_proxy_state = AsyncMock()
|
||
client.start_db_health_watchdog_task = AsyncMock()
|
||
client.health_check = AsyncMock(side_effect=health_check_error)
|
||
return client
|
||
|
||
|
||
async def _run_setup_prisma_client(mock_client):
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
|
||
with patch.object(proxy_server_module, "PrismaClient", return_value=mock_client):
|
||
result = await ProxyStartupEvent._setup_prisma_client(
|
||
database_url="postgresql://litellm:litellm@localhost:5432/litellm",
|
||
proxy_logging_obj=MagicMock(),
|
||
user_api_key_cache=DualCache(),
|
||
)
|
||
await asyncio.sleep(0.05)
|
||
return result
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_setup_prisma_client_retains_connected_client_when_startup_health_check_fails(
|
||
monkeypatch,
|
||
):
|
||
"""A transient failure of the startup ``SELECT 1`` must not discard a client
|
||
whose ``connect()`` already succeeded.
|
||
|
||
Discarding it assigns ``None`` to the module-level ``prisma_client`` for the
|
||
life of the process, so a database that came back a second later is never
|
||
used again until the proxy is restarted."""
|
||
monkeypatch.setenv("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", "False")
|
||
monkeypatch.setattr(
|
||
proxy_server_module,
|
||
"general_settings",
|
||
{"allow_requests_on_db_unavailable": True},
|
||
)
|
||
|
||
mock_client = _mock_startup_prisma_client(health_check_error=httpx.ReadTimeout("startup health check timed out"))
|
||
result = await _run_setup_prisma_client(mock_client)
|
||
|
||
assert mock_client.connect.await_count == 1
|
||
assert mock_client.health_check.await_count == 1
|
||
assert result is mock_client
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_setup_prisma_client_arms_health_watchdog_before_startup_health_check(
|
||
monkeypatch,
|
||
):
|
||
"""The health watchdog is the only thing that reconnects a dropped DB, so it
|
||
has to be armed before the startup health check can fail.
|
||
|
||
Armed after, the single failure it exists to recover from is exactly the one
|
||
that skips it, and recovery never happens."""
|
||
monkeypatch.setenv("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", "False")
|
||
monkeypatch.setattr(
|
||
proxy_server_module,
|
||
"general_settings",
|
||
{"allow_requests_on_db_unavailable": True},
|
||
)
|
||
|
||
mock_client = _mock_startup_prisma_client(health_check_error=httpx.ReadTimeout("startup health check timed out"))
|
||
call_order = MagicMock()
|
||
call_order.attach_mock(mock_client.start_db_health_watchdog_task, "watchdog")
|
||
call_order.attach_mock(mock_client.health_check, "health_check")
|
||
|
||
await _run_setup_prisma_client(mock_client)
|
||
|
||
assert mock_client.start_db_health_watchdog_task.await_count == 1
|
||
assert [call[0] for call in call_order.mock_calls] == ["watchdog", "health_check"]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_setup_prisma_client_raises_when_db_unavailable_is_not_allowed(monkeypatch):
|
||
"""Without ``allow_requests_on_db_unavailable`` a failed startup health check
|
||
must still hard-fail startup. Retaining the client is a fallback for
|
||
operators who opted into serving traffic without a database, never a way to
|
||
boot a proxy whose DB never answered."""
|
||
monkeypatch.setenv("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", "False")
|
||
monkeypatch.setattr(
|
||
proxy_server_module,
|
||
"general_settings",
|
||
{"allow_requests_on_db_unavailable": False},
|
||
)
|
||
|
||
mock_client = _mock_startup_prisma_client(health_check_error=httpx.ReadTimeout("startup health check timed out"))
|
||
with pytest.raises(httpx.ReadTimeout):
|
||
await _run_setup_prisma_client(mock_client)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_setup_prisma_client_returns_none_when_connect_itself_fails(monkeypatch):
|
||
"""Retaining only ever applies to a client that connected. If ``connect()``
|
||
failed there is no usable client and no watchdog to recover it, so the caller
|
||
must still get ``None``."""
|
||
monkeypatch.setenv("DISABLE_PRISMA_HEALTH_CHECK_ON_STARTUP", "False")
|
||
monkeypatch.setattr(
|
||
proxy_server_module,
|
||
"general_settings",
|
||
{"allow_requests_on_db_unavailable": True},
|
||
)
|
||
|
||
mock_client = _mock_startup_prisma_client(connect_error=httpx.ConnectError("connection refused"))
|
||
result = await _run_setup_prisma_client(mock_client)
|
||
|
||
assert result is None
|
||
assert mock_client.start_db_health_watchdog_task.await_count == 0
|
||
assert mock_client.health_check.await_count == 0
|
||
|
||
|
||
async def _run_scheduled_background_jobs():
|
||
from litellm.proxy.proxy_server import ProxyStartupEvent
|
||
from litellm.proxy.utils import ProxyLogging
|
||
|
||
mock_prisma_client = MagicMock()
|
||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(return_value=None)
|
||
|
||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||
mock_proxy_logging.slack_alerting_instance = MagicMock()
|
||
mock_proxy_logging.db_spend_update_writer = MagicMock()
|
||
mock_proxy_config = AsyncMock()
|
||
|
||
with (
|
||
patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config),
|
||
patch("litellm.proxy.proxy_server.store_model_in_db", True),
|
||
patch("litellm.proxy.proxy_server.get_secret_bool", return_value=True),
|
||
):
|
||
await ProxyStartupEvent.initialize_scheduled_background_jobs(
|
||
general_settings={},
|
||
prisma_client=mock_prisma_client,
|
||
proxy_budget_rescheduler_min_time=1,
|
||
proxy_budget_rescheduler_max_time=2,
|
||
proxy_batch_write_at=5,
|
||
proxy_logging_obj=mock_proxy_logging,
|
||
)
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
assert ps.scheduler is not None
|
||
return ps.scheduler
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_ptu_rollup_job_registered_at_startup(monkeypatch):
|
||
"""The PTU rollup cron is registered once an operator opts in; only models with PTU config accrue flat cost (asserted in test_ptu_flat_cost_rollup.py)."""
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
|
||
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import (
|
||
PTU_ROLLUP_JOB_ID,
|
||
)
|
||
|
||
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
|
||
|
||
scheduler = await _run_scheduled_background_jobs()
|
||
|
||
assert scheduler.get_job(PTU_ROLLUP_JOB_ID) is not None
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_ptu_rollup_job_hands_the_rollup_the_proxys_router(monkeypatch):
|
||
"""The rollup prices PTU deployments declared in config.yaml, which only the router
|
||
knows about. It takes the router as an argument, so nothing but this call site puts the
|
||
proxy's own router in front of it: without it that half of the feature is dead."""
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.spend_tracking import ptu_flat_cost_rollup
|
||
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
|
||
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import PTU_ROLLUP_JOB_ID
|
||
|
||
monkeypatch.setenv(PTU_COST_ATTRIBUTION_ENV_VAR, "true")
|
||
calls = []
|
||
monkeypatch.setattr(
|
||
ptu_flat_cost_rollup,
|
||
"run_scheduled_ptu_rollup",
|
||
AsyncMock(side_effect=lambda *args, **kwargs: calls.append(kwargs)),
|
||
)
|
||
|
||
scheduler = await _run_scheduled_background_jobs()
|
||
|
||
import litellm.proxy.proxy_server as ps
|
||
|
||
router = MagicMock()
|
||
monkeypatch.setattr(ps, "llm_router", router)
|
||
await scheduler.get_job(PTU_ROLLUP_JOB_ID).func()
|
||
|
||
assert [call["router"] for call in calls] == [router]
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_ptu_rollup_job_not_registered_without_opt_in(monkeypatch):
|
||
"""Without LITELLM_ENABLE_PTU_COST_ATTRIBUTION the rollup never runs, so no sentinel row
|
||
is ever written. This is the gate that keeps the whole feature inert by default."""
|
||
monkeypatch.delenv("STORE_MODEL_IN_DB", raising=False)
|
||
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
|
||
from litellm.proxy.spend_tracking.ptu_flat_cost_rollup import (
|
||
PTU_ROLLUP_JOB_ID,
|
||
)
|
||
|
||
monkeypatch.delenv(PTU_COST_ATTRIBUTION_ENV_VAR, raising=False)
|
||
|
||
scheduler = await _run_scheduled_background_jobs()
|
||
|
||
assert scheduler.get_job(PTU_ROLLUP_JOB_ID) is None
|
||
assert len(scheduler.get_jobs()) > 0
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_moderations_reraises_proxy_exception_unwrapped():
|
||
"""A 400 ProxyException from request validation must surface as-is,
|
||
not be re-wrapped into a code-500 ProxyException."""
|
||
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
||
|
||
exc = ProxyException(
|
||
message="Invalid type for 'metadata': expected an object, but got a string instead.",
|
||
type=ProxyErrorTypes.bad_request_error,
|
||
param="metadata",
|
||
code=400,
|
||
)
|
||
|
||
request = MagicMock()
|
||
request.body = AsyncMock(return_value=b'{"input": "hi", "metadata": "abc"}')
|
||
|
||
with (
|
||
patch.object(proxy_server_module, "add_litellm_data_to_request", new=AsyncMock(side_effect=exc)),
|
||
patch.object(proxy_server_module, "proxy_logging_obj") as mock_logging,
|
||
):
|
||
mock_logging.post_call_failure_hook = AsyncMock()
|
||
with pytest.raises(ProxyException) as exc_info:
|
||
await proxy_server_module.moderations(
|
||
request=request,
|
||
fastapi_response=MagicMock(),
|
||
user_api_key_dict=MagicMock(),
|
||
)
|
||
|
||
assert exc_info.value is exc
|
||
assert exc_info.value.code == "400"
|
||
assert exc_info.value.param == "metadata"
|
||
mock_logging.post_call_failure_hook.assert_awaited_once()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_agents_in_db_rebuilds_registry_under_agent_reconcile_lock(monkeypatch):
|
||
from litellm.proxy.agent_endpoints.agent_registry import (
|
||
AGENT_RECONCILE_LOCK,
|
||
global_agent_registry,
|
||
)
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
lock_states: list[bool] = []
|
||
|
||
async def fake_get_all_agents_from_db(prisma_client) -> list:
|
||
lock_states.append(AGENT_RECONCILE_LOCK.locked())
|
||
return []
|
||
|
||
def fake_load_agents_from_db_and_config(db_agents) -> None:
|
||
lock_states.append(AGENT_RECONCILE_LOCK.locked())
|
||
|
||
monkeypatch.setattr(global_agent_registry, "get_all_agents_from_db", fake_get_all_agents_from_db)
|
||
monkeypatch.setattr(global_agent_registry, "load_agents_from_db_and_config", fake_load_agents_from_db_and_config)
|
||
|
||
await ProxyConfig()._init_agents_in_db(prisma_client=MagicMock())
|
||
|
||
assert lock_states == [True, True]
|
||
assert not AGENT_RECONCILE_LOCK.locked()
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_guardrails_in_db_snapshots_and_reconciles_under_guardrail_reconcile_lock(monkeypatch):
|
||
from litellm.proxy.guardrails.guardrail_registry import (
|
||
GUARDRAIL_RECONCILE_LOCK,
|
||
IN_MEMORY_GUARDRAIL_HANDLER,
|
||
GuardrailRegistry,
|
||
)
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
lock_states: list[bool] = []
|
||
|
||
async def fake_get_all_guardrails_from_db(prisma_client) -> list:
|
||
lock_states.append(GUARDRAIL_RECONCILE_LOCK.locked())
|
||
return []
|
||
|
||
def fake_reconcile_db_guardrails(db_guardrail_ids) -> list:
|
||
lock_states.append(GUARDRAIL_RECONCILE_LOCK.locked())
|
||
return []
|
||
|
||
monkeypatch.setattr(GuardrailRegistry, "get_all_guardrails_from_db", fake_get_all_guardrails_from_db)
|
||
monkeypatch.setattr(IN_MEMORY_GUARDRAIL_HANDLER, "reconcile_db_guardrails", fake_reconcile_db_guardrails)
|
||
|
||
await ProxyConfig()._init_guardrails_in_db(prisma_client=MagicMock())
|
||
|
||
assert lock_states == [True, True]
|
||
assert not GUARDRAIL_RECONCILE_LOCK.locked()
|
||
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_prompts_in_db_reloads_rows_patched_on_another_worker(monkeypatch):
|
||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setattr(litellm, "callbacks", [])
|
||
|
||
def db_row(content: str) -> MagicMock:
|
||
row = MagicMock()
|
||
row.model_dump.return_value = {
|
||
"prompt_id": "greeting_sync",
|
||
"version": 1,
|
||
"environment": "development",
|
||
"created_by": None,
|
||
"litellm_params": json.dumps(
|
||
{
|
||
"prompt_id": "greeting_sync",
|
||
"prompt_integration": "dotprompt",
|
||
"prompt_data": {"content": content, "metadata": {}},
|
||
}
|
||
),
|
||
"prompt_info": json.dumps({"prompt_type": "db"}),
|
||
"created_at": None,
|
||
"updated_at": None,
|
||
}
|
||
return row
|
||
|
||
def served_content() -> str:
|
||
callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_sync.v1")
|
||
assert callback is not None
|
||
return callback.prompt_manager.get_prompt("greeting_sync").content
|
||
|
||
prisma_client = MagicMock()
|
||
try:
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[db_row("Begin every reply with AHOY")])
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
assert served_content() == "Begin every reply with AHOY"
|
||
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[db_row("Begin every reply with HOWDY")])
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
|
||
assert served_content() == "Begin every reply with HOWDY"
|
||
assert litellm.callbacks == [IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_sync.v1")]
|
||
finally:
|
||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_sync")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_prompts_in_db_syncs_remaining_rows_when_one_row_fails(monkeypatch):
|
||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setattr(litellm, "callbacks", [])
|
||
|
||
def db_row(prompt_id: str, integration: str) -> MagicMock:
|
||
row = MagicMock()
|
||
row.model_dump.return_value = {
|
||
"prompt_id": prompt_id,
|
||
"version": 1,
|
||
"environment": "development",
|
||
"created_by": None,
|
||
"litellm_params": json.dumps(
|
||
{
|
||
"prompt_id": prompt_id,
|
||
"prompt_integration": integration,
|
||
"prompt_data": {"content": "Begin every reply with AHOY", "metadata": {}},
|
||
}
|
||
),
|
||
"prompt_info": json.dumps({"prompt_type": "db"}),
|
||
"created_at": None,
|
||
"updated_at": None,
|
||
}
|
||
return row
|
||
|
||
prisma_client = MagicMock()
|
||
try:
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||
return_value=[db_row("broken_sync", "does_not_exist"), db_row("healthy_sync", "dotprompt")]
|
||
)
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("broken_sync.v1") is None
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("healthy_sync.v1") is not None
|
||
assert litellm.callbacks == [IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("healthy_sync.v1")]
|
||
finally:
|
||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("healthy_sync")
|
||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("broken_sync")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_prompts_in_db_serves_the_newest_row_when_environments_collide_on_a_versioned_id(monkeypatch):
|
||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setattr(litellm, "callbacks", [])
|
||
|
||
def db_row(environment: str, content: str, updated_at: datetime) -> MagicMock:
|
||
row = MagicMock()
|
||
row.model_dump.return_value = {
|
||
"prompt_id": "greeting_env",
|
||
"version": 1,
|
||
"environment": environment,
|
||
"created_by": None,
|
||
"litellm_params": json.dumps(
|
||
{
|
||
"prompt_id": "greeting_env",
|
||
"prompt_integration": "dotprompt",
|
||
"prompt_data": {"content": content, "metadata": {}},
|
||
}
|
||
),
|
||
"prompt_info": json.dumps({"prompt_type": "db"}),
|
||
"created_at": None,
|
||
"updated_at": updated_at,
|
||
}
|
||
return row
|
||
|
||
freshly_patched = db_row(
|
||
"production", "Begin every reply with HOWDY", datetime(2026, 8, 26, 12, 0, tzinfo=timezone.utc)
|
||
)
|
||
stale_sibling = db_row(
|
||
"development", "Begin every reply with AHOY", datetime(2026, 8, 26, 11, 0, tzinfo=timezone.utc)
|
||
)
|
||
|
||
prisma_client = MagicMock()
|
||
try:
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[freshly_patched, stale_sibling])
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
|
||
first_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_env.v1")
|
||
assert first_callback is not None
|
||
assert first_callback.prompt_manager.get_prompt("greeting_env").content == "Begin every reply with HOWDY"
|
||
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_env.v1") is first_callback
|
||
assert litellm.callbacks == [first_callback]
|
||
finally:
|
||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_env")
|
||
|
||
|
||
def _prompt_db_row(prompt_id: str, litellm_params: str) -> MagicMock:
|
||
row = MagicMock()
|
||
row.model_dump.return_value = {
|
||
"prompt_id": prompt_id,
|
||
"version": 1,
|
||
"environment": "development",
|
||
"created_by": None,
|
||
"litellm_params": litellm_params,
|
||
"prompt_info": json.dumps({"prompt_type": "db"}),
|
||
"created_at": None,
|
||
"updated_at": None,
|
||
}
|
||
return row
|
||
|
||
|
||
def _dotprompt_params(prompt_id: str) -> str:
|
||
return json.dumps(
|
||
{
|
||
"prompt_id": prompt_id,
|
||
"prompt_integration": "dotprompt",
|
||
"prompt_data": {"content": "Begin every reply with AHOY", "metadata": {}},
|
||
}
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_prompts_in_db_unloads_rows_deleted_on_another_worker(monkeypatch):
|
||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setattr(litellm, "callbacks", [])
|
||
|
||
prisma_client = MagicMock()
|
||
try:
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||
return_value=[_prompt_db_row("greeting_del", _dotprompt_params("greeting_del"))]
|
||
)
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_del.v1") is not None
|
||
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[])
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("greeting_del.v1") is None
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_del.v1") is None
|
||
assert litellm.callbacks == []
|
||
finally:
|
||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_del")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_prompts_in_db_keeps_config_prompts_when_their_id_has_no_db_row(monkeypatch):
|
||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
from litellm.types.prompts.init_prompts import PromptInfo, PromptLiteLLMParams, PromptSpec
|
||
|
||
monkeypatch.setattr(litellm, "callbacks", [])
|
||
|
||
config_prompt = PromptSpec(
|
||
prompt_id="greeting_cfg",
|
||
litellm_params=PromptLiteLLMParams(
|
||
prompt_id="greeting_cfg",
|
||
prompt_integration="dotprompt",
|
||
prompt_data={"content": "Begin every reply with AHOY", "metadata": {}},
|
||
),
|
||
prompt_info=PromptInfo(prompt_type="config"),
|
||
)
|
||
|
||
prisma_client = MagicMock()
|
||
try:
|
||
IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(prompt=config_prompt)
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(return_value=[])
|
||
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_cfg") is not None
|
||
assert len(litellm.callbacks) == 1
|
||
finally:
|
||
IN_MEMORY_PROMPT_REGISTRY.remove_prompt(prompt_id="greeting_cfg")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_prompts_in_db_keeps_the_in_memory_copy_when_a_row_fails_to_parse(monkeypatch):
|
||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
monkeypatch.setattr(litellm, "callbacks", [])
|
||
|
||
prisma_client = MagicMock()
|
||
try:
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||
return_value=[_prompt_db_row("greeting_broken", _dotprompt_params("greeting_broken"))]
|
||
)
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
loaded_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_broken.v1")
|
||
assert loaded_callback is not None
|
||
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(
|
||
return_value=[_prompt_db_row("greeting_broken", "this is not json")]
|
||
)
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_broken.v1") is loaded_callback
|
||
assert litellm.callbacks == [loaded_callback]
|
||
finally:
|
||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_broken")
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_init_prompts_in_db_keeps_a_prompt_created_while_the_sync_was_reading(monkeypatch):
|
||
from litellm.proxy.prompts.prompt_registry import IN_MEMORY_PROMPT_REGISTRY
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
from litellm.types.prompts.init_prompts import PromptInfo, PromptLiteLLMParams, PromptSpec
|
||
|
||
monkeypatch.setattr(litellm, "callbacks", [])
|
||
|
||
prisma_client = MagicMock()
|
||
try:
|
||
|
||
async def create_prompt_behind_the_select() -> list:
|
||
IN_MEMORY_PROMPT_REGISTRY.initialize_prompt(
|
||
prompt=PromptSpec(
|
||
prompt_id="greeting_race.v1",
|
||
litellm_params=PromptLiteLLMParams(
|
||
prompt_id="greeting_race",
|
||
prompt_integration="dotprompt",
|
||
prompt_data={"content": "Begin every reply with AHOY", "metadata": {}},
|
||
),
|
||
prompt_info=PromptInfo(prompt_type="db"),
|
||
)
|
||
)
|
||
return []
|
||
|
||
prisma_client.db.litellm_prompttable.find_many = AsyncMock(side_effect=create_prompt_behind_the_select)
|
||
await ProxyConfig()._init_prompts_in_db(prisma_client=prisma_client)
|
||
|
||
surviving_callback = IN_MEMORY_PROMPT_REGISTRY.get_prompt_callback_by_id("greeting_race.v1")
|
||
assert IN_MEMORY_PROMPT_REGISTRY.get_prompt_by_id("greeting_race.v1") is not None
|
||
assert surviving_callback is not None
|
||
assert litellm.callbacks == [surviving_callback]
|
||
finally:
|
||
IN_MEMORY_PROMPT_REGISTRY.delete_prompts_by_base_id("greeting_race")
|
||
|
||
|
||
class TestEmbeddingsFailureHookRequestData:
|
||
@pytest.mark.asyncio
|
||
async def test_failure_hook_gets_post_setup_data_with_logging_obj(self):
|
||
"""Request setup replaces the processor's data dict (adding the logging
|
||
object the failure hook needs to lift token usage from); the embeddings
|
||
exception handler must pass that replaced dict, not the raw request body
|
||
dict it was rebuilt from."""
|
||
from litellm.proxy._types import ProxyException
|
||
|
||
captured = {}
|
||
logging_obj_sentinel = MagicMock()
|
||
|
||
async def fake_process(self, **kwargs):
|
||
self.data = {**self.data, "litellm_logging_obj": logging_obj_sentinel}
|
||
captured["processor_data"] = self.data
|
||
raise RuntimeError("provider timeout")
|
||
|
||
with (
|
||
patch.object(
|
||
proxy_server_module,
|
||
"_read_request_body",
|
||
new=AsyncMock(return_value={"model": "my-embed", "input": "hello"}),
|
||
),
|
||
patch.object(
|
||
proxy_server_module.ProxyBaseLLMRequestProcessing,
|
||
"base_process_llm_request",
|
||
new=fake_process,
|
||
),
|
||
patch.object(proxy_server_module, "proxy_logging_obj") as mock_logging,
|
||
):
|
||
mock_logging.post_call_failure_hook = AsyncMock(return_value=None)
|
||
with pytest.raises(ProxyException):
|
||
await proxy_server_module.embeddings(
|
||
request=MagicMock(),
|
||
fastapi_response=MagicMock(),
|
||
user_api_key_dict=UserAPIKeyAuth(),
|
||
)
|
||
|
||
hook_request_data = mock_logging.post_call_failure_hook.await_args.kwargs["request_data"]
|
||
assert hook_request_data is captured["processor_data"]
|
||
assert hook_request_data["litellm_logging_obj"] is logging_obj_sentinel
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_authoritative_floor_spend_keeps_a_reset_marker_written_during_the_db_read():
|
||
"""A team-member spend reset writes the post-reset floor to the spend_db_floor marker
|
||
(auth_checks.invalidate_team_member_spend_state). A floor read already in flight when the
|
||
reset commits would otherwise cache its stale pre-reset DB value over the fresh marker,
|
||
letting a budget check raise the counter right back above the just-reset spend
|
||
(regression: PR #37971 Greptile finding)."""
|
||
from litellm.proxy.proxy_server import _authoritative_floor_spend
|
||
|
||
real_spend_counter_cache = DualCache()
|
||
counter_key = "spend:team_member:user-1:team-1"
|
||
marker_key = f"spend_db_floor:{counter_key}"
|
||
|
||
async def db_read_racing_with_a_reset(prisma_client, counter_key):
|
||
real_spend_counter_cache.in_memory_cache.set_cache(key=marker_key, value=0.0)
|
||
return 999.0
|
||
|
||
with (
|
||
patch.object( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
|
||
proxy_server_module, "spend_counter_cache", real_spend_counter_cache
|
||
),
|
||
patch.object( # test-quality-ok: the DB read must race the reset; no injectable seam for module-global prisma reads
|
||
proxy_server_module.SpendCounterReseed,
|
||
"from_db",
|
||
AsyncMock(side_effect=db_read_racing_with_a_reset),
|
||
),
|
||
):
|
||
result = await _authoritative_floor_spend(counter_key=counter_key)
|
||
|
||
assert result == 0.0
|
||
assert real_spend_counter_cache.in_memory_cache.get_cache(key=marker_key) == 0.0, (
|
||
"the in-flight DB read clobbered the post-reset floor marker with the stale pre-reset value"
|
||
)
|
||
|
||
|
||
@pytest.mark.asyncio
|
||
async def test_load_config_router_authorizes_fallback_targets_against_the_calling_key(tmp_path):
|
||
from litellm.proxy.auth.fallback_model_access import router_fallback_access_check
|
||
from litellm.proxy.proxy_server import ProxyConfig
|
||
|
||
config_file = tmp_path / "config.yaml"
|
||
config_file.write_text(
|
||
yaml.dump({"model_list": [{"model_name": "m", "litellm_params": {"model": "openai/m", "api_key": "k"}}]})
|
||
)
|
||
|
||
router, _, _ = await ProxyConfig().load_config(router=None, config_file_path=str(config_file))
|
||
|
||
assert router.fallback_access_check is router_fallback_access_check
|
||
|
||
|
||
def test_docs_redoc_openapi_are_reachable_by_default():
|
||
"""
|
||
LIT-6745: the interactive/machine-readable docs surfaces are on by
|
||
default (the customer-facing production toggle is opt-in, not opt-out).
|
||
"""
|
||
client = TestClient(app)
|
||
|
||
assert client.get("/redoc").status_code == 200
|
||
openapi_response = client.get("/openapi.json")
|
||
assert openapi_response.status_code == 200
|
||
assert "paths" in openapi_response.json()
|
||
|
||
|
||
def test_production_app_docs_urls_are_wired_to_the_real_env_helpers():
|
||
"""
|
||
LIT-6745: pins the actual `FastAPI(docs_url=..., redoc_url=..., openapi_url=...)`
|
||
construction in proxy_server.py to _get_docs_url/_get_redoc_url/_get_openapi_url,
|
||
so a hardcoded or drifted value at that call site fails this test even though
|
||
the helpers themselves are covered separately.
|
||
"""
|
||
from litellm.proxy import utils as proxy_utils
|
||
|
||
assert app.docs_url == proxy_utils._get_docs_url()
|
||
assert app.redoc_url == proxy_utils._get_redoc_url()
|
||
assert app.openapi_url == proxy_utils._get_openapi_url()
|
||
|
||
|
||
def _build_app_with_docs_env(monkeypatch, *, disabled: bool) -> FastAPI:
|
||
from litellm.proxy import utils as proxy_utils
|
||
from litellm.proxy.health_endpoints._health_endpoints import router as health_router
|
||
|
||
for flag in ("DOCS_URL", "REDOC_URL", "OPENAPI_URL"):
|
||
monkeypatch.delenv(flag, raising=False)
|
||
for flag in ("NO_DOCS", "NO_REDOC", "NO_OPENAPI"):
|
||
if disabled:
|
||
monkeypatch.setenv(flag, "True")
|
||
else:
|
||
monkeypatch.delenv(flag, raising=False)
|
||
|
||
# Mirrors the exact FastAPI() construction in proxy_server.py, so this
|
||
# exercises the real gating mechanism rather than a reimplementation of it.
|
||
app_under_test = FastAPI(
|
||
docs_url=proxy_utils._get_docs_url(),
|
||
redoc_url=proxy_utils._get_redoc_url(),
|
||
openapi_url=proxy_utils._get_openapi_url(),
|
||
)
|
||
app_under_test.include_router(health_router)
|
||
return app_under_test
|
||
|
||
|
||
def test_docs_endpoints_enabled_when_env_unset(monkeypatch):
|
||
app_under_test = _build_app_with_docs_env(monkeypatch, disabled=False)
|
||
assert app_under_test.docs_url == "/"
|
||
assert app_under_test.redoc_url == "/redoc"
|
||
assert app_under_test.openapi_url == "/openapi.json"
|
||
|
||
client = TestClient(app_under_test)
|
||
assert client.get(app_under_test.docs_url).status_code == 200
|
||
assert client.get(app_under_test.redoc_url).status_code == 200
|
||
assert client.get(app_under_test.openapi_url).status_code == 200
|
||
|
||
|
||
def test_no_docs_no_redoc_no_openapi_disable_every_documentation_surface(monkeypatch):
|
||
"""
|
||
LIT-6745: NO_DOCS, NO_REDOC and NO_OPENAPI must each 404 their surface
|
||
with no schema in the body, so a production/air-gapped deployment can
|
||
restrict every doc route consistently.
|
||
"""
|
||
app_under_test = _build_app_with_docs_env(monkeypatch, disabled=True)
|
||
assert app_under_test.docs_url is None
|
||
assert app_under_test.redoc_url is None
|
||
assert app_under_test.openapi_url is None
|
||
|
||
client = TestClient(app_under_test)
|
||
for route in ("/", "/redoc", "/openapi.json"):
|
||
response = client.get(route)
|
||
assert response.status_code == 404
|
||
assert "openapi" not in response.text.lower()
|
||
assert "paths" not in response.text.lower()
|
||
|
||
|
||
def test_disabling_docs_does_not_disable_other_routes(monkeypatch):
|
||
"""
|
||
LIT-6745: disabling the doc surfaces must not affect inference/management
|
||
routes, since NO_DOCS/NO_REDOC/NO_OPENAPI only remove the routes FastAPI
|
||
itself auto-registers for docs_url/redoc_url/openapi_url.
|
||
"""
|
||
app_under_test = _build_app_with_docs_env(monkeypatch, disabled=True)
|
||
client = TestClient(app_under_test)
|
||
|
||
assert client.get("/redoc").status_code == 404
|
||
assert client.get("/health/liveliness").status_code == 200
|