mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'upstream/litellm_internal_staging' into litellm_project_rate_limiting
This commit is contained in:
commit
dbf4f9637f
37 changed files with 1472 additions and 127 deletions
|
|
@ -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).
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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: {},
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
import React from "react";
|
||||
import { Switch } from "@tremor/react";
|
||||
import { Switch } from "antd";
|
||||
|
||||
interface TagFilteringToggleProps {
|
||||
enabled: boolean;
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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: [] }),
|
||||
);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue