Merge remote-tracking branch 'origin/litellm_internal_staging' into codex/budget-race-enforcement-greptile-fix

# Conflicts:
#	litellm/proxy/db/spend_counter_reseed.py
#	litellm/proxy/proxy_server.py
This commit is contained in:
user 2026-04-30 18:25:48 -07:00
commit 64fadc3b8e
29 changed files with 1986 additions and 1056 deletions

View file

@ -392,6 +392,7 @@ class DualCache(BaseCache):
value: float,
parent_otel_span: Optional[Span] = None,
local_only: bool = False,
refresh_ttl: bool = False,
**kwargs,
) -> Optional[float]:
"""
@ -399,6 +400,9 @@ class DualCache(BaseCache):
Value - float - the value you want to increment by
Refresh_ttl - bool - if True, resets the Redis TTL on every write.
Default False preserves window-style semantics.
Returns - the incremented value, or None if no cache backend is
available (in_memory_cache is None and Redis failed/is absent).
"""
@ -415,6 +419,7 @@ class DualCache(BaseCache):
value,
parent_otel_span=parent_otel_span,
ttl=kwargs.get("ttl", None),
refresh_ttl=refresh_ttl,
)
return result

View file

@ -824,6 +824,7 @@ class RedisCache(BaseCache):
value: float,
ttl: Optional[int] = None,
parent_otel_span: Optional[Span] = None,
refresh_ttl: bool = False,
) -> float:
from redis.asyncio import Redis
@ -834,11 +835,12 @@ class RedisCache(BaseCache):
try:
result = await _redis_client.incrbyfloat(name=key, amount=value)
if _used_ttl is not None:
# check if key already has ttl, if not -> set ttl
current_ttl = await _redis_client.ttl(key)
if current_ttl == -1:
# Key has no expiration
if refresh_ttl:
await _redis_client.expire(key, _used_ttl)
else:
current_ttl = await _redis_client.ttl(key)
if current_ttl == -1:
await _redis_client.expire(key, _used_ttl)
## LOGGING ##
end_time = time.time()

View file

@ -1425,6 +1425,7 @@ LITELLM_PROXY_ADMIN_NAME = "default_user_id"
LITELLM_CLI_SOURCE_IDENTIFIER = "litellm-cli"
LITELLM_CLI_SESSION_TOKEN_PREFIX = "litellm-session-token"
CLI_SSO_SESSION_CACHE_KEY_PREFIX = "cli_sso_session"
CLI_SSO_SESSION_TTL_SECONDS = 600
CLI_JWT_TOKEN_NAME = "cli-jwt-token"
# Support both CLI_JWT_EXPIRATION_HOURS and LITELLM_CLI_JWT_EXPIRATION_HOURS for backwards compatibility
CLI_JWT_EXPIRATION_HOURS = int(

View file

@ -2535,10 +2535,16 @@ class BaseLLMHTTPHandler:
},
)
delete_kwargs: Dict[str, Any] = {
"url": url,
"headers": headers,
"timeout": timeout,
}
if data:
delete_kwargs["json"] = data
try:
response = await async_httpx_client.delete(
url=url, headers=headers, json=data, timeout=timeout
)
response = await async_httpx_client.delete(**delete_kwargs)
except Exception as e:
raise self._handle_error(
@ -2619,10 +2625,16 @@ class BaseLLMHTTPHandler:
},
)
delete_kwargs: Dict[str, Any] = {
"url": url,
"headers": headers,
"timeout": timeout,
}
if data:
delete_kwargs["json"] = data
try:
response = sync_httpx_client.delete(
url=url, headers=headers, json=data, timeout=timeout
)
response = sync_httpx_client.delete(**delete_kwargs)
except Exception as e:
raise self._handle_error(

View file

@ -2161,8 +2161,8 @@ class PassThroughGenericEndpoint(LiteLLMPydanticObjectBase):
description="The USD cost per request to the target endpoint. This is used to calculate the cost of the request to the target endpoint.",
)
auth: bool = Field(
default=False,
description="Whether authentication is required for the pass-through endpoint. If True, requests to the endpoint will require a valid LiteLLM API key.",
default=True,
description="Whether authentication is required for the pass-through endpoint. Defaults to True so a pass-through silently created without an explicit value still requires a valid LiteLLM API key — set to False only if the endpoint is meant to be a public forwarder (e.g. an unauthenticated webhook target).",
)
guardrails: Optional[PassThroughGuardrailsConfig] = Field(
default=None,

View file

@ -473,7 +473,12 @@ async def check_api_key_for_custom_headers_or_pass_through_endpoints(
for endpoint in pass_through_endpoints:
if isinstance(endpoint, dict) and endpoint.get("path", "") == route:
## IF AUTH DISABLED
if endpoint.get("auth") is not True:
# Default to True: a config dict with no ``auth`` key
# otherwise produced an unauthenticated forwarder. The
# Pydantic ``PassThroughGenericEndpoint.auth`` default
# is also True, but raw config dicts skip that path —
# so this runtime check has to default to True too.
if endpoint.get("auth", True) is not True:
return UserAPIKeyAuth()
## IF AUTH ENABLED
### IF CUSTOM PARSER REQUIRED

View file

@ -313,23 +313,24 @@ sequenceDiagram
participant Proxy as LiteLLM Proxy
participant SSO as SSO Provider
CLI->>CLI: Generate key ID (sk-uuid)
CLI->>Browser: Open /sso/key/generate?source=litellm-cli&key=sk-uuid
CLI->>Proxy: POST /sso/cli/start
Proxy->>CLI: Return login_id, poll_secret, user_code
CLI->>Browser: Open /sso/key/generate?source=litellm-cli&key=login_id
Browser->>Proxy: GET /sso/key/generate?source=litellm-cli&key=sk-uuid
Proxy->>Proxy: Set cli_state = litellm-session-token:sk-uuid
Proxy->>SSO: Redirect with state=litellm-session-token:sk-uuid
Browser->>Proxy: GET /sso/key/generate?source=litellm-cli&key=login_id
Proxy->>Proxy: Set cli_state = litellm-session-token:login_id
Proxy->>SSO: Redirect with state=litellm-session-token:login_id
SSO->>Browser: Show login page
Browser->>SSO: User authenticates
SSO->>Proxy: Redirect to /sso/callback?state=litellm-session-token:sk-uuid
SSO->>Proxy: Redirect to /sso/callback?state=litellm-session-token:login_id
Proxy->>Proxy: Check if state starts with "litellm-session-token:"
Proxy->>Proxy: Generate API key with ID=sk-uuid
Proxy->>Browser: Show success page
Proxy->>Browser: Prompt for user_code
Browser->>Proxy: POST /sso/cli/complete/login_id
CLI->>Proxy: Poll /sso/cli/poll/sk-uuid
Proxy->>CLI: Return {"status": "ready", "key": "sk-uuid"}
CLI->>Proxy: Poll /sso/cli/poll/login_id with poll_secret header
Proxy->>CLI: Return {"status": "ready", "key": "jwt"}
CLI->>CLI: Save key to ~/.litellm/token.json
```
@ -343,13 +344,13 @@ The CLI provides three authentication commands:
### Authentication Flow Steps
1. **Generate Session ID**: CLI generates a unique key ID (`sk-{uuid}`)
2. **Open Browser**: CLI opens browser to `/sso/key/generate` with CLI source and key parameters
3. **SSO Redirect**: Proxy sets the formatted state (`litellm-session-token:sk-uuid`) as OAuth state parameter and redirects to SSO provider
1. **Start Session**: CLI creates a short-lived login session with `/sso/cli/start`
2. **Open Browser**: CLI opens browser to `/sso/key/generate` with CLI source and login ID parameters
3. **SSO Redirect**: Proxy sets the formatted state (`litellm-session-token:{login_id}`) as OAuth state parameter and redirects to SSO provider
4. **User Authentication**: User completes SSO authentication in browser
5. **Callback Processing**: SSO provider redirects back to proxy with state parameter
6. **Key Generation**: Proxy detects CLI login (state starts with "litellm-session-token:") and generates API key with pre-specified ID
7. **Polling**: CLI polls `/sso/cli/poll/{key_id}` endpoint until key is ready
6. **User Code Verification**: Browser confirms the verification code shown in the CLI
7. **Polling**: CLI polls `/sso/cli/poll/{login_id}` with the polling secret header until the JWT is ready
8. **Token Storage**: CLI saves the authentication token to `~/.litellm/token.json`
### Benefits of This Approach
@ -357,7 +358,7 @@ The CLI provides three authentication commands:
- **No Local Server**: No need to run a local callback server
- **Standard OAuth**: Uses OAuth 2.0 state parameter correctly
- **Remote Compatible**: Works with remote proxy servers
- **Secure**: Uses UUID session identifiers
- **Secure**: Keeps the polling secret out of the browser handoff
- **Simple Setup**: No additional OAuth redirect URL configuration needed
### Token Storage

View file

@ -5,6 +5,7 @@ import time
import webbrowser
from pathlib import Path
from typing import Any, Dict, List, Optional
from urllib.parse import urlencode
import click
import requests
@ -241,7 +242,7 @@ def prompt_team_selection(teams: List[Dict[str, Any]]) -> Optional[Dict[str, Any
def prompt_team_selection_fallback(
teams: List[Dict[str, Any]]
teams: List[Dict[str, Any]],
) -> Optional[Dict[str, Any]]:
"""Fallback team selection for non-interactive environments"""
if not teams:
@ -279,6 +280,7 @@ def prompt_team_selection_fallback(
def _poll_for_ready_data(
url: str,
*,
headers: Optional[Dict[str, str]] = None,
total_timeout: int = 300,
poll_interval: int = 2,
request_timeout: int = 10,
@ -291,7 +293,10 @@ def _poll_for_ready_data(
) -> Optional[Dict[str, Any]]:
for attempt in range(total_timeout // poll_interval):
try:
response = requests.get(url, timeout=request_timeout)
request_kwargs: Dict[str, Any] = {"timeout": request_timeout}
if headers is not None:
request_kwargs["headers"] = headers
response = requests.get(url, **request_kwargs)
if response.status_code == 200:
data = response.json()
status = data.get("status")
@ -346,7 +351,23 @@ def _normalize_teams(teams, team_details):
return []
def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
def _start_cli_sso_flow(base_url: str) -> Dict[str, Any]:
response = requests.post(f"{base_url}/sso/cli/start", timeout=10)
response.raise_for_status()
data = response.json()
required_fields = ("login_id", "poll_secret", "user_code")
if not all(isinstance(data.get(field), str) for field in required_fields):
raise ValueError("Invalid CLI SSO start response")
return data
def _get_cli_sso_poll_headers(poll_secret: str) -> Dict[str, str]:
return {"x-litellm-cli-poll-secret": poll_secret}
def _poll_for_authentication(
base_url: str, key_id: str, poll_secret: str
) -> Optional[dict]:
"""
Poll the server for authentication completion and handle team selection.
@ -356,6 +377,7 @@ def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
poll_url = f"{base_url}/sso/cli/poll/{key_id}"
data = _poll_for_ready_data(
poll_url,
headers=_get_cli_sso_poll_headers(poll_secret),
pending_message="Still waiting for authentication...",
)
if not data:
@ -373,6 +395,7 @@ def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
jwt_with_team = _handle_team_selection_during_polling(
base_url=base_url,
key_id=key_id,
poll_secret=poll_secret,
teams=normalized_teams,
)
@ -410,7 +433,7 @@ def _poll_for_authentication(base_url: str, key_id: str) -> Optional[dict]:
def _handle_team_selection_during_polling(
base_url: str, key_id: str, teams: List[Dict[str, Any]]
base_url: str, key_id: str, poll_secret: str, teams: List[Dict[str, Any]]
) -> Optional[str]:
"""
Handle team selection and re-poll with selected team_id.
@ -441,6 +464,7 @@ def _handle_team_selection_during_polling(
poll_url = f"{base_url}/sso/cli/poll/{key_id}?team_id={team_id}"
data = _poll_for_ready_data(
poll_url,
headers=_get_cli_sso_poll_headers(poll_secret),
pending_message="Still waiting for team authentication...",
other_status_message="Waiting for team authentication to complete...",
http_error_log_every=10,
@ -514,29 +538,24 @@ def _render_and_prompt_for_team_selection(teams: List[Dict[str, Any]]) -> Option
@click.pass_context
def login(ctx: click.Context):
"""Login to LiteLLM proxy using SSO authentication"""
from litellm._uuid import uuid
from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER
from litellm.proxy.client.cli.interface import show_commands
base_url = ctx.obj["base_url"]
# Check if we have an existing key to regenerate
existing_key = get_stored_api_key()
# Generate unique key ID for this login session
key_id = f"sk-{str(uuid.uuid4())}"
try:
# Construct SSO login URL with CLI source and pre-generated key
sso_url = f"{base_url}/sso/key/generate?source={LITELLM_CLI_SOURCE_IDENTIFIER}&key={key_id}"
cli_sso_flow = _start_cli_sso_flow(base_url=base_url)
key_id = cli_sso_flow["login_id"]
poll_secret = cli_sso_flow["poll_secret"]
user_code = cli_sso_flow["user_code"]
# If we have an existing key, include it as a parameter to the login endpoint
# The server will encode it in the OAuth state parameter for the SSO flow
if existing_key:
sso_url += f"&existing_key={existing_key}"
sso_url = f"{base_url}/sso/key/generate?" + urlencode(
{"source": LITELLM_CLI_SOURCE_IDENTIFIER, "key": key_id}
)
click.echo(f"Opening browser to: {sso_url}")
click.echo("Please complete the SSO authentication in your browser...")
click.echo(f"Verification code: {user_code}")
click.echo(f"Session ID: {key_id}")
# Open browser
@ -545,7 +564,9 @@ def login(ctx: click.Context):
# Poll for authentication completion
click.echo("Waiting for authentication...")
auth_result = _poll_for_authentication(base_url=base_url, key_id=key_id)
auth_result = _poll_for_authentication(
base_url=base_url, key_id=key_id, poll_secret=poll_secret
)
if auth_result:
api_key = auth_result["api_key"]

View file

@ -52,6 +52,37 @@ class ResetBudgetJob:
### RESET MULTI-WINDOW BUDGETS ###
await self.reset_budget_windows()
@staticmethod
async def _invalidate_spend_counter(counter_key: str) -> None:
"""Zero a spend counter so a DB-row reset takes effect immediately.
Call AFTER the DB write commits. Clearing Redis before the DB
commit opens a window where get_current_spend reads 0 from Redis
while the DB still holds the pre-reset value, allowing bypass.
"""
try:
from litellm.proxy.proxy_server import spend_counter_cache
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key, value=0.0, ttl=60
)
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_set_cache(
key=counter_key, value=0.0, ttl=60
)
except Exception as redis_err:
verbose_proxy_logger.warning(
"Failed to reset spend counter %s in Redis: %s. "
"Budget may be over-enforced until counter expires.",
counter_key,
redis_err,
)
except Exception as e:
verbose_proxy_logger.warning(
"Failed to reset spend counter %s: %s", counter_key, e
)
async def reset_budget_for_litellm_team_members(
self, budgets_to_reset: List[LiteLLM_BudgetTableFull]
):
@ -64,46 +95,30 @@ class ResetBudgetJob:
if budget.budget_id is not None
]
# Reset spend counters for affected team members.
# Reset Redis directly so a transient failure doesn't leave stale
# counters that get_current_spend would read as authoritative.
try:
from litellm.proxy.proxy_server import spend_counter_cache
memberships = await self.prisma_client.db.litellm_teammembership.find_many(
where={"budget_id": {"in": budget_ids}}
)
for m in memberships:
counter_key = f"spend:team_member:{m.user_id}:{m.team_id}"
# Always reset in-memory
spend_counter_cache.in_memory_cache.set_cache(
key=counter_key, value=0.0
)
# Explicitly reset Redis with warning on failure
if spend_counter_cache.redis_cache is not None:
try:
await spend_counter_cache.redis_cache.async_set_cache(
key=counter_key, value=0.0
)
except Exception as redis_err:
verbose_proxy_logger.warning(
"Failed to reset team member spend counter in Redis %s: %s. "
"Budget may be over-enforced until counter expires.",
counter_key,
redis_err,
)
except Exception as e:
memberships = []
verbose_proxy_logger.warning(
"Failed to reset team member spend counters: %s", e
"Failed to fetch team memberships for counter invalidation: %s", e
)
return await self.prisma_client.db.litellm_teammembership.update_many(
update_result = await self.prisma_client.db.litellm_teammembership.update_many(
where={"budget_id": {"in": budget_ids}},
data={
"spend": 0,
},
)
for m in memberships:
await self._invalidate_spend_counter(
f"spend:team_member:{m.user_id}:{m.team_id}"
)
return update_result
async def reset_budget_for_keys_linked_to_budgets(
self, budgets_to_reset: List[LiteLLM_BudgetTableFull]
):
@ -126,17 +141,36 @@ class ResetBudgetJob:
if not budget_ids:
return
return await self.prisma_client.db.litellm_verificationtoken.update_many(
where={
"budget_id": {"in": budget_ids},
"budget_duration": None, # only keys without their own reset schedule
"spend": {"gt": 0}, # only reset keys that have accumulated spend
},
data={
"spend": 0,
},
where_clause: dict = {
"budget_id": {"in": budget_ids},
"budget_duration": None, # only keys without their own reset schedule
"spend": {"gt": 0}, # only reset keys that have accumulated spend
}
try:
keys = await self.prisma_client.db.litellm_verificationtoken.find_many(
where=where_clause
)
except Exception as e:
keys = []
verbose_proxy_logger.warning(
"Failed to fetch keys for counter invalidation: %s", e
)
update_result = (
await self.prisma_client.db.litellm_verificationtoken.update_many(
where=where_clause,
data={
"spend": 0,
},
)
)
for k in keys:
await self._invalidate_spend_counter(f"spend:key:{k.token}")
return update_result
async def reset_budget_for_litellm_budget_table(self):
"""
Resets the budget for all LiteLLM End-Users (Customers), and Team Members if their budget has expired
@ -365,6 +399,10 @@ class ResetBudgetJob:
data_list=updated_keys,
table_name="key",
)
for k in updated_keys:
token = getattr(k, "token", None)
if token:
await self._invalidate_spend_counter(f"spend:key:{token}")
end_time = time.time()
if len(failed_keys) > 0: # If any keys failed to reset
@ -450,6 +488,12 @@ class ResetBudgetJob:
data_list=updated_users,
table_name="user",
)
for u in updated_users:
user_id = getattr(u, "user_id", None)
if user_id:
await self._invalidate_spend_counter(
f"spend:user:{user_id}"
)
end_time = time.time()
if len(failed_users) > 0: # If any users failed to reset
@ -541,6 +585,12 @@ class ResetBudgetJob:
data_list=updated_teams,
table_name="team",
)
for t in updated_teams:
team_id = getattr(t, "team_id", None)
if team_id:
await self._invalidate_spend_counter(
f"spend:team:{team_id}"
)
end_time = time.time()
if len(failed_teams) > 0: # If any teams failed to reset

View file

@ -0,0 +1,52 @@
"""Helpers for unauthenticated logo / favicon endpoints."""
import os
from typing import Optional, Tuple
from litellm._logging import verbose_proxy_logger
LOCAL_IMAGE_HEADER_BYTES = 512
def detect_local_image_media_type(header: bytes) -> Optional[str]:
"""Return a browser image media type for supported local image signatures."""
if header[0:8] == b"\x89PNG\r\n\x1a\n":
return "image/png"
if header[0:4] == b"GIF8" and header[5:6] == b"a":
return "image/gif"
if header[0:3] == b"\xff\xd8\xff":
return "image/jpeg"
if header[0:4] == b"RIFF" and header[8:12] == b"WEBP":
return "image/webp"
if header[0:4] in (b"\x00\x00\x01\x00", b"\x00\x00\x02\x00"):
return "image/x-icon"
return None
def resolve_validated_local_image_path(candidate: str) -> Optional[Tuple[str, str]]:
"""Resolve ``candidate`` only when it is an existing supported image file."""
if not candidate:
return None
try:
resolved = os.path.realpath(os.path.expanduser(candidate))
except (OSError, ValueError):
return None
if not os.path.isfile(resolved):
return None
try:
with open(resolved, "rb") as f:
header = f.read(LOCAL_IMAGE_HEADER_BYTES)
except OSError as exc:
verbose_proxy_logger.debug("Could not read local asset %r: %s", candidate, exc)
return None
media_type = detect_local_image_media_type(header)
if media_type is None:
verbose_proxy_logger.warning(
"Local asset %r is not a supported image file; falling back to default.",
candidate,
)
return None
return resolved, media_type

View file

@ -160,7 +160,9 @@ class SpendCounterReseed:
"""
lock = await SpendCounterReseed._get_lock(counter_key)
async with lock:
# Re-check after acquiring the lock - another waiter may have warmed it.
# Re-check after acquiring the lock. Skip in-memory on a clean
# Redis miss - in-memory is per-pod-stale.
redis_clean_miss = False
if spend_counter_cache.redis_cache is not None:
try:
val = await spend_counter_cache.redis_cache.async_get_cache(
@ -168,11 +170,13 @@ class SpendCounterReseed:
)
if val is not None:
return float(val)
redis_clean_miss = True
except Exception:
pass
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
if val is not None:
return float(val)
if not redis_clean_miss:
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
if val is not None:
return float(val)
db_spend = await SpendCounterReseed.from_db(prisma_client, counter_key)
if db_spend is None:
@ -184,6 +188,7 @@ class SpendCounterReseed:
await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=db_spend,
refresh_ttl=True,
)
)
spend_counter_cache.in_memory_cache.set_cache(
@ -192,7 +197,7 @@ class SpendCounterReseed:
)
else:
await spend_counter_cache.async_increment_cache(
key=counter_key, value=db_spend
key=counter_key, value=db_spend, refresh_ttl=True
)
except Exception:
verbose_proxy_logger.exception(

View file

@ -1,10 +1,6 @@
from datetime import datetime
from fastapi import APIRouter, Depends, Request, Response
from fastapi.responses import ORJSONResponse
from fastapi import APIRouter, Depends, HTTPException, Request, Response
from fastapi.responses import ORJSONResponse, StreamingResponse
import litellm
from litellm._uuid import uuid
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
@ -30,12 +26,17 @@ async def google_generate_content(
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import (
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
select_data_generator,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
@ -43,48 +44,33 @@ async def google_generate_content(
if "model" not in data:
data["model"] = model_name
# Extract generationConfig and pass it as config parameter
generation_config = data.pop("generationConfig", None)
if generation_config:
data["config"] = generation_config
# Add user authentication metadata for cost tracking
data = await add_litellm_data_to_request(
data=data,
request=request,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
general_settings=general_settings,
version=version,
)
# Create logging object with full request metadata so callbacks (e.g. S3) get user/trace_id
data["litellm_call_id"] = request.headers.get(
"x-litellm-call-id", str(uuid.uuid4())
)
logging_obj, data = litellm.utils.function_setup(
original_function="agenerate_content",
rules_obj=litellm.utils.Rules(),
start_time=datetime.now(),
**data,
)
data["litellm_logging_obj"] = logging_obj
# call router
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
response = await llm_router.agenerate_content(**data)
success_headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
response=response,
request_data=data,
request=request,
user_api_key_dict=user_api_key_dict,
logging_obj=logging_obj,
version=version,
proxy_logging_obj=proxy_logging_obj,
)
fastapi_response.headers.update(success_headers)
return response
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="agenerate_content",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=model_name,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
@router.post(
@ -101,73 +87,52 @@ async def google_stream_generate_content(
fastapi_response: Response,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
from litellm.proxy.proxy_server import (
general_settings,
llm_router,
proxy_config,
proxy_logging_obj,
select_data_generator,
user_api_base,
user_max_tokens,
user_model,
user_request_timeout,
user_temperature,
version,
)
data = await _read_request_body(request=request)
if "model" not in data:
data["model"] = model_name
data["stream"] = True
data["stream"] = True # enforce streaming for this endpoint
# Extract generationConfig and pass it as config parameter
generation_config = data.pop("generationConfig", None)
if generation_config:
data["config"] = generation_config
# Add user authentication metadata for cost tracking
data = await add_litellm_data_to_request(
data=data,
request=request,
user_api_key_dict=user_api_key_dict,
proxy_config=proxy_config,
general_settings=general_settings,
version=version,
)
# Create logging object with full request metadata so streaming END callbacks (e.g. S3) get user/trace_id
data["litellm_call_id"] = request.headers.get(
"x-litellm-call-id", str(uuid.uuid4())
)
logging_obj, data = litellm.utils.function_setup(
original_function="agenerate_content_stream",
rules_obj=litellm.utils.Rules(),
start_time=datetime.now(),
**data,
)
data["litellm_logging_obj"] = logging_obj
# call router
if llm_router is None:
raise HTTPException(status_code=500, detail="Router not initialized")
response = await llm_router.agenerate_content_stream(**data)
success_headers = await ProxyBaseLLMRequestProcessing.build_litellm_proxy_success_headers_from_llm_response(
response=response,
request_data=data,
request=request,
user_api_key_dict=user_api_key_dict,
logging_obj=logging_obj,
version=version,
proxy_logging_obj=proxy_logging_obj,
)
# Check if response is an async iterator (streaming response)
if response is not None and hasattr(response, "__aiter__"):
return StreamingResponse(
content=response,
media_type="text/event-stream",
headers=success_headers,
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
return await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
route_type="agenerate_content_stream",
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
general_settings=general_settings,
proxy_config=proxy_config,
select_data_generator=select_data_generator,
model=model_name,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
fastapi_response.headers.update(success_headers)
return response
@router.post(

View file

@ -13,7 +13,9 @@ import base64
import hashlib
import inspect
import os
import re
import secrets
from html import escape
from copy import deepcopy
from typing import (
TYPE_CHECKING,
@ -27,13 +29,13 @@ from typing import (
Union,
cast,
)
from urllib.parse import urlencode, urlparse
from urllib.parse import parse_qs, urlencode, urlparse
if TYPE_CHECKING:
import httpx
import jwt
from fastapi import APIRouter, Depends, HTTPException, Request, status
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi.responses import RedirectResponse
import litellm
@ -41,6 +43,9 @@ from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.caching import DualCache
from litellm.constants import (
CLI_SSO_SESSION_CACHE_KEY_PREFIX,
CLI_SSO_SESSION_TTL_SECONDS,
LITELLM_CLI_SOURCE_IDENTIFIER,
LITELLM_UI_SESSION_DURATION,
MAX_SPENDLOG_ROWS_TO_QUERY,
MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE,
@ -70,7 +75,7 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_user_object
from litellm.proxy.auth.auth_utils import _has_user_setup_sso
from litellm.proxy.auth.auth_utils import _get_request_ip_address, _has_user_setup_sso
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.admin_ui_utils import (
@ -123,6 +128,250 @@ router = APIRouter()
# Metadata fields (token_type, expires_in, scope) are intentionally kept so
# response convertors see the same fields in the PKCE path as in the non-PKCE path.
_OAUTH_TOKEN_FIELDS = frozenset({"access_token", "id_token", "refresh_token"})
_CLI_SSO_FLOW_CACHE_KEY_PREFIX = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:flow"
_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX = (
f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:start_rate_limit"
)
_CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS = 60
_CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS = 30
_CLI_SSO_USER_CODE_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789"
_CLI_SSO_LOGIN_ID_RE = re.compile(r"^cli-[A-Za-z0-9_-]{12,124}$")
def _hash_cli_sso_secret(secret: str) -> str:
return hashlib.sha256(secret.encode("utf-8")).hexdigest()
def _normalize_cli_sso_user_code(user_code: str) -> str:
return "".join(ch for ch in user_code.upper() if ch.isalnum())
def _generate_cli_sso_user_code() -> str:
user_code = "".join(secrets.choice(_CLI_SSO_USER_CODE_ALPHABET) for _ in range(8))
return f"{user_code[:4]}-{user_code[4:]}"
def _get_cli_sso_flow_cache_key(login_id: str) -> str:
return f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:{login_id}"
def _is_valid_cli_sso_login_id(login_id: Optional[str]) -> bool:
return isinstance(login_id, str) and bool(_CLI_SSO_LOGIN_ID_RE.fullmatch(login_id))
def _get_cli_sso_start_rate_limit_cache_key(
request: Request, use_x_forwarded_for: Optional[bool] = False
) -> str:
client_ip = (
_get_request_ip_address(
request=request, use_x_forwarded_for=use_x_forwarded_for
)
or "unknown"
)
client_ip_hash = _hash_cli_sso_secret(client_ip)
return f"{_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX}:{client_ip_hash}"
def _check_cli_sso_start_rate_limit(
request: Request,
cache: DualCache,
use_x_forwarded_for: Optional[bool] = False,
) -> None:
rate_limit_cache_key = _get_cli_sso_start_rate_limit_cache_key(
request=request, use_x_forwarded_for=use_x_forwarded_for
)
current_attempts = cache.increment_cache(
key=rate_limit_cache_key,
value=1,
ttl=_CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS,
)
if current_attempts > _CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS:
raise HTTPException(
status_code=429,
detail="Too many CLI login attempts. Try again later.",
)
def _get_cli_sso_flow_or_raise(login_id: Optional[str], cache: DualCache) -> dict:
if not _is_valid_cli_sso_login_id(login_id):
raise HTTPException(status_code=400, detail="Invalid CLI login session")
cache_key = _get_cli_sso_flow_cache_key(cast(str, login_id))
flow = cache.get_cache(key=cache_key)
if not isinstance(flow, dict) or "poll_secret_hash" not in flow:
raise HTTPException(status_code=400, detail="Invalid CLI login session")
return flow
def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None:
cache.set_cache(
key=_get_cli_sso_flow_cache_key(login_id),
value=flow,
ttl=CLI_SSO_SESSION_TTL_SECONDS,
)
def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool:
expected_poll_secret_hash = flow.get("poll_secret_hash")
if not isinstance(expected_poll_secret_hash, str) or not isinstance(
poll_secret, str
):
return False
supplied_poll_secret_hash = _hash_cli_sso_secret(poll_secret)
return secrets.compare_digest(supplied_poll_secret_hash, expected_poll_secret_hash)
def _render_cli_sso_verification_page(
verify_url: str, browser_complete_token: str
) -> str:
escaped_verify_url = escape(verify_url, quote=True)
escaped_browser_complete_token = escape(browser_complete_token, quote=True)
return f"""
<!doctype html>
<html>
<head>
<title>LiteLLM CLI Login</title>
<style>
body {{
font-family: -apple-system, BlinkMacSystemFont, "Segoe UI", sans-serif;
margin: 0;
min-height: 100vh;
display: flex;
align-items: center;
justify-content: center;
background: #f8fafc;
color: #0f172a;
}}
main {{
width: min(420px, calc(100vw - 32px));
background: #ffffff;
border: 1px solid #e2e8f0;
border-radius: 8px;
padding: 28px;
box-shadow: 0 12px 32px rgba(15, 23, 42, 0.08);
}}
h1 {{ font-size: 22px; margin: 0 0 12px; }}
p {{ line-height: 1.5; margin: 0 0 18px; color: #334155; }}
label {{ display: block; font-weight: 600; margin-bottom: 8px; }}
input {{
box-sizing: border-box;
width: 100%;
padding: 12px;
border: 1px solid #cbd5e1;
border-radius: 6px;
font-size: 20px;
letter-spacing: 0.08em;
text-transform: uppercase;
}}
button {{
margin-top: 16px;
width: 100%;
padding: 12px;
border: 0;
border-radius: 6px;
background: #0f172a;
color: #ffffff;
font-weight: 600;
cursor: pointer;
}}
</style>
</head>
<body>
<main>
<h1>Complete CLI Login</h1>
<p>Enter the verification code shown in your terminal to finish this login.</p>
<form method="post" action="{escaped_verify_url}">
<input type="hidden" name="browser_complete_token" value="{escaped_browser_complete_token}" />
<label for="user_code">Verification code</label>
<input id="user_code" name="user_code" autocomplete="one-time-code" required autofocus />
<button type="submit">Continue</button>
</form>
</main>
</body>
</html>
"""
@router.post("/sso/cli/start", tags=["experimental"], include_in_schema=False)
async def cli_sso_start(request: Request):
from litellm.proxy.proxy_server import general_settings, user_api_key_cache
_check_cli_sso_start_rate_limit(
request=request,
cache=user_api_key_cache,
use_x_forwarded_for=bool(
(general_settings or {}).get("use_x_forwarded_for", False)
),
)
login_id = f"cli-{secrets.token_urlsafe(24)}"
poll_secret = secrets.token_urlsafe(32)
user_code = _generate_cli_sso_user_code()
flow = {
"poll_secret_hash": _hash_cli_sso_secret(poll_secret),
"user_code_hash": _hash_cli_sso_secret(_normalize_cli_sso_user_code(user_code)),
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
return {
"login_id": login_id,
"poll_secret": poll_secret,
"user_code": user_code,
"expires_in": CLI_SSO_SESSION_TTL_SECONDS,
}
@router.post(
"/sso/cli/complete/{login_id}", tags=["experimental"], include_in_schema=False
)
async def cli_sso_complete(request: Request, login_id: str):
from fastapi.responses import HTMLResponse
from litellm.proxy.common_utils.html_forms.cli_sso_success import (
render_cli_sso_success_page,
)
from litellm.proxy.proxy_server import user_api_key_cache
flow = _get_cli_sso_flow_or_raise(login_id=login_id, cache=user_api_key_cache)
if not flow.get("sso_complete") or not flow.get("session_data"):
raise HTTPException(status_code=400, detail="CLI login is not ready")
body = (await request.body()).decode("utf-8")
form_values = parse_qs(body)
supplied_user_code = (form_values.get("user_code") or [""])[0]
supplied_browser_complete_token = (
form_values.get("browser_complete_token") or [""]
)[0]
supplied_user_code_hash = _hash_cli_sso_secret(
_normalize_cli_sso_user_code(supplied_user_code)
)
supplied_browser_complete_token_hash = _hash_cli_sso_secret(
supplied_browser_complete_token
)
expected_user_code_hash = flow.get("user_code_hash")
if not isinstance(expected_user_code_hash, str) or not secrets.compare_digest(
supplied_user_code_hash, expected_user_code_hash
):
raise HTTPException(status_code=400, detail="Invalid verification code")
expected_browser_complete_token_hash = flow.get("browser_complete_token_hash")
if not isinstance(
expected_browser_complete_token_hash, str
) or not secrets.compare_digest(
supplied_browser_complete_token_hash, expected_browser_complete_token_hash
):
raise HTTPException(status_code=400, detail="Invalid verification code")
flow["user_code_verified"] = True
_set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow)
html_content = render_cli_sso_success_page()
return HTMLResponse(content=html_content, status_code=200)
def normalize_email(email: Optional[str]) -> Optional[str]:
@ -333,6 +582,7 @@ async def google_login(
from litellm.proxy.proxy_server import (
premium_user,
prisma_client,
user_api_key_cache,
user_custom_ui_sso_sign_in_handler,
)
@ -382,14 +632,15 @@ async def google_login(
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
request=request,
sso_callback_route="sso/callback",
existing_key=existing_key,
)
# Store CLI key in state for OAuth flow
if source == LITELLM_CLI_SOURCE_IDENTIFIER:
_get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache)
# Store CLI login handle in state for OAuth flow
cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state(
source=source,
key=key,
existing_key=existing_key,
)
# check if user defined a custom auth sso sign in handler, if yes, use it
@ -1392,18 +1643,12 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
)
if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
# Extract the key ID and existing_key from the state
# State format: {PREFIX}:{key}:{existing_key} or {PREFIX}:{key}
state_parts = state.split(":", 2) # Split into max 3 parts
# State format: {PREFIX}:{login_id}
state_parts = state.split(":", 1)
key_id = state_parts[1] if len(state_parts) > 1 else None
existing_key = state_parts[2] if len(state_parts) > 2 else None
verbose_proxy_logger.info(
f"CLI SSO callback detected for key: {key_id}, existing_key: {existing_key}"
)
return await cli_sso_callback(
request=request, key=key_id, existing_key=existing_key, result=result
)
verbose_proxy_logger.info("CLI SSO callback detected")
return await cli_sso_callback(request=request, key=key_id, result=result)
# Control-plane cross-origin: read return_to from cookie.
# Starlette's cookie_parser already handles RFC 2109 unquoting.
@ -1424,13 +1669,10 @@ async def auth_callback(request: Request, state: Optional[str] = None): # noqa:
async def cli_sso_callback(
request: Request,
key: Optional[str] = None,
existing_key: Optional[str] = None,
result: Optional[Union[OpenID, dict]] = None,
):
"""CLI SSO callback - stores session info for JWT generation on polling"""
verbose_proxy_logger.info(
f"CLI SSO callback for key: {key}, existing_key: {existing_key}"
)
verbose_proxy_logger.info("CLI SSO callback")
from litellm.proxy.proxy_server import (
prisma_client,
@ -1438,11 +1680,7 @@ async def cli_sso_callback(
user_api_key_cache,
)
if not key or not key.startswith("sk-"):
raise HTTPException(
status_code=400,
detail="Invalid key parameter. Must be a valid key ID starting with 'sk-'",
)
flow = _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache)
if prisma_client is None:
raise HTTPException(
@ -1480,9 +1718,6 @@ async def cli_sso_callback(
status_code=500, detail="Failed to retrieve user information from SSO"
)
# Store session info in cache (10 min TTL)
from litellm.constants import CLI_SSO_SESSION_CACHE_KEY_PREFIX
# Get all teams from user_info - CLI will let user select which one
teams: List[str] = []
if hasattr(user_info, "teams") and user_info.teams:
@ -1523,21 +1758,25 @@ async def cli_sso_callback(
"team_details": team_details,
}
cache_key = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:{key}"
user_api_key_cache.set_cache(key=cache_key, value=session_data, ttl=600)
flow["session_data"] = session_data
flow["sso_complete"] = True
browser_complete_token = secrets.token_urlsafe(32)
flow["browser_complete_token_hash"] = _hash_cli_sso_secret(
browser_complete_token
)
_set_cli_sso_flow(login_id=cast(str, key), cache=user_api_key_cache, flow=flow)
verbose_proxy_logger.info(
f"Stored CLI SSO session for user: {user_info.user_id}, teams: {teams}, num_teams: {len(teams)}"
)
# Return success page
from fastapi.responses import HTMLResponse
from litellm.proxy.common_utils.html_forms.cli_sso_success import (
render_cli_sso_success_page,
verify_url = str(request.url_for("cli_sso_complete", login_id=key))
html_content = _render_cli_sso_verification_page(
verify_url=verify_url,
browser_complete_token=browser_complete_token,
)
html_content = render_cli_sso_success_page()
return HTMLResponse(content=html_content, status_code=200)
except Exception as e:
@ -1548,7 +1787,11 @@ async def cli_sso_callback(
@router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False)
async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
async def cli_poll_key(
key_id: str,
team_id: Optional[str] = None,
x_litellm_cli_poll_secret: Optional[str] = Header(default=None),
):
"""
CLI polling endpoint - retrieves session from cache and generates JWT.
@ -1557,22 +1800,25 @@ async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
2. Second poll (with team_id): Generates JWT with selected team and deletes session
Args:
key_id: The session key ID
key_id: The CLI login session ID
team_id: Optional team ID to assign to the JWT. If provided, must be one of user's teams.
"""
from litellm.constants import CLI_SSO_SESSION_CACHE_KEY_PREFIX
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
from litellm.proxy.proxy_server import user_api_key_cache
if not key_id.startswith("sk-"):
raise HTTPException(status_code=400, detail="Invalid key ID format")
try:
# Look up session in cache
cache_key = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:{key_id}"
session_data = user_api_key_cache.get_cache(key=cache_key)
flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache)
if not _verify_cli_sso_poll_secret(
flow=flow, poll_secret=x_litellm_cli_poll_secret
):
raise HTTPException(status_code=403, detail="Invalid CLI polling secret")
if session_data:
if not flow.get("sso_complete") or not flow.get("user_code_verified"):
return {"status": "pending"}
session_data = flow.get("session_data")
if isinstance(session_data, dict):
user_teams = session_data.get("teams", [])
user_team_details = session_data.get("team_details")
user_id = session_data["user_id"]
@ -1632,7 +1878,7 @@ async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
)
# Delete cache entry (single-use)
user_api_key_cache.delete_cache(key=cache_key)
user_api_key_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id))
verbose_proxy_logger.info(
f"CLI JWT generated for user: {user_id}, team: {team_id}"
@ -1650,6 +1896,8 @@ async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
else:
return {"status": "pending"}
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.error(f"Error polling for CLI JWT: {e}")
raise HTTPException(
@ -2393,20 +2641,15 @@ class SSOAuthenticationHandler:
This is used to authenticate through the CLI login flow.
The state parameter format is: {PREFIX}:{key}:{existing_key}
- If existing_key is provided, it's included in the state
The state parameter format is: {PREFIX}:{login_id}
- The state parameter is used to pass data through the OAuth flow without changing the callback URL
"""
from litellm.constants import (
LITELLM_CLI_SESSION_TOKEN_PREFIX,
LITELLM_CLI_SOURCE_IDENTIFIER,
)
if source == LITELLM_CLI_SOURCE_IDENTIFIER and key:
if existing_key:
return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}:{existing_key}"
else:
return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}"
return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}"
else:
return None

View file

@ -41,7 +41,6 @@ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.passthrough import BasePassthroughUtils
from litellm.proxy._types import (
CommonProxyErrors,
ConfigFieldInfo,
ConfigFieldUpdate,
LiteLLMRoutes,
@ -2325,12 +2324,10 @@ async def _register_pass_through_endpoint(
dependencies = None
if auth is not None and str(auth).lower() == "true":
if premium_user is not True:
raise ValueError(
"Error Setting Authentication on Pass Through Endpoint: {}".format(
CommonProxyErrors.not_premium_user.value
)
)
# Authentication on a pass-through endpoint used to be enterprise-only.
# That left OSS with no safe configuration: auth=True raised at startup
# unless the operator had a license. The safe option must always be free,
# and unauthenticated forwarding should require explicit opt-in.
dependencies = [Depends(user_api_key_auth)]
if path not in LiteLLMRoutes.openai_routes.value:
LiteLLMRoutes.openai_routes.value.append(path)

View file

@ -1798,12 +1798,16 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
3. Reseed from authoritative DB spend (counter expired, cross-pod stale)
4. Caller-supplied fallback (DB unavailable, cold start)
"""
# 1. Try Redis first (cross-pod authoritative)
# 1. Redis first (cross-pod authoritative). On clean miss, skip
# in-memory: per-pod in-memory only has this pod's writes, so it
# would mask cross-pod increments.
redis_clean_miss = False
if spend_counter_cache.redis_cache is not None:
try:
val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key)
if val is not None:
return float(val)
redis_clean_miss = True
except Exception as e:
verbose_proxy_logger.debug(
"get_current_spend: Redis read failed for %s, falling back to in-memory: %s",
@ -1811,10 +1815,11 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
e,
)
# 2. Fall back to in-memory counter (single-instance or Redis failure)
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
if val is not None:
return float(val)
# 2. In-memory only when Redis is unreachable.
if not redis_clean_miss:
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
if val is not None:
return float(val)
# 3. Reseed from DB - fallback_spend lags cross-pod, would allow bypass.
db_spend = await SpendCounterReseed.coalesced(
@ -2147,6 +2152,7 @@ async def _increment_spend_counter_cache(counter_key: str, increment: float):
current_value = await spend_counter_cache.redis_cache.async_increment(
key=counter_key,
value=increment,
refresh_ttl=True,
)
except Exception:
await _invalidate_spend_counter(counter_key=counter_key)
@ -2160,6 +2166,7 @@ async def _increment_spend_counter_cache(counter_key: str, increment: float):
return await spend_counter_cache.async_increment_cache(
key=counter_key,
value=increment,
refresh_ttl=True,
)
@ -12598,9 +12605,20 @@ async def claim_onboarding_link(data: InvitationClaim, request: Request):
@app.get("/get_logo_url", include_in_schema=False)
def get_logo_url():
"""Get the current logo URL from environment"""
"""Get the current logo URL from environment.
Only HTTP(S) URLs are returned those are intended to be loaded
directly by the browser from a public/internal CDN. Local file
paths set via ``UI_LOGO_PATH`` are NOT returned: they are admin-
only filesystem details, the dashboard falls back to ``/get_image``
which serves the file only when it is a supported image. Without
this filter, the unauthenticated endpoint would disclose internal
hostnames or filesystem paths to any caller.
"""
logo_path = os.getenv("UI_LOGO_PATH", "")
return {"logo_url": logo_path}
if logo_path.startswith(("http://", "https://")):
return {"logo_url": logo_path}
return {"logo_url": ""}
@app.get("/get_image", include_in_schema=False)
@ -12639,61 +12657,44 @@ async def get_image():
if assets_dir != current_dir and not os.path.exists(default_logo):
default_logo = default_site_logo
cache_dir = assets_dir if os.access(assets_dir, os.W_OK) else current_dir
cache_path = os.path.join(cache_dir, "cached_logo.jpg")
logo_path = os.getenv("UI_LOGO_PATH", default_logo)
verbose_proxy_logger.debug("Reading logo from path: %s", logo_path)
# If UI_LOGO_PATH points to a local file, serve it directly (skip cache)
from litellm.proxy.common_utils.static_asset_utils import (
resolve_validated_local_image_path,
)
if logo_path != default_logo and not logo_path.startswith(("http://", "https://")):
if os.path.exists(logo_path):
return FileResponse(logo_path, media_type="image/jpeg")
# Custom path doesn't exist — fall back to default
safe_logo = resolve_validated_local_image_path(logo_path)
if safe_logo is not None:
safe_logo_path, media_type = safe_logo
return FileResponse(safe_logo_path, media_type=media_type)
verbose_proxy_logger.warning(
f"UI_LOGO_PATH '{logo_path}' does not exist, falling back to default logo"
"UI_LOGO_PATH %r is not a supported image file or does not exist, "
"falling back to default logo",
logo_path,
)
logo_path = default_logo
# [OPTIMIZATION] For HTTP URLs and default logo, check if the cached image exists
if os.path.exists(cache_path):
return FileResponse(cache_path, media_type="image/jpeg")
# Check if the logo path is an HTTP/HTTPS URL
# Remote logo URLs are loaded by the browser. The proxy should not fetch
# arbitrary admin-configured URLs server-side.
if logo_path.startswith(("http://", "https://")):
try:
# Download the image and cache it
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
return RedirectResponse(url=logo_path)
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.UI,
params={"timeout": 5.0},
)
response = await async_client.get(logo_path)
if response.status_code == 200:
# Save the image to a local file
with open(cache_path, "wb") as f:
f.write(response.content)
# Return the cached image as a FileResponse
return FileResponse(cache_path, media_type="image/jpeg")
else:
# Handle the case when the image cannot be downloaded
return FileResponse(default_logo, media_type="image/jpeg")
except Exception as e:
# Handle any exceptions during the download (e.g., timeout, connection error)
verbose_proxy_logger.debug(f"Error downloading logo from {logo_path}: {e}")
return FileResponse(default_logo, media_type="image/jpeg")
else:
# Return the local image file if the logo path is not an HTTP/HTTPS URL
return FileResponse(logo_path, media_type="image/jpeg")
# Default logo (resolved from the bundled asset, not user-controlled).
safe_logo = resolve_validated_local_image_path(logo_path)
if safe_logo is not None:
safe_logo_path, media_type = safe_logo
return FileResponse(safe_logo_path, media_type=media_type)
return FileResponse(default_site_logo, media_type="image/jpeg")
@app.get("/get_favicon", include_in_schema=False)
async def get_favicon():
"""Get custom favicon for the admin UI."""
from fastapi.responses import Response
from litellm.proxy.common_utils.static_asset_utils import (
resolve_validated_local_image_path,
)
current_dir = os.path.dirname(os.path.abspath(__file__))
default_favicon = os.path.join(current_dir, "_experimental", "out", "favicon.ico")
@ -12706,42 +12707,17 @@ async def get_favicon():
raise HTTPException(status_code=404, detail="Default favicon not found")
if favicon_url.startswith(("http://", "https://")):
try:
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.UI,
params={"timeout": 5.0},
)
response = await async_client.get(favicon_url)
if response.status_code == 200:
content_type = response.headers.get("content-type", "image/x-icon")
return Response(
content=response.content,
media_type=content_type,
)
else:
verbose_proxy_logger.warning(
"Failed to fetch favicon from %s: status %s",
favicon_url,
response.status_code,
)
if os.path.exists(default_favicon):
return FileResponse(default_favicon, media_type="image/x-icon")
raise HTTPException(status_code=404, detail="Favicon not found")
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.debug(
"Error downloading favicon from %s: %s", favicon_url, e
)
if os.path.exists(default_favicon):
return FileResponse(default_favicon, media_type="image/x-icon")
raise HTTPException(status_code=404, detail="Favicon not found")
return RedirectResponse(url=favicon_url)
else:
if os.path.exists(favicon_url):
return FileResponse(favicon_url, media_type="image/x-icon")
safe_favicon = resolve_validated_local_image_path(favicon_url)
if safe_favicon is not None:
safe_favicon_path, media_type = safe_favicon
return FileResponse(safe_favicon_path, media_type=media_type)
verbose_proxy_logger.warning(
"LITELLM_FAVICON_URL %r is not a supported image file or does not "
"exist, falling back to default favicon",
favicon_url,
)
if os.path.exists(default_favicon):
return FileResponse(default_favicon, media_type="image/x-icon")
raise HTTPException(status_code=404, detail="Favicon not found")

View file

@ -6,13 +6,19 @@ from httpx import AsyncClient
from typing import Any, Optional, List, Literal
# The proxy strips client-supplied `mock_response` unless the calling key or
# team has this admin-metadata flag set. See `_UNTRUSTED_ROOT_CONTROL_FIELDS`
# in litellm/proxy/litellm_pre_call_utils.py.
_ALLOW_CLIENT_MOCK_METADATA = {"allow_client_mock_response": True}
async def generate_key(
session, models: Optional[List[str]] = None, team_id: Optional[str] = None
):
"""Helper function to generate a key with specific model access controls"""
url = "http://0.0.0.0:4000/key/generate"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {}
data: dict = {"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA)}
if models is not None:
data["models"] = models
if team_id is not None:
@ -25,7 +31,7 @@ async def generate_team(session, models: Optional[List[str]] = None):
"""Helper function to generate a team with specific model access"""
url = "http://0.0.0.0:4000/team/new"
headers = {"Authorization": "Bearer sk-1234", "Content-Type": "application/json"}
data = {}
data: dict = {"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA)}
if models is not None:
data["models"] = models
async with session.post(url, headers=headers, json=data) as response:
@ -111,7 +117,12 @@ async def test_model_access_update():
# Create initial key with restricted access
response = await client.post(
"/key/generate", json={"models": ["openai/gpt-4"]}, headers=headers
"/key/generate",
json={
"models": ["openai/gpt-4"],
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
key_data = response.json()
@ -214,7 +225,11 @@ async def test_team_model_access_update():
# Create initial team with restricted access
response = await client.post(
"/team/new",
json={"models": ["openai/gpt-4"], "name": "test-team"},
json={
"models": ["openai/gpt-4"],
"name": "test-team",
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
@ -223,7 +238,12 @@ async def test_team_model_access_update():
# Generate a key for this team
response = await client.post(
"/key/generate", json={"team_id": team_id}, headers=headers
"/key/generate",
json={
"team_id": team_id,
"metadata": dict(_ALLOW_CLIENT_MOCK_METADATA),
},
headers=headers,
)
assert response.status_code == 200
key = response.json()["key"]

View file

@ -1,6 +1,5 @@
import os
import sys
from unittest import mock
sys.path.insert(0, os.path.abspath("../.."))
@ -26,50 +25,30 @@ async def test_get_favicon_default():
@pytest.mark.asyncio
async def test_get_favicon_with_custom_url():
"""Test that get_favicon fetches from a custom URL."""
os.environ["LITELLM_FAVICON_URL"] = "https://example.com/favicon.ico"
async def test_get_favicon_with_custom_url(monkeypatch):
"""Test that get_favicon redirects browser-loaded custom URLs."""
monkeypatch.setenv("LITELLM_FAVICON_URL", "https://example.com/favicon.ico")
mock_response = mock.Mock()
mock_response.status_code = 200
mock_response.content = b"\x00\x00\x01\x00"
mock_response.headers = {"content-type": "image/x-icon"}
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
) as ac:
response = await ac.get("/get_favicon")
try:
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.return_value = mock_response
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
) as ac:
response = await ac.get("/get_favicon")
assert response.status_code == 200
assert response.headers["content-type"] == "image/x-icon"
finally:
os.environ.pop("LITELLM_FAVICON_URL", None)
assert response.status_code == 307
assert response.headers["location"] == "https://example.com/favicon.ico"
@pytest.mark.asyncio
async def test_get_favicon_url_error_fallback():
"""Test that get_favicon falls back to default on error."""
os.environ["LITELLM_FAVICON_URL"] = "https://invalid.com/favicon.ico"
async def test_get_favicon_remote_url_is_not_server_fetched(monkeypatch):
"""Test that get_favicon does not validate remote URLs server-side."""
monkeypatch.setenv("LITELLM_FAVICON_URL", "https://invalid.com/favicon.ico")
try:
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.side_effect = httpx.ConnectError("unreachable")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
) as ac:
response = await ac.get("/get_favicon")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app),
base_url="http://testserver",
) as ac:
response = await ac.get("/get_favicon")
assert response.status_code in [200, 404]
finally:
os.environ.pop("LITELLM_FAVICON_URL", None)
assert response.status_code == 307
assert response.headers["location"] == "https://invalid.com/favicon.ico"

View file

@ -5,85 +5,48 @@ from unittest import mock
# Standard path insertion
sys.path.insert(0, os.path.abspath("../.."))
import pytest
import httpx
import pytest
from litellm.proxy.proxy_server import app
@pytest.mark.asyncio
async def test_get_image_error_handling():
async def test_get_image_redirects_remote_logo_without_server_fetch(monkeypatch):
"""
Test that get_image handles network errors gracefully and doesn't hang.
Remote logo URLs should be loaded by the browser, not fetched by the proxy.
"""
# Set an unreachable URL
os.environ["UI_LOGO_PATH"] = "http://invalid-url-12345.com/logo.jpg"
monkeypatch.setenv("UI_LOGO_PATH", "http://invalid-url-12345.com/logo.jpg")
# Clear cache
parent_dir = os.path.dirname(
os.path.dirname(
app.__file__
if hasattr(app, "__file__")
else "litellm/proxy/proxy_server.py"
)
)
cache_path = os.path.join(parent_dir, "proxy", "cached_logo.jpg")
if os.path.exists(cache_path):
os.remove(cache_path)
# Mock AsyncHTTPHandler to simulate a timeout or connection error
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.side_effect = httpx.ConnectError("Network is unreachable")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app), base_url="http://testserver"
) as ac:
response = await ac.get("/get_image")
assert response.status_code == 200
assert response.headers["content-type"] == "image/jpeg"
assert response.status_code == 307
assert response.headers["location"] == "http://invalid-url-12345.com/logo.jpg"
mock_get.assert_not_called()
@pytest.mark.asyncio
async def test_get_image_cache_logic():
async def test_get_image_remote_logo_does_not_use_stale_cache(monkeypatch, tmp_path):
"""
Test that once cached, get_image doesn't hit the network.
A stale pre-fix cache file should not mask a configured remote logo URL.
"""
os.environ["UI_LOGO_PATH"] = "http://example.com/logo.jpg"
# Clear cache
parent_dir = os.path.dirname(
os.path.dirname(
app.__file__
if hasattr(app, "__file__")
else "litellm/proxy/proxy_server.py"
)
)
cache_path = os.path.join(parent_dir, "proxy", "cached_logo.jpg")
if os.path.exists(cache_path):
os.remove(cache_path)
# Mock response
mock_response = mock.Mock()
mock_response.status_code = 200
mock_response.content = b"fake image data"
monkeypatch.setenv("UI_LOGO_PATH", "http://example.com/logo.jpg")
monkeypatch.setenv("LITELLM_ASSETS_PATH", str(tmp_path))
(tmp_path / "cached_logo.jpg").write_bytes(b"\xff\xd8\xff cached logo")
with mock.patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get"
) as mock_get:
mock_get.return_value = mock_response
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=app), base_url="http://testserver"
) as ac:
# First call - should hit download logic
response1 = await ac.get("/get_image")
assert response1.status_code == 200
assert mock_get.call_count == 1
response = await ac.get("/get_image")
# Second call - should hit cache
response2 = await ac.get("/get_image")
assert response2.status_code == 200
# If cache works, mock_get shouldn't be called again
assert mock_get.call_count == 1
assert response.status_code == 307
assert response.headers["location"] == "http://example.com/logo.jpg"
mock_get.assert_not_called()

View file

@ -50,6 +50,50 @@ async def test_redis_cache_async_increment(namespace, monkeypatch, redis_no_ping
)
@pytest.mark.asyncio
async def test_redis_cache_async_increment_refresh_ttl_true_bumps_existing_ttl(
monkeypatch, redis_no_ping
):
"""With refresh_ttl=True, every increment should call expire() to bump
the TTL, even when the key already has a TTL (counter-style use)."""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_redis_instance.__aenter__.return_value = mock_redis_instance
mock_redis_instance.__aexit__.return_value = None
mock_redis_instance.ttl.return_value = 42 # key already has ~42s left
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
await redis_cache.async_increment(
key="spend:team_member:u:t", value=0.05, refresh_ttl=True
)
mock_redis_instance.expire.assert_awaited_once_with("spend:team_member:u:t", 60)
@pytest.mark.asyncio
async def test_redis_cache_async_increment_default_does_not_bump_existing_ttl(
monkeypatch, redis_no_ping
):
"""Default (refresh_ttl=False) preserves window-style semantics: TTL is
set only on first creation, never refreshed (used by rate-limit windows)."""
monkeypatch.setenv("REDIS_HOST", "https://my-test-host")
redis_cache = RedisCache()
mock_redis_instance = AsyncMock()
mock_redis_instance.__aenter__.return_value = mock_redis_instance
mock_redis_instance.__aexit__.return_value = None
mock_redis_instance.ttl.return_value = 42 # key already has ~42s left
with patch.object(
redis_cache, "init_async_client", return_value=mock_redis_instance
):
await redis_cache.async_increment(key="rate_limit:window", value=1)
mock_redis_instance.expire.assert_not_awaited()
@pytest.mark.asyncio
async def test_redis_client_init_with_socket_timeout(monkeypatch, redis_no_ping):
monkeypatch.setenv("REDIS_HOST", "my-fake-host")

View file

@ -1,3 +1,4 @@
import asyncio
import os
import sys
from unittest.mock import AsyncMock, Mock, patch
@ -8,6 +9,8 @@ import pytest
sys.path.insert(
0, os.path.abspath("../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
BaseLLMHTTPHandler,
_google_genai_streaming_hidden_params,
@ -103,7 +106,9 @@ def test_fingerprint_agentic_tools_is_deterministic():
tools_a = {"tool_calls": [{"id": "1", "input": {"q": "abc"}, "name": "web_search"}]}
tools_b = {"tool_calls": [{"name": "web_search", "input": {"q": "abc"}, "id": "1"}]}
assert handler._fingerprint_agentic_tools(tools_a) == handler._fingerprint_agentic_tools(tools_b)
assert handler._fingerprint_agentic_tools(
tools_a
) == handler._fingerprint_agentic_tools(tools_b)
@pytest.mark.asyncio
@ -350,3 +355,70 @@ def test_google_genai_streaming_hidden_params_model_info_and_router_fallback():
response_headers=httpx.Headers({}),
)
assert from_router["model_id"] == "router-model-id"
def _build_delete_response_mock(captured: dict):
"""Returns a fake httpx delete that records its kwargs."""
def _response() -> httpx.Response:
return httpx.Response(
status_code=200,
headers={"content-type": "application/json"},
content=b'{"id": "resp_x", "object": "response", "deleted": true}',
request=httpx.Request(method="DELETE", url="https://test.openai.azure.com"),
)
async def fake_async_delete(*args, **kwargs):
captured.update(kwargs)
return _response()
def fake_sync_delete(*args, **kwargs):
captured.update(kwargs)
return _response()
return fake_async_delete, fake_sync_delete
def test_async_delete_responses_omits_body_for_azure():
"""Azure responses DELETE rejects requests with any body. Verify the handler
does not pass `json=` to httpx when the transformer returns an empty dict."""
captured: dict = {}
fake_async_delete, _ = _build_delete_response_mock(captured)
async def run():
with patch.object(AsyncHTTPHandler, "delete", new=fake_async_delete):
await litellm.adelete_responses(
response_id="resp_xyz",
custom_llm_provider="azure",
api_base="https://test.openai.azure.com",
api_key="test-key",
api_version="2025-03-01-preview",
)
asyncio.run(run())
assert "json" not in captured
assert "data" not in captured
assert captured["url"].endswith(
"/openai/responses/resp_xyz?api-version=2025-03-01-preview"
)
def test_sync_delete_responses_omits_body_for_azure():
captured: dict = {}
_, fake_sync_delete = _build_delete_response_mock(captured)
with patch.object(HTTPHandler, "delete", new=fake_sync_delete):
litellm.delete_responses(
response_id="resp_xyz",
custom_llm_provider="azure",
api_base="https://test.openai.azure.com",
api_key="test-key",
api_version="2025-03-01-preview",
)
assert "json" not in captured
assert "data" not in captured
assert captured["url"].endswith(
"/openai/responses/resp_xyz?api-version=2025-03-01-preview"
)

View file

@ -6,11 +6,12 @@ This module tests the auth commands and their associated functionality.
import pytest
import requests
from unittest.mock import AsyncMock, patch, Mock, call
from unittest.mock import patch, Mock, call
from litellm.proxy.client.cli.commands.auth import (
_normalize_teams,
_poll_for_ready_data,
_poll_for_authentication,
_start_cli_sso_flow,
)
@ -57,6 +58,18 @@ async def test_normalize_teams_with_details_with_aliases():
]
@patch("litellm.proxy.client.cli.commands.auth.requests.post")
def test_start_cli_sso_flow_rejects_invalid_response(request_mock):
"""Test CLI SSO start rejects malformed server responses"""
response = Mock()
response.raise_for_status = Mock()
response.json.return_value = {"login_id": "cli-session", "user_code": "ABCD-EFGH"}
request_mock.return_value = response
with pytest.raises(ValueError, match="Invalid CLI SSO start response"):
_start_cli_sso_flow("https://litellm.com")
@pytest.mark.asyncio
@patch(
"litellm.proxy.client.cli.commands.auth.requests.get",
@ -195,10 +208,11 @@ async def test_poll_for_ready_connection_failure(sleep_mock, click_mock, request
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
async def test_poll_for_authentication_no_data(click_mock, poll_mock, handle_mock):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual is None
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_not_called()
@ -214,10 +228,11 @@ async def test_poll_for_authentication_no_data(click_mock, poll_mock, handle_moc
@patch("litellm.proxy.client.cli.commands.auth.click.echo")
async def test_poll_for_authentication_no_teams(click_mock, poll_mock, handle_mock):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual is None
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_not_called()
@ -243,7 +258,7 @@ async def test_poll_for_authentication_team_selection_success(
click_mock, poll_mock, handle_mock
):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual == {
"api_key": "jwt-123",
"user_id": "user-123",
@ -252,11 +267,13 @@ async def test_poll_for_authentication_team_selection_success(
}
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_called_once_with(
base_url="https://litellm.com",
key_id="key-123",
poll_secret="poll-secret",
teams=[
{"team_id": "1", "team_alias": None},
{"team_id": "2", "team_alias": None},
@ -283,15 +300,17 @@ async def test_poll_for_authentication_team_selection_cancelled(
click_mock, poll_mock, handle_mock
):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual is None
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_called_once_with(
base_url="https://litellm.com",
key_id="key-123",
poll_secret="poll-secret",
teams=[{"team_id": "team-1", "team_alias": None}],
)
click_mock.assert_called_once()
@ -314,7 +333,7 @@ async def test_poll_for_authentication_auto_assigned_team(
click_mock, poll_mock, handle_mock
):
"""Test poll_for_authentication function"""
actual = _poll_for_authentication("https://litellm.com", "key-123")
actual = _poll_for_authentication("https://litellm.com", "key-123", "poll-secret")
assert actual == {
"api_key": "jwt-456",
"user_id": "user-456",
@ -323,6 +342,7 @@ async def test_poll_for_authentication_auto_assigned_team(
}
poll_mock.assert_called_once_with(
"https://litellm.com/sso/cli/poll/key-123",
headers={"x-litellm-cli-poll-secret": "poll-secret"},
pending_message="Still waiting for authentication...",
)
handle_mock.assert_not_called()

View file

@ -1,17 +1,15 @@
import json
import os
import sys
import tempfile
import time
from pathlib import Path
from unittest.mock import MagicMock, Mock, mock_open, patch
from unittest.mock import Mock, mock_open, patch
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
import pytest
from click.testing import CliRunner
from litellm.proxy.client.cli.commands.auth import (
@ -26,6 +24,22 @@ from litellm.proxy.client.cli.commands.auth import (
)
def _mock_cli_sso_start_response(
login_id: str = "cli-session-uuid-456",
poll_secret: str = "poll-secret",
user_code: str = "ABCD-EFGH",
) -> Mock:
mock_response = Mock()
mock_response.status_code = 200
mock_response.json.return_value = {
"login_id": login_id,
"poll_secret": poll_secret,
"user_code": user_code,
}
mock_response.raise_for_status = Mock()
return mock_response
class TestTokenUtilities:
"""Test token file utility functions"""
@ -243,12 +257,15 @@ class TestLoginCommand:
with (
patch("webbrowser.open") as mock_browser,
patch(
"requests.post",
return_value=_mock_cli_sso_start_response(login_id="cli-test-uuid-123"),
) as mock_post,
patch("requests.get", return_value=mock_response) as mock_get,
patch("litellm.proxy.client.cli.commands.auth.save_token") as mock_save,
patch(
"litellm.proxy.client.cli.interface.show_commands"
) as mock_show_commands,
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -261,7 +278,13 @@ class TestLoginCommand:
mock_browser.assert_called_once()
call_args = mock_browser.call_args[0][0]
assert "https://test.example.com/sso/key/generate" in call_args
assert "sk-test-uuid-123" in call_args
assert "cli-test-uuid-123" in call_args
assert "Verification code: ABCD-EFGH" in result.output
mock_post.assert_called_once()
mock_get.assert_called()
assert mock_get.call_args.kwargs["headers"] == {
"x-litellm-cli-poll-secret": "poll-secret"
}
# Verify JWT was saved
mock_save.assert_called_once()
@ -284,9 +307,9 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", return_value=mock_response),
patch("time.sleep") as mock_sleep,
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
patch("time.sleep"),
):
# Mock time.sleep to avoid actual delays in tests
@ -306,9 +329,9 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", return_value=mock_response),
patch("time.sleep"),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -325,12 +348,12 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch(
"requests.get",
side_effect=requests.RequestException("Connection failed"),
),
patch("time.sleep"),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -345,8 +368,8 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", side_effect=KeyboardInterrupt),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -369,9 +392,9 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", return_value=mock_response),
patch("time.sleep"),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -386,8 +409,8 @@ class TestLoginCommand:
with (
patch("webbrowser.open"),
patch("requests.post", return_value=_mock_cli_sso_start_response()),
patch("requests.get", side_effect=ValueError("Invalid value")),
patch("litellm._uuid.uuid.uuid4", return_value="test-uuid-123"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -556,6 +579,12 @@ class TestCLIKeyRegenerationFlow:
# Simulate user selecting team #2 (team-beta)
with (
patch("webbrowser.open") as mock_browser,
patch(
"requests.post",
return_value=_mock_cli_sso_start_response(
login_id="cli-session-uuid-456"
),
),
patch(
"requests.get", side_effect=[mock_first_response, mock_second_response]
) as mock_get,
@ -563,7 +592,6 @@ class TestCLIKeyRegenerationFlow:
patch(
"litellm.proxy.client.cli.interface.show_commands"
) as mock_show_commands,
patch("litellm._uuid.uuid.uuid4", return_value="session-uuid-456"),
patch("click.prompt", return_value="2"),
): # User selects index 2
@ -585,8 +613,11 @@ class TestCLIKeyRegenerationFlow:
# First poll should be without team_id
first_poll_url = mock_get.call_args_list[0][0][0]
assert "sk-session-uuid-456" in first_poll_url
assert "cli-session-uuid-456" in first_poll_url
assert "team_id=" not in first_poll_url
assert mock_get.call_args_list[0].kwargs["headers"] == {
"x-litellm-cli-poll-secret": "poll-secret"
}
# Second poll should include team_id=team-beta
second_poll_url = mock_get.call_args_list[1][0][0]
@ -621,10 +652,15 @@ class TestCLIKeyRegenerationFlow:
with (
patch("webbrowser.open") as mock_browser,
patch(
"requests.post",
return_value=_mock_cli_sso_start_response(
login_id="cli-session-uuid-solo"
),
),
patch("requests.get", return_value=mock_response),
patch("litellm.proxy.client.cli.commands.auth.save_token") as mock_save,
patch("litellm.proxy.client.cli.interface.show_commands"),
patch("litellm._uuid.uuid.uuid4", return_value="session-uuid-solo"),
):
result = self.runner.invoke(login, obj=mock_context.obj)
@ -637,7 +673,7 @@ class TestCLIKeyRegenerationFlow:
call_args = mock_browser.call_args[0][0]
assert "https://test.example.com/sso/key/generate" in call_args
assert "source=litellm-cli" in call_args
assert "key=sk-session-uuid-solo" in call_args
assert "key=cli-session-uuid-solo" in call_args
# Verify JWT was saved
mock_save.assert_called_once()

View file

@ -1049,3 +1049,159 @@ def test_reset_budget_windows_query_error_does_not_break_team_path(monkeypatch):
asyncio.run(job.reset_budget_windows()) # must not raise
prisma_client.db.litellm_teamtable.update.assert_awaited_once()
# ---------------------------------------------------------------------------
# Counter invalidation on budget reset
# ---------------------------------------------------------------------------
def _make_counter_invalidation_job(monkeypatch):
"""Stub spend_counter_cache so we can observe invalidation calls."""
spend_counter_cache = MagicMock()
spend_counter_cache.in_memory_cache.set_cache = MagicMock()
spend_counter_cache.redis_cache = MagicMock()
spend_counter_cache.redis_cache.async_set_cache = AsyncMock()
fake_module = types.ModuleType("litellm.proxy.proxy_server")
fake_module.spend_counter_cache = spend_counter_cache
monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", fake_module)
return spend_counter_cache
def test_reset_budget_for_team_members_invalidates_redis_counter(monkeypatch):
"""Team-member budget reset clears the Redis spend counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
expired_budget = type("B", (), {"budget_id": "budget-1"})
membership = type(
"Membership",
(),
{"user_id": "alice", "team_id": "team-x", "budget_id": "budget-1"},
)
prisma_client = MagicMock()
prisma_client.db.litellm_teammembership.find_many = AsyncMock(
return_value=[membership]
)
prisma_client.db.litellm_teammembership.update_many = AsyncMock(
return_value={"count": 1}
)
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_litellm_team_members([expired_budget]))
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:team_member:alice:team-x", value=0.0, ttl=60
)
counter_cache.redis_cache.async_set_cache.assert_any_await(
key="spend:team_member:alice:team-x", value=0.0, ttl=60
)
def test_reset_budget_for_keys_invalidates_redis_counter(
reset_budget_job, mock_prisma_client, monkeypatch
):
"""Key budget reset must clear the Redis spend counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
now = datetime.now(timezone.utc)
mock_prisma_client.data["key"] = [
type(
"Key",
(),
{
"spend": 100.0,
"budget_duration": "30d",
"budget_reset_at": now,
"id": "key-1",
"token": "sk-abc",
},
)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_keys())
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:key:sk-abc", value=0.0, ttl=60
)
def test_reset_budget_for_users_invalidates_redis_counter(
reset_budget_job, mock_prisma_client, monkeypatch
):
"""User budget reset must clear the Redis spend counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
now = datetime.now(timezone.utc)
mock_prisma_client.data["user"] = [
type(
"User",
(),
{
"spend": 50.0,
"budget_duration": "7d",
"budget_reset_at": now,
"id": "user-1",
"user_id": "alice",
},
)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_users())
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:user:alice", value=0.0, ttl=60
)
def test_reset_budget_for_teams_invalidates_redis_counter(
reset_budget_job, mock_prisma_client, monkeypatch
):
"""Team budget reset must clear the Redis spend counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
now = datetime.now(timezone.utc)
mock_prisma_client.data["team"] = [
type(
"Team",
(),
{
"spend": 200.0,
"budget_duration": "1mo",
"budget_reset_at": now,
"id": "team-1",
"team_id": "team-x",
},
)
]
asyncio.run(reset_budget_job.reset_budget_for_litellm_teams())
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:team:team-x", value=0.0, ttl=60
)
def test_reset_budget_for_keys_linked_to_budgets_invalidates_redis_counter(monkeypatch):
"""Resetting keys via budget tier must clear each linked key's counter."""
counter_cache = _make_counter_invalidation_job(monkeypatch)
expired_budget = type("B", (), {"budget_id": "budget-1"})
linked_key = type("Key", (), {"token": "sk-linked"})
prisma_client = MagicMock()
prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[linked_key]
)
prisma_client.db.litellm_verificationtoken.update_many = AsyncMock(
return_value={"count": 1}
)
job = ResetBudgetJob(proxy_logging_obj=MagicMock(), prisma_client=prisma_client)
asyncio.run(job.reset_budget_for_keys_linked_to_budgets([expired_budget]))
counter_cache.in_memory_cache.set_cache.assert_any_call(
key="spend:key:sk-linked", value=0.0, ttl=60
)

View file

@ -0,0 +1,97 @@
"""
Unit tests for unauthenticated logo / favicon endpoint helpers.
Local image paths are an existing deployment workflow, so the helper keeps
arbitrary local image paths working while refusing non-image files like
``/etc/passwd`` or ``/proc/self/environ``.
"""
import os
import sys
import pytest
sys.path.insert(0, os.path.abspath("../../../.."))
from litellm.proxy.common_utils.static_asset_utils import (
detect_local_image_media_type,
resolve_validated_local_image_path,
)
@pytest.mark.parametrize(
("body", "media_type"),
[
(b"\x89PNG\r\n\x1a\nfake png body", "image/png"),
(b"GIF89a fake gif body", "image/gif"),
(b"\xff\xd8\xff fake jpeg body", "image/jpeg"),
(b"RIFF\x00\x00\x00\x00WEBP fake webp body", "image/webp"),
(b"\x00\x00\x01\x00 fake ico body", "image/x-icon"),
],
)
def test_detect_local_image_media_type_accepts_supported_images(body, media_type):
assert detect_local_image_media_type(body) == media_type
def test_detect_local_image_media_type_rejects_non_images():
assert detect_local_image_media_type(b"root:x:0:0:root:/root:/bin/bash") is None
class TestResolveValidatedLocalImagePath:
def test_returns_resolved_path_for_arbitrary_local_image(self, tmp_path):
logo = tmp_path / "logo.png"
logo.write_bytes(b"\x89PNG\r\n\x1a\nfake png body")
result = resolve_validated_local_image_path(str(logo))
assert result == (str(logo.resolve()), "image/png")
def test_rejects_etc_passwd(self):
result = resolve_validated_local_image_path("/etc/passwd")
assert result is None
def test_rejects_proc_self_environ(self):
result = resolve_validated_local_image_path("/proc/self/environ")
assert result is None
def test_rejects_symlink_pointing_to_non_image(self, tmp_path):
secret = tmp_path / "secret.txt"
secret.write_text("password=hunter2")
symlink = tmp_path / "logo.png"
os.symlink(str(secret), str(symlink))
result = resolve_validated_local_image_path(str(symlink))
assert result is None
def test_accepts_symlink_pointing_to_image(self, tmp_path):
logo = tmp_path / "real_logo.png"
logo.write_bytes(b"\x89PNG\r\n\x1a\nfake png body")
symlink = tmp_path / "logo.png"
os.symlink(str(logo), str(symlink))
result = resolve_validated_local_image_path(str(symlink))
assert result == (str(logo.resolve()), "image/png")
def test_rejects_path_traversal_to_non_image(self, tmp_path):
assets_dir = tmp_path / "assets"
assets_dir.mkdir()
secret = tmp_path / "secret.txt"
secret.write_text("nope")
traversal = str(assets_dir / ".." / "secret.txt")
result = resolve_validated_local_image_path(traversal)
assert result is None
def test_rejects_directory(self, tmp_path):
result = resolve_validated_local_image_path(str(tmp_path))
assert result is None
def test_rejects_nonexistent_file(self, tmp_path):
result = resolve_validated_local_image_path(str(tmp_path / "missing.jpg"))
assert result is None
def test_rejects_empty_path(self):
assert resolve_validated_local_image_path("") is None

View file

@ -4,7 +4,7 @@ Test to verify the Google GenAI proxy API endpoints
"""
import os
import sys
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@ -13,520 +13,171 @@ sys.path.insert(
) # Adds the parent directory to the system path
def test_google_generate_content_endpoint():
"""Test that the google_generate_content endpoint correctly routes requests"""
# Skip this test if we can't import the required modules due to missing dependencies
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
def _build_test_client():
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
from litellm.proxy.google_endpoints.endpoints import router as google_router
app = FastAPI()
app.include_router(google_router)
return TestClient(app)
def _patch_base_process(return_value=None):
"""Patch ProxyBaseLLMRequestProcessing.base_process_llm_request so endpoint
tests don't run the full pipeline. Returns the AsyncMock so callers can
inspect call args."""
if return_value is None:
return_value = {"test": "response"}
return patch(
"litellm.proxy.google_endpoints.endpoints.ProxyBaseLLMRequestProcessing.base_process_llm_request",
new_callable=AsyncMock,
return_value=return_value,
)
def test_google_generate_content_endpoint():
"""generateContent routes through ProxyBaseLLMRequestProcessing with the
agenerate_content route_type that pipeline runs pre_call_hook +
during_call_hook + post_call_success_hook for every guardrail callback."""
try:
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router (required for FastAPI 0.120+)
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock the router's agenerate_content method
with patch("litellm.proxy.proxy_server.llm_router") as mock_router:
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Send a request to the endpoint
with _patch_base_process() as mock_base:
response = client.post(
"/v1beta/models/test-model:generateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
)
# Verify the response
assert response.status_code == 200
assert response.json() == {"test": "response"}
# Verify that agenerate_content was called
mock_router.agenerate_content.assert_called_once()
mock_base.assert_called_once()
kwargs = mock_base.call_args.kwargs
assert kwargs["route_type"] == "agenerate_content"
assert kwargs["model"] == "test-model"
def test_google_stream_generate_content_endpoint():
"""Test that the google_stream_generate_content endpoint correctly routes streaming requests"""
# Skip this test if we can't import the required modules due to missing dependencies
"""streamGenerateContent must route through the same processor with the
streaming route_type so the guardrail pipeline runs."""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router (required for FastAPI 0.120+)
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock the router's agenerate_content_stream method to return a stream
async def mock_stream_generator():
yield 'data: {"test": "stream_chunk_1"}\n\n'
yield 'data: {"test": "stream_chunk_2"}\n\n'
yield "data: [DONE]\n\n"
with patch("litellm.proxy.proxy_server.llm_router") as mock_router:
mock_router.agenerate_content_stream = AsyncMock(
return_value=mock_stream_generator()
)
# Send a request to the endpoint
with (
_patch_base_process() as mock_base,
patch(
"litellm.proxy.google_endpoints.endpoints.ProxyBaseLLMRequestProcessing.__init__",
return_value=None,
) as mock_init,
):
response = client.post(
"/v1beta/models/test-model:streamGenerateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
)
# Verify the response
assert response.status_code == 200
mock_base.assert_called_once()
kwargs = mock_base.call_args.kwargs
assert kwargs["route_type"] == "agenerate_content_stream"
assert kwargs["model"] == "test-model"
# Verify that agenerate_content_stream was called with correct parameters
mock_router.agenerate_content_stream.assert_called_once()
call_args = mock_router.agenerate_content_stream.call_args
assert call_args[1]["stream"] is True
assert call_args[1]["model"] == "test-model"
assert call_args[1]["contents"] == [
# stream=True must be forced into the data the processor receives.
init_kwargs = mock_init.call_args.kwargs
assert init_kwargs["data"]["stream"] is True
assert init_kwargs["data"]["model"] == "test-model"
assert init_kwargs["data"]["contents"] == [
{"role": "user", "parts": [{"text": "Hello"}]}
]
def test_google_generate_content_with_cost_tracking_metadata():
"""Test that the google_generate_content endpoint includes user metadata for cost tracking"""
def test_google_generate_content_data_flows_through_processor():
"""The body the client sends must reach ProxyBaseLLMRequestProcessing
intact so the pipeline can apply guardrails to it."""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.google_endpoints.endpoints import router as google_router
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router (required for FastAPI 0.120+)
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
_patch_base_process(),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
"litellm.proxy.google_endpoints.endpoints.ProxyBaseLLMRequestProcessing.__init__",
return_value=None,
) as mock_init,
):
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to return data with metadata
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
# Simulate adding user metadata
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
"user_api_key_team_id": "test-team-id",
"user_api_key": "hashed-key",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request to the endpoint
response = client.post(
client.post(
"/v1beta/models/test-model:generateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
headers={"Authorization": "Bearer sk-test-key"},
)
# Verify the response
assert response.status_code == 200
# Verify that add_litellm_data_to_request was called
mock_add_data.assert_called_once()
# Verify that agenerate_content was called with metadata
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that litellm_metadata exists and contains user information
assert "litellm_metadata" in called_data
assert called_data["litellm_metadata"]["user_api_key_user_id"] == "test-user-id"
assert called_data["litellm_metadata"]["user_api_key_team_id"] == "test-team-id"
def test_google_stream_generate_content_with_cost_tracking_metadata():
"""Test that the google_stream_generate_content endpoint includes user metadata for cost tracking"""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router (required for FastAPI 0.120+)
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock the router's agenerate_content_stream method to return a stream
mock_stream = AsyncMock()
mock_stream.__aiter__ = lambda self: mock_stream
mock_stream.__anext__.side_effect = StopAsyncIteration
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
):
mock_router.agenerate_content_stream = AsyncMock(return_value=mock_stream)
# Mock add_litellm_data_to_request to return data with metadata
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
# Simulate adding user metadata
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
"user_api_key_team_id": "test-team-id",
"user_api_key": "hashed-key",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request to the endpoint
response = client.post(
"/v1beta/models/test-model:streamGenerateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
headers={"Authorization": "Bearer sk-test-key"},
)
# Verify the response
assert response.status_code == 200
# Verify that add_litellm_data_to_request was called
mock_add_data.assert_called_once()
# Verify that agenerate_content_stream was called with metadata
mock_router.agenerate_content_stream.assert_called_once()
call_args = mock_router.agenerate_content_stream.call_args
called_data = call_args[1]
# Verify that litellm_metadata exists and contains user information
assert "litellm_metadata" in called_data
assert called_data["litellm_metadata"]["user_api_key_user_id"] == "test-user-id"
assert called_data["litellm_metadata"]["user_api_key_team_id"] == "test-team-id"
# Verify stream is set to True
assert called_data["stream"] is True
def test_google_generate_content_with_system_instruction():
"""
Test that systemInstruction is correctly passed through from the endpoint to the router.
This test verifies the fix for systemInstruction being dropped when forwarding
requests to Vertex AI through the Google GenAI endpoint.
"""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
):
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to pass through data unchanged
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
return data
mock_add_data.side_effect = mock_add_litellm_data
# Define the systemInstruction to test
system_instruction = {"parts": [{"text": "Your name is Doodle."}]}
# Send a request with systemInstruction
response = client.post(
"/v1beta/models/gemini-2.5-pro:generateContent",
json={
"systemInstruction": system_instruction,
"contents": [
{"parts": [{"text": "What is your name?"}], "role": "user"}
],
},
headers={"Authorization": "Bearer sk-test-key"},
)
# Verify the response
assert response.status_code == 200
# Verify that agenerate_content was called
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that systemInstruction is present in the call arguments
assert "systemInstruction" in called_data
assert called_data["systemInstruction"] == system_instruction
assert (
called_data["systemInstruction"]["parts"][0]["text"]
== "Your name is Doodle."
)
# Verify contents are also present
assert "contents" in called_data
assert len(called_data["contents"]) == 1
assert called_data["contents"][0]["role"] == "user"
def test_google_generate_content_with_image_config():
"""
Test that imageConfig is correctly passed through from generationConfig to the router.
This test verifies that imageConfig parameters (aspectRatio, imageSize) are preserved
when forwarding requests to Google GenAI through the endpoint.
"""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
):
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to pass through data unchanged
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request with generationConfig containing imageConfig
response = client.post(
"/v1beta/models/gemini-3-pro-image-preview:generateContent",
json={
"contents": [
{
"role": "user",
"parts": [
{
"text": "Create a vibrant infographic about photosynthesis"
}
],
}
],
"contents": [{"role": "user", "parts": [{"text": "Hello"}]}],
"systemInstruction": {"parts": [{"text": "Your name is Doodle."}]},
"generationConfig": {
"responseModalities": ["TEXT", "IMAGE"],
"imageConfig": {"aspectRatio": "9:16", "imageSize": "4K"},
},
},
headers={"Authorization": "Bearer sk-test-key"},
)
# Verify the response
assert response.status_code == 200
# Verify that agenerate_content was called
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that config is present in the call arguments
assert "config" in called_data
# Verify that imageConfig is preserved in the config
assert "imageConfig" in called_data["config"]
assert called_data["config"]["imageConfig"]["aspectRatio"] == "9:16"
assert called_data["config"]["imageConfig"]["imageSize"] == "4K"
# Verify that responseModalities is also preserved
assert "responseModalities" in called_data["config"]
assert called_data["config"]["responseModalities"] == ["TEXT", "IMAGE"]
# Verify contents are also present
assert "contents" in called_data
assert len(called_data["contents"]) == 1
assert called_data["contents"][0]["role"] == "user"
data = mock_init.call_args.kwargs["data"]
assert data["model"] == "test-model"
assert data["contents"] == [{"role": "user", "parts": [{"text": "Hello"}]}]
assert data["systemInstruction"] == {
"parts": [{"text": "Your name is Doodle."}]
}
# generationConfig arrives intact here; the rename to `config` is
# done downstream in route_request (see test_route_llm_request).
assert data["generationConfig"]["responseModalities"] == ["TEXT", "IMAGE"]
assert data["generationConfig"]["imageConfig"]["aspectRatio"] == "9:16"
def test_google_generate_content_metadata_and_trace_id_callbacks():
"""Test that google_generate_content sets litellm_call_id and logging_obj for callbacks (e.g. S3, Langfuse)"""
def test_google_generate_content_forwards_call_id_header():
"""The endpoint must forward the x-litellm-call-id header to the processor
so the helper can stamp it on the logging object. Trace continuity from
client callbacks (S3, Langfuse, etc.) depends on this header surviving
the hop through these endpoints."""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
# Create a FastAPI app and include the router
app = FastAPI()
app.include_router(google_router)
# Create a test client
client = TestClient(app)
# Mock all required proxy server dependencies
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
):
mock_router.agenerate_content = AsyncMock(return_value={"test": "response"})
# Mock add_litellm_data_to_request to return data with metadata
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
# Simulate adding user metadata
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
# Send a request to the endpoint with x-litellm-call-id header
test_call_id = "test-custom-call-id"
response = client.post(
with _patch_base_process() as mock_base:
client.post(
"/v1beta/models/test-model:generateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
headers={
"Authorization": "Bearer sk-test-key",
"x-litellm-call-id": test_call_id,
},
headers={"x-litellm-call-id": "trace-abc-123"},
)
assert response.status_code == 200
mock_router.agenerate_content.assert_called_once()
call_args = mock_router.agenerate_content.call_args
called_data = call_args[1]
# Verify that the litellm_logging_obj got assigned in the final called_data to router
assert "litellm_logging_obj" in called_data
assert "litellm_call_id" in called_data
assert called_data["litellm_call_id"] == test_call_id
forwarded_request = mock_base.call_args.kwargs["request"]
assert forwarded_request.headers.get("x-litellm-call-id") == "trace-abc-123"
def test_google_stream_generate_content_metadata_and_trace_id_callbacks():
"""Test that google_stream_generate_content sets litellm_call_id and logging_obj for callbacks"""
def test_google_count_tokens_unchanged():
"""countTokens has its own path and isn't affected by the pipeline change."""
try:
from fastapi import FastAPI
from fastapi.testclient import TestClient
from litellm.proxy.google_endpoints.endpoints import router as google_router
client = _build_test_client()
except ImportError as e:
pytest.skip(f"Skipping test due to missing dependency: {e}")
app = FastAPI()
app.include_router(google_router)
client = TestClient(app)
fake_response = MagicMock()
fake_response.original_response = {
"totalTokens": 7,
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 7}],
}
fake_response.total_tokens = 7
mock_stream = AsyncMock()
mock_stream.__aiter__ = lambda self: mock_stream
mock_stream.__anext__.side_effect = StopAsyncIteration
with (
patch("litellm.proxy.proxy_server.llm_router") as mock_router,
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config") as mock_proxy_config,
patch("litellm.proxy.proxy_server.version", "1.0.0"),
patch(
"litellm.proxy.litellm_pre_call_utils.add_litellm_data_to_request"
) as mock_add_data,
with patch(
"litellm.proxy.proxy_server.token_counter",
new_callable=AsyncMock,
return_value=fake_response,
):
mock_router.agenerate_content_stream = AsyncMock(return_value=mock_stream)
async def mock_add_litellm_data(
data, request, user_api_key_dict, proxy_config, general_settings, version
):
data["litellm_metadata"] = {
"user_api_key_user_id": "test-user-id",
}
return data
mock_add_data.side_effect = mock_add_litellm_data
test_call_id = "test-custom-stream-call-id"
response = client.post(
"/v1beta/models/test-model:streamGenerateContent",
json={"contents": [{"role": "user", "parts": [{"text": "Hello stream"}]}]},
headers={
"Authorization": "Bearer sk-test-key",
"x-litellm-call-id": test_call_id,
},
"/v1beta/models/test-model:countTokens",
json={"contents": [{"role": "user", "parts": [{"text": "Hello"}]}]},
)
assert response.status_code == 200
mock_router.agenerate_content_stream.assert_called_once()
call_args = mock_router.agenerate_content_stream.call_args
called_data = call_args[1]
assert "litellm_logging_obj" in called_data
assert "litellm_call_id" in called_data
assert called_data["litellm_call_id"] == test_call_id
body = response.json()
assert body["totalTokens"] == 7

View file

@ -2,9 +2,9 @@ import asyncio
import json
import os
import sys
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import HTTPException, Request
@ -25,7 +25,6 @@ from litellm.proxy.management_endpoints.ui_sso import (
SSOAuthenticationHandler,
_setup_team_mappings,
_sync_user_role_from_jwt_role_map,
determine_role_from_groups,
normalize_email,
process_sso_jwt_access_token,
)
@ -1471,13 +1470,13 @@ class TestAuthCallbackRouting:
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
# Test CLI state detection logic
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-test123"
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-test1234567890"
# This mimics the logic in auth_callback
if cli_state and cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"):
# Extract the key ID from the state
# Extract the login ID from the state
key_id = cli_state.split(":", 1)[1]
assert key_id == "sk-test123"
assert key_id == "cli-test1234567890"
else:
assert False, "CLI state should have been detected"
@ -1510,13 +1509,13 @@ class TestGoogleLoginCLIIntegration:
# Test the CLI state generation logic used in google_login
source = "litellm-cli"
key = "sk-test123"
key = "cli-test1234567890"
cli_state = SSOAuthenticationHandler._get_cli_state(source=source, key=key)
assert cli_state is not None
assert cli_state.startswith("litellm-session-token:")
assert "sk-test123" in cli_state
assert "cli-test1234567890" in cli_state
def test_google_login_no_cli_state_when_missing_params(self):
"""Test that google_login doesn't generate CLI state when CLI parameters are missing"""
@ -1526,8 +1525,8 @@ class TestGoogleLoginCLIIntegration:
test_cases = [
(None, None),
("litellm-cli", None),
(None, "sk-test123"),
("wrong-source", "sk-test123"),
(None, "cli-test1234567890"),
("wrong-source", "cli-test1234567890"),
]
for source, key in test_cases:
@ -1634,19 +1633,19 @@ class TestSSOStateHandling:
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source="litellm-cli", key="sk-test123"
source="litellm-cli", key="cli-test1234567890"
)
assert state is not None
assert state.startswith("litellm-session-token:")
assert "sk-test123" in state
assert "cli-test1234567890" in state
def test_get_cli_state_invalid_source(self):
"""Test generating CLI state with invalid source"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source="invalid_source", key="sk-test123"
source="invalid_source", key="cli-test1234567890"
)
assert state is None
@ -1663,40 +1662,40 @@ class TestSSOStateHandling:
"""Test generating CLI state without source"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(source=None, key="sk-test123")
state = SSOAuthenticationHandler._get_cli_state(
source=None, key="cli-test1234567890"
)
assert state is None
def test_get_cli_state_with_existing_key(self):
"""Test generating CLI state with existing_key embedded in state parameter"""
def test_get_cli_state_ignores_existing_key(self):
"""Test CLI state does not embed an existing key"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source="litellm-cli",
key="sk-new-key-123",
key="cli-new-key-1234567890",
existing_key="sk-existing-key-456",
)
assert state is not None
assert state.startswith("litellm-session-token:")
assert "sk-new-key-123" in state
assert "sk-existing-key-456" in state
# Verify the format: {PREFIX}:{key}:{existing_key}
assert state == "litellm-session-token:sk-new-key-123:sk-existing-key-456"
assert "cli-new-key-1234567890" in state
assert "sk-existing-key-456" not in state
assert state == "litellm-session-token:cli-new-key-1234567890"
def test_get_cli_state_without_existing_key(self):
"""Test generating CLI state without existing_key"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
state = SSOAuthenticationHandler._get_cli_state(
source="litellm-cli", key="sk-new-key-789", existing_key=None
source="litellm-cli", key="cli-new-key-789123456", existing_key=None
)
assert state is not None
assert state.startswith("litellm-session-token:")
assert "sk-new-key-789" in state
# Verify the format: {PREFIX}:{key} (no third part)
assert state == "litellm-session-token:sk-new-key-789"
assert "cli-new-key-789123456" in state
assert state == "litellm-session-token:cli-new-key-789123456"
assert state.count(":") == 1 # Only one colon separator
@ -1708,44 +1707,37 @@ class TestStateRouting:
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
# Test CLI state format
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-test123"
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-test1234567890"
assert cli_state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:")
# Test extraction of key from state
key_id = cli_state.split(":", 1)[1]
assert key_id == "sk-test123"
assert key_id == "cli-test1234567890"
def test_cli_state_parsing_with_existing_key(self):
"""Test parsing CLI state with existing_key embedded"""
def test_cli_state_parsing_uses_single_login_id(self):
"""Test parsing CLI state with a single login ID"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
# State format: {PREFIX}:{key}:{existing_key}
cli_state = (
f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-key-456:sk-existing-key-789"
)
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-new-key-456123"
# Parse as done in auth_callback
state_parts = cli_state.split(":", 2) # Split into max 3 parts
state_parts = cli_state.split(":", 1)
key_id = state_parts[1] if len(state_parts) > 1 else None
existing_key = state_parts[2] if len(state_parts) > 2 else None
assert key_id == "sk-new-key-456"
assert existing_key == "sk-existing-key-789"
assert key_id == "cli-new-key-456123"
def test_cli_state_parsing_without_existing_key(self):
"""Test parsing CLI state without existing_key"""
def test_cli_state_parsing_without_extra_segments(self):
"""Test parsing CLI state uses a single login ID"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
# State format: {PREFIX}:{key}
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-key-999"
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-new-key-999123"
# Parse as done in auth_callback
state_parts = cli_state.split(":", 2) # Split into max 3 parts
state_parts = cli_state.split(":", 1)
key_id = state_parts[1] if len(state_parts) > 1 else None
existing_key = state_parts[2] if len(state_parts) > 2 else None
assert key_id == "sk-new-key-999"
assert existing_key is None
assert key_id == "cli-new-key-999123"
def test_non_cli_state_detection(self):
"""Test detection of non-CLI state parameters"""
@ -2007,6 +1999,178 @@ class TestCustomUISSO:
class TestCLIKeyRegenerationFlow:
"""Test the end-to-end CLI key regeneration flow"""
def test_cli_sso_login_id_validation_restricts_charset(self):
"""Test CLI SSO login IDs only allow the generated character set"""
from litellm.proxy.management_endpoints.ui_sso import (
_is_valid_cli_sso_login_id,
)
assert _is_valid_cli_sso_login_id("cli-test_1234567890")
assert not _is_valid_cli_sso_login_id("cli-session")
assert not _is_valid_cli_sso_login_id("cli-test\n1234567890")
assert not _is_valid_cli_sso_login_id("cli-test\x001234567890")
assert not _is_valid_cli_sso_login_id("sk-test1234567890")
@pytest.mark.asyncio
async def test_cli_sso_start_creates_bound_flow(self):
"""Test CLI SSO start creates a polling secret bound flow"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
_normalize_cli_sso_user_code,
cli_sso_start,
)
mock_request = MagicMock(spec=Request)
mock_request.client = SimpleNamespace(host="127.0.0.1")
mock_request.headers = {}
mock_cache = MagicMock()
mock_cache.increment_cache.return_value = 1
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
result = await cli_sso_start(request=mock_request)
assert result["login_id"].startswith("cli-")
assert result["poll_secret"]
assert result["user_code"]
mock_cache.increment_cache.assert_called_once()
assert mock_cache.increment_cache.call_args.kwargs["ttl"] == 60
mock_cache.set_cache.assert_called_once()
flow_data = mock_cache.set_cache.call_args.kwargs["value"]
assert flow_data["poll_secret_hash"] == _hash_cli_sso_secret(
result["poll_secret"]
)
assert flow_data["user_code_hash"] == _hash_cli_sso_secret(
_normalize_cli_sso_user_code(result["user_code"])
)
assert flow_data["poll_secret_hash"] != result["poll_secret"]
assert flow_data["user_code_hash"] != result["user_code"]
@pytest.mark.asyncio
async def test_cli_sso_start_rate_limits_by_client_ip(self):
"""Test CLI SSO start enforces a coarse per-client rate limit"""
from litellm.proxy.management_endpoints.ui_sso import cli_sso_start
mock_request = MagicMock(spec=Request)
mock_request.client = SimpleNamespace(host="127.0.0.1")
mock_request.headers = {}
mock_cache = MagicMock()
mock_cache.increment_cache.return_value = 31
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
with pytest.raises(HTTPException) as exc_info:
await cli_sso_start(request=mock_request)
assert exc_info.value.status_code == 429
mock_cache.set_cache.assert_not_called()
@pytest.mark.asyncio
async def test_cli_sso_complete_verifies_user_code(self):
"""Test CLI SSO complete marks a session as verified"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
_normalize_cli_sso_user_code,
cli_sso_complete,
)
mock_request = MagicMock(spec=Request)
mock_request.body = AsyncMock(
return_value=b"user_code=ABCD-EFGH&browser_complete_token=browser-token"
)
mock_cache = MagicMock()
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"user_code_hash": _hash_cli_sso_secret(
_normalize_cli_sso_user_code("ABCD-EFGH")
),
"browser_complete_token_hash": _hash_cli_sso_secret("browser-token"),
"sso_complete": True,
"user_code_verified": False,
"session_data": {"user_id": "test-user-123"},
}
with (
patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache),
patch(
"litellm.proxy.common_utils.html_forms.cli_sso_success.render_cli_sso_success_page",
return_value="<html>Success</html>",
),
):
result = await cli_sso_complete(
request=mock_request, login_id="cli-session-4567890"
)
assert result.status_code == 200
flow_data = mock_cache.set_cache.call_args.kwargs["value"]
assert flow_data["user_code_verified"] is True
@pytest.mark.asyncio
async def test_cli_sso_complete_requires_callback_token(self):
"""Test CLI SSO complete requires the callback-delivered token"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
_normalize_cli_sso_user_code,
cli_sso_complete,
)
mock_request = MagicMock(spec=Request)
mock_request.body = AsyncMock(return_value=b"user_code=ABCD-EFGH")
mock_cache = MagicMock()
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"user_code_hash": _hash_cli_sso_secret(
_normalize_cli_sso_user_code("ABCD-EFGH")
),
"browser_complete_token_hash": _hash_cli_sso_secret("browser-token"),
"sso_complete": True,
"user_code_verified": False,
"session_data": {"user_id": "test-user-123"},
}
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
with pytest.raises(HTTPException) as exc_info:
await cli_sso_complete(
request=mock_request, login_id="cli-session-4567890"
)
assert exc_info.value.status_code == 400
mock_cache.set_cache.assert_not_called()
@pytest.mark.asyncio
async def test_cli_sso_complete_waits_for_callback_before_token_checks(self):
"""Test CLI SSO complete returns not-ready before verification checks"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
_normalize_cli_sso_user_code,
cli_sso_complete,
)
mock_request = MagicMock(spec=Request)
mock_request.body = AsyncMock(
return_value=b"user_code=ABCD-EFGH&browser_complete_token=browser-token"
)
mock_cache = MagicMock()
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"user_code_hash": _hash_cli_sso_secret(
_normalize_cli_sso_user_code("ABCD-EFGH")
),
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
with pytest.raises(HTTPException) as exc_info:
await cli_sso_complete(
request=mock_request, login_id="cli-session-4567890"
)
assert exc_info.value.status_code == 400
assert exc_info.value.detail == "CLI login is not ready"
mock_request.body.assert_not_awaited()
mock_cache.set_cache.assert_not_called()
@pytest.mark.asyncio
async def test_cli_sso_callback_stores_session(self):
"""Test CLI SSO callback stores session data in cache for JWT generation"""
@ -2017,7 +2181,7 @@ class TestCLIKeyRegenerationFlow:
mock_request = MagicMock(spec=Request)
# Test data
session_key = "sk-session-456"
session_key = "cli-session-4567890"
# Mock user info
mock_user_info = LiteLLM_UserTable(
@ -2032,6 +2196,16 @@ class TestCLIKeyRegenerationFlow:
# Mock cache
mock_cache = MagicMock()
mock_cache.get_cache.return_value = {
"poll_secret_hash": "poll-secret-hash",
"user_code_hash": "user-code-hash",
"sso_complete": False,
"user_code_verified": False,
"session_data": None,
}
mock_request.url_for.return_value = (
"https://test.example.com/sso/cli/complete/cli-session-4567890"
)
with (
patch(
@ -2049,7 +2223,6 @@ class TestCLIKeyRegenerationFlow:
result = await cli_sso_callback(
request=mock_request,
key=session_key,
existing_key=None,
result=mock_sso_result,
)
@ -2062,14 +2235,18 @@ class TestCLIKeyRegenerationFlow:
assert session_key in call_args.kwargs["key"]
# Verify session data structure
session_data = call_args.kwargs["value"]
flow_data = call_args.kwargs["value"]
session_data = flow_data["session_data"]
assert flow_data["sso_complete"] is True
assert flow_data["user_code_verified"] is False
assert isinstance(flow_data["browser_complete_token_hash"], str)
assert session_data["user_id"] == "test-user-123"
assert session_data["user_role"] == "internal_user"
assert session_data["teams"] == ["team1", "team2"]
assert session_data["models"] == ["gpt-4"]
# Verify TTL
assert call_args.kwargs["ttl"] == 600 # 10 minutes
assert call_args.kwargs["ttl"] == 600
assert result.status_code == 200
# Verify response contains success message (response is HTML)
@ -2078,10 +2255,13 @@ class TestCLIKeyRegenerationFlow:
@pytest.mark.asyncio
async def test_cli_poll_key_returns_teams_for_selection(self):
"""Test CLI poll endpoint returns teams for user selection when multiple teams exist"""
from litellm.proxy.management_endpoints.ui_sso import cli_poll_key
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
# Test data
session_key = "sk-session-789"
session_key = "cli-session-789123"
session_data = {
"user_id": "test-user-456",
"user_role": "internal_user",
@ -2091,11 +2271,20 @@ class TestCLIKeyRegenerationFlow:
# Mock cache
mock_cache = MagicMock()
mock_cache.get_cache.return_value = session_data
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
# Act - First poll without team_id
result = await cli_poll_key(key_id=session_key, team_id=None)
result = await cli_poll_key(
key_id=session_key,
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
# Assert - should return teams list for selection
assert result["status"] == "ready"
@ -2108,16 +2297,72 @@ class TestCLIKeyRegenerationFlow:
mock_cache.delete_cache.assert_not_called()
@pytest.mark.asyncio
async def test_auth_callback_routes_to_cli_with_existing_key(self):
"""Test that auth_callback properly routes CLI requests and extracts existing_key from state parameter"""
async def test_cli_poll_key_requires_poll_secret(self):
"""Test CLI poll endpoint rejects callers without the polling secret"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
mock_cache = MagicMock()
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": {
"user_id": "test-user-456",
"user_role": "internal_user",
"teams": [],
"models": ["gpt-4"],
},
}
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
with pytest.raises(HTTPException) as exc_info:
await cli_poll_key(key_id="cli-session-789123", team_id=None)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_cli_poll_key_waits_for_user_code_verification(self):
"""Test CLI poll endpoint stays pending until user code verification"""
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
mock_cache = MagicMock()
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": False,
"session_data": {
"user_id": "test-user-456",
"user_role": "internal_user",
"teams": [],
"models": ["gpt-4"],
},
}
with patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache):
result = await cli_poll_key(
key_id="cli-session-789123",
team_id=None,
x_litellm_cli_poll_secret="poll-secret",
)
assert result == {"status": "pending"}
@pytest.mark.asyncio
async def test_auth_callback_routes_to_cli(self):
"""Test that auth_callback properly routes CLI requests"""
from litellm.constants import LITELLM_CLI_SESSION_TOKEN_PREFIX
from litellm.proxy.management_endpoints.ui_sso import auth_callback
# Mock request (no query params needed - existing_key is in state)
# Mock request
mock_request = MagicMock(spec=Request)
# CLI state with existing_key embedded: {PREFIX}:{key}:{existing_key}
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:sk-new-session-key-456:sk-existing-cli-key-123"
cli_state = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-new-session-key-456"
# Mock the CLI callback and required proxy server components
mock_result = {"user_id": "test-user", "email": "test@example.com"}
@ -2142,16 +2387,14 @@ class TestCLIKeyRegenerationFlow:
# Act
await auth_callback(request=mock_request, state=cli_state)
# Assert - existing_key should be extracted from state parameter
mock_cli_callback.assert_called_once_with(
request=mock_request,
key="sk-new-session-key-456",
existing_key="sk-existing-cli-key-123",
key="cli-new-session-key-456",
result=mock_result,
)
def test_get_redirect_url_does_not_include_existing_key_in_url(self):
"""Test that redirect URL generation does NOT include existing_key in URL (uses state parameter instead)"""
"""Test that redirect URL generation does NOT include existing_key in URL"""
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Mock request
@ -2194,10 +2437,13 @@ class TestCLIKeyRegenerationFlow:
async def test_cli_poll_key_generates_jwt_with_team(self):
"""Test CLI poll endpoint generates JWT when team_id is provided"""
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.management_endpoints.ui_sso import cli_poll_key
from litellm.proxy.management_endpoints.ui_sso import (
_hash_cli_sso_secret,
cli_poll_key,
)
# Test data
session_key = "sk-session-999"
session_key = "cli-session-999123"
selected_team = "team-b"
session_data = {
"user_id": "test-user-789",
@ -2217,7 +2463,12 @@ class TestCLIKeyRegenerationFlow:
# Mock cache
mock_cache = MagicMock()
mock_cache.get_cache.return_value = session_data
mock_cache.get_cache.return_value = {
"poll_secret_hash": _hash_cli_sso_secret("poll-secret"),
"sso_complete": True,
"user_code_verified": True,
"session_data": session_data,
}
mock_jwt_token = "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.token"
@ -2235,7 +2486,11 @@ class TestCLIKeyRegenerationFlow:
)
# Act - Second poll with team_id
result = await cli_poll_key(key_id=session_key, team_id=selected_team)
result = await cli_poll_key(
key_id=session_key,
team_id=selected_team,
x_litellm_cli_poll_secret="poll-secret",
)
# Assert - should return JWT
assert result["status"] == "ready"
@ -2901,7 +3156,7 @@ class TestGetGenericSSORedirectParams:
from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler
# Arrange
cli_state = "litellm-session-token:sk-test123"
cli_state = "litellm-session-token:cli-test1234567890"
with patch.dict(os.environ, {"GENERIC_CLIENT_STATE": "env_state_value"}):
# Act

View file

@ -0,0 +1,136 @@
"""
Regression tests for the pass-through endpoint auth-default fix
(GHSA-7h34-mmrh-6g58).
Two failures the fix closes:
1. ``PassThroughGenericEndpoint.auth`` defaulted to ``False`` an
admin who added a pass-through to ``general_settings`` without
explicitly setting ``auth: true`` shipped an unauthenticated
forwarder.
2. Setting ``auth: true`` was rejected at startup unless the operator
had a LiteLLM Enterprise license, leaving OSS deployments with no
safe configuration.
The fix flips the default to ``True`` (safe-by-default) and removes
the enterprise gate so OSS operators can register an authenticated
pass-through. The runtime check in ``user_api_key_auth.py`` also now
defaults to ``True`` so a config dict (raw, not Pydantic) without an
``auth`` key still requires authentication.
"""
import os
import sys
from unittest.mock import AsyncMock, MagicMock
import pytest
from fastapi import FastAPI
sys.path.insert(0, os.path.abspath("../../../.."))
from litellm.proxy._types import PassThroughGenericEndpoint
from litellm.proxy.auth.user_api_key_auth import (
check_api_key_for_custom_headers_or_pass_through_endpoints,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
_register_pass_through_endpoint,
)
def test_passthrough_auth_defaults_to_true():
# Regression: an admin who configures a pass-through without setting
# auth explicitly used to ship an unauthenticated forwarder. The
# default is now safe.
endpoint = PassThroughGenericEndpoint(
path="/canary-forwarder",
target="https://postman-echo.com/get",
)
assert endpoint.auth is True
def test_passthrough_auth_can_still_be_explicitly_disabled():
# Operators who genuinely need an unauthenticated forwarder (e.g.
# public webhook receiver) can opt in explicitly.
endpoint = PassThroughGenericEndpoint(
path="/public-webhook",
target="https://example.com/webhook",
auth=False,
)
assert endpoint.auth is False
@pytest.mark.asyncio
async def test_register_passthrough_with_auth_true_works_for_oss(monkeypatch):
# Regression: setting ``auth: true`` used to raise at startup
# unless ``premium_user`` was True, leaving OSS with no safe
# configuration.
app = MagicMock(spec=FastAPI)
visited: set = set()
endpoint = PassThroughGenericEndpoint(
path="/forwarder",
target="https://example.com",
auth=True,
)
# Should not raise; OSS premium_user=False is allowed to use auth=True.
await _register_pass_through_endpoint(
endpoint=endpoint,
app=app,
premium_user=False,
visited_endpoints=visited,
)
@pytest.mark.asyncio
async def test_runtime_check_treats_missing_auth_key_as_authenticated():
# The runtime dispatch in user_api_key_auth pulls
# pass_through_endpoints from general_settings as raw dicts (the
# Pydantic default never applies). A dict without an ``auth`` key
# must default to "authenticated" — without this, the previous
# behaviour (``endpoint.get("auth") is not True`` -> True -> empty
# auth) ships an unauthenticated forwarder.
request = MagicMock()
request.headers = {}
raw_endpoint_no_auth_key = {
"path": "/forwarder",
"target": "https://example.com",
# ``auth`` deliberately omitted
}
result = await check_api_key_for_custom_headers_or_pass_through_endpoints(
request=request,
route="/forwarder",
pass_through_endpoints=[raw_endpoint_no_auth_key],
api_key="sk-1234",
)
# Result is the api_key string (auth is REQUIRED for this endpoint
# — flow continues to normal key validation), NOT an empty
# ``UserAPIKeyAuth()`` (which was the unauthenticated-forwarder
# bug).
assert result == "sk-1234"
@pytest.mark.asyncio
async def test_runtime_check_explicit_auth_false_still_skips_validation():
# Operators who explicitly set ``auth: False`` get the legacy
# behaviour — an empty UserAPIKeyAuth, no key required.
from litellm.proxy._types import UserAPIKeyAuth
request = MagicMock()
request.headers = {}
raw_endpoint_auth_false = {
"path": "/public-webhook",
"target": "https://example.com",
"auth": False,
}
result = await check_api_key_for_custom_headers_or_pass_through_endpoints(
request=request,
route="/public-webhook",
pass_through_endpoints=[raw_endpoint_auth_false],
api_key="",
)
assert isinstance(result, UserAPIKeyAuth)

View file

@ -457,6 +457,59 @@ def test_fallback_login_has_no_deprecation_banner(client_no_auth):
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(
@ -3980,7 +4033,7 @@ async def test_get_image_root_case_uses_current_dir(monkeypatch):
@pytest.mark.asyncio
async def test_get_image_custom_local_logo_bypasses_cache(monkeypatch):
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.
@ -3989,11 +4042,11 @@ async def test_get_image_custom_local_logo_bypasses_cache(monkeypatch):
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 unittest.mock import patch
from litellm.proxy.proxy_server import get_image
monkeypatch.setenv("UI_LOGO_PATH", "/app/custom_logo.jpg")
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)
@ -4004,8 +4057,6 @@ async def test_get_image_custom_local_logo_bypasses_cache(monkeypatch):
return MagicMock()
with (
patch("litellm.proxy.proxy_server.os.path.exists", return_value=True),
patch("litellm.proxy.proxy_server.os.access", return_value=True),
patch(
"litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response
),
@ -4015,25 +4066,27 @@ async def test_get_image_custom_local_logo_bypasses_cache(monkeypatch):
assert (
len(calls_to_file_response) == 1
), "FileResponse should be called exactly once"
assert calls_to_file_response[0] == "/app/custom_logo.jpg", (
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_still_uses_cache(monkeypatch):
async def test_get_image_default_logo_ignores_stale_cache(monkeypatch, tmp_path):
"""
Test that when UI_LOGO_PATH is NOT set (default logo), the cache
optimization still works cached_logo.jpg is returned if it exists.
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.delenv("LITELLM_ASSETS_PATH", raising=False)
monkeypatch.setenv("LITELLM_ASSETS_PATH", str(tmp_path))
calls_to_file_response = []
@ -4042,8 +4095,6 @@ async def test_get_image_default_logo_still_uses_cache(monkeypatch):
return MagicMock()
with (
patch("litellm.proxy.proxy_server.os.path.exists", return_value=True),
patch("litellm.proxy.proxy_server.os.access", return_value=True),
patch(
"litellm.proxy.proxy_server.FileResponse", side_effect=fake_file_response
),
@ -4054,24 +4105,26 @@ async def test_get_image_default_logo_still_uses_cache(monkeypatch):
len(calls_to_file_response) == 1
), "FileResponse should be called exactly once"
served_path = calls_to_file_response[0]
assert served_path.endswith(
"cached_logo.jpg"
), f"Expected cached_logo.jpg for default logo, got {served_path}"
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):
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 cache/default logo instead of failing.
get_image falls through to the default logo instead of failing.
"""
from unittest.mock import patch
from litellm.proxy.proxy_server import get_image
monkeypatch.setenv("UI_LOGO_PATH", "/app/nonexistent_logo.jpg")
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.delenv("LITELLM_ASSETS_PATH", raising=False)
monkeypatch.setenv("LITELLM_ASSETS_PATH", str(tmp_path))
calls_to_file_response = []
@ -4079,17 +4132,7 @@ async def test_get_image_custom_logo_missing_falls_through_to_default(monkeypatc
calls_to_file_response.append(path)
return MagicMock()
def exists_side_effect(path):
# The custom logo does NOT exist; cache and default DO exist
if path == "/app/nonexistent_logo.jpg":
return False
return True
with (
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.FileResponse", side_effect=fake_file_response
),
@ -4100,28 +4143,29 @@ async def test_get_image_custom_logo_missing_falls_through_to_default(monkeypatc
len(calls_to_file_response) == 1
), "FileResponse should be called exactly once"
served_path = calls_to_file_response[0]
assert (
served_path != "/app/nonexistent_logo.jpg"
assert served_path != str(
custom_logo_path
), "Should not attempt to serve a non-existent custom logo"
assert served_path.endswith(
"cached_logo.jpg"
), f"Expected fallback to cached_logo.jpg, got {served_path}"
assert served_path.endswith("logo.jpg")
@pytest.mark.asyncio
async def test_get_image_custom_logo_missing_no_cache_serves_default(monkeypatch):
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.
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
monkeypatch.setenv("UI_LOGO_PATH", "/app/nonexistent_logo.jpg")
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.delenv("LITELLM_ASSETS_PATH", raising=False)
monkeypatch.setenv("LITELLM_ASSETS_PATH", str(tmp_path))
calls_to_file_response = []
@ -4129,19 +4173,7 @@ async def test_get_image_custom_logo_missing_no_cache_serves_default(monkeypatch
calls_to_file_response.append(path)
return MagicMock()
def exists_side_effect(path):
# Neither the custom logo nor the cache exist
if path == "/app/nonexistent_logo.jpg":
return False
if "cached_logo.jpg" in path:
return False
return True
with (
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.FileResponse", side_effect=fake_file_response
),
@ -4152,8 +4184,8 @@ async def test_get_image_custom_logo_missing_no_cache_serves_default(monkeypatch
len(calls_to_file_response) == 1
), "FileResponse should be called exactly once"
served_path = calls_to_file_response[0]
assert (
served_path != "/app/nonexistent_logo.jpg"
assert served_path != str(
custom_logo_path
), "Should not attempt to serve a non-existent custom logo"
assert served_path.endswith(
"logo.jpg"
@ -5625,39 +5657,34 @@ class TestLazyFeaturesNotImportedAtStartup:
"""
def test_heavy_modules_absent_at_startup(self):
# Force a fresh `proxy_server` import in a subprocess so other tests
# in this run (which may have triggered lazy loads via the TestClient)
# don't pollute the result.
import subprocess
# 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
check = (
"import sys; "
"from litellm.proxy.proxy_server import app; " # noqa: F401
"heavy = ["
"'litellm.proxy._experimental.mcp_server.rest_endpoints',"
"'litellm.proxy._experimental.mcp_server.server',"
"'litellm.proxy.management_endpoints.config_override_endpoints',"
"'litellm.proxy.guardrails.guardrail_endpoints',"
"'litellm.proxy.openai_evals_endpoints.endpoints',"
"]; "
"still_present = [m for m in heavy if m in sys.modules]; "
"print('PRESENT_AT_STARTUP:', still_present)"
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}"
)
result = subprocess.run(
[sys.executable, "-c", check],
capture_output=True,
text=True,
timeout=120,
)
# Last non-empty line of stdout (skip warnings printed before)
out_lines = [
line for line in result.stdout.strip().splitlines() if line.strip()
]
report = next((line for line in out_lines if "PRESENT_AT_STARTUP" in line), "")
assert report, f"no report emitted (stderr: {result.stderr[-500:]})"
assert (
"PRESENT_AT_STARTUP: []" in report
), f"expected no heavy modules at startup, got: {report}"
class TestLazyFeatureMiddleware:
@ -5835,3 +5862,90 @@ class TestLazyFeatureMiddleware:
assert attempts == [
"called"
], f"failing register_fn should be invoked once, not on every request; got {attempts}"
@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} "
f"(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

View file

@ -239,3 +239,55 @@ async def test_route_request_with_router_settings_override_preserves_existing():
assert call_kwargs["num_retries"] == 10
# Key/team timeout should be applied since not in request
assert call_kwargs["timeout"] == 30
@pytest.mark.parametrize(
"route_type", ["agenerate_content", "agenerate_content_stream"]
)
@pytest.mark.asyncio
async def test_route_request_maps_generation_config_for_google_routes(route_type):
"""For Google generate_content routes, route_request must rename
`generationConfig` (Google's wire format) to `config` (the kwarg the
router method expects). Without this mapping the request reaches the
LLM with the field under the wrong name and the config is dropped."""
data = {
"model": "gemini-2.5-flash",
"contents": [{"role": "user", "parts": [{"text": "Hello"}]}],
"generationConfig": {
"responseModalities": ["TEXT", "IMAGE"],
"imageConfig": {"aspectRatio": "9:16", "imageSize": "4K"},
},
}
llm_router = MagicMock()
getattr(llm_router, route_type).return_value = "ok"
await route_request(data, llm_router, None, route_type)
call_kwargs = getattr(llm_router, route_type).call_args[1]
assert "generationConfig" not in call_kwargs
assert "config" in call_kwargs
assert call_kwargs["config"]["responseModalities"] == ["TEXT", "IMAGE"]
assert call_kwargs["config"]["imageConfig"]["aspectRatio"] == "9:16"
assert call_kwargs["config"]["imageConfig"]["imageSize"] == "4K"
@pytest.mark.parametrize(
"route_type", ["agenerate_content", "agenerate_content_stream"]
)
@pytest.mark.asyncio
async def test_route_request_preserves_existing_config_for_google_routes(route_type):
"""If the caller already supplies `config`, route_request must not
overwrite it with `generationConfig`."""
data = {
"model": "gemini-2.5-flash",
"contents": [{"role": "user", "parts": [{"text": "Hello"}]}],
"config": {"existing": True},
"generationConfig": {"shouldNotWin": True},
}
llm_router = MagicMock()
getattr(llm_router, route_type).return_value = "ok"
await route_request(data, llm_router, None, route_type)
call_kwargs = getattr(llm_router, route_type).call_args[1]
assert call_kwargs["config"] == {"existing": True}