Merge remote-tracking branch 'upstream/litellm_internal_staging' into litellm_project_rate_limiting

This commit is contained in:
shivam 2026-04-18 13:26:24 -07:00
commit dbf4f9637f
No known key found for this signature in database
37 changed files with 1472 additions and 127 deletions

View file

@ -110,7 +110,7 @@ LiteLLM is a unified interface for 100+ LLM providers with two main components:
- When wiring a new UI entity type to an existing backend endpoint, verify the backend API contract (single value vs. array, required vs. optional params) and ensure the UI controls match — e.g., use a single-select dropdown when the backend accepts a single value, not a multi-select
### UI Component Library
- **Always use `antd` for new UI components** — we are migrating off of `@tremor/react`. Do not introduce new `Badge`, `Text`, `Card`, `Grid`, `Title`, or other imports from `@tremor/react` in any new or modified file. Use `antd` equivalents: `Tag` for labels, plain `<span>`/`<div>` with Tailwind classes (or `Typography.Text`) for text, `Card` from `antd`, etc. Note that `antd` has no `"yellow"` Tag color — use `"gold"` for amber/yellow.
- **Always use `antd` for new UI components** — we are migrating off of `@tremor/react`. Do not introduce new `Badge`, `Text`, `Card`, `Grid`, `Title`, or other imports from `@tremor/react` in any new or modified file. Use `antd` equivalents: `Tag` for labels, `Typography.Text` / `Typography.Title` / `Typography.Paragraph` for textual content (avoid plain text-only `<span>`, `<p>`, `<h*>` when Typography fits), and `Card` from `antd`. Note that `antd` has no `"yellow"` Tag color — use `"gold"` for amber/yellow.
### MCP OAuth / OpenAPI Transport Mapping
- `TRANSPORT.OPENAPI` is a UI-only concept. The backend only accepts `"http"`, `"sse"`, or `"stdio"`. Always map it to `"http"` before any API call (including pre-OAuth temp-session calls).

View file

@ -47,7 +47,7 @@ spec:
{{- toYaml .Values.podSecurityContext | nindent 8 }}
{{- with .Values.extraInitContainers }}
initContainers:
{{- toYaml . | nindent 8 }}
{{- tpl (toYaml .) $ | nindent 8 }}
{{- end }}
containers:
- name: {{ include "litellm.name" . }}
@ -212,7 +212,7 @@ spec:
{{- toYaml . | nindent 12 }}
{{- end }}
{{- with .Values.extraContainers }}
{{- toYaml . | nindent 8 }}
{{- tpl (toYaml .) $ | nindent 8 }}
{{- end }}
volumes:
{{ if .Values.securityContext.readOnlyRootFilesystem }}

View file

@ -37,7 +37,7 @@ spec:
serviceAccountName: {{ include "litellm.migrationServiceAccountName" . }}
{{- with .Values.migrationJob.extraInitContainers }}
initContainers:
{{- toYaml . | nindent 8 }}
{{- tpl (toYaml .) $ | nindent 8 }}
{{- end }}
containers:
- name: prisma-migrations
@ -96,7 +96,7 @@ spec:
{{- toYaml . | nindent 12 }}
{{- end }}
{{- with .Values.migrationJob.extraContainers }}
{{- toYaml . | nindent 8 }}
{{- tpl (toYaml .) $ | nindent 8 }}
{{- end }}
{{- with .Values.volumes }}
volumes:

View file

@ -319,3 +319,61 @@ tests:
asserts:
- notExists:
path: spec.minReadySeconds
- it: should work with extraInitContainers
template: deployment.yaml
set:
extraInitContainers:
- name: init-test
image: busybox:latest
command: ["echo", "hello"]
asserts:
- contains:
path: spec.template.spec.initContainers
content:
name: init-test
image: busybox:latest
command: ["echo", "hello"]
- it: should support tpl in extraInitContainers
template: deployment.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
tag: test
extraInitContainers:
- name: init-tpl
image: "{{ .Values.image.repository }}:{{ .Values.image.tag }}"
command: ["echo", "hello"]
asserts:
- contains:
path: spec.template.spec.initContainers
content:
name: init-tpl
image: "ghcr.io/berriai/litellm-database:test"
command: ["echo", "hello"]
- it: should work with extraContainers
template: deployment.yaml
set:
extraContainers:
- name: sidecar
image: busybox:latest
asserts:
- contains:
path: spec.template.spec.containers
content:
name: sidecar
image: busybox:latest
- it: should support tpl in extraContainers
template: deployment.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
tag: test
extraContainers:
- name: sidecar-tpl
image: "{{ .Values.image.repository }}:{{ .Values.image.tag }}"
asserts:
- contains:
path: spec.template.spec.containers
content:
name: sidecar-tpl
image: "ghcr.io/berriai/litellm-database:test"

View file

@ -188,3 +188,69 @@ tests:
- equal:
path: spec.template.spec.serviceAccountName
value: pre-existing-sa
- it: should work with extraInitContainers
template: migrations-job.yaml
set:
migrationJob:
enabled: true
extraInitContainers:
- name: init-test
image: busybox:latest
command: ["echo", "hello"]
asserts:
- contains:
path: spec.template.spec.initContainers
content:
name: init-test
image: busybox:latest
command: ["echo", "hello"]
- it: should support tpl in extraInitContainers
template: migrations-job.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
tag: test
migrationJob:
enabled: true
extraInitContainers:
- name: init-tpl
image: "{{ .Values.image.repository }}:{{ .Values.image.tag }}"
command: ["echo", "hello"]
asserts:
- contains:
path: spec.template.spec.initContainers
content:
name: init-tpl
image: "ghcr.io/berriai/litellm-database:test"
command: ["echo", "hello"]
- it: should work with extraContainers
template: migrations-job.yaml
set:
migrationJob:
enabled: true
extraContainers:
- name: sidecar
image: busybox:latest
asserts:
- contains:
path: spec.template.spec.containers
content:
name: sidecar
image: busybox:latest
- it: should support tpl in extraContainers
template: migrations-job.yaml
set:
image:
repository: ghcr.io/berriai/litellm-database
tag: test
migrationJob:
enabled: true
extraContainers:
- name: sidecar-tpl
image: "{{ .Values.image.repository }}:{{ .Values.image.tag }}"
asserts:
- contains:
path: spec.template.spec.containers
content:
name: sidecar-tpl
image: "ghcr.io/berriai/litellm-database:test"

View file

@ -487,7 +487,8 @@ router_settings:
| AZURE_STORAGE_CLIENT_ID | The Application Client ID to use for Authentication to Azure Blob Storage logging
| AZURE_STORAGE_CLIENT_SECRET | The Application Client Secret to use for Authentication to Azure Blob Storage logging
| AZURE_VECTOR_STORE_COST_PER_GB_PER_DAY | Cost per GB per day for Azure Vector Store service
| BACKGROUND_HEALTH_CHECK_MAX_TOKENS | Optional global default for `max_tokens` on proxy background health checks when a model has no `health_check_max_tokens`. If unset, non-wildcard models default to 1. Applies to wildcard routes when set. Default is unset
| BACKGROUND_HEALTH_CHECK_MAX_TOKENS | Optional global default for `max_tokens` on proxy background health checks when a model has no `health_check_max_tokens`. If unset, non-wildcard models default to 5. Applies to wildcard routes when set. Default is unset
| BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING | For **non-wildcard** reasoning models (`supports_reasoning(model)=true`), this takes precedence over `BACKGROUND_HEALTH_CHECK_MAX_TOKENS` when set. If unset, reasoning models fall back to `BACKGROUND_HEALTH_CHECK_MAX_TOKENS` (if set) or default behavior. Wildcard routes ignore this. Default is unset
| BATCH_STATUS_POLL_INTERVAL_SECONDS | Interval in seconds for polling batch status. Default is 3600 (1 hour)
| BATCH_STATUS_POLL_MAX_ATTEMPTS | Maximum number of attempts for polling batch status. Default is 24 (for 24 hours)
| BEDROCK_MAX_POLICY_SIZE | Maximum size for Bedrock policy. Default is 75

View file

@ -338,7 +338,7 @@ model_list:
## Health Check Max Tokens
By default, health checks use `max_tokens=1` to minimize cost and latency. For wildcard models, the default is `max_tokens=10`.
By default, health checks use `max_tokens=5` to balance reliability with low cost and latency. For wildcard models, the default is `max_tokens=10`.
You can override this per-model by setting `health_check_max_tokens` in the `model_info` section of your config.yaml.
@ -352,6 +352,30 @@ model_list:
health_check_max_tokens: 5 # 👈 OVERRIDE HEALTH CHECK MAX TOKENS
```
### Reasoning vs non-reasoning defaults
Reasoning models (per `supports_reasoning` in the model map) often need a higher health-check `max_tokens` because providers count reasoning tokens toward the completion budget. You can set **separate** limits without listing every model:
**Per deployment (`model_info`)** — used when `health_check_max_tokens` is not set. Ignored for wildcard routes (`*` in `litellm_params.model`, i.e. the deployment model string; not `health_check_model`).
```yaml
model_list:
- model_name: openai-stack
litellm_params:
model: openai/gpt-5-nano
api_key: os.environ/OPENAI_API_KEY
model_info:
health_check_max_tokens_reasoning: 128
health_check_max_tokens_non_reasoning: 1
```
**Global (environment)**:
- `BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING` — for non-wildcard reasoning models, this value takes precedence when set
- `BACKGROUND_HEALTH_CHECK_MAX_TOKENS` — global fallback for all models (including wildcard routes)
If neither is set, non-wildcard models default to `5` and wildcard routes omit `max_tokens`.
## `/health/readiness`
Unprotected endpoint for checking if proxy is ready to accept requests

View file

@ -35,6 +35,17 @@ By default, LiteLLM strips `x-api-key` from client requests for security. Settin
:::
:::tip Configure via UI instead of config.yaml
You can also complete this setup from the LiteLLM admin UI:
- Add the model via **Models → Add Model**, leaving the **API Key** field blank.
- Enable the toggle at **Settings → UI Settings → "Forward LLM provider auth headers"**.
Both UI actions write to the database and override `config.yaml` at runtime.
:::
## Step 2: Create a LiteLLM Virtual Key
Create a virtual key in the LiteLLM UI or via API.

View file

@ -1360,6 +1360,25 @@ try:
)
except (ValueError, TypeError):
BACKGROUND_HEALTH_CHECK_MAX_TOKENS = None
_background_health_check_max_tokens_reasoning_env = os.getenv(
"BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING"
)
try:
_raw_background_health_check_max_tokens_reasoning = (
_background_health_check_max_tokens_reasoning_env.strip()
if _background_health_check_max_tokens_reasoning_env is not None
else ""
)
BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING: Optional[int] = (
int(_raw_background_health_check_max_tokens_reasoning)
if _raw_background_health_check_max_tokens_reasoning
else None
)
except (ValueError, TypeError):
BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING = None
LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME = "litellm-internal-health-check"
LITTELM_CLI_SERVICE_ACCOUNT_NAME = "litellm-cli"
LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME = "litellm_internal_jobs"

View file

@ -684,6 +684,9 @@ def generic_cost_per_token( # noqa: PLR0915
- cache_creation
- image_tokens
)
# Clamp to zero: inconsistent streaming usage
if text_tokens < 0:
text_tokens = 0
prompt_tokens_details["text_tokens"] = text_tokens
(

View file

@ -12,6 +12,7 @@ import asyncio
import re
import time
from typing import Any, Dict, Optional
from urllib.parse import quote
import httpx
@ -55,10 +56,87 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
"""
Get supported OCR parameters for Azure Document Intelligence.
Azure DI has minimal optional parameters compared to Mistral OCR.
Most Mistral-specific params are ignored during transformation.
Azure DI exposes a `pages` query parameter on the analyze endpoint
(1-based, e.g. "1-3,5,7-9"). To keep the public request shape
aligned with Mistral OCR, callers pass `pages` using Mistral
semantics — a list of 0-based integers — or a pre-formatted
Azure-style string. Other Mistral-specific params (e.g.
`include_image_base64`) are not supported by Azure DI and are
ignored during transformation.
"""
return []
return ["pages"]
def map_ocr_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
) -> dict:
"""
Map OCR params to Azure DI format.
Translates Mistral-style `pages` (list[int], 0-based) into Azure's
`pages` query string (1-based, e.g. "1,2,3" or "1-3,5"). A raw
string that already matches Azure's format is passed through
unchanged.
"""
pages = non_default_params.get("pages")
if pages is None:
return optional_params
normalized = self._normalize_pages_param(pages)
if normalized:
optional_params["pages"] = normalized
return optional_params
@staticmethod
def _normalize_pages_param(pages: Any) -> str:
"""
Convert a caller-provided `pages` value to Azure DI's query-string
form. Azure expects 1-based page numbers, grammar: `^(\\d+(-\\d+)?)(,\\s*(\\d+(-\\d+)?))*$`.
Accepted inputs:
- list[int]: Mistral-style 0-based indices. Converted to 1-based
and joined (e.g. [0,1,2] -> "1,2,3").
- list[str]: tokens like "1" or "3-5". Validated, joined as-is
(treated as Azure-native, i.e. 1-based).
- str: already in Azure format. Validated and whitespace-stripped.
"""
pages_pattern = re.compile(r"^\s*\d+(-\d+)?(\s*,\s*\d+(-\d+)?)*\s*$")
if isinstance(pages, str):
if not pages_pattern.match(pages):
raise ValueError(
f"Invalid `pages` string for Azure Document Intelligence: "
f"{pages!r}. Expected format like '1-3,5,7-9'."
)
return pages.replace(" ", "")
if isinstance(pages, list):
if len(pages) == 0:
return ""
if any(isinstance(p, bool) for p in pages):
raise ValueError("`pages` must be integers, not booleans")
if all(isinstance(p, int) for p in pages):
if any(p < 0 for p in pages):
raise ValueError(
"`pages` integers must be >= 0 (Mistral 0-based indices)"
)
# Mistral 0-based -> Azure 1-based.
return ",".join(str(p + 1) for p in sorted(set(pages)))
if all(isinstance(p, str) for p in pages):
joined = ",".join(p.strip() for p in pages)
if not pages_pattern.match(joined):
raise ValueError(
f"Invalid `pages` list for Azure Document Intelligence: "
f"{pages!r}. Expected tokens like '1' or '3-5'."
)
return joined
raise ValueError(
"`pages` must be a list[int] (0-based, Mistral-style) or a "
"string like '1-3,5,7-9'."
)
def validate_environment(
self,
@ -142,7 +220,18 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
# Azure Document Intelligence analyze endpoint
# Note: API version 2024-11-30+ uses /documentintelligence/ (not /formrecognizer/)
return f"{api_base}/documentintelligence/documentModels/{model_id}:analyze?api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}"
url = (
f"{api_base}/documentintelligence/documentModels/{model_id}:analyze"
f"?api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}"
)
# Azure DI accepts `pages` as a query param (1-based, e.g. "1-3,5").
# `optional_params` has already been normalized in `map_ocr_params`.
pages = optional_params.get("pages") if optional_params else None
if pages:
url += f"&pages={quote(str(pages), safe=',-')}"
return url
def _extract_base64_from_data_uri(self, data_uri: str) -> str:
"""
@ -234,8 +323,9 @@ class AzureDocumentIntelligenceOCRConfig(BaseOCRConfig):
data["urlSource"] = document_url
verbose_logger.debug("Using urlSource for Azure Document Intelligence")
# Azure DI doesn't support most Mistral-specific params
# Ignore pages, include_image_base64, etc.
# Azure DI: `pages` is a query param (wired in get_complete_url),
# not a body field. Other Mistral-specific params (e.g.
# include_image_base64, image_limit) are unsupported and ignored.
return OCRRequestData(data=data, files=None)

View file

@ -591,11 +591,14 @@ class AmazonAnthropicClaudeMessagesConfig(
"""
Bedrock invoke does not return SSE formatted data. This function is a wrapper to ensure litellm chunks are SSE formatted.
Bedrock's Anthropic-compatible streaming puts cache usage fields
(cache_creation_input_tokens, cache_read_input_tokens) only on
message_stop, not on message_start or message_delta. Claude Code's
SDK only merges usage from message_delta, so we promote those fields
from message_stop onto message_delta before yielding.
Bedrock's Anthropic-compatible streaming usually puts cache usage fields
(cache_creation_input_tokens, cache_read_input_tokens) on message_stop.
Some deployments (including GovCloud) emit the cache breakdown only on
``message_start.message.usage``; ``message_delta`` / ``message_stop`` then
repeat uncached ``input_tokens`` only. We promote cache fields from
``message_stop`` onto ``message_delta``, and when those are absent we
merge them from ``message_start`` so logging/cost sees a consistent usage
object (fixes negative input costs: LIT-2411).
"""
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
BaseAnthropicMessagesStreamingIterator,
@ -611,6 +614,27 @@ class AmazonAnthropicClaudeMessagesConfig(
async for chunk in handler.async_sse_wrapper(patched_stream):
yield chunk
@staticmethod
def _merge_message_start_cache_into_delta_usage(
delta_usage: Dict[str, Any],
start_usage: Optional[Dict[str, Any]],
) -> None:
"""
Copy cache breakdown from message_start onto message_delta usage when
those keys are missing on the delta (GovCloud / some Bedrock streams).
"""
if not start_usage:
return
for field in ("cache_creation_input_tokens", "cache_read_input_tokens"):
if field not in delta_usage:
val = start_usage.get(field)
if val is not None:
delta_usage[field] = val
if "cache_creation" not in delta_usage:
cc = start_usage.get("cache_creation")
if cc is not None:
delta_usage["cache_creation"] = cc
@staticmethod
async def _promote_message_stop_usage(
completion_stream: AsyncIterator[
@ -618,20 +642,13 @@ class AmazonAnthropicClaudeMessagesConfig(
],
) -> AsyncIterator[Union[bytes, GenericStreamingChunk, ModelResponseStream, dict]]:
"""
Promote cache usage fields from message_stop onto message_delta.
Bedrock reports input_tokens (uncached only) on message_start, and
the full breakdown (input_tokens, cache_creation_input_tokens,
cache_read_input_tokens) only on message_stop. Claude Code's SDK
merges usage from message_start and message_delta but ignores
message_stop. This method buffers message_delta and, when
message_stop arrives with cache usage, merges those fields into the
message_delta usage. input_tokens is kept as the uncached-only
count; downstream calculate_usage adds cache tokens to
prompt_tokens.
Promote cache usage fields onto message_delta from message_stop (and,
when stop lacks them, from message_start). Ensures the final usage
chunk that logging/cost sees is always self-consistent.
"""
_CACHE_FIELDS = ("cache_creation_input_tokens", "cache_read_input_tokens")
pending_delta = None
pending_delta: Optional[Dict[str, Any]] = None
start_usage_snapshot: Optional[Dict[str, Any]] = None
async for chunk in completion_stream:
if not isinstance(chunk, dict):
@ -643,8 +660,19 @@ class AmazonAnthropicClaudeMessagesConfig(
chunk_type = chunk.get("type")
if chunk_type == "message_start":
msg: Dict[str, Any] = cast(Dict[str, Any], chunk.get("message") or {})
u = msg.get("usage")
if isinstance(u, dict):
start_usage_snapshot = dict(u)
if pending_delta is not None:
yield pending_delta
pending_delta = None
yield chunk
continue
if chunk_type == "message_delta":
pending_delta = chunk
pending_delta = cast(Dict[str, Any], chunk)
continue
if chunk_type == "message_stop" and pending_delta is not None:
@ -661,6 +689,10 @@ class AmazonAnthropicClaudeMessagesConfig(
raw_input if isinstance(raw_input, int) else 0
)
AmazonAnthropicClaudeMessagesConfig._merge_message_start_cache_into_delta_usage(
delta_usage, start_usage_snapshot
)
if delta_usage:
pending_delta["usage"] = delta_usage # type: ignore[arg-type]
@ -676,6 +708,12 @@ class AmazonAnthropicClaudeMessagesConfig(
yield chunk
if pending_delta is not None:
delta_usage = dict(pending_delta.get("usage") or {})
AmazonAnthropicClaudeMessagesConfig._merge_message_start_cache_into_delta_usage(
delta_usage, start_usage_snapshot
)
if delta_usage:
pending_delta["usage"] = delta_usage # type: ignore[arg-type]
yield pending_delta

View file

@ -4,6 +4,7 @@ import time
from datetime import datetime, timezone
from typing import Any, List, Literal, Optional, Union
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import (
LiteLLM_BudgetTableFull,
@ -165,16 +166,35 @@ class ResetBudgetJob:
table_name="budget",
)
budget_ids_to_reset = [
budget.budget_id
for budget in budgets_to_reset
if budget.budget_id is not None
]
endusers_to_reset = await self.prisma_client.get_data(
table_name="enduser",
query_type="find_all",
budget_id_list=[
budget.budget_id
for budget in budgets_to_reset
if budget.budget_id is not None
],
budget_id_list=budget_ids_to_reset,
)
# Also reset end users with no budget_id (NULL) who use the
# default budget via litellm.max_end_user_budget_id. These
# users are enforced in-memory but never had budget_id
# persisted, so the query above misses them.
if (
litellm.max_end_user_budget_id is not None
and litellm.max_end_user_budget_id in budget_ids_to_reset
):
default_budget_endusers = (
await self._get_endusers_with_no_budget_id()
)
if default_budget_endusers:
if endusers_to_reset is None:
endusers_to_reset = default_budget_endusers
else:
endusers_to_reset.extend(default_budget_endusers)
await self.reset_budget_for_litellm_team_members(
budgets_to_reset=budgets_to_reset
)
@ -282,6 +302,23 @@ class ResetBudgetJob:
)
verbose_proxy_logger.exception("Failed to reset budget for endusers: %s", e)
async def _get_endusers_with_no_budget_id(
self,
) -> List[LiteLLM_EndUserTable]:
"""
Fetch end users that have no explicit budget_id set (NULL) and have
accumulated spend > 0. These are implicitly-created end users that
rely on the default budget (litellm.max_end_user_budget_id) applied
in-memory during auth checks.
"""
rows = await self.prisma_client.db.litellm_endusertable.find_many(
where={
"budget_id": None,
"spend": {"gt": 0},
},
)
return [LiteLLM_EndUserTable(**row.dict()) for row in rows]
async def reset_budget_for_litellm_keys(self):
"""
Resets the budget for all the litellm keys

View file

@ -13,6 +13,7 @@ import litellm
logger = logging.getLogger(__name__)
from litellm.constants import (
BACKGROUND_HEALTH_CHECK_MAX_TOKENS,
BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING,
DEFAULT_HEALTH_CHECK_PROMPT,
HEALTH_CHECK_TIMEOUT_SECONDS,
)
@ -292,6 +293,69 @@ def build_deployment_health_states(
return states
def _deployment_model_string_for_health_check(litellm_params: dict) -> str:
"""Deployment model from litellm_params (before Bedrock rewrite).
Used for reasoning vs non-reasoning max_tokens and wildcard detection only.
Does not use ``health_check_model``; that override applies later to the request.
"""
return litellm_params.get("model") or ""
def _health_check_deployment_is_wildcard(litellm_params: dict) -> bool:
return "*" in _deployment_model_string_for_health_check(litellm_params)
def _resolve_health_check_max_tokens(model_info: dict, litellm_params: dict) -> Optional[int]:
"""
Pick max_tokens for the health check request.
Priority:
1. model_info.health_check_max_tokens (explicit override)
2. For non-wildcard routes: health_check_max_tokens_reasoning / _non_reasoning
from model_info based on litellm.supports_reasoning(litellm_params["model"])
3. For non-wildcard reasoning routes: BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING
from env (if set)
4. BACKGROUND_HEALTH_CHECK_MAX_TOKENS (global, any route including wildcards)
5. Non-wildcard default: 5
6. Wildcard and nothing from (1)(4): leave unset (caller omits max_tokens)
"""
explicit = model_info.get("health_check_max_tokens", None)
if explicit is not None:
return int(explicit)
is_wildcard = _health_check_deployment_is_wildcard(litellm_params)
deployment_model = _deployment_model_string_for_health_check(litellm_params)
if not is_wildcard:
try:
is_reasoning = litellm.supports_reasoning(deployment_model)
except Exception:
is_reasoning = False
tokens_reasoning = model_info.get("health_check_max_tokens_reasoning", None)
tokens_non_reasoning = model_info.get(
"health_check_max_tokens_non_reasoning", None
)
if tokens_reasoning is not None or tokens_non_reasoning is not None:
if is_reasoning and tokens_reasoning is not None:
return int(tokens_reasoning)
if not is_reasoning and tokens_non_reasoning is not None:
return int(tokens_non_reasoning)
if (
is_reasoning
and BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING is not None
):
return int(BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING)
if BACKGROUND_HEALTH_CHECK_MAX_TOKENS is not None:
return int(BACKGROUND_HEALTH_CHECK_MAX_TOKENS)
if not is_wildcard:
return 5
return None
def _update_litellm_params_for_health_check(
model_info: dict, litellm_params: dict
) -> dict:
@ -304,15 +368,9 @@ def _update_litellm_params_for_health_check(
- for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID
"""
litellm_params["messages"] = _get_random_llm_message()
_health_check_max_tokens = model_info.get("health_check_max_tokens", None)
if _health_check_max_tokens is not None:
litellm_params["max_tokens"] = _health_check_max_tokens
elif BACKGROUND_HEALTH_CHECK_MAX_TOKENS is not None:
litellm_params["max_tokens"] = BACKGROUND_HEALTH_CHECK_MAX_TOKENS
elif "*" not in (
model_info.get("health_check_model") or litellm_params.get("model") or ""
):
litellm_params["max_tokens"] = 1
_resolved_max_tokens = _resolve_health_check_max_tokens(model_info, litellm_params)
if _resolved_max_tokens is not None:
litellm_params["max_tokens"] = _resolved_max_tokens
_health_check_model = model_info.get("health_check_model", None)
if _health_check_model is not None:

View file

@ -214,12 +214,22 @@
"provider_display_name": "Anthropic",
"litellm_provider": "anthropic",
"credential_fields": [
{
"key": "api_base",
"label": "Upstream API Base",
"placeholder": "https://api.anthropic.com",
"tooltip": "Optional. Where the proxy forwards requests upstream. Leave blank to use Anthropic's public API. Set this only for private Anthropic deployments or reverse proxies. Do NOT set this to your LiteLLM proxy URL — that causes a recursive loop.",
"required": false,
"field_type": "text",
"options": null,
"default_value": null
},
{
"key": "api_key",
"label": "API Key",
"placeholder": "sk-",
"tooltip": null,
"required": true,
"tooltip": "Leave empty for BYOK (bring-your-own-key) flows, where clients forward their own Anthropic key via the x-api-key header. Requires the 'Forward LLM provider auth headers' UI setting to be enabled.",
"required": false,
"field_type": "password",
"options": null,
"default_value": null

View file

@ -97,7 +97,24 @@ class UISettings(BaseModel):
forward_client_headers_to_llm_api: bool = Field(
default=False,
description="If enabled, forwards client headers (e.g. Authorization) to the LLM API. Required for Claude Code with Max subscription.",
description=(
"Forwards client headers (Authorization, anthropic-beta, and x-* "
"custom headers) to the upstream LLM. Enable for Claude Code with a "
"Max subscription (forwards the OAuth token) or to pass custom/tracing "
"headers through to the provider. Independent of the BYOK toggle — "
"enable only the one(s) you need."
),
)
forward_llm_provider_auth_headers: bool = Field(
default=False,
description=(
"Forwards provider auth headers (x-api-key, x-goog-api-key, api-key, "
"ocp-apim-subscription-key) to the upstream LLM, overriding any "
"deployment-configured key for that request. Enable for Claude Code "
"BYOK (clients bring their own API key). Independent of the "
"client-headers toggle — enable only the one(s) you need."
),
)
enable_projects_ui: bool = Field(
@ -149,6 +166,7 @@ ALLOWED_UI_SETTINGS_FIELDS = {
"enabled_ui_pages_internal_users",
"require_auth_for_public_ai_hub",
"forward_client_headers_to_llm_api",
"forward_llm_provider_auth_headers",
"enable_projects_ui",
"disable_agents_for_internal_users",
"allow_agents_for_team_admins",
@ -162,6 +180,7 @@ ALLOWED_UI_SETTINGS_FIELDS = {
# general_settings at runtime (on both read and write).
_RUNTIME_GENERAL_SETTINGS_FLAGS = [
"forward_client_headers_to_llm_api",
"forward_llm_provider_auth_headers",
"disable_agents_for_internal_users",
"allow_agents_for_team_admins",
"disable_vector_stores_for_internal_users",

View file

@ -10,6 +10,10 @@ import os
import pytest
from base_ocr_unit_tests import BaseOCRTest
from litellm.constants import AZURE_DOCUMENT_INTELLIGENCE_API_VERSION
from litellm.llms.azure_ai.ocr.document_intelligence.transformation import (
AzureDocumentIntelligenceOCRConfig,
)
class TestAzureDocumentIntelligenceOCR(BaseOCRTest):
@ -42,3 +46,126 @@ class TestAzureDocumentIntelligenceOCR(BaseOCRTest):
"api_key": api_key,
"api_base": endpoint,
}
class TestAzureDocumentIntelligencePagesParam:
"""
Unit tests for the Mistral-compatible `pages` parameter translation to
Azure Document Intelligence's `pages` query string.
These tests exercise the transformation layer directly and do not
require Azure credentials or a network call.
"""
@pytest.fixture
def cfg(self) -> AzureDocumentIntelligenceOCRConfig:
return AzureDocumentIntelligenceOCRConfig()
def test_get_supported_ocr_params_includes_pages(self, cfg):
assert cfg.get_supported_ocr_params("prebuilt-layout") == ["pages"]
def test_map_ocr_params_mistral_zero_based_int_list(self, cfg):
mapped = cfg.map_ocr_params({"pages": [0, 1, 2]}, {}, "prebuilt-layout")
assert mapped == {"pages": "1,2,3"}
def test_map_ocr_params_dedupes_and_sorts(self, cfg):
mapped = cfg.map_ocr_params({"pages": [2, 0, 0, 1]}, {}, "prebuilt-layout")
assert mapped == {"pages": "1,2,3"}
def test_map_ocr_params_empty_list_omits_pages(self, cfg):
mapped = cfg.map_ocr_params({"pages": []}, {}, "prebuilt-layout")
assert mapped == {}
def test_map_ocr_params_azure_native_string_range(self, cfg):
mapped = cfg.map_ocr_params({"pages": "3-9"}, {}, "prebuilt-layout")
assert mapped == {"pages": "3-9"}
def test_map_ocr_params_azure_native_string_with_spaces_stripped(self, cfg):
mapped = cfg.map_ocr_params({"pages": "1-3, 5"}, {}, "prebuilt-layout")
assert mapped == {"pages": "1-3,5"}
def test_map_ocr_params_list_of_string_tokens(self, cfg):
mapped = cfg.map_ocr_params({"pages": ["1", "3-5"]}, {}, "prebuilt-layout")
assert mapped == {"pages": "1,3-5"}
def test_map_ocr_params_invalid_string_raises(self, cfg):
with pytest.raises(ValueError, match="Invalid `pages` string"):
cfg.map_ocr_params({"pages": "a,b"}, {}, "prebuilt-layout")
def test_map_ocr_params_negative_index_raises(self, cfg):
with pytest.raises(ValueError, match="must be >= 0"):
cfg.map_ocr_params({"pages": [-1]}, {}, "prebuilt-layout")
def test_map_ocr_params_bool_list_raises(self, cfg):
with pytest.raises(ValueError, match="must be integers, not booleans"):
cfg.map_ocr_params({"pages": [True, False]}, {}, "prebuilt-layout")
def test_map_ocr_params_unsupported_type_raises(self, cfg):
with pytest.raises(ValueError):
cfg.map_ocr_params({"pages": 5}, {}, "prebuilt-layout")
def test_get_complete_url_appends_pages_query(self, cfg):
url = cfg.get_complete_url(
api_base="https://example.cognitiveservices.azure.com/",
model="azure_ai/doc-intelligence/prebuilt-layout",
optional_params={"pages": "1-3,5"},
)
assert (
f"api-version={AZURE_DOCUMENT_INTELLIGENCE_API_VERSION}" in url
), url
assert "pages=1-3,5" in url, url
assert "/documentintelligence/documentModels/prebuilt-layout:analyze" in url
def test_get_complete_url_no_pages_when_optional_params_empty(self, cfg):
url = cfg.get_complete_url(
api_base="https://example.cognitiveservices.azure.com",
model="prebuilt-layout",
optional_params={},
)
assert "pages=" not in url
def test_transform_ocr_request_does_not_put_pages_in_body(self, cfg):
req = cfg.transform_ocr_request(
model="prebuilt-layout",
document={
"type": "document_url",
"document_url": "https://example.com/x.pdf",
},
optional_params={"pages": "1,2,3"},
headers={},
)
assert req.data is not None
assert "pages" not in req.data
assert req.data.get("urlSource") == "https://example.com/x.pdf"
def test_end_to_end_mistral_shape_to_azure_query(self, cfg):
"""
Caller sends Mistral-style `pages: [2,3,4,5,6,7,8]` (0-based,
meaning human pages 3-9). LiteLLM should turn that into Azure's
`&pages=3,4,5,6,7,8,9` on the analyze URL, and the body should
still only contain urlSource.
"""
non_default_params = {"pages": [2, 3, 4, 5, 6, 7, 8]}
optional_params = cfg.map_ocr_params(
non_default_params=non_default_params,
optional_params={},
model="prebuilt-layout",
)
url = cfg.get_complete_url(
api_base="https://example.cognitiveservices.azure.com",
model="prebuilt-layout",
optional_params=optional_params,
)
req = cfg.transform_ocr_request(
model="prebuilt-layout",
document={
"type": "document_url",
"document_url": "https://example.com/x.pdf",
},
optional_params=optional_params,
headers={},
)
assert "pages=3,4,5,6,7,8,9" in url
assert req.data == {"urlSource": "https://example.com/x.pdf"}

View file

@ -618,6 +618,147 @@ async def test_promote_message_stop_usage_preserves_message_delta_output_tokens(
assert delta_out["usage"]["input_tokens"] == 3
@pytest.mark.asyncio
async def test_promote_message_start_cache_when_message_stop_omits_cache_fields():
"""
GovCloud / some Bedrock streams put cache_read only on message_start; delta and
stop repeat uncached input_tokens only. Merging start cache onto message_delta
avoids inconsistent usage and negative input costs (LIT-2411).
"""
cfg = AmazonAnthropicClaudeMessagesConfig()
async def _stream(): # type: ignore[return-type]
yield {
"type": "message_start",
"message": {
"id": "msg_1",
"type": "message",
"role": "assistant",
"content": [],
"model": "claude-sonnet-4-5-20250929",
"usage": {
"input_tokens": 10,
"cache_creation_input_tokens": 0,
"cache_read_input_tokens": 22167,
"cache_creation": {
"ephemeral_5m_input_tokens": 0,
"ephemeral_1h_input_tokens": 0,
},
"output_tokens": 4,
},
},
}
yield {
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"input_tokens": 10, "output_tokens": 181},
}
yield {"type": "message_stop", "usage": {"input_tokens": 10, "output_tokens": 181}}
merged: list[dict] = []
async for chunk in cfg._promote_message_stop_usage(_stream()):
if isinstance(chunk, dict):
merged.append(chunk)
delta_chunks = [c for c in merged if c.get("type") == "message_delta"]
assert len(delta_chunks) == 1
u = delta_chunks[0]["usage"]
assert u["input_tokens"] == 10
assert u["output_tokens"] == 181
assert u["cache_read_input_tokens"] == 22167
assert u["cache_creation_input_tokens"] == 0
@pytest.mark.asyncio
async def test_unified_bedrock_messages_cache_on_start_only_never_negative_cost():
"""
Regression guard for LIT-2411:
If cache usage is present only on message_start (and omitted from
message_delta/message_stop), final reconstructed usage + cost must still
be consistent and non-negative.
"""
from litellm import completion_cost
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
cfg = AmazonAnthropicClaudeMessagesConfig()
async def _stream(): # type: ignore[return-type]
yield {
"type": "message_start",
"message": {
"id": "msg_bdrk_01WuFzkDbE9KWgiWakMRNKcA",
"type": "message",
"role": "assistant",
"content": [],
"model": "claude-sonnet-4-5-20250929",
"stop_reason": None,
"stop_sequence": None,
"usage": {
"input_tokens": 10,
"cache_creation_input_tokens": 0,
"cache_read_input_tokens": 22167,
"cache_creation": {
"ephemeral_5m_input_tokens": 0,
"ephemeral_1h_input_tokens": 0,
},
"output_tokens": 4,
},
},
}
yield {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}
yield {
"type": "content_block_delta",
"index": 0,
"delta": {"type": "text_delta", "text": "Hello from regression test"},
}
yield {"type": "content_block_stop", "index": 0}
yield {
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 181, "input_tokens": 10},
}
yield {"type": "message_stop", "usage": {"input_tokens": 10, "output_tokens": 181}}
logging_obj = LiteLLMLoggingObj(
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
messages=[{"role": "user", "content": "Hi"}],
stream=True,
call_type="chat",
start_time=datetime.now(),
litellm_call_id="test_cache_on_start_only_never_negative_cost",
function_id="test_cache_on_start_only_never_negative_cost",
)
collected: list[bytes] = []
async for sse in cfg.bedrock_sse_wrapper(
completion_stream=_stream(),
litellm_logging_obj=logging_obj,
request_body={"model": "anthropic.claude-3-5-sonnet-20240620-v1:0"},
):
collected.append(sse)
built = AnthropicPassthroughLoggingHandler._build_complete_streaming_response(
all_chunks=collected,
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
litellm_logging_obj=Mock(),
)
assert built.usage is not None
assert built.usage.prompt_tokens == 22177
assert built.usage.completion_tokens == 181
assert built.usage.cache_creation_input_tokens == 0
assert built.usage.cache_read_input_tokens == 22167
cost = completion_cost(
completion_response=built,
model="bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0",
custom_llm_provider="bedrock",
)
assert cost > 0
assert cost == pytest.approx(0.0093951, rel=0, abs=1e-9)
@pytest.mark.asyncio
async def test_unified_bedrock_messages_sse_usage_and_cost_claude_sonnet_46():
"""

View file

@ -36,10 +36,24 @@ class MockLiteLLMVerificationToken:
return {"count": 1}
class MockLiteLLMEndUserTable:
def __init__(self):
self.find_many_calls: List[Dict[str, Any]] = []
self._find_many_results: List[Any] = []
def set_find_many_results(self, results: List[Any]):
self._find_many_results = results
async def find_many(self, where: Dict[str, Any]) -> List[Any]:
self.find_many_calls.append({"where": where})
return self._find_many_results
class MockDB:
def __init__(self):
self.litellm_teammembership = MockLiteLLMTeamMembership()
self.litellm_verificationtoken = MockLiteLLMVerificationToken()
self.litellm_endusertable = MockLiteLLMEndUserTable()
class MockPrismaClient:
@ -599,3 +613,174 @@ def test_budget_table_reset_also_resets_linked_keys(
)
assert calls[0]["where"]["budget_id"] == {"in": ["7d-budget-tier"]}
assert calls[0]["data"]["spend"] == 0
def test_reset_budget_resets_endusers_with_null_budget_id(
reset_budget_job, mock_prisma_client
):
"""
When litellm.max_end_user_budget_id is configured and that budget is
being reset, end users with budget_id=NULL should also have their spend
reset. These users were implicitly created and have no budget_id persisted,
but are enforced against the default budget in-memory.
"""
import litellm
now = datetime.now(timezone.utc)
default_budget_id = "default-enduser-budget"
litellm.max_end_user_budget_id = default_budget_id
# Budget that is due for reset — matches the default end user budget
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 50.0,
"budget_duration": "1d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": default_budget_id,
"created_at": now - timedelta(days=1),
},
)
# End user WITH explicit budget_id (found by the normal budget_id_list query)
enduser_with_budget = type(
"LiteLLM_EndUserTable",
(),
{
"spend": 30.0,
"litellm_budget_table": test_budget,
"user_id": "enduser-explicit",
},
)
# End user WITHOUT budget_id (NULL) — should also be reset
enduser_no_budget_row = type(
"EndUserRow",
(),
{
"spend": 25.0,
"user_id": "enduser-implicit",
"budget_id": None,
"alias": None,
"allowed_model_region": None,
"default_model": None,
"blocked": False,
"object_permission_id": None,
"object_permission": None,
"litellm_budget_table": None,
"dict": lambda self=None: {
"spend": 25.0,
"user_id": "enduser-implicit",
"blocked": False,
"alias": None,
"allowed_model_region": None,
"default_model": None,
"litellm_budget_table": None,
"object_permission_id": None,
"object_permission": None,
},
},
)
mock_prisma_client.data["budget"] = [test_budget]
mock_prisma_client.data["enduser"] = [enduser_with_budget]
# Set up the DB mock for NULL-budget-id end users
mock_prisma_client.db.litellm_endusertable.set_find_many_results(
[enduser_no_budget_row]
)
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
# Both end users should have been reset
updated = mock_prisma_client.updated_data["enduser"]
assert len(updated) == 2, (
f"Expected 2 endusers reset (1 explicit + 1 implicit), got {len(updated)}"
)
user_ids = {u.user_id for u in updated}
assert "enduser-explicit" in user_ids
assert "enduser-implicit" in user_ids
for u in updated:
assert u.spend == 0.0, f"Expected spend=0 for {u.user_id}, got {u.spend}"
# Verify find_many was called to fetch NULL-budget-id end users
find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls
assert len(find_many_calls) == 1
assert find_many_calls[0]["where"] == {"budget_id": None, "spend": {"gt": 0}}
litellm.max_end_user_budget_id = None
def test_reset_budget_skips_null_budget_id_endusers_when_default_not_configured(
reset_budget_job, mock_prisma_client
):
"""
When litellm.max_end_user_budget_id is NOT configured, end users with
budget_id=NULL should NOT be fetched or reset.
"""
import litellm
now = datetime.now(timezone.utc)
litellm.max_end_user_budget_id = None
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 50.0,
"budget_duration": "1d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "some-budget",
"created_at": now - timedelta(days=1),
},
)
mock_prisma_client.data["budget"] = [test_budget]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
# Should NOT have queried for NULL-budget-id end users
find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls
assert len(find_many_calls) == 0
litellm.max_end_user_budget_id = None
def test_reset_budget_skips_null_budget_id_endusers_when_default_not_in_reset_list(
reset_budget_job, mock_prisma_client
):
"""
When litellm.max_end_user_budget_id IS configured but the corresponding
budget is NOT in the budgets-to-reset list (not yet expired), end users
with budget_id=NULL should NOT be reset.
"""
import litellm
now = datetime.now(timezone.utc)
litellm.max_end_user_budget_id = "default-budget-not-expired"
# A different budget that IS expiring (not the default one)
test_budget = type(
"LiteLLM_BudgetTableFull",
(),
{
"max_budget": 50.0,
"budget_duration": "1d",
"budget_reset_at": now - timedelta(hours=1),
"budget_id": "other-budget",
"created_at": now - timedelta(days=1),
},
)
mock_prisma_client.data["budget"] = [test_budget]
asyncio.run(reset_budget_job.reset_budget_for_litellm_budget_table())
# Should NOT have queried for NULL-budget-id end users
find_many_calls = mock_prisma_client.db.litellm_endusertable.find_many_calls
assert len(find_many_calls) == 0
litellm.max_end_user_budget_id = None

View file

@ -143,6 +143,51 @@ def test_azure_provider_fields_include_entra_id():
assert fields_by_key["client_secret"]["required"] is False
def test_anthropic_provider_fields_support_byok():
"""
The Anthropic provider form must allow BYOK:
- api_key is optional (not required) so admins can create models without a key
- api_key has a non-null tooltip explaining the BYOK use case
"""
app_instance = FastAPI()
app_instance.include_router(router)
test_client = TestClient(app_instance)
response = test_client.get("/public/providers/fields")
assert response.status_code == 200
providers = response.json()
anthropic = next((p for p in providers if p["provider"] == "Anthropic"), None)
assert anthropic is not None, "Anthropic provider entry not found"
fields_by_key = {f["key"]: f for f in anthropic["credential_fields"]}
assert "api_key" in fields_by_key
assert fields_by_key["api_key"]["required"] is False, (
"Anthropic api_key must be optional so admins can configure BYOK models "
"without entering a key. See BYOK tutorial."
)
assert fields_by_key["api_key"].get("tooltip"), (
"Anthropic api_key must have a tooltip explaining the BYOK use case."
)
assert "api_base" in fields_by_key, (
"Anthropic provider form must expose api_base so cloud customers "
"can override the upstream URL without env var access."
)
api_base_field = fields_by_key["api_base"]
assert api_base_field["required"] is False
assert api_base_field["field_type"] == "text"
assert api_base_field.get("tooltip"), (
"api_base should have a tooltip explaining it is optional."
)
# UI forms render fields in credential_fields order; api_base should come first
# so an admin sees the URL override before the key field.
field_order = [f["key"] for f in anthropic["credential_fields"]]
assert field_order.index("api_base") < field_order.index("api_key"), (
"api_base must appear before api_key in credential_fields (matches AI21 and ANTHROPIC_TEXT convention)."
)
def test_public_model_hub_with_healthy_model():
"""Test that health information is populated for a healthy model"""
app = FastAPI()

View file

@ -4,20 +4,25 @@ import pytest
from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers
from litellm.proxy import health_check as hc_module
from litellm.proxy.health_check import _update_litellm_params_for_health_check
from litellm.proxy.health_check import (
_resolve_health_check_max_tokens,
_update_litellm_params_for_health_check,
)
@pytest.mark.asyncio
async def test_update_litellm_params_max_tokens_default():
async def test_update_litellm_params_max_tokens_default(monkeypatch):
"""
Test that max_tokens defaults to 1 for non-wildcard models.
Test that max_tokens defaults to 5 for non-wildcard models.
"""
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", None)
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING", None)
model_info = {}
litellm_params = {"model": "gpt-4"}
updated_params = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated_params["max_tokens"] == 1
assert updated_params["max_tokens"] == 5
@pytest.mark.asyncio
@ -126,3 +131,97 @@ async def test_global_env_var_applies_to_wildcard_models(monkeypatch):
updated_params = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated_params["max_tokens"] == 15
def test_resolve_health_check_max_tokens_reasoning_specific_model_info():
model_info = {
"health_check_max_tokens_reasoning": 64,
"health_check_max_tokens_non_reasoning": 2,
}
litellm_params = {"model": "openai/gpt-4o"}
with patch.object(hc_module.litellm, "supports_reasoning", return_value=False):
assert _resolve_health_check_max_tokens(model_info, litellm_params) == 2
with patch.object(hc_module.litellm, "supports_reasoning", return_value=True):
assert _resolve_health_check_max_tokens(model_info, litellm_params) == 64
def test_explicit_health_check_max_tokens_beats_reasoning_specific():
model_info = {
"health_check_max_tokens": 9,
"health_check_max_tokens_reasoning": 64,
"health_check_max_tokens_non_reasoning": 2,
}
litellm_params = {"model": "openai/gpt-4o"}
with patch.object(hc_module.litellm, "supports_reasoning", return_value=True):
assert _resolve_health_check_max_tokens(model_info, litellm_params) == 9
def test_reasoning_specific_falls_through_when_wrong_branch_only(monkeypatch):
"""Only non-reasoning key set but model is reasoning → fall back to default 5."""
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", None)
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING", None)
model_info = {"health_check_max_tokens_non_reasoning": 3}
litellm_params = {"model": "openai/o1"}
with patch.object(hc_module.litellm, "supports_reasoning", return_value=True):
assert _resolve_health_check_max_tokens(model_info, litellm_params) == 5
@pytest.mark.asyncio
async def test_background_split_env_reasoning_vs_non_reasoning(monkeypatch):
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", None)
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING", 50)
model_info = {}
litellm_params = {"model": "azure/gpt-4"}
with patch.object(hc_module.litellm, "supports_reasoning", return_value=False):
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated["max_tokens"] == 5
litellm_params2 = {"model": "openai/o1"}
with patch.object(hc_module.litellm, "supports_reasoning", return_value=True):
updated2 = _update_litellm_params_for_health_check(model_info, litellm_params2)
assert updated2["max_tokens"] == 50
@pytest.mark.asyncio
async def test_reasoning_env_precedence_over_global(monkeypatch):
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", 10)
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING", 20)
model_info = {}
litellm_params = {"model": "openai/gpt-5.4"}
with patch.object(hc_module.litellm, "supports_reasoning", return_value=True):
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated["max_tokens"] == 20
@pytest.mark.asyncio
async def test_non_reasoning_uses_global_when_reasoning_env_set(monkeypatch):
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", 10)
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING", 20)
model_info = {}
litellm_params = {"model": "azure/gpt-4"}
with patch.object(hc_module.litellm, "supports_reasoning", return_value=False):
updated = _update_litellm_params_for_health_check(model_info, litellm_params)
assert updated["max_tokens"] == 10
def test_wildcard_ignores_reasoning_split_model_info(monkeypatch):
"""Wildcard routes do not use reasoning/non-reasoning model_info split."""
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS", None)
monkeypatch.setattr(hc_module, "BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING", None)
model_info = {
"health_check_max_tokens_reasoning": 99,
"health_check_max_tokens_non_reasoning": 7,
}
litellm_params = {"model": "openai/*"}
assert _resolve_health_check_max_tokens(model_info, litellm_params) is None

View file

@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import Request
from starlette.datastructures import Headers
import litellm
from litellm.proxy._types import TeamCallbackMetadata, UserAPIKeyAuth
@ -23,6 +24,7 @@ from litellm.proxy.litellm_pre_call_utils import (
add_guardrails_from_policy_engine,
add_litellm_data_to_request,
check_if_token_is_service_account,
clean_headers,
)
from litellm.types.utils import CredentialItem
@ -2487,3 +2489,79 @@ def test_resolve_provider_hint_from_model_name():
"azure/gpt-4", config, None, pre_alias_model_name="gpt-4", provider="azure"
)
assert result == "azure-cred"
def test_clean_headers_preserves_x_api_key_when_byok_enabled():
"""
Regression test: when forward_llm_provider_auth_headers=True,
clean_headers() must preserve the client-supplied x-api-key header
so it can be forwarded to the upstream Anthropic API (BYOK flow).
"""
headers = Headers(
{
"x-api-key": "sk-ant-api03-client-key",
"x-litellm-api-key": "sk-proxy-virtual-key",
"content-type": "application/json",
}
)
result = clean_headers(
headers=headers,
litellm_key_header_name="x-litellm-api-key",
forward_llm_provider_auth_headers=True,
authenticated_with_header="x-litellm-api-key",
)
# x-api-key must be preserved for BYOK
assert result.get("x-api-key") == "sk-ant-api03-client-key"
# x-litellm-api-key must NOT leak to the upstream
assert "x-litellm-api-key" not in result
def test_clean_headers_strips_x_api_key_when_byok_disabled():
"""
Regression test: with forward_llm_provider_auth_headers=False (default),
x-api-key must be stripped so proxy-configured keys are not overridden
by a client-supplied one.
"""
headers = Headers(
{
"x-api-key": "sk-ant-api03-client-key",
"x-litellm-api-key": "sk-proxy-virtual-key",
}
)
result = clean_headers(
headers=headers,
litellm_key_header_name="x-litellm-api-key",
forward_llm_provider_auth_headers=False,
authenticated_with_header="x-litellm-api-key",
)
assert "x-api-key" not in result
def test_clean_headers_strips_x_api_key_when_byok_enabled_but_x_api_key_was_auth_header():
"""
Anti-replay regression: even when forward_llm_provider_auth_headers=True,
if the client authenticated TO the proxy using x-api-key (i.e., the proxy
key arrived as x-api-key), clean_headers() must NOT forward that header
upstream. Otherwise a proxy-auth key would leak to the LLM provider.
"""
headers = Headers(
{
"x-api-key": "sk-proxy-auth-key-masquerading-as-anthropic-key",
"content-type": "application/json",
}
)
result = clean_headers(
headers=headers,
litellm_key_header_name="x-litellm-api-key",
forward_llm_provider_auth_headers=True,
authenticated_with_header="x-api-key",
)
# Even with BYOK enabled, x-api-key must be stripped when it was used
# as the LiteLLM auth header (anti-replay guard).
assert "x-api-key" not in result

View file

@ -1021,6 +1021,83 @@ class TestProxySettingEndpoints:
assert "unsupported_flag" not in stored_settings
assert stored_settings["disable_model_add_for_internal_users"] is False
def test_update_ui_settings_persists_forward_llm_provider_auth_headers(
self, mock_auth, monkeypatch
):
"""BYOK flag must be allowlisted and persisted to litellm_uisettings."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
mock_prisma = MagicMock()
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
payload = {"forward_llm_provider_auth_headers": True}
try:
response = client.patch("/update/ui_settings", json=payload)
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
data = response.json()
assert data["status"] == "success"
assert data["settings"]["forward_llm_provider_auth_headers"] is True
assert mock_prisma.db.litellm_uisettings.upsert.called
call_args = mock_prisma.db.litellm_uisettings.upsert.call_args
create_data = call_args.kwargs["data"]["create"]
stored_settings = json.loads(create_data["ui_settings"])
assert stored_settings["forward_llm_provider_auth_headers"] is True
def test_update_ui_settings_syncs_forward_llm_provider_auth_headers_to_general_settings(
self, mock_auth, monkeypatch
):
"""BYOK flag must be synced into general_settings dict so the request path sees it."""
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
mock_user_auth = UserAPIKeyAuth(
user_id="test-user-123",
user_role=LitellmUserRoles.PROXY_ADMIN,
)
app.dependency_overrides[user_api_key_auth] = lambda: mock_user_auth
monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", True)
# Reset general_settings so the test is hermetic
general_settings: dict = {}
monkeypatch.setattr(
"litellm.proxy.proxy_server.general_settings", general_settings
)
mock_prisma = MagicMock()
mock_prisma.db.litellm_uisettings.upsert = AsyncMock()
mock_prisma.db.litellm_uisettings.find_unique = AsyncMock(return_value=None)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
payload = {"forward_llm_provider_auth_headers": True}
try:
response = client.patch("/update/ui_settings", json=payload)
finally:
app.dependency_overrides.clear()
assert response.status_code == 200
assert general_settings.get("forward_llm_provider_auth_headers") is True
def test_get_sso_settings_from_database(
self, mock_proxy_config, mock_auth, monkeypatch
):

View file

@ -17,6 +17,8 @@ export default function UISettings() {
const disableTeamAdminDeleteProperty = schema?.properties?.disable_team_admin_delete_team_user;
const requireAuthForPublicAIHubProperty = schema?.properties?.require_auth_for_public_ai_hub;
const forwardClientHeadersProperty = schema?.properties?.forward_client_headers_to_llm_api;
const forwardLLMProviderAuthHeadersProperty =
schema?.properties?.forward_llm_provider_auth_headers;
const enableProjectsUIProperty = schema?.properties?.enable_projects_ui;
const enabledPagesProperty = schema?.properties?.enabled_ui_pages_internal_users;
const disableAgentsProperty = schema?.properties?.disable_agents_for_internal_users;
@ -84,6 +86,20 @@ export default function UISettings() {
);
};
const handleToggleForwardLLMProviderAuthHeaders = (checked: boolean) => {
updateSettings(
{ forward_llm_provider_auth_headers: checked },
{
onSuccess: () => {
NotificationManager.success("UI settings updated successfully");
},
onError: (error) => {
NotificationManager.fromBackend(error);
},
},
);
};
const handleToggleEnableProjectsUI = (checked: boolean) => {
updateSettings(
{ enable_projects_ui: checked },
@ -279,7 +295,27 @@ export default function UISettings() {
<Typography.Text strong>Forward client headers to LLM API</Typography.Text>
<Typography.Text type="secondary">
{forwardClientHeadersProperty?.description ??
"If enabled, forwards client headers (e.g. Authorization) to the LLM API. Required for Claude Code with Max subscription."}
"Forwards client headers (Authorization, anthropic-beta, and x-* custom headers) to the upstream LLM. Enable for Claude Code with a Max subscription (forwards the OAuth token) or to pass custom/tracing headers through to the provider. Independent of the BYOK toggle — enable only the one(s) you need."}
</Typography.Text>
</Space>
</Space>
<Space align="start" size="middle">
<Switch
checked={Boolean(values.forward_llm_provider_auth_headers)}
disabled={isUpdating}
loading={isUpdating}
onChange={handleToggleForwardLLMProviderAuthHeaders}
aria-label={
forwardLLMProviderAuthHeadersProperty?.description ??
"Forward LLM provider auth headers"
}
/>
<Space direction="vertical" size={4}>
<Typography.Text strong>Forward LLM provider auth headers</Typography.Text>
<Typography.Text type="secondary">
{forwardLLMProviderAuthHeadersProperty?.description ??
"Forwards provider auth headers (x-api-key, x-goog-api-key, api-key, ocp-apim-subscription-key) to the upstream LLM, overriding any deployment-configured key for that request. Enable for Claude Code BYOK (clients bring their own API key). Independent of the client-headers toggle — enable only the one(s) you need."}
</Typography.Text>
</Space>
</Space>

View file

@ -1,6 +1,5 @@
import React, { useState } from "react";
import { Card, Title, Text, TextInput } from "@tremor/react";
import { List, Empty, Spin, Checkbox } from "antd";
import { Card, List, Empty, Spin, Input, Typography } from "antd";
import { ExperimentOutlined, SearchOutlined } from "@ant-design/icons";
import GuardrailTestPanel from "./GuardrailTestPanel";
import { applyGuardrail } from "../networking";
@ -117,18 +116,18 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
return (
<div className="w-full h-[calc(100vh-200px)]">
<Card className="h-full">
<Card className="h-full" styles={{ body: { padding: 0, height: "100%" } }}>
<div className="flex h-full">
{/* Left Sidebar - Guardrails List */}
<div className="w-1/4 border-r border-gray-200 flex flex-col overflow-hidden">
<div className="p-4 border-b border-gray-200">
<div className="mb-3">
<Title className="text-lg font-semibold mb-3">Guardrails</Title>
<TextInput
icon={SearchOutlined}
<h3 className="text-lg font-semibold mb-3">Guardrails</h3>
<Input
prefix={<SearchOutlined />}
placeholder="Search guardrails..."
value={searchQuery}
onValueChange={setSearchQuery}
onChange={(e) => setSearchQuery(e.target.value)}
/>
</div>
</div>
@ -156,24 +155,14 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
toggleGuardrailSelection(guardrail.guardrail_name);
}
}}
className={`cursor-pointer hover:bg-gray-50 transition-colors px-4 ${
style={{ paddingLeft: 24, paddingRight: 16 }}
className={`cursor-pointer hover:bg-gray-50 transition-colors ${
selectedGuardrails.has(guardrail.guardrail_name || "")
? "bg-blue-50 border-l-4 border-l-blue-500"
: "border-l-4 border-l-transparent"
}`}
>
<List.Item.Meta
avatar={
<Checkbox
checked={selectedGuardrails.has(guardrail.guardrail_name || "")}
onClick={(e) => {
e.stopPropagation();
if (guardrail.guardrail_name) {
toggleGuardrailSelection(guardrail.guardrail_name);
}
}}
/>
}
title={
<div className="flex items-center space-x-2">
<ExperimentOutlined className="text-gray-400" />
@ -206,29 +195,31 @@ const GuardrailTestPlayground: React.FC<GuardrailTestPlaygroundProps> = ({
</div>
<div className="p-3 border-t border-gray-200 bg-gray-50">
<Text className="text-xs text-gray-600">
<Typography.Text className="text-xs text-gray-600">
{selectedGuardrails.size} of {filteredGuardrails.length} selected
</Text>
</Typography.Text>
</div>
</div>
{/* Right Panel - Test Area */}
<div className="w-3/4 flex flex-col bg-white">
<div className="p-4 border-b border-gray-200 flex justify-between items-center">
<Title className="text-xl font-semibold mb-0">Guardrail Testing Playground</Title>
<Typography.Title level={2} className="text-xl font-semibold mb-0">
Guardrail Testing Playground
</Typography.Title>
</div>
<div className="flex-1 overflow-auto p-4">
{selectedGuardrails.size === 0 ? (
<div className="h-full flex flex-col items-center justify-center text-gray-400">
<ExperimentOutlined style={{ fontSize: "48px", marginBottom: "16px" }} />
<Text className="text-lg font-medium text-gray-600 mb-2">
<Typography.Paragraph className="text-lg font-medium text-gray-600 mb-2">
Select Guardrails to Test
</Text>
<Text className="text-center text-gray-500 max-w-md">
</Typography.Paragraph>
<Typography.Paragraph className="text-center text-gray-500 max-w-md">
Choose one or more guardrails from the left sidebar to start testing and
comparing results.
</Text>
</Typography.Paragraph>
</div>
) : (
<div className="h-full">

View file

@ -27,10 +27,6 @@ const MCPPermissionManagement: React.FC<MCPPermissionManagementProps> = ({
// Set initial values when mcpServer changes
useEffect(() => {
if (mcpServer) {
// Set extra_headers if they exist
if (mcpServer.extra_headers) {
form.setFieldValue("extra_headers", mcpServer.extra_headers);
}
if (mcpServer.static_headers) {
const staticHeaders = Object.entries(mcpServer.static_headers).map(([header, value]) => ({
header,

View file

@ -189,6 +189,7 @@ const MCPServerEdit: React.FC<MCPServerEditProps> = ({
...mcpServer,
transport: effectiveTransport,
static_headers: initialStaticHeaders,
extra_headers: mcpServer.extra_headers || [],
oauth_flow_type: mcpServer.token_url ? OAUTH_FLOW.M2M : OAUTH_FLOW.INTERACTIVE,
token_validation_json: mcpServer.token_validation
? JSON.stringify(mcpServer.token_validation, null, 2)

View file

@ -1,5 +1,5 @@
import React from "react";
import { TextInput } from "@tremor/react";
import { Input } from "antd";
interface routingStrategyArgs {
ttl?: number;
@ -42,7 +42,7 @@ const LatencyBasedConfiguration: React.FC<LatencyBasedConfigurationProps> = ({
<p className="text-xs text-gray-500 mt-0.5 mb-2">
{paramExplanation[param] || ""}
</p>
<TextInput
<Input
name={param}
defaultValue={typeof value === "object" ? JSON.stringify(value, null, 2) : value?.toString()}
className="font-mono text-sm w-full"

View file

@ -1,5 +1,5 @@
import React from "react";
import { TextInput } from "@tremor/react";
import { Input } from "antd";
interface ReliabilityRetriesSectionProps {
routerSettings: { [key: string]: any };
@ -36,7 +36,7 @@ const ReliabilityRetriesSection: React.FC<ReliabilityRetriesSectionProps> = ({
<p className="text-xs text-gray-500 mt-0.5 mb-2">
{routerFieldsMetadata[param]?.field_description || ""}
</p>
<TextInput
<Input
name={param}
defaultValue={
value === null || value === undefined || value === "null"

View file

@ -4,25 +4,30 @@ import userEvent from "@testing-library/user-event";
import RouterSettingsForm from "./RouterSettingsForm";
import type { RouterSettingsFormValue } from "./RouterSettingsForm";
// Use the same antd mock as RoutingStrategySelector to keep things consistent
vi.mock("antd", () => ({
Select: Object.assign(
({ value, onChange, children }: any) => (
<select
data-testid="strategy-select"
value={value ?? ""}
onChange={(e) => onChange(e.target.value)}
>
{children}
</select>
),
{
Option: ({ value, children }: any) => (
<option value={value}>{children}</option>
// Override antd Select (complex to drive in JSDOM) while preserving the rest
// of antd (Switch, Button, etc.) so nested components render normally.
vi.mock("antd", async (importOriginal) => {
const actual = await importOriginal<typeof import("antd")>();
return {
...actual,
Select: Object.assign(
({ value, onChange, children }: any) => (
<select
data-testid="strategy-select"
value={value ?? ""}
onChange={(e) => onChange(e.target.value)}
>
{children}
</select>
),
}
),
}));
{
Option: ({ value, children }: any) => (
<option value={value}>{children}</option>
),
}
),
};
});
const defaultValue: RouterSettingsFormValue = {
routerSettings: {},

View file

@ -90,6 +90,7 @@ describe("TagFilteringToggle", () => {
await user.click(screen.getByRole("switch"));
expect(onToggle).toHaveBeenCalledWith(true);
expect(onToggle).toHaveBeenCalledTimes(1);
expect(onToggle.mock.calls[0][0]).toBe(true);
});
});

View file

@ -1,5 +1,5 @@
import React from "react";
import { Switch } from "@tremor/react";
import { Switch } from "antd";
interface TagFilteringToggleProps {
enabled: boolean;

View file

@ -3,24 +3,28 @@ import { renderWithProviders, screen, waitFor } from "../../../tests/test-utils"
import userEvent from "@testing-library/user-event";
import RouterSettings from "./index";
vi.mock("antd", () => ({
Select: Object.assign(
({ value, onChange, children }: any) => (
<select
data-testid="strategy-select"
value={value ?? ""}
onChange={(e) => onChange(e.target.value)}
>
{children}
</select>
),
{
Option: ({ value, children }: any) => (
<option value={value}>{children}</option>
vi.mock("antd", async (importOriginal) => {
const actual = await importOriginal<typeof import("antd")>();
return {
...actual,
Select: Object.assign(
({ value, onChange, children }: any) => (
<select
data-testid="strategy-select"
value={value ?? ""}
onChange={(e) => onChange(e.target.value)}
>
{children}
</select>
),
}
),
}));
{
Option: ({ value, children }: any) => (
<option value={value}>{children}</option>
),
}
),
};
});
vi.mock("@/components/networking", () => ({
getCallbacksCall: vi.fn(),

View file

@ -1,4 +1,4 @@
import { Button } from "@tremor/react";
import { Button } from "antd";
import React, { useEffect, useState } from "react";
import NotificationsManager from "../molecules/notifications_manager";
import { getCallbacksCall, getRouterSettingsCall, setCallbacksCall } from "../networking";
@ -191,10 +191,10 @@ const RouterSettings: React.FC<RouterSettingsProps> = ({ accessToken, userRole,
{/* Actions - Sticky at bottom */}
<div className="border-t border-gray-200 pt-6 flex justify-end gap-3">
<Button variant="secondary" size="sm" onClick={() => window.location.reload()} className="text-sm">
<Button onClick={() => window.location.reload()}>
Reset
</Button>
<Button size="sm" onClick={handleSaveChanges} className="text-sm font-medium">
<Button type="primary" onClick={handleSaveChanges}>
Save Changes
</Button>
</div>

View file

@ -7,8 +7,20 @@ import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { useResetKeySpend } from "@/app/(dashboard)/hooks/keys/useResetKeySpend";
import { KeyResponse, Team } from "../key_team_helpers/key_list";
import { keyUpdateCall } from "../networking";
import KeyInfoView from "./key_info_view";
const editViewMocks = vi.hoisted(() => ({
onSubmit: undefined as ((v: Record<string, any>) => Promise<void>) | undefined,
}));
vi.mock("./key_edit_view", () => ({
KeyEditView: ({ onSubmit }: { onSubmit: (v: Record<string, any>) => Promise<void> }) => {
editViewMocks.onSubmit = onSubmit;
return <div data-testid="key-edit-view-stub" />;
},
}));
vi.mock("@/app/(dashboard)/hooks/useTeams", () => ({
default: vi.fn(),
}));
@ -680,4 +692,89 @@ describe("KeyInfoView", () => {
});
});
});
describe("premium metadata payload normalization", () => {
const enterEditMode = async (keyData: KeyResponse) => {
vi.mocked(useAuthorized).mockReturnValue({
...baseUseAuthorizedMock,
userId: "proxy-admin-user",
userRole: "proxy_admin",
});
renderWithProviders(
<KeyInfoView
keyData={keyData}
onClose={() => {}}
keyId="test-key-id"
onKeyDataUpdate={() => {}}
teams={[]}
/>,
);
await userEvent.click(screen.getByRole("tab", { name: /settings/i }));
await userEvent.click(screen.getByRole("button", { name: /edit settings/i }));
await waitFor(() => expect(editViewMocks.onSubmit).toBeDefined());
};
beforeEach(() => {
editViewMocks.onSubmit = undefined;
vi.mocked(keyUpdateCall).mockClear();
vi.mocked(keyUpdateCall).mockResolvedValue({});
});
it("should drop an empty policies field when the key previously had no policies", async () => {
// Reproduces the real bug: after a successful /key/update, the response echoes
// top-level `policies: []` into client state. Without stripping, the next save
// resends `[]` and trips the premium gate in prepare_metadata_fields.
const keyData: KeyResponse = {
...MOCK_KEY_DATA,
user_id: "proxy-admin-user",
metadata: {},
policies: [],
} as KeyResponse;
await enterEditMode(keyData);
await editViewMocks.onSubmit!({ key: keyData.token, token: keyData.token, policies: [] });
expect(keyUpdateCall).toHaveBeenCalledWith(
expect.anything(),
expect.not.objectContaining({ policies: expect.anything() }),
);
});
it("should keep an empty policies field when the key previously had policies set", async () => {
// Premium users must still be able to clear existing policies by sending `[]`.
const keyData: KeyResponse = {
...MOCK_KEY_DATA,
user_id: "proxy-admin-user",
metadata: { policies: ["existing-policy"] },
policies: ["existing-policy"],
} as KeyResponse;
await enterEditMode(keyData);
await editViewMocks.onSubmit!({ key: keyData.token, token: keyData.token, policies: [] });
expect(keyUpdateCall).toHaveBeenCalledWith(
expect.anything(),
expect.objectContaining({ policies: [] }),
);
});
it("should keep an empty policies field when the previous value lives only at the top level of keyData", async () => {
// Defensive: some premium fields may be present at the top level but not
// mirrored into metadata. A genuine clear must still be forwarded.
const keyData: KeyResponse = {
...MOCK_KEY_DATA,
user_id: "proxy-admin-user",
metadata: {},
policies: ["existing-policy"],
} as KeyResponse;
await enterEditMode(keyData);
await editViewMocks.onSubmit!({ key: keyData.token, token: keyData.token, policies: [] });
expect(keyUpdateCall).toHaveBeenCalledWith(
expect.anything(),
expect.objectContaining({ policies: [] }),
);
});
});
});

View file

@ -34,6 +34,21 @@ interface KeyInfoViewProps {
backButtonText?: string;
}
// Must stay in sync with LiteLLM_ManagementEndpoint_MetadataFields_Premium
// in litellm/proxy/_types.py — limited to fields the key-edit form submits.
const PREMIUM_METADATA_FIELDS = [
"policies",
"guardrails",
"prompts",
"tags",
"allowed_passthrough_routes",
] as const;
const isEmptyValue = (v: unknown): boolean =>
v == null ||
(Array.isArray(v) && v.length === 0) ||
(typeof v === "string" && v.trim() === "");
/**
* ─────────────────────────────────────────────────────────────────────────
* @deprecated
@ -146,6 +161,19 @@ export default function KeyInfoView({
delete formValues.prompts;
}
// Drop premium metadata fields that are empty AND were empty before.
// The /key/update response echoes defaults like `policies: []` back into
// state; without this, the next save resends `[]` and trips the premium
// gate in prepare_metadata_fields for non-premium users.
for (const field of PREMIUM_METADATA_FIELDS) {
const previousValue =
(currentKeyData.metadata as Record<string, unknown> | undefined)?.[field] ??
(currentKeyData as unknown as Record<string, unknown>)[field];
if (isEmptyValue(formValues[field]) && isEmptyValue(previousValue)) {
delete formValues[field];
}
}
// Handle max budget empty string
formValues.max_budget = mapEmptyStringToNull(formValues.max_budget);