Merge remote-tracking branch 'origin/litellm_internal_staging' into HEAD

This commit is contained in:
mateo-berri 2026-09-01 11:45:45 -07:00
commit 4ce2b4d5b4
90 changed files with 1789 additions and 1003 deletions

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.62"
version = "0.1.63"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.62"
version = "0.1.63"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -7,6 +7,10 @@ metadata:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: backend
spec:
{{- with .Values.backend.strategy }}
strategy:
{{- toYaml . | nindent 4 }}
{{- end }}
selector:
matchLabels:
{{- include "litellm.backend.selectorLabels" . | nindent 6 }}

View file

@ -7,6 +7,10 @@ metadata:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: gateway
spec:
{{- with .Values.gateway.strategy }}
strategy:
{{- toYaml . | nindent 4 }}
{{- end }}
selector:
matchLabels:
{{- include "litellm.gateway.selectorLabels" . | nindent 6 }}

View file

@ -7,6 +7,8 @@
#
# Running this pre-upgrade closes the window where new application pods would
# otherwise serve traffic against the previous release's unmigrated schema.
# Argo CD users can swap the Helm hook for a PreSync hook through
# `migrationJob.hooks`, which re-runs the Job on every sync.
apiVersion: batch/v1
kind: Job
metadata:
@ -14,10 +16,18 @@ metadata:
labels:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: migrations
{{- if or .Values.migrationJob.hooks.helm.enabled .Values.migrationJob.hooks.argocd.enabled }}
annotations:
{{- if .Values.migrationJob.hooks.helm.enabled }}
helm.sh/hook: pre-install,pre-upgrade
helm.sh/hook-delete-policy: before-hook-creation
helm.sh/hook-weight: "0"
helm.sh/hook-weight: {{ .Values.migrationJob.hooks.helm.weight | default "0" | quote }}
{{- end }}
{{- if .Values.migrationJob.hooks.argocd.enabled }}
argocd.argoproj.io/hook: PreSync
argocd.argoproj.io/hook-delete-policy: BeforeHookCreation
{{- end }}
{{- end }}
spec:
backoffLimit: {{ .Values.migrationJob.backoffLimit }}
ttlSecondsAfterFinished: {{ .Values.migrationJob.ttlSecondsAfterFinished }}

View file

@ -7,6 +7,10 @@ metadata:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: ui
spec:
{{- with .Values.ui.strategy }}
strategy:
{{- toYaml . | nindent 4 }}
{{- end }}
selector:
matchLabels:
{{- include "litellm.ui.selectorLabels" . | nindent 6 }}

View file

@ -0,0 +1,63 @@
suite: test migrations Job hook annotations
templates:
- migrations-job.yaml
values:
- ./values/required.yaml
tests:
- it: runs as a Helm pre-install / pre-upgrade hook by default
asserts:
- equal:
path: metadata.annotations["helm.sh/hook"]
value: pre-install,pre-upgrade
- equal:
path: metadata.annotations["helm.sh/hook-delete-policy"]
value: before-hook-creation
- equal:
path: metadata.annotations["helm.sh/hook-weight"]
value: "0"
- notExists:
path: metadata.annotations["argocd.argoproj.io/hook"]
- it: adds the Argo CD PreSync hook when asked
set:
migrationJob.hooks.argocd.enabled: true
asserts:
- equal:
path: metadata.annotations["argocd.argoproj.io/hook"]
value: PreSync
- equal:
path: metadata.annotations["argocd.argoproj.io/hook-delete-policy"]
value: BeforeHookCreation
- it: drops the Helm hook so Argo CD owns the Job
set:
migrationJob.hooks.argocd.enabled: true
migrationJob.hooks.helm.enabled: false
asserts:
- equal:
path: metadata.annotations["argocd.argoproj.io/hook"]
value: PreSync
- notExists:
path: metadata.annotations["helm.sh/hook"]
- notExists:
path: metadata.annotations["helm.sh/hook-delete-policy"]
- notExists:
path: metadata.annotations["helm.sh/hook-weight"]
- it: renders an ordinary Job when both hooks are disabled
set:
migrationJob.hooks.helm.enabled: false
asserts:
- notExists:
path: metadata.annotations
- equal:
path: kind
value: Job
- it: honours a custom Helm hook weight
set:
migrationJob.hooks.helm.weight: "-5"
asserts:
- equal:
path: metadata.annotations["helm.sh/hook-weight"]
value: "-5"

View file

@ -0,0 +1,66 @@
suite: test rolling update strategy on the component deployments
templates:
- gateway/deployment.yaml
- gateway/configmap.yaml
- backend/deployment.yaml
- ui/deployment.yaml
values:
- ./values/required.yaml
tests:
- it: leaves the strategy to Kubernetes defaults when unset
asserts:
- notExists:
path: spec.strategy
- it: renders the configured strategy on each deployment
set:
gateway.strategy:
type: RollingUpdate
rollingUpdate:
maxUnavailable: 0
maxSurge: 1
backend.strategy:
type: RollingUpdate
rollingUpdate:
maxUnavailable: "25%"
maxSurge: 2
ui.strategy:
type: Recreate
asserts:
- equal:
path: spec.strategy
value:
type: RollingUpdate
rollingUpdate:
maxUnavailable: 0
maxSurge: 1
template: gateway/deployment.yaml
- equal:
path: spec.strategy
value:
type: RollingUpdate
rollingUpdate:
maxUnavailable: 25%
maxSurge: 2
template: backend/deployment.yaml
- equal:
path: spec.strategy
value:
type: Recreate
template: ui/deployment.yaml
- it: keeps a component on the cluster default when only another one sets a strategy
set:
gateway.strategy:
type: Recreate
asserts:
- equal:
path: spec.strategy.type
value: Recreate
template: gateway/deployment.yaml
- notExists:
path: spec.strategy
template: backend/deployment.yaml
- notExists:
path: spec.strategy
template: ui/deployment.yaml

View file

@ -75,6 +75,22 @@ serviceAccounts:
# generate` — the migration engine doesn't need the generated client.
migrationJob:
enabled: true
# Which controller is responsible for running the Job.
#
# `helm.enabled` renders the Helm pre-install / pre-upgrade hook, so the Job
# runs whenever `helm upgrade` sees a change to apply. `argocd.enabled`
# renders an Argo CD PreSync hook instead, which runs the Job on every sync
# even when the rendered manifests are unchanged: the way to re-run
# migrations on demand from a GitOps pipeline. Turning the Helm hook off
# while the Argo CD hook is on leaves the Job out of Helm's own upgrade
# path, which is what Argo CD users want since Argo, not Helm, applies the
# manifests.
hooks:
helm:
enabled: true
weight: "0"
argocd:
enabled: false
backoffLimit: 4
ttlSecondsAfterFinished: 120
# Wall-clock budget for the whole Job, shared across every `backoffLimit`
@ -257,6 +273,15 @@ gateway:
initialDelaySeconds: 5
periodSeconds: 10
timeoutSeconds: 10
# Rolling update tuning for the gateway Deployment. Empty by default, so
# Kubernetes applies its own RollingUpdate defaults (25% maxSurge /
# 25% maxUnavailable). Example, for a surge-only rollout behind a load
# balancer that must never lose capacity:
# type: RollingUpdate
# rollingUpdate:
# maxUnavailable: 0
# maxSurge: 1
strategy: {}
# Optional startupProbe. Empty by default, so existing installs are unchanged
# and liveness/readiness apply from container start. Set it to gate
# liveness/readiness until a slow cold start finishes — a high failureThreshold
@ -369,6 +394,8 @@ backend:
initialDelaySeconds: 5
periodSeconds: 10
timeoutSeconds: 10
# Same shape as gateway.strategy.
strategy: {}
# Optional startupProbe; same shape as gateway.startupProbe. Empty by default.
startupProbe: {}
hpa:
@ -433,6 +460,8 @@ ui:
httpGet: { path: /, port: http }
initialDelaySeconds: 2
periodSeconds: 10
# Same shape as gateway.strategy.
strategy: {}
# Optional startupProbe; same shape as gateway.startupProbe. Empty by default.
startupProbe: {}
hpa:

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.91"
version = "0.4.92"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.91"
version = "0.4.92"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -1910,12 +1910,15 @@ def ocr_cost(
if credits is not None and cost_per_credit is not None:
return cost_per_credit * credits, 0.0
ocr_cost_per_page: float | None = None
if model_info is not None:
ocr_cost_per_page = model_info.get("ocr_cost_per_page")
ocr_cost_per_page: Final = model_info.get("ocr_cost_per_page") if model_info is not None else None
annotation_cost_per_page: Final = model_info.get("annotation_cost_per_page") if model_info is not None else None
annotation_rate: Final = annotation_cost_per_page if annotation_cost_per_page is not None else ocr_cost_per_page
pages_processed: Final = response.usage_info.pages_processed
if pages_processed is None:
annotation_pages: Final = response.usage_info.pages_processed_annotation or 0
has_billable_annotation_pages: Final = annotation_rate is not None and annotation_pages > 0
if pages_processed is None and not has_billable_annotation_pages:
if cost_per_credit is not None or ocr_cost_per_page is None:
# Surface missing usage data instead of silently under-reporting
# cost. The previous behavior raised ValueError; we now return 0.0
@ -1931,7 +1934,7 @@ def ocr_cost(
return 0.0, 0.0
raise ValueError("OCR response pages_processed is None")
if ocr_cost_per_page is None:
if ocr_cost_per_page is None and not has_billable_annotation_pages:
# No per-page pricing configured. Either the model is on credit-based
# pricing (and credits weren't returned, so the credit branch above did
# not match) or the model has no OCR pricing entry at all. Surface a
@ -1947,8 +1950,9 @@ def ocr_cost(
)
return 0.0, 0.0
total_ocr_processing_cost: Final[float] = ocr_cost_per_page * pages_processed
return total_ocr_processing_cost, 0.0
ocr_pages_cost: Final = (ocr_cost_per_page or 0.0) * (pages_processed or 0)
annotation_pages_cost: Final = (annotation_rate or 0.0) * annotation_pages
return ocr_pages_cost + annotation_pages_cost, 0.0
def vector_store_search_cost(

View file

@ -416,25 +416,15 @@ class WebSearchInterceptionLogger(CustomLogger):
if not tools:
return None
is_responses_call: Final = call_type in (CallTypes.responses, CallTypes.aresponses)
has_websearch: Final = (
any(is_web_search_tool_responses(tool) for tool in tools)
if is_responses_call
else any(is_web_search_tool(tool) for tool in tools)
)
if call_type in (CallTypes.responses, CallTypes.aresponses):
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
# Check if any tool is a web search tool (native or already LiteLLM standard)
has_websearch: Final = any(is_web_search_tool(t) for t in tools)
if not has_websearch:
return None
if self.search_tool_name:
try:
from litellm.proxy.proxy_server import llm_router
except ImportError:
llm_router = None
self._select_search_tool_from_router(llm_router=llm_router)
if is_responses_call:
return self._convert_responses_tools(kwargs=kwargs, tools=tools)
verbose_logger.debug("WebSearchInterception: Converting native web_search tools to LiteLLM standard")
# If the client sent an Anthropic-native web_search_* tool, mark the
@ -1641,36 +1631,34 @@ class WebSearchInterceptionLogger(CustomLogger):
return None
def _select_search_tool_from_router(self, llm_router: object) -> "_SearchToolConfig | None":
search_tools: Final = list(getattr(llm_router, "search_tools", []) or [])
if llm_router is None or not hasattr(llm_router, "search_tools"):
return None
search_tools: Final = tuple(getattr(llm_router, "search_tools", None) or ())
return self._select_search_tool_from_list(search_tools=search_tools, source="router")
def _select_search_tool_from_list(
self,
search_tools: list[_SearchToolConfig],
search_tools: Sequence[_SearchToolConfig],
source: str,
) -> "_SearchToolConfig | None":
if self.search_tool_name:
matching_tools = [tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name]
if not matching_tools:
raise ValueError(f"Configured search tool '{self.search_tool_name}' was not found")
selected_tool: Final = matching_tools[0]
litellm_params: Final = selected_tool.get("litellm_params")
selected_search_provider: Final = (
litellm_params.get("search_provider") if isinstance(litellm_params, Mapping) else None
matching_tools: Final = tuple(
tool for tool in search_tools if tool.get("search_tool_name") == self.search_tool_name
)
if not isinstance(selected_search_provider, str) or not selected_search_provider.strip():
raise ValueError(
f"Configured search tool '{self.search_tool_name}' does not define a valid search provider"
if matching_tools:
search_provider = (matching_tools[0].get("litellm_params", {}) or {}).get("search_provider")
verbose_logger.debug(
"WebSearchInterception: Found search tool '%s' from %s with provider '%s'",
self.search_tool_name,
source,
search_provider,
)
return matching_tools[0]
verbose_logger.debug(
"WebSearchInterception: Found search tool '%s' from %s with provider '%s'",
"WebSearchInterception: Search tool '%s' not found in %s, falling back to first available or perplexity",
self.search_tool_name,
source,
selected_search_provider,
)
return selected_tool
if search_tools:
first_tool: Final = search_tools[0]

View file

@ -75,6 +75,7 @@ class OCRUsageInfo(LiteLLMPydanticObjectBase):
"""Usage information from OCR response."""
pages_processed: int | None = None
pages_processed_annotation: int | None = None
credits: float | None = None
doc_size_bytes: int | None = None

View file

@ -4,13 +4,12 @@ VLLM is a superset of OpenAI's `embedding` endpoint.
## `encoding_format`
For OpenAI-compatible embedding calls (including `openai/...` with a custom `api_base` pointing at vLLM), LiteLLM resolves `encoding_format` when it is not set on the request:
For OpenAI-compatible embedding calls (including `openai/...` with a custom `api_base` pointing at vLLM), LiteLLM resolves `encoding_format` when it is not set on the request. `hosted_vllm/...` models use a separate handler that never adds the field on its own, so this resolution applies to the `openai/...`-style routes only:
1. Explicit value on the embedding call (`encoding_format=...`).
2. Model config (`litellm_params.encoding_format` on the proxy `model_list` entry).
3. Environment variable `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT` (e.g. in `.env` or container env).
4. Default **`float`**.
That avoids forwarding `encoding_format=None` to the provider/SDK where some servers behave poorly.
If none of those is set, or the winning value is the literal string `none`, the field is omitted from the upstream request entirely (LiteLLM also bypasses the OpenAI SDK's own base64 default), so OpenAI-compatible servers that reject `encoding_format` keep working.
To pass provider-specific parameters, see [provider-specific params](https://docs.litellm.ai/docs/completion/provider_specific_params).
To pass provider-specific parameters, see [provider-specific params](https://docs.litellm.ai/docs/completion/provider_specific_params).

View file

@ -12,9 +12,14 @@ if TYPE_CHECKING:
import openai
from openai import AsyncOpenAI, OpenAI
from openai._base_client import make_request_options
from openai._constants import RAW_RESPONSE_HEADER
from openai._legacy_response import LegacyAPIResponse
from openai._types import RequestOptions
from openai.types import CreateEmbeddingResponse
from openai.types.beta.assistant_deleted import AssistantDeleted
from openai.types.file_deleted import FileDeleted
from pydantic import BaseModel
from pydantic import BaseModel, TypeAdapter
from typing_extensions import overload
import litellm
@ -329,6 +334,28 @@ class OpenAIChatCompletionResponseIterator(BaseModelResponseIterator):
raise e
_EXTRA_HEADERS_ADAPTER: Final = TypeAdapter(dict[str, str] | None)
_EXTRA_QUERY_ADAPTER: Final = TypeAdapter(dict[str, object] | None)
_NO_EXTRA_HEADERS: Final[Mapping[str, str]] = types.MappingProxyType({})
_SDK_OPTION_KEYS: Final = frozenset(("extra_headers", "extra_query", "extra_body"))
def _embedding_request_without_sdk_defaults(
data: Mapping[str, object], timeout: float | httpx.Timeout
) -> tuple[Mapping[str, object], RequestOptions]:
body: Final = { # mutable-ok: the SDK json-encodes the body and needs a plain dict
k: v for k, v in data.items() if k not in _SDK_OPTION_KEYS
}
extra_headers: Final = _EXTRA_HEADERS_ADAPTER.validate_python(data.get("extra_headers")) or _NO_EXTRA_HEADERS
options: Final = make_request_options(
extra_headers=types.MappingProxyType({**extra_headers, RAW_RESPONSE_HEADER: "true"}),
extra_query=_EXTRA_QUERY_ADAPTER.validate_python(data.get("extra_query")),
extra_body=data.get("extra_body"),
timeout=timeout,
)
return body, options
class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
def __init__(self) -> None:
super().__init__()
@ -1177,19 +1204,15 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
data: dict,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
):
"""
Helper to:
- call embeddings.create.with_raw_response when litellm.return_response_headers is True
- call embeddings.create by default
"""
try:
raw_response = await openai_aclient.embeddings.with_raw_response.create(**data, timeout=timeout)
headers: Final = dict(raw_response.headers)
response: Final = raw_response.parse()
return headers, response
except Exception as e:
raise e
) -> LegacyAPIResponse[CreateEmbeddingResponse]:
if "encoding_format" not in data:
body, options = _embedding_request_without_sdk_defaults(data, timeout)
bypass_response: Final = await openai_aclient.post(
"/embeddings", body=body, options=options, cast_to=CreateEmbeddingResponse
)
assert isinstance(bypass_response, LegacyAPIResponse)
return bypass_response
return await openai_aclient.embeddings.with_raw_response.create(**data, timeout=timeout)
@track_llm_api_timing()
def make_sync_openai_embedding_request(
@ -1198,20 +1221,15 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
data: dict,
timeout: float | httpx.Timeout,
logging_obj: LiteLLMLoggingObj,
):
"""
Helper to:
- call embeddings.create.with_raw_response when litellm.return_response_headers is True
- call embeddings.create by default
"""
try:
raw_response = openai_client.embeddings.with_raw_response.create(**data, timeout=timeout)
headers: Final = dict(raw_response.headers)
response: Final = raw_response.parse()
return headers, response
except Exception as e:
raise e
) -> LegacyAPIResponse[CreateEmbeddingResponse]:
if "encoding_format" not in data:
body, options = _embedding_request_without_sdk_defaults(data, timeout)
bypass_response: Final = openai_client.post(
"/embeddings", body=body, options=options, cast_to=CreateEmbeddingResponse
)
assert isinstance(bypass_response, LegacyAPIResponse)
return bypass_response
return openai_client.embeddings.with_raw_response.create(**data, timeout=timeout)
async def aembedding(
self,
@ -1236,14 +1254,15 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client=client,
shared_session=shared_session,
)
headers, response = await self.make_openai_embedding_request(
raw_response: Final = await self.make_openai_embedding_request(
openai_aclient=openai_aclient,
data=data,
timeout=timeout,
logging_obj=logging_obj,
)
headers: Final = dict(raw_response.headers)
logging_obj.model_call_details["response_headers"] = headers
stringified_response: Final = response.model_dump()
stringified_response: Final = raw_response.parse().model_dump()
## LOGGING
logging_obj.post_call(
input=input,
@ -1335,13 +1354,14 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
)
## embedding CALL
headers: dict | None = None
headers, sync_embedding_response = self.make_sync_openai_embedding_request(
raw_response: Final = self.make_sync_openai_embedding_request(
openai_client=openai_client,
data=data,
timeout=timeout,
logging_obj=logging_obj,
)
headers: Final = dict(raw_response.headers)
sync_embedding_response: Final = raw_response.parse()
## LOGGING
logging_obj.model_call_details["response_headers"] = headers

View file

@ -6292,18 +6292,15 @@ def embedding(
if headers is not None and headers != {}:
optional_params["extra_headers"] = headers
if encoding_format is not None:
optional_params["encoding_format"] = encoding_format
requested_encoding_format: Final = (
encoding_format
or optional_params.get("encoding_format")
or get_secret_str("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT")
)
if requested_encoding_format is None or requested_encoding_format.strip().lower() == "none":
optional_params.pop("encoding_format", None)
else:
env_fmt: Final = get_secret_str("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT")
if env_fmt is not None and env_fmt.strip().lower() == "none":
optional_params.pop("encoding_format", None)
else:
_default_fmt: Final = optional_params.get("encoding_format") or env_fmt or "float"
if _default_fmt.strip().lower() == "none":
optional_params.pop("encoding_format", None)
else:
optional_params["encoding_format"] = _default_fmt
optional_params["encoding_format"] = requested_encoding_format
api_version = None

View file

@ -554,6 +554,7 @@
"supports_vision": true
},
"amazon.nova-sonic-v1:0": {
"deprecation_date": "2026-09-14",
"input_cost_per_audio_token": 3.4e-06,
"input_cost_per_token": 6e-08,
"litellm_provider": "bedrock",
@ -3045,6 +3046,7 @@
"prompt_cache_min_tokens": 2048
},
"azure_ai/claude-fable-5": {
"deprecation_date": "2027-12-05",
"supports_mid_conversation_system": true,
"input_cost_per_token": 1e-05,
"output_cost_per_token": 5e-05,
@ -3078,6 +3080,7 @@
"prompt_cache_min_tokens": 512
},
"azure_ai/claude-opus-5": {
"deprecation_date": "2027-07-08",
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"input_cost_per_token": 5e-06,
@ -3110,6 +3113,7 @@
"prompt_cache_min_tokens": 512
},
"azure_ai/claude-opus-4-8": {
"deprecation_date": "2027-09-01",
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"input_cost_per_token": 5e-06,
@ -3188,6 +3192,7 @@
"prompt_cache_min_tokens": 1024
},
"azure_ai/claude-sonnet-5": {
"deprecation_date": "2027-06-30",
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_1hr": 4e-06,
@ -12287,6 +12292,7 @@
"supports_tool_choice": true
},
"cerebras/zai-glm-4.7": {
"deprecation_date": "2026-08-17",
"input_cost_per_token": 2.25e-06,
"litellm_provider": "cerebras",
"max_input_tokens": 128000,
@ -15101,6 +15107,62 @@
"supports_tool_choice": true,
"supports_vision": true
},
"databricks/databricks-deepseek-v4-flash-0731": {
"cache_creation_input_token_cost": 1.4e-07,
"cache_read_input_token_cost": 2.8e-08,
"input_cost_per_token": 1.4e-07,
"input_dbu_cost_per_token": 2e-06,
"litellm_provider": "databricks",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"metadata": {
"notes": "Input/output cost per token is dbu cost * $0.070. Billing reads the per-token dollar fields; the '*_dbu_cost_per_token' fields are the published Databricks rates, kept for reference. Context/max output are the DeepSeek-published model limits (1M context, 384K max output)."
},
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"output_dbu_cost_per_token": 4e-06,
"source": "https://www.databricks.com/product/pricing/foundation-model-serving",
"supported_modalities": [
"text"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": false
},
"databricks/databricks-deepseek-v4-pro-0813": {
"cache_creation_input_token_cost": 1.31999e-06,
"cache_read_input_token_cost": 1.3202e-07,
"input_cost_per_token": 1.31999e-06,
"input_dbu_cost_per_token": 1.8857e-05,
"litellm_provider": "databricks",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"metadata": {
"notes": "Input/output cost per token is dbu cost * $0.070. Billing reads the per-token dollar fields; the '*_dbu_cost_per_token' fields are the published Databricks rates, kept for reference. Context/max output are the DeepSeek-published model limits (1M context, 384K max output)."
},
"mode": "chat",
"output_cost_per_token": 3.95997e-06,
"output_dbu_cost_per_token": 5.6571e-05,
"source": "https://www.databricks.com/product/pricing/foundation-model-serving",
"supported_modalities": [
"text"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": false
},
"databricks/databricks-gemini-2-5-flash": {
"cache_creation_input_token_cost": 3.0002e-07,
"cache_read_input_token_cost": 3.0002e-08,
@ -20631,6 +20693,7 @@
"supports_image_size": false
},
"gemini-live-2.5-flash-native-audio": {
"deprecation_date": "2026-12-13",
"input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai-language-models",
@ -23852,8 +23915,10 @@
"max_input_tokens": 1024,
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.15,
"source": "https://ai.google.dev/gemini-api/docs/video",
"output_cost_per_second": 0.1,
"output_cost_per_second_1080p": 0.12,
"output_cost_per_second_4k": 0.3,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
"text"
],
@ -23867,7 +23932,8 @@
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.4,
"source": "https://ai.google.dev/gemini-api/docs/video",
"output_cost_per_second_4k": 0.6,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
"text"
],
@ -23895,8 +23961,10 @@
"max_input_tokens": 1024,
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.15,
"source": "https://ai.google.dev/gemini-api/docs/video",
"output_cost_per_second": 0.1,
"output_cost_per_second_1080p": 0.12,
"output_cost_per_second_4k": 0.3,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
"text"
],
@ -23910,7 +23978,8 @@
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.4,
"source": "https://ai.google.dev/gemini-api/docs/video",
"output_cost_per_second_4k": 0.6,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
"text"
],
@ -43440,7 +43509,8 @@
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.4,
"source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/veo/3-1-generate",
"output_cost_per_second_4k": 0.6,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text"
],
@ -43453,8 +43523,10 @@
"max_input_tokens": 1024,
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.15,
"source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/veo/3-1-generate",
"output_cost_per_second": 0.1,
"output_cost_per_second_1080p": 0.12,
"output_cost_per_second_4k": 0.3,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text"
],
@ -43469,7 +43541,8 @@
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.4,
"source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/veo/3-1-generate",
"output_cost_per_second_4k": 0.6,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text"
],
@ -43483,8 +43556,10 @@
"max_input_tokens": 1024,
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.15,
"source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/veo/3-1-generate",
"output_cost_per_second": 0.1,
"output_cost_per_second_1080p": 0.12,
"output_cost_per_second_4k": 0.3,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text"
],
@ -52231,14 +52306,14 @@
"supports_vision": true
},
"fireworks_ai/deepseek-v4-flash-0731": {
"cache_read_input_token_cost": 2.8e-08,
"input_cost_per_token": 1.4e-07,
"cache_read_input_token_cost": 7e-09,
"input_cost_per_token": 2.2e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"output_cost_per_token": 6.6e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -55124,5 +55199,55 @@
"max_tokens": 40960,
"mode": "embedding",
"source": "https://docs.fireworks.ai/serverless/pricing"
},
"zai/glm-5.2": {
"cache_creation_input_token_cost": 0,
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "zai",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"source": "https://docs.z.ai/guides/overview/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"together_ai/Qwen/Qwen3.8-Flash": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "together_ai",
"max_input_tokens": 1000000,
"max_tokens": 1000000,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.together.ai/docs/serverless-models"
},
"cerebras/gemma-4-31b": {
"input_cost_per_token": 9.9e-07,
"litellm_provider": "cerebras",
"max_input_tokens": 131072,
"max_output_tokens": 40960,
"max_tokens": 40960,
"mode": "chat",
"output_cost_per_token": 1.49e-06,
"source": "https://api.cerebras.ai/public/v1/models/gemma-4-31b",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"elevenlabs/scribe_v2": {
"input_cost_per_second": 6.11e-05,
"litellm_provider": "elevenlabs",
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://elevenlabs.io/pricing/api",
"supported_endpoints": [
"/v1/audio/transcriptions"
]
}
}

View file

@ -286,7 +286,7 @@ async def _resolve_mcp_server_identifiers_to_ids(
return resolved
def _rewrite_object_permission_mcp_servers(
def _drop_stale_object_permission_mcp_servers(
object_permission: ObjectPermissionDict,
identifier_to_server_ids: dict[str, set[str]],
) -> None:
@ -294,16 +294,18 @@ def _rewrite_object_permission_mcp_servers(
if not isinstance(mcp_servers, list):
return
normalized_servers: Final[list[str]] = []
for identifier in mcp_servers:
if identifier == SpecialMCPServerNames.no_mcp_servers.value:
normalized_servers.append(SpecialMCPServerNames.no_mcp_servers.value)
continue
normalized_servers.extend(sorted(identifier_to_server_ids.get(identifier, [])))
object_permission["mcp_servers"] = _dedupe_preserving_order(normalized_servers)
# Persist original identifiers, never resolved ids: shared-DB multi-region
# instances each expand a name/alias to their own local server id at read
# time. Only entries resolving to nothing (deleted servers, typos) drop.
kept_servers: Final = [
identifier
for identifier in mcp_servers
if identifier == SpecialMCPServerNames.no_mcp_servers.value or identifier_to_server_ids.get(identifier)
]
object_permission["mcp_servers"] = _dedupe_preserving_order(kept_servers)
def _rewrite_object_permission_mcp_tool_permissions(
def _drop_stale_object_permission_mcp_tool_permissions(
object_permission: ObjectPermissionDict,
identifier_to_server_ids: dict[str, set[str]],
) -> None:
@ -311,31 +313,25 @@ def _rewrite_object_permission_mcp_tool_permissions(
if not isinstance(mcp_tool_permissions, dict):
return
normalized_tool_permissions: Final[dict[str, list[str]]] = {}
for identifier, tools in mcp_tool_permissions.items():
if not isinstance(tools, list):
tools = []
for server_id in sorted(identifier_to_server_ids.get(identifier, [])):
normalized_tool_permissions.setdefault(server_id, [])
normalized_tool_permissions[server_id].extend(tools)
object_permission["mcp_tool_permissions"] = {
server_id: _dedupe_preserving_order(tools) for server_id, tools in normalized_tool_permissions.items()
identifier: _dedupe_preserving_order(tools if isinstance(tools, list) else [])
for identifier, tools in mcp_tool_permissions.items()
if identifier_to_server_ids.get(identifier)
}
def _rewrite_object_permission_mcp_identifiers(
def _drop_stale_object_permission_mcp_identifiers(
object_permission: ObjectPermissionDict | None,
identifier_to_server_ids: dict[str, set[str]],
) -> None:
if not object_permission or not isinstance(object_permission, dict):
return
_rewrite_object_permission_mcp_servers(
_drop_stale_object_permission_mcp_servers(
object_permission=object_permission,
identifier_to_server_ids=identifier_to_server_ids,
)
_rewrite_object_permission_mcp_tool_permissions(
_drop_stale_object_permission_mcp_tool_permissions(
object_permission=object_permission,
identifier_to_server_ids=identifier_to_server_ids,
)
@ -615,7 +611,7 @@ async def validate_key_mcp_servers_against_team(
"validate_key_mcp_servers_against_team: ignoring stale MCP server identifiers (no longer in registry or DB): %s",
sorted(stale_identifiers),
)
_rewrite_object_permission_mcp_identifiers(
_drop_stale_object_permission_mcp_identifiers(
object_permission=object_permission,
identifier_to_server_ids=identifier_to_server_ids,
)

View file

@ -6,6 +6,7 @@ pass/fail actions (allow, block, next, modify_response) and data forwarding.
"""
import time
from collections.abc import Sequence
from typing import Any, Final, Literal
import litellm
@ -114,11 +115,7 @@ class PipelineExecutor:
# Handle terminal actions
if action == "allow":
return PipelineExecutionResult(
terminal_action="allow",
step_results=step_results,
modified_data=working_data if working_data != data else None,
)
return _allow_result(step_results=step_results, working_data=working_data, request_data=data)
if action == "block":
return PipelineExecutionResult(
@ -138,11 +135,7 @@ class PipelineExecutor:
# action == "next" → continue to next step
# Ran out of steps without a terminal action → default allow
return PipelineExecutionResult(
terminal_action="allow",
step_results=step_results,
modified_data=working_data if working_data != data else None,
)
return _allow_result(step_results=step_results, working_data=working_data, request_data=data)
@staticmethod
async def _run_step(
@ -251,6 +244,45 @@ class PipelineExecutor:
return None
def _allow_result(
step_results: Sequence[PipelineStepResult],
working_data: dict, # mutable-ok: same request-payload shape as execute_steps' data
request_data: dict, # mutable-ok: same request-payload shape as execute_steps' data
) -> PipelineExecutionResult:
"""Build the terminal-allow result, propagating pipeline modifications without the per-step guardrail override."""
restored: Final = _restore_request_guardrails(working_data, request_data)
return PipelineExecutionResult(
terminal_action="allow",
step_results=list(step_results), # mutable-ok: PipelineExecutionResult field is a list
modified_data=restored if restored != request_data else None,
)
def _restore_request_guardrails(
working_data: dict, # mutable-ok: same request-payload shape as execute_steps' data
request_data: dict, # mutable-ok: same request-payload shape as execute_steps' data
) -> dict: # mutable-ok: merged back into the request dict, which downstream code mutates
"""
Restore the request's own metadata["guardrails"] activation list.
_run_step overrides it to [step.guardrail] so should_run_guardrail() allows each
step; letting that override escape via modified_data permanently drops every
independently activated guardrail from later lifecycle stages (post_call, etc.).
"""
working_metadata: Final = working_data.get("metadata")
if not isinstance(working_metadata, dict):
return working_data
request_metadata: Final = request_data.get("metadata")
original_guardrails: Final = request_metadata.get("guardrails") if isinstance(request_metadata, dict) else None
stripped: Final = {k: v for k, v in working_metadata.items() if k != "guardrails"} # mutable-ok: request dict
if original_guardrails is not None:
restored: Final = {**stripped, "guardrails": original_guardrails} # mutable-ok: request dict
return {**working_data, "metadata": restored} # mutable-ok: request dict
if not stripped and not isinstance(request_metadata, dict):
return {k: v for k, v in working_data.items() if k != "metadata"} # mutable-ok: request dict
return {**working_data, "metadata": stripped} # mutable-ok: request dict
def _pipeline_action_for_outcome(step: PipelineStep, outcome: str) -> str:
"""
Map pipeline step outcome to the configured action.

View file

@ -45,6 +45,21 @@ def is_custom_tool_call(tool_name: str, custom_tool_names: set[str]) -> bool:
return tool_name in custom_tool_names
def serialize_tool_call_arguments(raw_arguments: object, default: str = "") -> str:
"""Render tool call arguments as the JSON string tool-call schemas require.
Arguments normally arrive already JSON-encoded, but clients and providers
also send the decoded object. ``str()`` on a dict yields a Python repr with
single quotes, which every downstream JSON parser rejects with errors like
"Expecting ',' delimiter".
"""
if isinstance(raw_arguments, str):
return raw_arguments or default
if raw_arguments is None:
return default
return json.dumps(raw_arguments, default=str)
def unwrap_custom_tool_arguments(arguments: str) -> str:
"""Extract the raw content string from JSON-wrapped arguments.

View file

@ -8,6 +8,7 @@ from litellm.main import stream_chunk_builder
from litellm.responses.litellm_completion_transformation.custom_tools import (
build_tool_call_item_kwargs,
extract_custom_tool_names,
serialize_tool_call_arguments,
)
from litellm.responses.litellm_completion_transformation.transformation import (
LiteLLMCompletionResponsesConfig,
@ -213,10 +214,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
fn_args_delta = ""
if isinstance(fn, dict):
fn_name = str(fn.get("name") or "")
fn_args_delta = str(fn.get("arguments") or "")
fn_args_delta = serialize_tool_call_arguments(fn.get("arguments"))
else:
fn_name = str(getattr(fn, "name", "") or "")
fn_args_delta = str(getattr(fn, "arguments", "") or "")
fn_args_delta = serialize_tool_call_arguments(getattr(fn, "arguments", ""))
tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name)
output_index = self._get_or_assign_tool_output_index(call_id)
@ -284,10 +285,10 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator):
fn_args = ""
if isinstance(fn, dict):
fn_name = str(fn.get("name") or "")
fn_args = str(fn.get("arguments") or "")
fn_args = serialize_tool_call_arguments(fn.get("arguments"))
else:
fn_name = str(getattr(fn, "name", "") or "")
fn_args = str(getattr(fn, "arguments", "") or "")
fn_args = serialize_tool_call_arguments(getattr(fn, "arguments", ""))
tool_name, tool_namespace = self._responses_namespace_tool_call_fields(fn_name)
# Track if this is a new tool call that wasn't streamed

View file

@ -93,6 +93,7 @@ from .custom_tools import (
convert_custom_tool_to_function_tool,
extract_custom_tool_names,
is_custom_tool_call,
serialize_tool_call_arguments,
unwrap_custom_tool_arguments,
validated_allowed_callers,
)
@ -1010,7 +1011,7 @@ class LiteLLMCompletionResponsesConfig:
type=cast(Literal["function"], tool_use_type),
function=ChatCompletionToolCallFunctionChunk(
name=str(function.get("name", "")),
arguments=str(function.get("arguments", "{}")),
arguments=serialize_tool_call_arguments(function.get("arguments"), "{}"),
),
index=index,
)
@ -1539,7 +1540,7 @@ class LiteLLMCompletionResponsesConfig:
type=cast(Literal["function"], _tool_use_definition.get("type") or "function"),
function=ChatCompletionToolCallFunctionChunk(
name=function.get("name") or "",
arguments=str(function.get("arguments") or ""),
arguments=serialize_tool_call_arguments(function.get("arguments")),
),
index=0,
)
@ -1589,7 +1590,7 @@ class LiteLLMCompletionResponsesConfig:
type="function",
function=ChatCompletionToolCallFunctionChunk(
name=f"{namespace}__{raw_name}" if qualify else raw_name,
arguments=str(raw_arguments or ""),
arguments=serialize_tool_call_arguments(raw_arguments),
),
index=0,
)
@ -2024,7 +2025,7 @@ class LiteLLMCompletionResponsesConfig:
function_definition = tool.function
tool_name = function_definition.name or ""
tool_id = tool.id or ""
tool_arguments = function_definition.get("arguments") or ""
tool_arguments = serialize_tool_call_arguments(function_definition.get("arguments"))
# Check if this is a custom tool
if is_custom_tool_call(tool_name, custom_tool_names):
@ -2559,7 +2560,7 @@ class LiteLLMCompletionResponsesConfig:
type="function",
function=Function(
name=tool_call.get("name") or "",
arguments=tool_call.get("arguments") or "",
arguments=serialize_tool_call_arguments(tool_call.get("arguments")),
),
)

View file

@ -3561,10 +3561,10 @@ def get_optional_params_embeddings(
non_default_params=non_default_params, optional_params={}, kwargs=kwargs
)
elif custom_llm_provider == "vertex_ai" or custom_llm_provider == "gemini":
# OpenAI SDKs (and litellm's own client) send encoding_format="float"
# by default; float lists are exactly what the vertex API returns, so
# the param is a no-op — don't reject the provider default. Other
# values (e.g. "base64") stay on the unsupported-param path below.
# OpenAI SDKs send encoding_format="float" by default; float lists are
# exactly what the vertex API returns, so the param is a no-op and the
# provider default is not rejected. Other values (e.g. "base64") stay
# on the unsupported-param path below.
if non_default_params.get("encoding_format") == "float":
non_default_params.pop("encoding_format")
supported_params = get_supported_openai_params(

View file

@ -554,6 +554,7 @@
"supports_vision": true
},
"amazon.nova-sonic-v1:0": {
"deprecation_date": "2026-09-14",
"input_cost_per_audio_token": 3.4e-06,
"input_cost_per_token": 6e-08,
"litellm_provider": "bedrock",
@ -3045,6 +3046,7 @@
"prompt_cache_min_tokens": 2048
},
"azure_ai/claude-fable-5": {
"deprecation_date": "2027-12-05",
"supports_mid_conversation_system": true,
"input_cost_per_token": 1e-05,
"output_cost_per_token": 5e-05,
@ -3078,6 +3080,7 @@
"prompt_cache_min_tokens": 512
},
"azure_ai/claude-opus-5": {
"deprecation_date": "2027-07-08",
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"input_cost_per_token": 5e-06,
@ -3110,6 +3113,7 @@
"prompt_cache_min_tokens": 512
},
"azure_ai/claude-opus-4-8": {
"deprecation_date": "2027-09-01",
"supports_mid_conversation_system": true,
"supports_adaptive_thinking": true,
"input_cost_per_token": 5e-06,
@ -3188,6 +3192,7 @@
"prompt_cache_min_tokens": 1024
},
"azure_ai/claude-sonnet-5": {
"deprecation_date": "2027-06-30",
"supports_mid_conversation_system": true,
"cache_creation_input_token_cost": 2.5e-06,
"cache_creation_input_token_cost_above_1hr": 4e-06,
@ -12287,6 +12292,7 @@
"supports_tool_choice": true
},
"cerebras/zai-glm-4.7": {
"deprecation_date": "2026-08-17",
"input_cost_per_token": 2.25e-06,
"litellm_provider": "cerebras",
"max_input_tokens": 128000,
@ -15101,6 +15107,62 @@
"supports_tool_choice": true,
"supports_vision": true
},
"databricks/databricks-deepseek-v4-flash-0731": {
"cache_creation_input_token_cost": 1.4e-07,
"cache_read_input_token_cost": 2.8e-08,
"input_cost_per_token": 1.4e-07,
"input_dbu_cost_per_token": 2e-06,
"litellm_provider": "databricks",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"metadata": {
"notes": "Input/output cost per token is dbu cost * $0.070. Billing reads the per-token dollar fields; the '*_dbu_cost_per_token' fields are the published Databricks rates, kept for reference. Context/max output are the DeepSeek-published model limits (1M context, 384K max output)."
},
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"output_dbu_cost_per_token": 4e-06,
"source": "https://www.databricks.com/product/pricing/foundation-model-serving",
"supported_modalities": [
"text"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": false
},
"databricks/databricks-deepseek-v4-pro-0813": {
"cache_creation_input_token_cost": 1.31999e-06,
"cache_read_input_token_cost": 1.3202e-07,
"input_cost_per_token": 1.31999e-06,
"input_dbu_cost_per_token": 1.8857e-05,
"litellm_provider": "databricks",
"max_input_tokens": 1000000,
"max_output_tokens": 393216,
"max_tokens": 393216,
"metadata": {
"notes": "Input/output cost per token is dbu cost * $0.070. Billing reads the per-token dollar fields; the '*_dbu_cost_per_token' fields are the published Databricks rates, kept for reference. Context/max output are the DeepSeek-published model limits (1M context, 384K max output)."
},
"mode": "chat",
"output_cost_per_token": 3.95997e-06,
"output_dbu_cost_per_token": 5.6571e-05,
"source": "https://www.databricks.com/product/pricing/foundation-model-serving",
"supported_modalities": [
"text"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true,
"supports_vision": false
},
"databricks/databricks-gemini-2-5-flash": {
"cache_creation_input_token_cost": 3.0002e-07,
"cache_read_input_token_cost": 3.0002e-08,
@ -20631,6 +20693,7 @@
"supports_image_size": false
},
"gemini-live-2.5-flash-native-audio": {
"deprecation_date": "2026-12-13",
"input_cost_per_audio_token": 3e-06,
"input_cost_per_token": 5e-07,
"litellm_provider": "vertex_ai-language-models",
@ -23852,8 +23915,10 @@
"max_input_tokens": 1024,
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.15,
"source": "https://ai.google.dev/gemini-api/docs/video",
"output_cost_per_second": 0.1,
"output_cost_per_second_1080p": 0.12,
"output_cost_per_second_4k": 0.3,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
"text"
],
@ -23867,7 +23932,8 @@
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.4,
"source": "https://ai.google.dev/gemini-api/docs/video",
"output_cost_per_second_4k": 0.6,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
"text"
],
@ -23895,8 +23961,10 @@
"max_input_tokens": 1024,
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.15,
"source": "https://ai.google.dev/gemini-api/docs/video",
"output_cost_per_second": 0.1,
"output_cost_per_second_1080p": 0.12,
"output_cost_per_second_4k": 0.3,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
"text"
],
@ -23910,7 +23978,8 @@
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.4,
"source": "https://ai.google.dev/gemini-api/docs/video",
"output_cost_per_second_4k": 0.6,
"source": "https://ai.google.dev/gemini-api/docs/pricing",
"supported_modalities": [
"text"
],
@ -43440,7 +43509,8 @@
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.4,
"source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/veo/3-1-generate",
"output_cost_per_second_4k": 0.6,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text"
],
@ -43453,8 +43523,10 @@
"max_input_tokens": 1024,
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.15,
"source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/veo/3-1-generate",
"output_cost_per_second": 0.1,
"output_cost_per_second_1080p": 0.12,
"output_cost_per_second_4k": 0.3,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text"
],
@ -43469,7 +43541,8 @@
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.4,
"source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/veo/3-1-generate",
"output_cost_per_second_4k": 0.6,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text"
],
@ -43483,8 +43556,10 @@
"max_input_tokens": 1024,
"max_tokens": 1024,
"mode": "video_generation",
"output_cost_per_second": 0.15,
"source": "https://docs.cloud.google.com/gemini-enterprise-agent-platform/models/veo/3-1-generate",
"output_cost_per_second": 0.1,
"output_cost_per_second_1080p": 0.12,
"output_cost_per_second_4k": 0.3,
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing",
"supported_modalities": [
"text"
],
@ -52231,14 +52306,14 @@
"supports_vision": true
},
"fireworks_ai/deepseek-v4-flash-0731": {
"cache_read_input_token_cost": 2.8e-08,
"input_cost_per_token": 1.4e-07,
"cache_read_input_token_cost": 7e-09,
"input_cost_per_token": 2.2e-07,
"litellm_provider": "fireworks_ai",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 2.8e-07,
"output_cost_per_token": 6.6e-07,
"source": "https://docs.fireworks.ai/serverless/pricing",
"supports_function_calling": true,
"supports_reasoning": true,
@ -55124,5 +55199,55 @@
"max_tokens": 40960,
"mode": "embedding",
"source": "https://docs.fireworks.ai/serverless/pricing"
},
"zai/glm-5.2": {
"cache_creation_input_token_cost": 0,
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 1.4e-06,
"litellm_provider": "zai",
"max_input_tokens": 1000000,
"max_output_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.4e-06,
"source": "https://docs.z.ai/guides/overview/pricing",
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"together_ai/Qwen/Qwen3.8-Flash": {
"input_cost_per_token": 1.5e-07,
"litellm_provider": "together_ai",
"max_input_tokens": 1000000,
"max_tokens": 1000000,
"mode": "chat",
"output_cost_per_token": 4.7e-07,
"source": "https://docs.together.ai/docs/serverless-models"
},
"cerebras/gemma-4-31b": {
"input_cost_per_token": 9.9e-07,
"litellm_provider": "cerebras",
"max_input_tokens": 131072,
"max_output_tokens": 40960,
"max_tokens": 40960,
"mode": "chat",
"output_cost_per_token": 1.49e-06,
"source": "https://api.cerebras.ai/public/v1/models/gemma-4-31b",
"supports_function_calling": true,
"supports_parallel_function_calling": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true
},
"elevenlabs/scribe_v2": {
"input_cost_per_second": 6.11e-05,
"litellm_provider": "elevenlabs",
"mode": "audio_transcription",
"output_cost_per_second": 0.0,
"source": "https://elevenlabs.io/pricing/api",
"supported_endpoints": [
"/v1/audio/transcriptions"
]
}
}

View file

@ -1,6 +1,6 @@
[project]
name = "litellm"
version = "1.100.0"
version = "1.101.0"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.15"
@ -67,8 +67,8 @@ proxy = [
"azure-identity>=1.25.2,<2.0",
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.28.1,<2.0",
"litellm-proxy-extras==0.4.91",
"litellm-enterprise==0.1.62",
"litellm-proxy-extras==0.4.92",
"litellm-enterprise==0.1.63",
"RestrictedPython>=8.5,<9.0",
"rich>=13.9.4,<14.0",
"InquirerPy>=0.3.4,<1.0",
@ -319,7 +319,7 @@ members = ["enterprise", "litellm-proxy-extras"]
profile = "black"
[tool.commitizen]
version = "1.100.0"
version = "1.101.0"
version_files = [
"pyproject.toml:^version",
]

View file

@ -9,7 +9,7 @@
"limit": 809
},
"ANN201": {
"limit": 2002
"limit": 2001
},
"ANN202": {
"limit": 841
@ -240,10 +240,10 @@
"limit": 96
},
"TRY201": {
"limit": 405
"limit": 403
},
"TRY203": {
"limit": 113
"limit": 111
},
"TRY300": {
"limit": 855

View file

@ -3,7 +3,7 @@
"limit": 733
},
"TQ002": {
"limit": 742
"limit": 741
},
"TQ003": {
"limit": 62
@ -21,6 +21,6 @@
"limit": 117
},
"TQ008": {
"limit": 11139
"limit": 11135
}
}

View file

@ -30,6 +30,34 @@ def _key(client: McpClient, resources: ResourceManager, *, mcp_servers: list[str
return key
class TestMcpKeyGrantByAlias:
def test_alias_grant_persists_verbatim_and_lists_tools(
self,
client: McpClient,
resources: ResourceManager,
) -> None:
"""A key granted an MCP server by its alias must store the alias, not the
resolved server_id: in a shared-DB multi-region deployment each instance
derives a different id for the same config server, so only the alias
grants access on every region. The same key must still see the server's
tools, proving the alias grant is honored at request time."""
server_id = register_datadog_mcp(client, resources)
client.await_registered(server_id)
alias = next(row.alias for row in client.registered_servers() if row.server_id == server_id)
assert alias, f"registered server {server_id} has no alias to grant by"
key = _key(client, resources, mcp_servers=[alias])
stored = client.proxy.key_info(key).object_permission
assert stored is not None and stored.mcp_servers == [alias], (
f"alias grant was rewritten before persisting (expected [{alias!r}]): "
f"{stored.mcp_servers if stored else None}. A stored server_id is region-local "
f"and breaks the grant on every other instance sharing this database"
)
_ = client.await_tool(key, server_id, SEARCH_LOGS_TOOL)
class TestMcpKeyWithoutAccessIsDenied:
@pytest.mark.covers("mcp.list_tools.api_key.denied_without_permission")
def test_list_tools_denied_without_permission(

View file

@ -114,6 +114,7 @@ class KeyInfo(BaseModel):
budget_id: str | None = None
litellm_budget_table: LiteLLMBudgetTable | None = None
budget_limits: list[BudgetWindowState] | None = None
object_permission: ObjectPermission | None = None
class KeyInfoResponse(BaseModel):

View file

@ -5,6 +5,7 @@ from io import BytesIO
from unittest.mock import AsyncMock
import httpx
import litellm
from litellm import completion, embedding
import pytest
@ -92,44 +93,54 @@ async def test_litellm_gateway_from_sdk_embedding(is_async):
litellm.set_verbose = True
litellm._turn_on_debug()
captured_bodies = []
def handler(request: httpx.Request) -> httpx.Response:
captured_bodies.append(json.loads(request.content))
return httpx.Response(
200,
json={
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"model": "my-vllm-model",
"usage": {"prompt_tokens": 2, "total_tokens": 2},
},
)
if is_async:
from openai import AsyncOpenAI
openai_client = AsyncOpenAI(api_key="fake-key")
mock_method = AsyncMock()
patch_target = openai_client.embeddings.create
openai_client = AsyncOpenAI(
api_key="fake-key",
http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)),
)
response = await litellm.aembedding(
model="litellm_proxy/my-vllm-model",
input="Hello world",
client=openai_client,
api_base="my-custom-api-base",
)
else:
from openai import OpenAI
openai_client = OpenAI(api_key="fake-key")
mock_method = MagicMock()
patch_target = openai_client.embeddings.create
openai_client = OpenAI(
api_key="fake-key",
http_client=httpx.Client(transport=httpx.MockTransport(handler)),
)
response = litellm.embedding(
model="litellm_proxy/my-vllm-model",
input="Hello world",
client=openai_client,
api_base="my-custom-api-base",
)
with patch.object(patch_target.__self__, patch_target.__name__, new=mock_method):
try:
if is_async:
await litellm.aembedding(
model="litellm_proxy/my-vllm-model",
input="Hello world",
client=openai_client,
api_base="my-custom-api-base",
)
else:
litellm.embedding(
model="litellm_proxy/my-vllm-model",
input="Hello world",
client=openai_client,
api_base="my-custom-api-base",
)
except Exception as e:
print(e)
request_body = captured_bodies[0]
print("Request body - {}".format(request_body))
mock_method.assert_called_once()
print("Call KWARGS - {}".format(mock_method.call_args.kwargs))
assert "Hello world" == mock_method.call_args.kwargs["input"]
assert "my-vllm-model" == mock_method.call_args.kwargs["model"]
assert "Hello world" == request_body["input"]
assert "my-vllm-model" == request_body["model"]
assert "encoding_format" not in request_body
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
@pytest.mark.parametrize("is_async", [False, True])

View file

@ -63,27 +63,39 @@ def test_embedding_nvidia_nim():
litellm.set_verbose = True
from openai import OpenAI
captured_bodies = []
def handler(request: httpx.Request) -> httpx.Response:
captured_bodies.append(json.loads(request.content))
return httpx.Response(
200,
json={
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"model": "nvidia/nv-embedqa-e5-v5",
"usage": {"prompt_tokens": 6, "total_tokens": 6},
},
)
client = OpenAI(
api_key="fake-api-key",
http_client=httpx.Client(transport=httpx.MockTransport(handler)),
)
with patch.object(client.embeddings.with_raw_response, "create") as mock_client:
try:
litellm.embedding(
model="nvidia_nim/nvidia/nv-embedqa-e5-v5",
input="What is the meaning of life?",
input_type="passage",
dimensions=1024,
client=client,
)
except Exception as e:
print(e)
mock_client.assert_called_once()
request_body = mock_client.call_args.kwargs
print("request_body: ", request_body)
assert request_body["input"] == "What is the meaning of life?"
assert request_body["model"] == "nvidia/nv-embedqa-e5-v5"
assert request_body["extra_body"]["input_type"] == "passage"
assert request_body["dimensions"] == 1024
response = litellm.embedding(
model="nvidia_nim/nvidia/nv-embedqa-e5-v5",
input="What is the meaning of life?",
input_type="passage",
dimensions=1024,
client=client,
)
request_body = captured_bodies[0]
print("request_body: ", request_body)
assert request_body["input"] == "What is the meaning of life?"
assert request_body["model"] == "nvidia/nv-embedqa-e5-v5"
assert request_body["input_type"] == "passage"
assert request_body["dimensions"] == 1024
assert "encoding_format" not in request_body
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
def test_chat_completion_nvidia_nim_with_tools():

View file

@ -3,6 +3,8 @@ import os
import re
import traceback
import httpx
import openai
import pytest
from dotenv import load_dotenv
@ -1255,56 +1257,42 @@ def test_jina_ai_img_embeddings(input_data, expected_payload_input):
assert sent_data["input"] == expected_payload_input
def test_encoding_format_defaults_to_float_for_openai_sdk(monkeypatch):
def test_encoding_format_omitted_by_default_for_openai_sdk(monkeypatch):
"""
When encoding_format is not provided, LiteLLM sends `float` for OpenAI-path embeddings.
When encoding_format is not provided, LiteLLM leaves it out of the upstream request.
Optional global override: `LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT`.
"""
monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False)
with patch(
"litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client"
) as mock_get_client:
# Create a mock client instance
mock_client_instance = MagicMock()
mock_get_client.return_value = mock_client_instance
captured_bodies = []
# Mock the embeddings.with_raw_response.create method
mock_response = MagicMock()
mock_response.parse.return_value = MagicMock(
model_dump=lambda: {
"data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}],
"model": "text-embedding-ada-002",
def handler(request: httpx.Request) -> httpx.Response:
captured_bodies.append(json.loads(request.content))
return httpx.Response(
200,
json={
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"model": "text-embedding-ada-002",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
}
)
mock_response.headers = {}
mock_client_instance.embeddings.with_raw_response.create.return_value = (
mock_response
},
)
# Call the embedding function without encoding_format
response = embedding(
model="text-embedding-ada-002",
input="Hello world",
)
client = openai.OpenAI(
api_key="sk-test", http_client=httpx.Client(transport=httpx.MockTransport(handler))
)
# Get the call arguments to verify what was sent to OpenAI SDK
call_args = mock_client_instance.embeddings.with_raw_response.create.call_args
assert (
call_args is not None
), "OpenAI SDK embeddings.create should have been called"
response = embedding(
model="text-embedding-ada-002",
input="Hello world",
api_key="sk-test",
client=client,
)
call_kwargs = call_args[1] # Get kwargs
assert "encoding_format" in call_kwargs
assert (
call_kwargs["encoding_format"] == "float"
), "encoding_format should default to float when not provided by user"
print("✅ PASS: encoding_format='float' is correctly passed to OpenAI SDK")
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
assert "encoding_format" not in captured_bodies[0], (
"encoding_format should be omitted from the upstream request when not provided by user"
)
def test_encoding_format_explicit_value_preserved():

View file

@ -5,7 +5,7 @@ import traceback
from typing import Any
import httpx
from openai import AsyncOpenAI, AuthenticationError, BadRequestError, OpenAIError, RateLimitError
from openai import AsyncAzureOpenAI, AsyncOpenAI, AuthenticationError, AzureOpenAI, BadRequestError, OpenAIError, RateLimitError
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
@ -895,7 +895,12 @@ def _pre_call_utils(
):
if call_type == "embedding":
data["input"] = "Hello world!"
mapped_target: Any = client.embeddings.with_raw_response
if isinstance(client, (AzureOpenAI, AsyncAzureOpenAI)):
mapped_target: Any = client.embeddings.with_raw_response
patched_attr = "create"
else:
mapped_target = client
patched_attr = "post"
if sync_mode:
original_function = litellm.embedding
else:
@ -905,6 +910,7 @@ def _pre_call_utils(
if streaming is True:
data["stream"] = True
mapped_target = client.chat.completions.with_raw_response # type: ignore
patched_attr = "create"
if sync_mode:
original_function = litellm.completion
else:
@ -914,12 +920,13 @@ def _pre_call_utils(
if streaming is True:
data["stream"] = True
mapped_target = client.completions.with_raw_response # type: ignore
patched_attr = "create"
if sync_mode:
original_function = litellm.text_completion
else:
original_function = litellm.atext_completion
return data, original_function, mapped_target
return data, original_function, mapped_target, patched_attr
def _pre_call_utils_httpx(
@ -1003,7 +1010,7 @@ async def test_exception_with_headers(sync_mode, provider, model, call_type, str
)
data = {"model": model}
data, original_function, mapped_target = _pre_call_utils(
data, original_function, mapped_target, patched_attr = _pre_call_utils(
call_type=call_type,
data=data,
client=openai_client,
@ -1049,7 +1056,7 @@ async def test_exception_with_headers(sync_mode, provider, model, call_type, str
with patch.object(
mapped_target,
"create",
patched_attr,
side_effect=_return_exception,
):
new_retry_after_mock_client = MagicMock(return_value=-1)

View file

@ -2032,8 +2032,8 @@ def test_router_dynamic_cooldown_correct_retry_after_time():
raise exception
with patch.object(
openai_client.embeddings.with_raw_response,
"create",
openai_client,
"post",
side_effect=_return_exception,
):
new_retry_after_mock_client = MagicMock(return_value=-1)

View file

@ -937,6 +937,14 @@ async def test_pre_request_hook_modifies_request_body():
print("✅ WebSearchInterceptionLogger initialized")
mock_router = MagicMock()
mock_router.search_tools = [
{
"search_tool_name": "test-search-tool",
"litellm_params": {"search_provider": "tavily"},
}
]
# Track what actually gets sent to the API
captured_request = {}
@ -987,6 +995,9 @@ async def test_pre_request_hook_modifies_request_body():
with patch(
"litellm.llms.anthropic.experimental_pass_through.messages.handler.anthropic_messages_handler",
side_effect=mock_anthropic_messages_handler,
), patch( # test-quality-ok: the hook imports this process-global router at call time; no injection seam exists to register search_tools
"litellm.proxy.proxy_server.llm_router",
mock_router,
):
print(

View file

@ -14,7 +14,7 @@ from litellm.integrations.websearch_interception.handler import (
)
from litellm.llms.base_llm.search.transformation import SearchResponse
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, ProxyException, UserAPIKeyAuth
from litellm.types.utils import CallTypes, LlmProviders
from litellm.types.utils import LlmProviders
def test_initialize_from_proxy_config():
@ -230,124 +230,6 @@ async def test_execute_search_passes_selected_search_tool_litellm_params(monkeyp
assert forwarded_kwargs["max_retries"] == 2
@pytest.mark.asyncio
@pytest.mark.parametrize(
("search_tools", "error"),
[
pytest.param(None, "was not found", id="router-not-configured"),
pytest.param(
[{"search_tool_name": "other-search", "litellm_params": {"search_provider": "tavily"}}],
"was not found",
id="requested-tool-not-configured",
),
pytest.param(
[{"search_tool_name": "parallel-search", "litellm_params": "not-a-mapping"}],
"does not define a valid search provider",
id="invalid-parameters",
),
pytest.param(
[{"search_tool_name": "parallel-search", "litellm_params": {}}],
"does not define a valid search provider",
id="missing-provider",
),
pytest.param(
[{"search_tool_name": "parallel-search", "litellm_params": {"search_provider": " "}}],
"does not define a valid search provider",
id="whitespace-provider",
),
pytest.param(
[{"search_tool_name": "parallel-search", "litellm_params": {"search_provider": 123}}],
"does not define a valid search provider",
id="invalid-provider",
),
],
)
async def test_execute_search_rejects_invalid_explicit_search_tool(monkeypatch, search_tools, error):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(search_tool_name="parallel-search")
router = None if search_tools is None else MagicMock(search_tools=search_tools)
mock_asearch = AsyncMock()
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(ValueError, match=f"Configured search tool 'parallel-search' {error}"):
await logger._execute_search("what is litellm")
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_execute_search_honors_explicit_parallel_search_tool(monkeypatch):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(search_tool_name="parallel-search")
router = MagicMock(
search_tools=[
{
"search_tool_name": "other-search",
"litellm_params": {"search_provider": "tavily", "api_key": "other-key"},
},
{
"search_tool_name": "parallel-search",
"litellm_params": {"search_provider": "parallel_ai", "api_key": "parallel-key"},
},
],
)
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
await logger._execute_search("what is litellm")
mock_asearch.assert_awaited_once_with(
query="what is litellm",
search_provider="parallel_ai",
api_key="parallel-key",
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("search_tools", "expected_search_kwargs"),
[
pytest.param(None, {"search_provider": "perplexity"}, id="router-not-configured"),
pytest.param(
[
{
"search_tool_name": "first-search",
"litellm_params": {"search_provider": "tavily", "api_key": "first-key"},
},
{
"search_tool_name": "parallel-search",
"litellm_params": {"search_provider": "parallel_ai", "api_key": "parallel-key"},
},
],
{"search_provider": "tavily", "api_key": "first-key"},
id="first-configured-tool",
),
],
)
async def test_execute_search_preserves_implicit_provider_selection(monkeypatch, search_tools, expected_search_kwargs):
import litellm
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger()
router = None if search_tools is None else MagicMock(search_tools=search_tools)
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
monkeypatch.setattr(proxy_server, "llm_router", router)
monkeypatch.setattr(litellm, "asearch", mock_asearch)
await logger._execute_search("what is litellm")
mock_asearch.assert_awaited_once_with(query="what is litellm", **expected_search_kwargs)
@pytest.mark.asyncio
async def test_execute_search_attributes_spend_to_the_calling_key(monkeypatch):
"""An intercepted search is billed and logged against the key that made the LLM request.
@ -515,72 +397,6 @@ async def test_execute_search_enforces_team_search_tool_permission(monkeypatch):
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
("call_type", "web_search_tool"),
[
pytest.param(
CallTypes.acompletion,
{"type": "web_search_20250305", "name": "web_search"},
id="chat-completion",
),
pytest.param(CallTypes.responses, {"type": "web_search"}, id="responses"),
pytest.param(CallTypes.aresponses, {"type": "web_search"}, id="async-responses"),
pytest.param(
CallTypes.anthropic_messages,
{"type": "web_search_20250305", "name": "web_search"},
id="anthropic-messages",
),
],
)
async def test_deployment_hook_dispatcher_propagates_missing_explicit_search_tool(
monkeypatch, call_type, web_search_tool
):
import litellm
from litellm.proxy import proxy_server
from litellm.utils import async_pre_call_deployment_hook
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="parallel-search")
mock_asearch = AsyncMock()
kwargs = {
"model": "bedrock/claude-sonnet-4",
"tools": [web_search_tool],
"custom_llm_provider": "bedrock",
}
monkeypatch.setattr(
proxy_server,
"llm_router",
MagicMock(search_tools=[{"search_tool_name": "other-search", "litellm_params": {"search_provider": "tavily"}}]),
)
monkeypatch.setattr(litellm, "callbacks", [logger])
monkeypatch.setattr(litellm, "asearch", mock_asearch)
with pytest.raises(ValueError, match="Configured search tool 'parallel-search' was not found"):
await async_pre_call_deployment_hook(kwargs=kwargs, call_type=call_type.value)
assert kwargs["tools"] == [web_search_tool]
mock_asearch.assert_not_awaited()
@pytest.mark.asyncio
async def test_deployment_hook_skips_explicit_tool_validation_for_non_search_responses(monkeypatch):
from litellm.proxy import proxy_server
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="parallel-search")
monkeypatch.setattr(proxy_server, "llm_router", MagicMock(search_tools=[]))
result = await logger.async_pre_call_deployment_hook(
kwargs={
"tools": [{"type": "function", "name": "calculator"}],
"custom_llm_provider": "bedrock",
},
call_type=CallTypes.aresponses,
)
assert result is None
@pytest.mark.asyncio
async def test_async_pre_call_deployment_hook_provider_from_top_level_kwargs():
"""Test that async_pre_call_deployment_hook finds custom_llm_provider at top-level kwargs.

View file

@ -62,6 +62,8 @@ PUBLISHED_DBU_PER_MILLION: Final = {
"databricks/databricks-gemini-2-5-pro": ("22.321", "178.571", "22.321", "2.232"),
"databricks/databricks-gemini-2-5-flash": ("5.357", "44.643", "5.357", "0.536"),
"databricks/databricks-kimi-k3": ("42.857", "214.286", "42.857", "4.286"),
"databricks/databricks-deepseek-v4-flash-0731": ("2.000", "4.000", "2.000", "0.400"),
"databricks/databricks-deepseek-v4-pro-0813": ("18.857", "56.571", "18.857", "1.886"),
"databricks/databricks-glm-5-2": ("20.000", "62.857", "20.000", "3.714"),
}
PROMOTIONAL_DISCOUNT: Final = 0.80

View file

@ -24,6 +24,9 @@ OCR3_MODEL = "mistral/mistral-ocr-2512"
OCR3_COST_PER_PAGE = 0.002
OCR3_ANNOTATION_COST_PER_PAGE = 0.003
AZURE_DOC_AI_MODEL = "azure_ai/mistral-document-ai-2512"
AZURE_DOC_AI_COST_PER_PAGE = 0.003
def _ocr_response(model: str, pages_processed: int) -> OCRResponse:
return OCRResponse(
@ -33,6 +36,14 @@ def _ocr_response(model: str, pages_processed: int) -> OCRResponse:
)
def _annotated_ocr_response(model: str, pages_processed: int | None, annotation_pages: int) -> OCRResponse:
return OCRResponse(
pages=[],
model=model,
usage_info=OCRUsageInfo(pages_processed=pages_processed, pages_processed_annotation=annotation_pages),
)
@pytest.mark.parametrize("model", ["mistral-ocr-4-0", "mistral-ocr-latest"])
def test_model_info_ocr4_price(model: str) -> None:
info = litellm.get_model_info(model=f"mistral/{model}", custom_llm_provider="mistral")
@ -79,3 +90,46 @@ def test_ocr3_cost_scales_with_pages(local_model_cost_map, pages_processed: int)
call_type="ocr",
)
assert cost == pytest.approx(OCR3_COST_PER_PAGE * pages_processed)
def test_ocr3_bills_ocr_and_annotation_pages_at_their_own_rates(local_model_cost_map) -> None:
cost = completion_cost(
completion_response=_annotated_ocr_response("mistral-ocr-2512", 2, 3),
model=OCR3_MODEL,
custom_llm_provider="mistral",
call_type="ocr",
)
assert cost == pytest.approx(2 * OCR3_COST_PER_PAGE + 3 * OCR3_ANNOTATION_COST_PER_PAGE)
def test_ocr3_bills_annotation_only_response(local_model_cost_map) -> None:
cost = completion_cost(
completion_response=_annotated_ocr_response("mistral-ocr-2512", 0, 3),
model=OCR3_MODEL,
custom_llm_provider="mistral",
call_type="ocr",
)
assert cost == pytest.approx(3 * OCR3_ANNOTATION_COST_PER_PAGE)
def test_ocr3_bills_annotation_pages_when_pages_processed_missing(local_model_cost_map) -> None:
cost = completion_cost(
completion_response=_annotated_ocr_response("mistral-ocr-2512", None, 4),
model=OCR3_MODEL,
custom_llm_provider="mistral",
call_type="ocr",
)
assert cost == pytest.approx(4 * OCR3_ANNOTATION_COST_PER_PAGE)
def test_azure_doc_ai_annotation_pages_fall_back_to_ocr_rate(local_model_cost_map) -> None:
info = litellm.get_model_info(model=AZURE_DOC_AI_MODEL, custom_llm_provider="azure_ai")
assert info.get("annotation_cost_per_page") is None
assert info["ocr_cost_per_page"] == AZURE_DOC_AI_COST_PER_PAGE
cost = completion_cost(
completion_response=_annotated_ocr_response("mistral-document-ai-2512", 0, 1),
model=AZURE_DOC_AI_MODEL,
custom_llm_provider="azure_ai",
call_type="ocr",
)
assert cost == pytest.approx(AZURE_DOC_AI_COST_PER_PAGE)

View file

@ -2,6 +2,30 @@ import os
import pytest
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
@pytest.fixture(autouse=True)
def _hermetic_mcp_server_registry():
"""Restore the singleton ``global_mcp_server_manager``'s registry state around every
test, so entries seeded by one test never leak into another on a shared shard."""
saved_registry = dict(global_mcp_server_manager.registry)
saved_config_servers = dict(global_mcp_server_manager.config_mcp_servers)
saved_tool_mapping = dict(global_mcp_server_manager.tool_name_to_mcp_server_name_mapping)
saved_oauth_slots = global_mcp_server_manager._oauth_discovery_slots
try:
yield
finally:
global_mcp_server_manager.registry.clear()
global_mcp_server_manager.registry.update(saved_registry)
global_mcp_server_manager.config_mcp_servers.clear()
global_mcp_server_manager.config_mcp_servers.update(saved_config_servers)
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.clear()
global_mcp_server_manager.tool_name_to_mcp_server_name_mapping.update(saved_tool_mapping)
global_mcp_server_manager._oauth_discovery_slots = saved_oauth_slots
@pytest.fixture(autouse=True)
def _hermetic_server_root_path():

View file

@ -23,6 +23,12 @@ MGMT_MODULE = "litellm.proxy.management_endpoints.mcp_management_endpoints"
@contextlib.contextmanager
def _env_and_reload(**env):
saved = {key: os.environ.get(key) for key in env}
utils_module = importlib.import_module(UTILS_MODULE)
mgmt_module = importlib.import_module(MGMT_MODULE)
# Restore pre-reload module attributes afterwards instead of reloading again:
# a reload re-creates the module's classes, breaking exception identity for
# modules that imported them earlier
snapshots = {module: dict(vars(module)) for module in (utils_module, mgmt_module)}
def _apply_env(values):
for key, value in values.items():
@ -32,8 +38,8 @@ def _env_and_reload(**env):
os.environ[key] = value
def _reload():
utils = importlib.reload(importlib.import_module(UTILS_MODULE))
mgmt = importlib.reload(importlib.import_module(MGMT_MODULE))
utils = importlib.reload(utils_module)
mgmt = importlib.reload(mgmt_module)
return utils, mgmt
try:
@ -41,7 +47,10 @@ def _env_and_reload(**env):
yield _reload()
finally:
_apply_env(saved)
_reload()
for module, snapshot in snapshots.items():
for key in [key for key in vars(module) if key not in snapshot]:
delattr(module, key)
vars(module).update(snapshot)
def test_defaults_used_when_env_unset():

View file

@ -1,11 +1,9 @@
import json
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException
from unittest.mock import AsyncMock, MagicMock, patch
from litellm.proxy._types import (
LiteLLM_ObjectPermissionBase,
LiteLLM_ObjectPermissionTable,
@ -13,10 +11,10 @@ from litellm.proxy._types import (
SpecialMCPServerName,
)
from litellm.proxy.management_helpers.object_permission_utils import (
_drop_stale_object_permission_mcp_servers,
_extract_requested_mcp_access_groups,
_extract_requested_mcp_server_ids,
_resolve_team_allowed_mcp_servers,
_rewrite_object_permission_mcp_servers,
_set_object_permission,
enforce_all_proxy_mcp_servers_grant_is_admin_only,
validate_key_mcp_servers_against_team,
@ -153,10 +151,10 @@ def test_extract_requested_mcp_server_ids_excludes_no_mcp_servers_sentinel():
assert _extract_requested_mcp_server_ids(obj_perm) == {"server-1"}
def test_rewrite_object_permission_mcp_servers_preserves_sentinel():
obj_perm = {"mcp_servers": ["no-mcp-servers", "alias-1"]}
_rewrite_object_permission_mcp_servers(obj_perm, {"alias-1": {"server-1"}})
assert obj_perm["mcp_servers"] == ["no-mcp-servers", "server-1"]
def test_drop_stale_object_permission_mcp_servers_preserves_sentinel_and_alias():
obj_perm = {"mcp_servers": ["no-mcp-servers", "alias-1", "gone-id"]}
_drop_stale_object_permission_mcp_servers(obj_perm, {"alias-1": {"server-1"}, "gone-id": set()})
assert obj_perm["mcp_servers"] == ["no-mcp-servers", "alias-1"]
@pytest.mark.asyncio
@ -692,9 +690,10 @@ async def test_validate_mcp_server_alias_outside_team_scope_raises(
new_callable=AsyncMock,
return_value=[],
)
async def test_validate_mcp_server_alias_is_normalized_before_save(
mock_access_groups, mock_allow_all
):
async def test_validate_mcp_server_alias_persists_verbatim(mock_access_groups, mock_allow_all):
"""Regression for the multi-region shared-DB setup: an alias grant must be
stored as the alias, so every instance can expand it to its own local id.
Rewriting to this instance's server_id breaks access on the other region."""
team_obj = _make_team_obj(mcp_servers=["allowed-server-id"])
object_permission = {
"mcp_servers": ["allowed-alias"],
@ -706,8 +705,27 @@ async def test_validate_mcp_server_alias_is_normalized_before_save(
team_obj=team_obj,
)
assert object_permission["mcp_servers"] == ["allowed-server-id"]
assert object_permission["mcp_tool_permissions"] == {"allowed-server-id": ["tool1"]}
assert object_permission["mcp_servers"] == ["allowed-alias"]
assert object_permission["mcp_tool_permissions"] == {"Allowed Server": ["tool1"]}
def test_alias_grant_expands_on_other_region_after_save():
"""Cross-region flow: the west instance saves an alias grant (its resolver maps
the alias to west's hash-derived id), then the central instance, whose registry
maps the same alias to a different id, expands the persisted grant. Rewriting
to west's id at save time is exactly the regression this guards against."""
west_mgr = _make_mock_mcp_manager(servers=[_make_mock_mcp_server("west-id", alias="github-mcp")])
central_mgr = _make_mock_mcp_manager(servers=[_make_mock_mcp_server("central-id", alias="github-mcp")])
object_permission = {"mcp_servers": ["github-mcp"]}
_drop_stale_object_permission_mcp_servers(object_permission, {"github-mcp": {"west-id"}})
assert object_permission["mcp_servers"] == ["github-mcp"]
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
expand = MCPServerManager.expand_permission_list
assert expand(west_mgr, object_permission["mcp_servers"]) == ["west-id"]
assert expand(central_mgr, object_permission["mcp_servers"]) == ["central-id"]
@pytest.mark.asyncio

View file

@ -749,6 +749,106 @@ async def test_single_step_pipeline_allow(monkeypatch):
assert guard.calls == 1
@pytest.mark.asyncio
async def test_allow_restores_independent_guardrails_list(monkeypatch):
"""
Request activates an independent guardrail; an unrelated pipeline runs and allows.
Expected: no modified_data escapes, so the request's guardrails list survives
and the independent guardrail still runs at later lifecycle stages (post_call).
Regression: LIT-6587 (pipeline clobbered the list with its last step's guardrail).
"""
pipeline_guard = AlwaysPassGuardrail(guardrail_name="input-scan")
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[PipelineStep(guardrail="input-scan", on_fail="block", on_pass="allow")],
)
monkeypatch.setattr(litellm, "callbacks", [pipeline_guard])
data = {
"messages": [{"role": "user", "content": "clean content"}],
"metadata": {"guardrails": ["independent-output-guard"]},
}
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data=data,
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="input-pipeline-policy",
)
assert pipeline_guard.calls == 1
assert result.terminal_action == "allow"
propagated = result.modified_data or data
assert propagated["metadata"]["guardrails"] == ["independent-output-guard"]
assert data["metadata"]["guardrails"] == ["independent-output-guard"]
@pytest.mark.asyncio
async def test_allow_does_not_leak_guardrails_into_bare_request(monkeypatch):
"""A request without metadata must not gain a metadata.guardrails list from the pipeline."""
pipeline_guard = AlwaysPassGuardrail(guardrail_name="input-scan")
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[PipelineStep(guardrail="input-scan", on_fail="block", on_pass="allow")],
)
monkeypatch.setattr(litellm, "callbacks", [pipeline_guard])
data = {"messages": [{"role": "user", "content": "clean content"}]}
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data=data,
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="input-pipeline-policy",
)
assert result.terminal_action == "allow"
propagated = result.modified_data or data
assert "guardrails" not in propagated.get("metadata", {})
assert "metadata" not in data
@pytest.mark.asyncio
async def test_data_forwarding_keeps_changes_and_restores_guardrails_list(monkeypatch):
"""A pass_data pipeline's modifications propagate while the request's guardrails list is restored."""
pii_guard = PiiMaskingGuardrail(guardrail_name="pii-masker")
content_guard = ContentCheckGuardrail(guardrail_name="content-check")
pipeline = GuardrailPipeline(
mode="pre_call",
steps=[
PipelineStep(guardrail="pii-masker", on_fail="block", on_pass="next", pass_data=True),
PipelineStep(guardrail="content-check", on_fail="block", on_pass="allow"),
],
)
monkeypatch.setattr(litellm, "callbacks", [pii_guard, content_guard])
data = {
"messages": [{"role": "user", "content": "Hello John Smith"}],
"metadata": {"guardrails": ["independent-output-guard"]},
}
result = await PipelineExecutor.execute_steps(
steps=pipeline.steps,
mode=pipeline.mode,
data=data,
user_api_key_dict=MagicMock(),
call_type="completion",
policy_name="pii-then-safety",
)
assert result.terminal_action == "allow"
assert result.modified_data is not None
assert result.modified_data["messages"][0]["content"] == "Hello [REDACTED]"
assert result.modified_data["metadata"]["guardrails"] == ["independent-output-guard"]
@pytest.mark.asyncio
async def test_step_results_include_duration(monkeypatch):
"""Step results should include timing information."""

View file

@ -14,6 +14,8 @@ from litellm.proxy.spend_tracking.savings import (
from litellm.router import Router
from litellm.types.utils import Usage
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
def _anthropic_costs(model: str) -> tuple[float, float]:
info = litellm.get_model_info(model=model, custom_llm_provider="anthropic")

View file

@ -985,6 +985,55 @@ class TestFunctionCallTransformation:
assert result[0]["tool_calls"][0]["function"]["arguments"] == "{}"
def test_function_call_transformation_json_encodes_object_arguments(self):
"""A decoded arguments object must be JSON-encoded, not str()'d.
Clients and providers sometimes send `arguments` as an object rather
than a JSON string; `str()` on a dict produces a Python repr with
single quotes, which downstream JSON parsers reject with errors like
"Expecting ',' delimiter".
"""
function_call_item = {
"type": "function_call",
"name": "shell",
"arguments": {"command": "ls", "timeout": 30, "flags": ["-l", "-a"]},
"call_id": "call_123",
"id": "call_123",
"status": "completed",
}
result = LiteLLMCompletionResponsesConfig._transform_responses_api_function_call_to_chat_completion_message(
function_call=function_call_item
)
arguments = result[0].get("tool_calls", [])[0].get("function", {}).get("arguments")
assert json.loads(arguments) == {"command": "ls", "timeout": 30, "flags": ["-l", "-a"]}
assert "'" not in arguments
def test_create_tool_call_chunk_json_encodes_object_arguments(self):
"""Cached tool_call definitions with object arguments stay valid JSON."""
chunk = LiteLLMCompletionResponsesConfig._create_tool_call_chunk(
tool_use_definition={
"id": "call_456",
"type": "function",
"function": {"name": "shell", "arguments": {"command": "ls"}},
},
tool_call_id="call_456",
index=0,
)
assert json.loads(chunk["function"]["arguments"]) == {"command": "ls"}
def test_create_tool_call_chunk_keeps_empty_arguments_default(self):
"""Missing arguments still fall back to an empty JSON object."""
chunk = LiteLLMCompletionResponsesConfig._create_tool_call_chunk(
tool_use_definition={"id": "call_789", "type": "function", "function": {"name": "shell"}},
tool_call_id="call_789",
index=0,
)
assert chunk["function"]["arguments"] == "{}"
def test_complete_input_transformation_with_function_calls(self):
"""Test the complete transformation with the exact input from the issue"""
test_input = [

View file

@ -10,6 +10,7 @@ before response.completed, and that every event of a bridged stream carries the
spend tracking stores, so a follow-up previous_response_id still finds the conversation.
"""
import json
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -523,3 +524,36 @@ async def test_streaming_response_id_falls_back_when_upstream_yields_nothing():
assert response_ids
assert len(set(response_ids)) == 1
assert response_ids[0].startswith("resp_")
def test_object_tool_call_arguments_stream_as_valid_json():
"""A provider that sends decoded object arguments must still stream valid JSON.
`str()` on a dict yields a Python repr with single quotes, which clients
parsing function_call_arguments reject with errors like
"Expecting ',' delimiter".
"""
iterator = LiteLLMCompletionStreamingIterator(
model="test-model",
litellm_custom_stream_wrapper=AsyncMock(),
request_input="Test input",
responses_api_request={},
)
iterator._queue_tool_call_delta_events(
[
{
"index": 0,
"id": "call_obj",
"type": "function",
"function": {"name": "shell", "arguments": {"command": "ls", "flags": ["-l"]}},
}
]
)
streamed_arguments = "".join(
evt.delta
for evt in iterator._pending_tool_events
if evt.type == ResponsesAPIStreamEvents.FUNCTION_CALL_ARGUMENTS_DELTA
)
assert json.loads(streamed_arguments) == {"command": "ls", "flags": ["-l"]}

View file

@ -84,3 +84,41 @@ def test_bare_fireworks_ids_resolve_through_prefixed_entries():
assert info["output_cost_per_token"] == pytest.approx(expected["output_cost_per_token"])
assert info["max_input_tokens"] == expected["max_input_tokens"]
assert info["max_output_tokens"] == expected["max_output_tokens"]
TWIN_PINNED_PRICES = {
"deepseek-v4-flash-0731": {
"input_cost_per_token": 2.2e-07,
"cache_read_input_token_cost": 7e-09,
"output_cost_per_token": 6.6e-07,
},
}
def test_deepseek_v4_flash_0731_twins_pin_published_pricing(model_data):
"""Both 0731 entries carry the price published at docs.fireworks.ai/serverless/pricing."""
for bare_suffix, expected in TWIN_PINNED_PRICES.items():
for key in (
f"fireworks_ai/{bare_suffix}",
f"fireworks_ai/accounts/fireworks/models/{bare_suffix}",
):
entry = model_data[key]
for field, value in expected.items():
assert entry[field] == pytest.approx(value), f"{key}.{field}"
def test_fireworks_account_prefixed_twins_agree_on_price(model_data):
"""Every accounts/fireworks/models/X entry prices identically to its bare fireworks_ai/X twin."""
prefix = "fireworks_ai/accounts/fireworks/models/"
pairs_checked = 0
for key, entry in model_data.items():
if not key.startswith(prefix):
continue
bare_key = f"fireworks_ai/{key[len(prefix):]}"
bare_entry = model_data.get(bare_key)
if bare_entry is None:
continue
pairs_checked += 1
for field in sorted({f for f in (*entry, *bare_entry) if "cost" in f}):
assert entry.get(field) == bare_entry.get(field), f"{key} vs {bare_key}: {field}"
assert pairs_checked >= 20

View file

@ -1,124 +1,121 @@
from unittest.mock import MagicMock, patch
import json
from typing import Final
import httpx
import pytest
import respx
from litellm import embedding
import litellm
@pytest.mark.parametrize(
"set_env, env_value, expected",
[
(False, None, "float"),
(True, "base64", "base64"),
],
)
def test_openai_embedding_encoding_format_default(
monkeypatch, set_env, env_value, expected
):
monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False)
if set_env:
monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_value)
mock_response = MagicMock()
mock_response.parse.return_value = MagicMock(
model_dump=lambda: {
"data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}],
"model": "text-embedding-ada-002",
"object": "list",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
}
def _mock_openai_embedding_route(respx_mock: respx.MockRouter) -> respx.Route:
return respx_mock.post("https://api.openai.com/v1/embeddings").mock(
return_value=httpx.Response(
200,
json={
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 2, "total_tokens": 2},
},
)
)
mock_response.headers = {}
with patch(
"litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client"
) as mock_get_client:
mock_client_instance = MagicMock()
mock_get_client.return_value = mock_client_instance
mock_client_instance.embeddings.with_raw_response.create.return_value = (
mock_response
)
embedding(
model="text-embedding-ada-002",
input="Hello world",
)
@pytest.fixture(autouse=True)
def clear_default_encoding_format_env(monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", raising=False)
call_kwargs = (
mock_client_instance.embeddings.with_raw_response.create.call_args[1]
)
assert call_kwargs["encoding_format"] == expected
def test_embedding_openai_omits_encoding_format_when_client_omits_it(respx_mock: respx.MockRouter) -> None:
mock_route: Final = _mock_openai_embedding_route(respx_mock)
response: Final = litellm.embedding(model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test")
request_body: Final = json.loads(mock_route.calls.last.request.read())
assert "encoding_format" not in request_body
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
def test_embedding_openai_forwards_explicit_encoding_format(respx_mock: respx.MockRouter) -> None:
mock_route: Final = _mock_openai_embedding_route(respx_mock)
litellm.embedding(
model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test", encoding_format="base64"
)
request_body: Final = json.loads(mock_route.calls.last.request.read())
assert request_body["encoding_format"] == "base64"
def test_embedding_openai_explicit_encoding_format_wins_over_env_var(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", "float")
mock_route: Final = _mock_openai_embedding_route(respx_mock)
litellm.embedding(
model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test", encoding_format="base64"
)
request_body: Final = json.loads(mock_route.calls.last.request.read())
assert request_body["encoding_format"] == "base64"
@pytest.mark.parametrize("env_value", ["float", "base64"])
def test_embedding_openai_env_var_sets_default_encoding_format(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, env_value: str
) -> None:
monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_value)
mock_route: Final = _mock_openai_embedding_route(respx_mock)
litellm.embedding(model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test")
request_body: Final = json.loads(mock_route.calls.last.request.read())
assert request_body["encoding_format"] == env_value
@pytest.mark.parametrize("env_none", ["none", "NONE", " none "])
def test_openai_embedding_encoding_format_env_none_omits_param(
monkeypatch, env_none
):
"""LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT=none omits encoding_format (provider default)."""
def test_embedding_openai_env_none_omits_encoding_format(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch, env_none: str
) -> None:
monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", env_none)
mock_route: Final = _mock_openai_embedding_route(respx_mock)
mock_response = MagicMock()
mock_response.parse.return_value = MagicMock(
model_dump=lambda: {
"data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}],
"model": "text-embedding-ada-002",
"object": "list",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
}
litellm.embedding(model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test")
request_body: Final = json.loads(mock_route.calls.last.request.read())
assert "encoding_format" not in request_body
@pytest.mark.asyncio
async def test_aembedding_openai_omits_encoding_format_when_client_omits_it(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
mock_route: Final = _mock_openai_embedding_route(respx_mock)
response: Final = await litellm.aembedding(model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test")
request_body: Final = json.loads(mock_route.calls.last.request.read())
assert "encoding_format" not in request_body
assert response.data[0]["embedding"] == [0.1, 0.2, 0.3]
def test_embedding_openai_omitted_encoding_format_maps_provider_errors(
respx_mock: respx.MockRouter, monkeypatch: pytest.MonkeyPatch
) -> None:
respx_mock.post("https://api.openai.com/v1/embeddings").mock(
return_value=httpx.Response(
429,
headers={"retry-after": "42", "x-should-retry": "false"},
json={"error": {"message": "rate limited", "type": "rate_limit_error"}},
)
)
mock_response.headers = {}
with patch(
"litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client"
) as mock_get_client:
mock_client_instance = MagicMock()
mock_get_client.return_value = mock_client_instance
mock_client_instance.embeddings.with_raw_response.create.return_value = (
mock_response
with pytest.raises(litellm.RateLimitError) as exc_info:
litellm.embedding(
model="openai/text-embedding-3-small", input=["hello"], api_key="sk-test", max_retries=0
)
embedding(
model="text-embedding-ada-002",
input="Hello world",
)
call_kwargs = (
mock_client_instance.embeddings.with_raw_response.create.call_args[1]
)
assert "encoding_format" not in call_kwargs
def test_openai_embedding_encoding_format_explicit_overrides_env(monkeypatch):
"""Request `encoding_format` wins over LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT."""
monkeypatch.setenv("LITELLM_DEFAULT_EMBEDDING_ENCODING_FORMAT", "float")
mock_response = MagicMock()
mock_response.parse.return_value = MagicMock(
model_dump=lambda: {
"data": [{"embedding": [0.1, 0.2, 0.3], "index": 0}],
"model": "text-embedding-ada-002",
"object": "list",
"usage": {"prompt_tokens": 1, "total_tokens": 1},
}
)
mock_response.headers = {}
with patch(
"litellm.llms.openai.openai.OpenAIChatCompletion._get_openai_client"
) as mock_get_client:
mock_client_instance = MagicMock()
mock_get_client.return_value = mock_client_instance
mock_client_instance.embeddings.with_raw_response.create.return_value = (
mock_response
)
embedding(
model="text-embedding-ada-002",
input="Hello world",
encoding_format="base64",
)
call_kwargs = (
mock_client_instance.embeddings.with_raw_response.create.call_args[1]
)
assert call_kwargs["encoding_format"] == "base64"
assert int(exc_info.value.litellm_response_headers["retry-after"]) == 42

View file

@ -532,6 +532,41 @@ class TestVideoGeneration:
assert abs(cost_for("runwayml/seedance2_5", "480p", 8.0) - 1.6) < 0.001
assert abs(cost_for("runwayml/gen4.5", None, 8.0) - 0.96) < 0.001
def test_completion_cost_veo_31_tiers_pin_published_rates(self, monkeypatch):
"""The gemini and vertex_ai veo 3.1 entries bill Google's published per-second tier rates."""
from litellm.cost_calculator import completion_cost
local_map_path = os.path.join(
os.path.dirname(__file__), "..", "..", "model_prices_and_context_window.json"
)
with open(local_map_path, "r") as f:
monkeypatch.setattr(litellm, "model_cost", json.load(f))
def cost_for(model: str, provider: str, resolution: str | None, duration: float) -> float:
mock_response = MagicMock()
mock_response.usage = {
"duration_seconds": duration,
**({"video_resolution": resolution} if resolution else {}),
}
type(mock_response)._hidden_params = {}
return completion_cost(
completion_response=mock_response,
model=model,
call_type="create_video",
custom_llm_provider=provider,
)
for provider in ("gemini", "vertex_ai"):
for suffix in ("generate-preview", "generate-001"):
standard = f"{provider}/veo-3.1-{suffix}"
fast = f"{provider}/veo-3.1-fast-{suffix}"
assert abs(cost_for(standard, provider, None, 8.0) - 3.2) < 1e-6
assert abs(cost_for(standard, provider, "1080p", 8.0) - 3.2) < 1e-6
assert abs(cost_for(standard, provider, "4k", 8.0) - 4.8) < 1e-6
assert abs(cost_for(fast, provider, "720p", 8.0) - 0.8) < 1e-6
assert abs(cost_for(fast, provider, "1080p", 8.0) - 0.96) < 1e-6
assert abs(cost_for(fast, provider, "4k", 8.0) - 2.4) < 1e-6
def test_video_generation_with_files(self):
"""Test video generation with file uploads."""
config = OpenAIVideoConfig()

View file

@ -4,5 +4,8 @@
"complexity": { "max": 140, "target": 80 },
"max-depth": { "max": 70, "target": 30 },
"local/no-large-inline-object-arg": { "max": 559, "target": 300 },
"local/no-long-condition-chain": { "max": 265, "target": 120 }
"local/no-long-condition-chain": { "max": 265, "target": 120 },
"testing-library/no-container": { "max": 133, "target": 50 },
"testing-library/no-node-access": { "max": 716, "target": 500 },
"testing-library/prefer-screen-queries": { "max": 18, "target": 18 }
}

View file

@ -104,10 +104,13 @@ const eslintConfig = [
plugins: { "testing-library": testingLibrary, "jest-dom": jestDom },
rules: {
"testing-library/await-async-queries": "error",
"testing-library/no-container": "warn",
"testing-library/no-node-access": "warn",
"testing-library/no-wait-for-multiple-assertions": "error",
"testing-library/no-wait-for-side-effects": "error",
"testing-library/prefer-find-by": "error",
"testing-library/prefer-presence-queries": "error",
"testing-library/prefer-screen-queries": "warn",
"jest-dom/prefer-checked": "error",
"jest-dom/prefer-empty": "error",
"jest-dom/prefer-enabled-disabled": "error",

View file

@ -13,17 +13,17 @@ describe("APIReferenceView", () => {
it("uses the API doc base url when provided", () => {
const apiDocUrl = "https://docs.litellm.test";
const { getAllByTestId } = render(<APIReferenceView proxySettings={{ LITELLM_UI_API_DOC_BASE_URL: apiDocUrl }} />);
render(<APIReferenceView proxySettings={{ LITELLM_UI_API_DOC_BASE_URL: apiDocUrl }} />);
const codeBlocks = getAllByTestId(codeBlockTestId);
const codeBlocks = screen.getAllByTestId(codeBlockTestId);
expect(codeBlocks[0]).toHaveTextContent(new RegExp(apiDocUrl));
});
it("falls back to the proxy base url when the docs url is missing", () => {
const proxyUrl = "https://proxy.litellm.test";
const { getAllByTestId } = render(<APIReferenceView proxySettings={{ PROXY_BASE_URL: proxyUrl }} />);
render(<APIReferenceView proxySettings={{ PROXY_BASE_URL: proxyUrl }} />);
const codeBlocks = getAllByTestId(codeBlockTestId);
const codeBlocks = screen.getAllByTestId(codeBlockTestId);
expect(codeBlocks[0]).toHaveTextContent(new RegExp(proxyUrl));
});
@ -31,7 +31,7 @@ describe("APIReferenceView", () => {
const apiDocUrl = "https://docs-preferred.litellm.test";
const proxyUrl = "https://proxy-backup.litellm.test";
const { getAllByTestId } = render(
render(
<APIReferenceView
proxySettings={{
LITELLM_UI_API_DOC_BASE_URL: apiDocUrl,
@ -40,7 +40,7 @@ describe("APIReferenceView", () => {
/>,
);
const codeBlocks = getAllByTestId(codeBlockTestId);
const codeBlocks = screen.getAllByTestId(codeBlockTestId);
const renderedCode = codeBlocks[0].textContent ?? "";
expect(renderedCode).toContain(apiDocUrl);
expect(renderedCode).not.toContain(proxyUrl);

View file

@ -1,12 +1,10 @@
import { describe, expect, it } from "vitest";
import RedisTypeSelector from "./RedisTypeSelector";
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
describe("RedisTypeSelector", () => {
it("should render the component", () => {
const { getAllByText } = render(
<RedisTypeSelector redisType="redis" redisTypeDescriptions={{}} onTypeChange={() => {}} />,
);
expect(getAllByText(/Redis/i).length).toBeGreaterThan(0);
render(<RedisTypeSelector redisType="redis" redisTypeDescriptions={{}} onTypeChange={() => {}} />);
expect(screen.getAllByText(/Redis/i).length).toBeGreaterThan(0);
});
});

View file

@ -1,4 +1,4 @@
import { fireEvent, render } from "@testing-library/react";
import { fireEvent, render, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import type { DailyData, KeyMetricWithMetadata, SpendMetrics } from "@/components/UsagePage/types";
@ -79,25 +79,25 @@ const renderWith = (results: DailyData[], overrides: Partial<DailyActivityRange>
describe("CacheLeakageCard", () => {
it("ranks leaking keys by uncached prompt tokens and shows cache hit ratio", () => {
const { getByText, getByLabelText } = renderWith([
renderWith([
dayWithKeys("2026-07-12", {
"hash-caching": key("caching-key", { prompt_tokens: 1000, cache_read_input_tokens: 900 }),
"hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }),
}),
]);
expect(getByText("leaky-key")).toBeInTheDocument();
expect(getByText("0.0%")).toBeInTheDocument();
expect(getByText("90.0%")).toBeInTheDocument();
expect(screen.getByText("leaky-key")).toBeInTheDocument();
expect(screen.getByText("0.0%")).toBeInTheDocument();
expect(screen.getByText("90.0%")).toBeInTheDocument();
[
"Input tokens you sent in this range that weren't served from or written to the cache",
"Share of your input tokens that were served from the cache",
"About how much you'd save if this uncached input used prompt caching. Estimated as uncached input tokens times what your cached traffic already nets per cached token (realized cache savings, after write premiums, ÷ cache read and write tokens). Blank when caching is not currently saving anything overall.",
].forEach((info) => expect(getByLabelText(info)).toBeInTheDocument());
].forEach((info) => expect(screen.getByLabelText(info)).toBeInTheDocument());
});
it("sorts by the clicked column, worst cache hit rate first", () => {
const { getAllByRole, getByText } = renderWith([
renderWith([
dayWithKeys("2026-07-12", {
"hash-a": key("alpha", {
prompt_tokens: 10000,
@ -111,48 +111,48 @@ describe("CacheLeakageCard", () => {
}),
}),
]);
const firstDataRow = () => getAllByRole("row")[1];
const firstDataRow = () => screen.getAllByRole("row")[1];
expect(firstDataRow()).toHaveTextContent("alpha");
fireEvent.click(getByText("Cache hit rate"));
fireEvent.click(screen.getByText("Cache hit rate"));
expect(firstDataRow()).toHaveTextContent("bravo");
fireEvent.click(getByText("Cache hit rate"));
fireEvent.click(screen.getByText("Cache hit rate"));
expect(firstDataRow()).toHaveTextContent("alpha");
});
it("switches to the model view and lists only Anthropic models", () => {
const { getByText, queryByText } = renderWith([
renderWith([
dayWithModels("2026-07-12", {
"claude-sonnet-5": { prompt_tokens: 5000, cache_read_input_tokens: 0 },
"gpt-4o": { prompt_tokens: 8000, cache_read_input_tokens: 0 },
}),
]);
fireEvent.click(getByText("By model"));
fireEvent.click(screen.getByText("By model"));
expect(getByText("Cache leakage by model")).toBeInTheDocument();
expect(getByText("claude-sonnet-5")).toBeInTheDocument();
expect(queryByText("gpt-4o")).not.toBeInTheDocument();
expect(screen.getByText("Cache leakage by model")).toBeInTheDocument();
expect(screen.getByText("claude-sonnet-5")).toBeInTheDocument();
expect(screen.queryByText("gpt-4o")).not.toBeInTheDocument();
});
it("shows an empty state when no key used tokens in the range", () => {
const { getByText, queryByRole } = renderWith([dayWithKeys("2026-07-12", {})]);
renderWith([dayWithKeys("2026-07-12", {})]);
expect(getByText("No key usage in this range.")).toBeInTheDocument();
expect(queryByRole("table")).not.toBeInTheDocument();
expect(screen.getByText("No key usage in this range.")).toBeInTheDocument();
expect(screen.queryByRole("table")).not.toBeInTheDocument();
});
it("tells the user the table is still filling in while fallback pages stream", () => {
const day = dayWithKeys("2026-07-12", {
"hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }),
});
const { getByText, getByRole } = renderWith([day], { isFetchingMore: true });
renderWith([day], { isFetchingMore: true });
expect(getByRole("table")).toBeInTheDocument();
expect(screen.getByRole("table")).toBeInTheDocument();
expect(
getByText("Data is still loading; rows and totals will update as the rest of the range arrives."),
screen.getByText("Data is still loading; rows and totals will update as the rest of the range arrives."),
).toBeInTheDocument();
});
@ -160,10 +160,10 @@ describe("CacheLeakageCard", () => {
const day = dayWithKeys("2026-07-12", {
"hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }),
});
const { queryByText } = renderWith([day], { loading: true });
renderWith([day], { loading: true });
expect(
queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."),
screen.queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."),
).not.toBeInTheDocument();
});
@ -171,10 +171,10 @@ describe("CacheLeakageCard", () => {
const day = dayWithKeys("2026-07-12", {
"hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }),
});
const { queryByText } = renderWith([day]);
renderWith([day]);
expect(
queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."),
screen.queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."),
).not.toBeInTheDocument();
});
});

View file

@ -1,5 +1,5 @@
import React from "react";
import { fireEvent, render, waitFor } from "@testing-library/react";
import { fireEvent, render, waitFor, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
@ -56,7 +56,7 @@ describe("CostOptimizationView daily activity", () => {
useAuthorizedMock.mockReturnValue({ accessToken: "test-token", userId: "u1", userRole: "proxy_admin" });
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
const { getByRole, getByTestId, findByTestId, queryByText } = render(
render(
<QueryClientProvider client={queryClient}>
<CostOptimizationView accessToken="test-token" userId="u1" userRole="proxy_admin" />
</QueryClientProvider>,
@ -64,12 +64,12 @@ describe("CostOptimizationView daily activity", () => {
await waitFor(() => expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalledTimes(1));
fireEvent.click(getByRole("tab", { name: "Prompt Caching" }));
await findByTestId("caching-settings");
fireEvent.click(screen.getByRole("tab", { name: "Prompt Caching" }));
await screen.findByTestId("caching-settings");
expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalledTimes(1);
expect(mockUserDailyActivityCall).not.toHaveBeenCalled();
expect(queryByText(/Currently fetching spend data/)).not.toBeInTheDocument();
expect(screen.queryByText(/Currently fetching spend data/)).not.toBeInTheDocument();
});
it("shows the fetch-progress banner while the paginated fallback streams pages in", async () => {
@ -84,13 +84,13 @@ describe("CostOptimizationView daily activity", () => {
useAuthorizedMock.mockReturnValue({ accessToken: "test-token", userId: "u1", userRole: "proxy_admin" });
const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } });
const { findByText, getByRole } = render(
render(
<QueryClientProvider client={queryClient}>
<CostOptimizationView accessToken="test-token" userId="u1" userRole="proxy_admin" />
</QueryClientProvider>,
);
expect(await findByText(/Currently fetching spend data: fetched 1 \/ 3 pages/)).toBeInTheDocument();
expect(getByRole("button", { name: "Stop" })).toBeInTheDocument();
expect(await screen.findByText(/Currently fetching spend data: fetched 1 \/ 3 pages/)).toBeInTheDocument();
expect(screen.getByRole("button", { name: "Stop" })).toBeInTheDocument();
});
});

View file

@ -1,5 +1,5 @@
import React from "react";
import { fireEvent, render } from "@testing-library/react";
import { fireEvent, render, screen } from "@testing-library/react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
@ -45,32 +45,32 @@ describe("CostOptimizationView", () => {
});
it("renders the standard page header with the sidebar's Cost Optimization icon", () => {
const { container, getByRole, getByText } = renderView();
const { container } = renderView();
expect(getByRole("heading", { level: 1, name: "Cost Optimization" })).toBeInTheDocument();
expect(getByText(/Track and configure the mechanisms that save you money/)).toBeInTheDocument();
expect(screen.getByRole("heading", { level: 1, name: "Cost Optimization" })).toBeInTheDocument();
expect(screen.getByText(/Track and configure the mechanisms that save you money/)).toBeInTheDocument();
expect(container.querySelector(".lucide-piggy-bank")).not.toBeNull();
});
it("renders the four cost-optimization tabs", () => {
const { getByText } = renderView();
renderView();
expect(getByText("Overall")).toBeInTheDocument();
expect(getByText("Prompt Compression")).toBeInTheDocument();
expect(getByText("Prompt Caching")).toBeInTheDocument();
expect(getByText("Auto-Router")).toBeInTheDocument();
expect(screen.getByText("Overall")).toBeInTheDocument();
expect(screen.getByText("Prompt Compression")).toBeInTheDocument();
expect(screen.getByText("Prompt Caching")).toBeInTheDocument();
expect(screen.getByText("Auto-Router")).toBeInTheDocument();
});
it("defaults to the Overall tab and switches the active tab on click", () => {
const { getByRole } = renderView();
renderView();
expect(getByRole("tab", { name: "Overall" })).toHaveAttribute("aria-selected", "true");
expect(getByRole("tab", { name: "Prompt Compression" })).toHaveAttribute("aria-selected", "false");
expect(screen.getByRole("tab", { name: "Overall" })).toHaveAttribute("aria-selected", "true");
expect(screen.getByRole("tab", { name: "Prompt Compression" })).toHaveAttribute("aria-selected", "false");
fireEvent.click(getByRole("tab", { name: "Prompt Compression" }));
fireEvent.click(screen.getByRole("tab", { name: "Prompt Compression" }));
expect(getByRole("tab", { name: "Overall" })).toHaveAttribute("aria-selected", "false");
expect(getByRole("tab", { name: "Prompt Compression" })).toHaveAttribute("aria-selected", "true");
expect(screen.getByRole("tab", { name: "Overall" })).toHaveAttribute("aria-selected", "false");
expect(screen.getByRole("tab", { name: "Prompt Compression" })).toHaveAttribute("aria-selected", "true");
});
// Unlike the other three pages in this cleanup, Cost Optimization keeps its
@ -80,21 +80,21 @@ describe("CostOptimizationView", () => {
// are proxy-admin-only, so those are what disappear.
describe("proxy-admin-only tabs", () => {
it.each(["Internal User", "Internal Viewer", "Org Admin"])("shows %s the Overall tab only", (userRole) => {
const { getByRole, queryByRole } = renderView(userRole);
renderView(userRole);
expect(getByRole("tab", { name: "Overall" })).toBeInTheDocument();
expect(queryByRole("tab", { name: "Prompt Compression" })).not.toBeInTheDocument();
expect(queryByRole("tab", { name: "Prompt Caching" })).not.toBeInTheDocument();
expect(queryByRole("tab", { name: "Auto-Router" })).not.toBeInTheDocument();
expect(screen.getByRole("tab", { name: "Overall" })).toBeInTheDocument();
expect(screen.queryByRole("tab", { name: "Prompt Compression" })).not.toBeInTheDocument();
expect(screen.queryByRole("tab", { name: "Prompt Caching" })).not.toBeInTheDocument();
expect(screen.queryByRole("tab", { name: "Auto-Router" })).not.toBeInTheDocument();
});
it("never mounts the panels behind the admin-only endpoints for an internal user", () => {
const { getByTestId, queryByTestId } = renderView("Internal User");
renderView("Internal User");
expect(getByTestId("usage-tab")).toBeInTheDocument();
expect(queryByTestId("compression-tab")).not.toBeInTheDocument();
expect(queryByTestId("caching-tab")).not.toBeInTheDocument();
expect(queryByTestId("autorouter-benchmarks-tab")).not.toBeInTheDocument();
expect(screen.getByTestId("usage-tab")).toBeInTheDocument();
expect(screen.queryByTestId("compression-tab")).not.toBeInTheDocument();
expect(screen.queryByTestId("caching-tab")).not.toBeInTheDocument();
expect(screen.queryByTestId("autorouter-benchmarks-tab")).not.toBeInTheDocument();
});
});
});

View file

@ -1,4 +1,4 @@
import { render, waitFor } from "@testing-library/react";
import { render, waitFor, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
const mockGetGeneralSettingsCall = vi.fn();
@ -37,10 +37,10 @@ describe("PromptCachingTab", () => {
cancelled: false,
cancel: vi.fn(),
};
const { getByTestId } = render(<PromptCachingTab accessToken="test-token" activity={activity} />);
render(<PromptCachingTab accessToken="test-token" activity={activity} />);
expect(getByTestId("caching-settings")).toBeInTheDocument();
expect(getByTestId("cache-leakage-card")).toBeInTheDocument();
expect(screen.getByTestId("caching-settings")).toBeInTheDocument();
expect(screen.getByTestId("cache-leakage-card")).toBeInTheDocument();
await waitFor(() => expect(mockCacheLeakageCard).toHaveBeenCalledWith(expect.objectContaining({ activity })));
});
});

View file

@ -1,4 +1,4 @@
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import type { ToolSpendResponse } from "@/components/networking";
@ -152,13 +152,13 @@ describe("UsageTab", () => {
gateway_injected_caching_savings_spend: 0.006,
compression_saved_tokens: 100000,
};
const { getByText } = renderWith([day("2026-07-12", firstDay), day("2026-07-13", secondDay)]);
renderWith([day("2026-07-12", firstDay), day("2026-07-13", secondDay)]);
expect(getByText("$0.1500")).toBeInTheDocument();
expect(getByText("$0.1400")).toBeInTheDocument();
expect(getByText("$0.0100")).toBeInTheDocument();
expect(getByText("$0.0160")).toBeInTheDocument();
expect(getByText("140,000 tokens compressed")).toBeInTheDocument();
expect(screen.getByText("$0.1500")).toBeInTheDocument();
expect(screen.getByText("$0.1400")).toBeInTheDocument();
expect(screen.getByText("$0.0100")).toBeInTheDocument();
expect(screen.getByText("$0.0160")).toBeInTheDocument();
expect(screen.getByText("140,000 tokens compressed")).toBeInTheDocument();
});
const twoDays = () => [
@ -167,11 +167,11 @@ describe("UsageTab", () => {
];
it("opens on a running total anchored at $0 at the start of the range", () => {
const { getByTestId } = renderWith(twoDays());
renderWith(twoDays());
// Cumulative prepends a synthetic $0 point at the range start (Jul 1) so the
// line rises from zero rather than floating; the daily running totals follow.
const series = readSeries(getByTestId("area-chart"));
const series = readSeries(screen.getByTestId("area-chart"));
expect(series).toHaveLength(3);
expect(series[0]).toMatchObject({ date: "Jul 1", Compression: 0, "Prompt caching": 0 });
expect(series[1]).toMatchObject({ Compression: 0.04, "Prompt caching": 0.006 });
@ -183,12 +183,12 @@ describe("UsageTab", () => {
// The original complaint: a one-day range plotted a single floating dot. The
// synthetic start anchor gives the line a zero origin to climb from.
const oneDay = new Date(2026, 6, 24);
const { getByTestId } = renderWith(
[day("2026-07-24", { compression_savings_spend: 0.2, gateway_injected_caching_savings_spend: 0.05 })],
{ from: oneDay, to: oneDay },
);
renderWith([day("2026-07-24", { compression_savings_spend: 0.2, gateway_injected_caching_savings_spend: 0.05 })], {
from: oneDay,
to: oneDay,
});
const series = readSeries(getByTestId("area-chart"));
const series = readSeries(screen.getByTestId("area-chart"));
expect(series).toHaveLength(2);
expect(series[0]).toMatchObject({ date: "Jul 24", Compression: 0, "Prompt caching": 0 });
expect(series[1]).toMatchObject({ date: "Jul 24", Compression: 0.2, "Prompt caching": 0.05 });
@ -202,49 +202,49 @@ describe("UsageTab", () => {
day("2026-07-13", { gateway_injected_caching_savings_spend: 0.1 }),
day("2026-07-12", { gateway_injected_caching_savings_spend: 0.04 }),
];
const { getByTestId, getByRole } = renderWith(newestFirst);
renderWith(newestFirst);
// The $0 anchor leads, then the days climb oldest to newest.
const cumulative = readSeries(getByTestId("area-chart"));
const cumulative = readSeries(screen.getByTestId("area-chart"));
expect(cumulative.map((p: { date: string }) => p.date)).toEqual(["Jul 1", "Jul 12", "Jul 13"]);
expect(cumulative[1]["Prompt caching"]).toBeCloseTo(0.04, 5);
expect(cumulative[2]["Prompt caching"]).toBeCloseTo(0.14, 5);
expect(cumulative[2]["Prompt caching"]).toBeGreaterThan(cumulative[1]["Prompt caching"]);
await userEvent.click(getByRole("tab", { name: "Per day" }));
const perDay = readSeries(getByTestId("bar-chart"));
await userEvent.click(screen.getByRole("tab", { name: "Per day" }));
const perDay = readSeries(screen.getByTestId("bar-chart"));
expect(perDay.map((p: { date: string }) => p.date)).toEqual(["Jul 12", "Jul 13"]);
});
it("draws bars of the raw per-interval readings on the other tab", async () => {
const { getByRole, getByTestId, queryByTestId } = renderWith(twoDays());
renderWith(twoDays());
// Cumulative opens on the area line.
expect(getByTestId("area-chart")).toBeInTheDocument();
expect(screen.getByTestId("area-chart")).toBeInTheDocument();
await userEvent.click(getByRole("tab", { name: "Per day" }));
await userEvent.click(screen.getByRole("tab", { name: "Per day" }));
// Per day switches to a bar chart of the unaccumulated daily savings, with no
// synthetic anchor prepended.
expect(queryByTestId("area-chart")).not.toBeInTheDocument();
const series = readSeries(getByTestId("bar-chart"));
expect(screen.queryByTestId("area-chart")).not.toBeInTheDocument();
const series = readSeries(screen.getByTestId("bar-chart"));
expect(series).toHaveLength(2);
expect(series[0]).toMatchObject({ Compression: 0.04, "Prompt caching": 0.006 });
expect(series[1]).toMatchObject({ Compression: 0.1, "Prompt caching": 0.01 });
});
it("says what the line means and over what range", async () => {
const { getByText, getByRole } = renderWith(twoDays());
renderWith(twoDays());
expect(getByText("Running total saved · Jul 1 – Jul 14 (UTC)")).toBeInTheDocument();
await userEvent.click(getByRole("tab", { name: "Per day" }));
expect(getByText("Saved per day · Jul 1 – Jul 14 (UTC)")).toBeInTheDocument();
expect(screen.getByText("Running total saved · Jul 1 – Jul 14 (UTC)")).toBeInTheDocument();
await userEvent.click(screen.getByRole("tab", { name: "Per day" }));
expect(screen.getByText("Saved per day · Jul 1 – Jul 14 (UTC)")).toBeInTheDocument();
});
it("builds the per-driver donut from the range totals, not the running total", () => {
const { getByTestId } = renderWith(twoDays());
renderWith(twoDays());
const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
const slices = JSON.parse(screen.getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
expect(slices).toEqual([
{ driver: "Compression", color: "emerald", usd: expect.closeTo(0.14, 5) },
{ driver: "Prompt caching", color: "blue", usd: expect.closeTo(0.016, 5) },
@ -252,9 +252,9 @@ describe("UsageTab", () => {
});
it("omits a driver slice when that driver has no savings", () => {
const { getByTestId } = renderWith([day("2026-07-12", { compression_savings_spend: 0.04 })]);
renderWith([day("2026-07-12", { compression_savings_spend: 0.04 })]);
const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
const slices = JSON.parse(screen.getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
expect(slices).toEqual([{ driver: "Compression", color: "emerald", usd: expect.closeTo(0.04, 5) }]);
});
@ -262,7 +262,7 @@ describe("UsageTab", () => {
// Stacking sums the series into one bar. Auto-router savings go negative when a
// model switch pays for a cold cache, and that segment would be drawn below the
// axis while the rest of the bar still read as the day's total.
const { getByRole, getByTestId } = renderWith([
renderWith([
day("2026-07-12", {
compression_savings_spend: 0.1,
gateway_injected_caching_savings_spend: 0.02,
@ -270,8 +270,8 @@ describe("UsageTab", () => {
}),
]);
await userEvent.click(getByRole("tab", { name: "Per day" }));
const bars = getByTestId("bar-chart");
await userEvent.click(screen.getByRole("tab", { name: "Per day" }));
const bars = screen.getByTestId("bar-chart");
expect(bars).toHaveAttribute("data-stack", "false");
expect(readSeries(bars)[0]).toMatchObject({ "Auto-router": -0.05 });
});
@ -281,10 +281,10 @@ describe("UsageTab", () => {
// per day"). Hand-rolled rows made it compete with the legend and the toggle for
// width, so the header grew a line on one tab and the chart moved with it. CardHeader
// sizes the action column to its content and gives the rest to the title column.
const { getByRole, getByTestId, container } = renderWith(twoDays());
const { container } = renderWith(twoDays());
const header = () => {
const legend = getByTestId("chart-legend");
const legend = screen.getByTestId("chart-legend");
const action = legend.closest('[data-slot="card-action"]') as HTMLElement;
const cardHeader = action.parentElement as HTMLElement;
const description = cardHeader.querySelector('[data-slot="card-description"]') as HTMLElement;
@ -295,12 +295,12 @@ describe("UsageTab", () => {
expect(before.action).toBeTruthy();
expect(before.description).toBeTruthy();
// the toggle rides in the same action slot as the legend, so neither moves alone
expect(before.action.contains(getByRole("tablist"))).toBe(true);
expect(before.action.contains(screen.getByRole("tablist"))).toBe(true);
// the subtitle lives outside that slot, so its length cannot reposition the controls
expect(before.action.contains(before.description)).toBe(false);
expect(before.description).toHaveTextContent(/Running total saved/);
await userEvent.click(getByRole("tab", { name: "Per day" }));
await userEvent.click(screen.getByRole("tab", { name: "Per day" }));
const after = header();
expect(after.action).toBe(before.action);
@ -314,7 +314,7 @@ describe("UsageTab", () => {
// Switching models leaves the new one with a cold cache, so a route can cost more
// than the baseline would have. A negative slice is meaningless in a donut, but the
// total has to keep the loss or the page can only ever report good news.
const { getByText, getByTestId } = renderWith([
renderWith([
day("2026-07-12", {
compression_savings_spend: 0.1,
gateway_injected_caching_savings_spend: 0.02,
@ -322,16 +322,16 @@ describe("UsageTab", () => {
}),
]);
expect(getByText("$0.0700")).toBeInTheDocument();
expect(getByText("-$0.0500")).toBeInTheDocument();
expect(screen.getByText("$0.0700")).toBeInTheDocument();
expect(screen.getByText("-$0.0500")).toBeInTheDocument();
const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
const slices = JSON.parse(screen.getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
expect(slices.map((d: { driver: string }) => d.driver)).toEqual(["Compression", "Prompt caching"]);
expect(getByTestId("donut-chart")).toHaveAttribute("data-label", "$0.1200");
expect(screen.getByTestId("donut-chart")).toHaveAttribute("data-label", "$0.1200");
});
it("carries auto-router savings into the summary card, donut slice, and cumulative series", () => {
const { getByText, getByTestId } = renderWith([
renderWith([
day("2026-07-12", {
compression_savings_spend: 0.04,
gateway_injected_caching_savings_spend: 0.006,
@ -345,11 +345,11 @@ describe("UsageTab", () => {
]);
// Total saved now sums three drivers, and the auto-router card carries its own total.
expect(getByText("$0.2260")).toBeInTheDocument();
expect(getByText("$0.0700")).toBeInTheDocument();
expect(screen.getByText("$0.2260")).toBeInTheDocument();
expect(screen.getByText("$0.0700")).toBeInTheDocument();
// The driver donut gains a third slice priced from the range totals.
const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
const slices = JSON.parse(screen.getByTestId("donut-chart").getAttribute("data-slices") ?? "[]");
expect(slices).toEqual([
{ driver: "Compression", color: "emerald", usd: expect.closeTo(0.14, 5) },
{ driver: "Prompt caching", color: "blue", usd: expect.closeTo(0.016, 5) },
@ -357,7 +357,7 @@ describe("UsageTab", () => {
]);
// And the cumulative line accumulates the auto-router series alongside the others.
const series = readSeries(getByTestId("area-chart"));
const series = readSeries(screen.getByTestId("area-chart"));
expect(series[2]["Auto-router"]).toBeCloseTo(0.07, 5);
});
@ -371,9 +371,9 @@ describe("UsageTab", () => {
start_date: "2026-07-12",
end_date: "2026-07-12",
};
const { findAllByTestId } = renderWith([day("2026-07-12", {})], { toolSpend });
renderWith([day("2026-07-12", {})], { toolSpend });
const bars = await findAllByTestId("bar-chart");
const bars = await screen.findAllByTestId("bar-chart");
const series = JSON.parse(bars[0].getAttribute("data-series") ?? "[]");
expect(series[0]).toMatchObject({ tool_name: "search", spend: 4.0 });
// The 64px bar cap is this card's opt-in; the shared BarChart must not cap
@ -391,14 +391,16 @@ describe("UsageTab", () => {
start_date: "2026-07-12",
end_date: "2026-07-12",
};
const { findAllByTestId, getAllByTestId } = renderWith([day("2026-07-12", {})], { toolSpend });
renderWith([day("2026-07-12", {})], { toolSpend });
const bars = await findAllByTestId("bar-chart");
const bars = await screen.findAllByTestId("bar-chart");
const [totalByTool, dailyByTool] = bars.slice(-2);
expect(dailyByTool).toHaveAttribute("data-show-legend", "false");
expect(totalByTool).toHaveAttribute("data-colors", dailyByTool.getAttribute("data-colors"));
const toolLegends = getAllByTestId("chart-legend").filter((legend) => legend.textContent === "search,read_file");
const toolLegends = screen
.getAllByTestId("chart-legend")
.filter((legend) => legend.textContent === "search,read_file");
expect(toolLegends).toHaveLength(1);
});
@ -415,23 +417,23 @@ describe("UsageTab", () => {
it.each(["Internal User", "Internal Viewer", "Org Admin"])(
"hides the card and never calls the endpoint for %s",
async (userRole) => {
const { queryByText, getByTestId } = renderWith([day("2026-07-12", { compression_savings_spend: 0.04 })], {
renderWith([day("2026-07-12", { compression_savings_spend: 0.04 })], {
toolSpend,
userRole,
});
// Liveness gate: the daily-activity charts still render for this role,
// so the absence below is the gate, not an empty tab.
expect(getByTestId("donut-chart")).toBeInTheDocument();
expect(queryByText("Spend by tool")).not.toBeInTheDocument();
expect(screen.getByTestId("donut-chart")).toBeInTheDocument();
expect(screen.queryByText("Spend by tool")).not.toBeInTheDocument();
await vi.waitFor(() => expect(mockGetToolSpend).not.toHaveBeenCalled());
},
);
it("keeps the card and the endpoint call for an admin", async () => {
const { findByText } = renderWith([day("2026-07-12", { compression_savings_spend: 0.04 })], { toolSpend });
renderWith([day("2026-07-12", { compression_savings_spend: 0.04 })], { toolSpend });
expect(await findByText("Spend by tool")).toBeInTheDocument();
expect(await screen.findByText("Spend by tool")).toBeInTheDocument();
expect(mockGetToolSpend).toHaveBeenCalled();
});
});

View file

@ -1,5 +1,5 @@
import * as networking from "@/components/networking";
import { fireEvent, render, waitFor, within } from "@testing-library/react";
import { fireEvent, render, waitFor, within, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { afterEach, describe, expect, it, vi } from "vitest";
import GuardrailInfoView from "./guardrail_info";
@ -65,21 +65,19 @@ describe("Guardrail Info", () => {
vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({});
const { getAllByText, getByText } = render(
<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />,
);
render(<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />);
// Wait for the loading to complete and data to be rendered
await waitFor(() => {
// The guardrail name appears in multiple places (title and settings tab)
const elements = getAllByText("Test Guardrail");
const elements = screen.getAllByText("Test Guardrail");
expect(elements.length).toBeGreaterThan(0);
});
// Verify other key elements are present
expect(getByText("Back to Guardrails")).toBeInTheDocument();
expect(getByText("Overview")).toBeInTheDocument();
expect(getByText("Settings")).toBeInTheDocument();
expect(screen.getByText("Back to Guardrails")).toBeInTheDocument();
expect(screen.getByText("Overview")).toBeInTheDocument();
expect(screen.getByText("Settings")).toBeInTheDocument();
});
it("should render a tag-based mode object rather than crashing the detail view", async () => {
@ -105,11 +103,9 @@ describe("Guardrail Info", () => {
vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({});
const { findAllByText } = render(
<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />,
);
render(<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />);
expect(await findAllByText("pre_call, post_call (tag-based)")).not.toHaveLength(0);
expect(await screen.findAllByText("pre_call, post_call (tag-based)")).not.toHaveLength(0);
});
it("should render the provider logo from the bundled guardrail logo map", async () => {
@ -135,11 +131,9 @@ describe("Guardrail Info", () => {
vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({});
const { findByAltText } = render(
<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />,
);
render(<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />);
const logo = await findByAltText("Presidio PII logo");
const logo = await screen.findByAltText("Presidio PII logo");
expect(logo).toHaveAttribute("src", expect.stringContaining("microsoft_azure.svg"));
});
@ -167,25 +161,27 @@ describe("Guardrail Info", () => {
vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({});
const { getByText, findByText, container } = render(
const { container } = render(
<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />,
);
await waitFor(() => {
expect(getByText("Settings")).toBeInTheDocument();
expect(screen.getByText("Settings")).toBeInTheDocument();
});
// Click the Settings tab
fireEvent.click(getByText("Settings"));
fireEvent.click(screen.getByText("Settings"));
// Wait for the Settings panel to render
await waitFor(() => {
expect(getByText("Guardrail Settings")).toBeInTheDocument();
expect(screen.getByText("Guardrail Settings")).toBeInTheDocument();
});
await userEvent.hover(within(container).getByRole("img", { name: "Config guardrail details" }));
expect(await findByText("Guardrail is defined in the config file and cannot be edited.")).toBeInTheDocument();
expect(
await screen.findByText("Guardrail is defined in the config file and cannot be edited."),
).toBeInTheDocument();
});
it("should render the guardrail info", async () => {
@ -216,12 +212,10 @@ describe("Guardrail Info", () => {
vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({});
const { getByText } = render(
<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />,
);
render(<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />);
await waitFor(() => {
expect(getByText("PII Entity Configuration")).toBeInTheDocument();
expect(screen.getByText("PII Entity Configuration")).toBeInTheDocument();
});
});
it("should handle content filter updates correctly", async () => {
@ -251,30 +245,28 @@ describe("Guardrail Info", () => {
vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({});
vi.mocked(networking.updateGuardrailCall).mockResolvedValue({ status: "success" });
const { getByText, getByLabelText } = render(
<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />,
);
render(<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />);
await waitFor(() => {
expect(getByText("Settings")).toBeInTheDocument();
expect(screen.getByText("Settings")).toBeInTheDocument();
});
// Go to Settings tab
fireEvent.click(getByText("Settings"));
fireEvent.click(screen.getByText("Settings"));
await waitFor(() => {
expect(getByText("Guardrail Settings")).toBeInTheDocument();
expect(screen.getByText("Guardrail Settings")).toBeInTheDocument();
});
// Enter Edit Mode
fireEvent.click(getByText("Edit Settings"));
fireEvent.click(screen.getByText("Edit Settings"));
// Modify Guardrail Name to force an update
const nameInput = getByLabelText("Guardrail Name");
const nameInput = screen.getByLabelText("Guardrail Name");
fireEvent.change(nameInput, { target: { value: "Updated Name" } });
// Save with only name change
const saveButton = getByText("Save Changes");
const saveButton = screen.getByText("Save Changes");
fireEvent.click(saveButton);
await waitFor(() => {
@ -300,16 +292,16 @@ describe("Guardrail Info", () => {
// Enter Edit Mode again to make changes
await waitFor(() => {
expect(getByText("Edit Settings")).toBeInTheDocument();
expect(screen.getByText("Edit Settings")).toBeInTheDocument();
});
fireEvent.click(getByText("Edit Settings"));
fireEvent.click(screen.getByText("Edit Settings"));
// Now modify the values using the mock button
const simulateChangeButton = getByText("Simulate Change");
const simulateChangeButton = screen.getByText("Simulate Change");
fireEvent.click(simulateChangeButton);
// Save again
fireEvent.click(getByText("Save Changes"));
fireEvent.click(screen.getByText("Save Changes"));
await waitFor(() => {
expect(networking.updateGuardrailCall).toHaveBeenCalled();
@ -339,12 +331,10 @@ describe("Guardrail Info", () => {
});
vi.mocked(networking.getGuardrailProviderSpecificParams).mockResolvedValue({});
const { findByRole, getByRole, getByText } = render(
<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />,
);
render(<GuardrailInfoView guardrailId="123" onClose={() => {}} accessToken="123" isAdmin={true} />);
expect(await findByRole("tab", { name: "Overview" })).toHaveAttribute("aria-selected", "true");
expect(getByRole("tab", { name: "Settings" })).toHaveAttribute("aria-selected", "false");
expect(getByText("Guardrail Settings")).toBeInTheDocument();
expect(await screen.findByRole("tab", { name: "Overview" })).toHaveAttribute("aria-selected", "true");
expect(screen.getByRole("tab", { name: "Settings" })).toHaveAttribute("aria-selected", "false");
expect(screen.getByText("Guardrail Settings")).toBeInTheDocument();
});
});

View file

@ -1,4 +1,4 @@
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import { describe, it, expect } from "vitest";
import { CategoryFilter, QuickActions, PiiEntityList } from "./pii_components";
import type { PiiEntityCategory } from "@/components/guardrails/types";
@ -6,25 +6,21 @@ import type { PiiEntityCategory } from "@/components/guardrails/types";
describe("CategoryFilter", () => {
it("should render", () => {
const emptyCategories: PiiEntityCategory[] = [];
const { getByText } = render(
<CategoryFilter categories={emptyCategories} selectedCategories={[]} onChange={() => {}} />,
);
expect(getByText("Filter by category")).toBeInTheDocument();
render(<CategoryFilter categories={emptyCategories} selectedCategories={[]} onChange={() => {}} />);
expect(screen.getByText("Filter by category")).toBeInTheDocument();
});
});
describe("QuickActions", () => {
it("should render", () => {
const { getByText } = render(
<QuickActions onSelectAll={() => {}} onUnselectAll={() => {}} hasSelectedEntities={false} />,
);
expect(getByText("Quick Actions")).toBeInTheDocument();
render(<QuickActions onSelectAll={() => {}} onUnselectAll={() => {}} hasSelectedEntities={false} />);
expect(screen.getByText("Quick Actions")).toBeInTheDocument();
});
});
describe("PiiEntityList", () => {
it("should render", () => {
const { getByText } = render(
render(
<PiiEntityList
entities={[]}
selectedEntities={[]}
@ -35,6 +31,6 @@ describe("PiiEntityList", () => {
entityToCategoryMap={new Map()}
/>,
);
expect(getByText("No PII types match your filter criteria")).toBeInTheDocument();
expect(screen.getByText("No PII types match your filter criteria")).toBeInTheDocument();
});
});

View file

@ -1,10 +1,10 @@
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import { describe, it, expect } from "vitest";
import PiiConfiguration from "./pii_configuration";
describe("PiiConfiguration", () => {
it("should render", () => {
const { getByText } = render(
render(
<PiiConfiguration
entities={[]}
actions={[]}
@ -15,6 +15,6 @@ describe("PiiConfiguration", () => {
entityCategories={[]}
/>,
);
expect(getByText("Configure PII Protection")).toBeInTheDocument();
expect(screen.getByText("Configure PII Protection")).toBeInTheDocument();
});
});

View file

@ -45,7 +45,7 @@ describe("MCPServers", () => {
vi.mocked(networking.fetchMCPServers).mockResolvedValue([]);
const queryClient = createQueryClient();
const { getByText } = render(
render(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} />
</QueryClientProvider>,
@ -53,11 +53,11 @@ describe("MCPServers", () => {
// Wait for the component to load and check if title renders
await waitFor(() => {
expect(getByText("MCP Servers")).toBeInTheDocument();
expect(screen.getByText("MCP Servers")).toBeInTheDocument();
});
// Verify the title is rendered
expect(getByText("MCP Servers")).toBeInTheDocument();
expect(screen.getByText("MCP Servers")).toBeInTheDocument();
});
it("should render mocked MCP servers data in the table", async () => {
@ -96,7 +96,7 @@ describe("MCPServers", () => {
vi.mocked(networking.fetchMCPServers).mockResolvedValue(mockServers);
const queryClient = createQueryClient();
const { getByText, getAllByText } = render(
render(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} />
</QueryClientProvider>,
@ -104,19 +104,19 @@ describe("MCPServers", () => {
// Wait for the component to load
await waitFor(() => {
expect(getByText("MCP Servers")).toBeInTheDocument();
expect(screen.getByText("MCP Servers")).toBeInTheDocument();
});
// Wait for the mocked data to render in the table
await waitFor(() => {
expect(getByText("Test Server 1")).toBeInTheDocument();
expect(screen.getByText("Test Server 1")).toBeInTheDocument();
});
// Verify the mocked server data is rendered in the table
expect(getByText("Test Server 1")).toBeInTheDocument();
expect(getByText("Test Server 2")).toBeInTheDocument();
expect(getAllByText("test-server-1").length).toBeGreaterThan(0);
expect(getAllByText("test-server-2").length).toBeGreaterThan(0);
expect(screen.getByText("Test Server 1")).toBeInTheDocument();
expect(screen.getByText("Test Server 2")).toBeInTheDocument();
expect(screen.getAllByText("test-server-1").length).toBeGreaterThan(0);
expect(screen.getAllByText("test-server-2").length).toBeGreaterThan(0);
// Verify the API was called
// Note: useMCPServers uses useAuthorized() internally, which returns "123" from global mock
@ -168,7 +168,7 @@ describe("MCPServers", () => {
vi.mocked(networking.fetchMCPServerHealth).mockResolvedValue(mockHealthStatuses);
const queryClient = createQueryClient();
const { getByText } = render(
render(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} />
</QueryClientProvider>,
@ -176,7 +176,7 @@ describe("MCPServers", () => {
// Wait for the component to load
await waitFor(() => {
expect(getByText("MCP Servers")).toBeInTheDocument();
expect(screen.getByText("MCP Servers")).toBeInTheDocument();
});
// Verify the health check API was called (without a server ID filter — the hook always
@ -211,7 +211,7 @@ describe("MCPServers", () => {
);
const queryClient = createQueryClient();
const { getByText } = render(
render(
<QueryClientProvider client={queryClient}>
<MCPServers {...defaultProps} />
</QueryClientProvider>,
@ -219,7 +219,7 @@ describe("MCPServers", () => {
// Wait for the component to load
await waitFor(() => {
expect(getByText("MCP Servers")).toBeInTheDocument();
expect(screen.getByText("MCP Servers")).toBeInTheDocument();
});
// Verify that health check was initiated

View file

@ -1,5 +1,5 @@
/* @vitest-environment jsdom */
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import PriceDataManagementTab from "./PriceDataManagementTab";
@ -11,7 +11,7 @@ vi.mock("@/app/(dashboard)/hooks/models/useModelCostMap", () => ({
describe("PriceDataManagementTab", () => {
it("renders its content standalone, without a tab-panel ancestor", () => {
const { getByText } = render(<PriceDataManagementTab />);
expect(getByText("Price Data Management")).toBeInTheDocument();
render(<PriceDataManagementTab />);
expect(screen.getByText("Price Data Management")).toBeInTheDocument();
});
});

View file

@ -1,6 +1,6 @@
/* @vitest-environment jsdom */
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import ModelsAndEndpointsPage from "./page";
@ -64,48 +64,48 @@ describe("ModelsAndEndpointsPage", () => {
});
it("renders the admin tab bar and the All Models panel by default", () => {
const { getByRole, getByTestId } = renderPage();
expect(getByRole("tab", { name: "All Models" })).toBeInTheDocument();
expect(getByRole("tab", { name: "LLM Credentials" })).toBeInTheDocument();
expect(getByRole("tab", { name: "Health Status" })).toBeInTheDocument();
expect(getByTestId("panel-all-models")).toBeInTheDocument();
renderPage();
expect(screen.getByRole("tab", { name: "All Models" })).toBeInTheDocument();
expect(screen.getByRole("tab", { name: "LLM Credentials" })).toBeInTheDocument();
expect(screen.getByRole("tab", { name: "Health Status" })).toBeInTheDocument();
expect(screen.getByTestId("panel-all-models")).toBeInTheDocument();
});
it("switches tabs in-memory, mounting only the active panel", async () => {
const user = userEvent.setup();
const { getByRole, getByTestId, queryByTestId } = renderPage();
await user.click(getByRole("tab", { name: "Health Status" }));
expect(getByTestId("panel-health")).toBeInTheDocument();
expect(queryByTestId("panel-all-models")).not.toBeInTheDocument();
renderPage();
await user.click(screen.getByRole("tab", { name: "Health Status" }));
expect(screen.getByTestId("panel-health")).toBeInTheDocument();
expect(screen.queryByTestId("panel-all-models")).not.toBeInTheDocument();
});
it("renders the model detail overlay from the ?model drill-in and hides the tabs", () => {
detailState.modelId = "abc-123";
const { getByTestId, queryByRole } = renderPage();
expect(getByTestId("model-info")).toHaveTextContent("model:abc-123");
expect(queryByRole("tab", { name: "All Models" })).not.toBeInTheDocument();
renderPage();
expect(screen.getByTestId("model-info")).toHaveTextContent("model:abc-123");
expect(screen.queryByRole("tab", { name: "All Models" })).not.toBeInTheDocument();
});
it("renders the team detail overlay from the ?team drill-in", () => {
detailState.teamId = "team-9";
const { getByTestId } = renderPage();
expect(getByTestId("team-info")).toHaveTextContent("team:team-9");
renderPage();
expect(screen.getByTestId("team-info")).toHaveTextContent("team:team-9");
});
it("hides admin-only tabs for a non-admin user", () => {
mockUseAuthorized.mockReturnValue(NON_ADMIN);
const { queryByRole } = renderPage();
expect(queryByRole("tab", { name: "LLM Credentials" })).not.toBeInTheDocument();
expect(queryByRole("tab", { name: "Health Status" })).not.toBeInTheDocument();
renderPage();
expect(screen.queryByRole("tab", { name: "LLM Credentials" })).not.toBeInTheDocument();
expect(screen.queryByRole("tab", { name: "Health Status" })).not.toBeInTheDocument();
});
// Auto-routers are excluded from the All Models table, so this tab is their home: the only
// place in the product to list, create, edit or delete one.
describe("Auto-Routers tab", () => {
it("sits third, after All Models and Add Model", () => {
const { getAllByRole } = renderPage();
renderPage();
const tabs = getAllByRole("tab").map((tab) => tab.textContent);
const tabs = screen.getAllByRole("tab").map((tab) => tab.textContent);
expect(tabs[0]).toContain("All Models");
expect(tabs[1]).toBe("Add Model");
expect(tabs[2]).toContain("Auto-Routers");
@ -115,17 +115,17 @@ describe("ModelsAndEndpointsPage", () => {
it("renders its panel when selected", async () => {
const user = userEvent.setup();
const { getByRole, getByTestId } = renderPage();
renderPage();
await user.click(getByRole("tab", { name: /Auto-Routers/ }));
expect(getByTestId("panel-auto-routers")).toBeInTheDocument();
await user.click(screen.getByRole("tab", { name: /Auto-Routers/ }));
expect(screen.getByTestId("panel-auto-routers")).toBeInTheDocument();
});
it("is hidden from non-admins, who cannot write models", () => {
mockUseAuthorized.mockReturnValue(NON_ADMIN);
const { queryByRole } = renderPage();
renderPage();
expect(queryByRole("tab", { name: /Auto-Routers/ })).not.toBeInTheDocument();
expect(screen.queryByRole("tab", { name: /Auto-Routers/ })).not.toBeInTheDocument();
});
});
});

View file

@ -93,13 +93,13 @@ describe("ChatMessageBubble", () => {
])("should paint the $role surface from theme tokens, not fixed colours", ({ role, bubble, avatar }) => {
render(<ChatMessageBubble {...defaultProps} message={{ role, content: "Hello" }} />);
const header = screen.getByText(role).closest("div") as HTMLElement;
const surface = header.parentElement as HTMLElement;
const surface = screen.getByTestId("message-surface");
const avatarEl = screen.getByTestId("message-avatar");
expect(surface).toHaveClass(...bubble);
expect(surface).not.toHaveAttribute("style");
expect(header.firstElementChild).toHaveClass(avatar);
expect(header.firstElementChild).not.toHaveAttribute("style");
expect(avatarEl).toHaveClass(avatar);
expect(avatarEl).not.toHaveAttribute("style");
});
it("should show model badge for assistant messages when model is provided", () => {

View file

@ -46,6 +46,7 @@ function ChatMessageBubble({
return (
<div className={`mb-4 min-w-0 ${isUser ? "text-right" : "text-left"}`}>
<div
data-testid="message-surface"
className={`inline-block min-w-0 max-w-[92%] overflow-hidden rounded-lg border p-3 text-left text-card-foreground shadow-xs sm:max-w-[85%] sm:px-4 ${
isUser ? "border-info/20 bg-info/10" : "border-border bg-card"
}`}
@ -53,6 +54,7 @@ function ChatMessageBubble({
{/* Header: role icon + name + model badge */}
<div className="mb-1.5 flex min-w-0 items-center gap-2">
<div
data-testid="message-avatar"
className={`flex items-center justify-center w-6 h-6 rounded-full mr-1 ${
isUser ? "bg-info/20" : "bg-muted"
}`}

View file

@ -90,21 +90,19 @@ beforeEach(() => {
describe("CompareUI", () => {
it("should render", () => {
const { getByTestId } = render(<CompareUI accessToken="test-token" disabledPersonalKeyCreation={false} />);
expect(getByTestId("comparison-panel-1")).toBeInTheDocument();
expect(getByTestId("comparison-panel-2")).toBeInTheDocument();
expect(getByTestId("message-input")).toBeInTheDocument();
render(<CompareUI accessToken="test-token" disabledPersonalKeyCreation={false} />);
expect(screen.getByTestId("comparison-panel-1")).toBeInTheDocument();
expect(screen.getByTestId("comparison-panel-2")).toBeInTheDocument();
expect(screen.getByTestId("message-input")).toBeInTheDocument();
});
it("adds a comparison when Add Comparison button is clicked", async () => {
const user = userEvent.setup();
const { container, getByTestId } = render(
<CompareUI accessToken="test-token" disabledPersonalKeyCreation={false} />,
);
const { container } = render(<CompareUI accessToken="test-token" disabledPersonalKeyCreation={false} />);
// Verify initial state: 2 comparison panels
expect(getByTestId("comparison-panel-1")).toBeInTheDocument();
expect(getByTestId("comparison-panel-2")).toBeInTheDocument();
expect(screen.getByTestId("comparison-panel-1")).toBeInTheDocument();
expect(screen.getByTestId("comparison-panel-2")).toBeInTheDocument();
let comparisonPanels = container.querySelectorAll('[data-testid^="comparison-panel-"]');
expect(comparisonPanels).toHaveLength(2);
@ -117,15 +115,13 @@ describe("CompareUI", () => {
});
// Verify the original 2 panels are still there
expect(getByTestId("comparison-panel-1")).toBeInTheDocument();
expect(getByTestId("comparison-panel-2")).toBeInTheDocument();
expect(screen.getByTestId("comparison-panel-1")).toBeInTheDocument();
expect(screen.getByTestId("comparison-panel-2")).toBeInTheDocument();
});
it("should handle image upload and send message with attachment", async () => {
const user = userEvent.setup();
const { getByTestId, queryByTestId } = render(
<CompareUI accessToken="test-token" disabledPersonalKeyCreation={false} />,
);
render(<CompareUI accessToken="test-token" disabledPersonalKeyCreation={false} />);
const file = new File(["test content"], "test-image.png", { type: "image/png" });
@ -138,13 +134,13 @@ describe("CompareUI", () => {
}
await waitFor(() => {
expect(getByTestId("has-attachment")).toBeInTheDocument();
expect(screen.getByTestId("has-attachment")).toBeInTheDocument();
});
const textarea = getByTestId("message-textarea");
const textarea = screen.getByTestId("message-textarea");
fireEvent.change(textarea, { target: { value: "Describe this image" } });
const sendButton = getByTestId("send-button");
const sendButton = screen.getByTestId("send-button");
expect(sendButton).toBeEnabled();
await user.click(sendButton);

View file

@ -1,4 +1,4 @@
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import type { MessageType } from "@/components/chat_ui/types";
import { MessageDisplay } from "./MessageDisplay";
@ -39,9 +39,9 @@ describe("MessageDisplay", () => {
model: "gpt-4",
},
];
const { getByText } = render(<MessageDisplay messages={messages} isLoading={false} />);
expect(getByText("Hello")).toBeInTheDocument();
expect(getByText("Hi there!")).toBeInTheDocument();
render(<MessageDisplay messages={messages} isLoading={false} />);
expect(screen.getByText("Hello")).toBeInTheDocument();
expect(screen.getByText("Hi there!")).toBeInTheDocument();
});
it("displays user and assistant messages with proper grouping and shows loading state", () => {
@ -64,13 +64,13 @@ describe("MessageDisplay", () => {
},
},
];
const { getByText, getByTestId } = render(<MessageDisplay messages={messages} isLoading={false} />);
expect(getByText("You")).toBeInTheDocument();
expect(getByText("What is 2+2?")).toBeInTheDocument();
expect(getByText("gpt-4")).toBeInTheDocument();
expect(getByText("calculator")).toBeInTheDocument();
expect(getByText("2+2 equals 4")).toBeInTheDocument();
expect(getByTestId("response-metrics")).toBeInTheDocument();
render(<MessageDisplay messages={messages} isLoading={false} />);
expect(screen.getByText("You")).toBeInTheDocument();
expect(screen.getByText("What is 2+2?")).toBeInTheDocument();
expect(screen.getByText("gpt-4")).toBeInTheDocument();
expect(screen.getByText("calculator")).toBeInTheDocument();
expect(screen.getByText("2+2 equals 4")).toBeInTheDocument();
expect(screen.getByTestId("response-metrics")).toBeInTheDocument();
});
it("should display image attachment in user message", () => {
@ -86,10 +86,10 @@ describe("MessageDisplay", () => {
model: "gpt-4",
},
];
const { getByTestId, getByText } = render(<MessageDisplay messages={messages} isLoading={false} />);
expect(getByText("What is in this image? [Image attached]")).toBeInTheDocument();
expect(getByTestId("chat-image-renderer")).toBeInTheDocument();
const image = getByTestId("chat-image-renderer").querySelector("img");
render(<MessageDisplay messages={messages} isLoading={false} />);
expect(screen.getByText("What is in this image? [Image attached]")).toBeInTheDocument();
expect(screen.getByTestId("chat-image-renderer")).toBeInTheDocument();
const image = screen.getByTestId("chat-image-renderer").querySelector("img");
expect(image).toHaveAttribute("src", "blob:test-image-url");
});
});

View file

@ -5,6 +5,7 @@ import { describe, it, expect, vi, beforeEach } from "vitest";
import { CreateUserButton } from "./CreateUserButton";
import * as networking from "./networking";
import { toast } from "@/lib/toast";
import { expectControlBesideLabel } from "../../tests/fieldOrientation";
vi.mock("./networking", () => ({
userCreateCall: vi.fn(),
@ -294,6 +295,20 @@ describe("CreateUserButton", () => {
});
});
it("lays the send invitation email checkbox out beside its label", async () => {
const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never });
renderWithProviders(<CreateUserButton {...defaultProps} />);
await waitFor(() => {
expect(screen.getByRole("button", { name: /\+ invite user/i })).toBeInTheDocument();
});
await user.click(screen.getByRole("button", { name: /\+ invite user/i }));
const dialog = screen.getByRole("dialog", { name: /invite user/i });
expectControlBesideLabel(within(dialog).getByRole("checkbox"));
});
describe("organizations", () => {
it("should send organizations list in POST body when organizations are selected", async () => {
const { useOrganizations } = await import("@/app/(dashboard)/hooks/organizations/useOrganizations");

View file

@ -270,7 +270,7 @@ export const CreateUserButton: React.FC<CreateuserProps> = ({
);
const sendInviteEmailField = (
<FormField control={form.control} name="send_invite_email" label="Send invitation email">
<FormField control={form.control} name="send_invite_email" label="Send invitation email" orientation="horizontal">
{({ id, value, onChange, onBlur }) => (
<Checkbox id={id} checked={value} onCheckedChange={onChange} onBlur={onBlur} />
)}

View file

@ -10,6 +10,7 @@
*/
import { describe, it, expect, vi, beforeEach } from "vitest";
import { screen } from "@testing-library/react";
import { renderWithProviders } from "../../../tests/test-utils";
import userEvent from "@testing-library/user-event";
import EntityUsageExportModal from "./EntityUsageExportModal";
@ -73,13 +74,13 @@ describe("EntityUsageExportModal", () => {
const user = userEvent.setup();
const { handleExportCSV } = await import("./utils");
const { getByRole } = renderWithProviders(<EntityUsageExportModal {...baseProps} />);
renderWithProviders(<EntityUsageExportModal {...baseProps} />);
// Default primary action reflects CSV export
expect(getByRole("button", { name: /Export CSV/i })).toBeInTheDocument();
expect(screen.getByRole("button", { name: /Export CSV/i })).toBeInTheDocument();
// Click export
await user.click(getByRole("button", { name: /Export CSV/i }));
await user.click(screen.getByRole("button", { name: /Export CSV/i }));
// Verifies export function was invoked with correct parameters
expect(handleExportCSV).toHaveBeenCalledWith(baseProps.spendData, "daily", "Tag", "tag", {});
@ -97,14 +98,14 @@ describe("EntityUsageExportModal", () => {
const user = userEvent.setup();
const { handleExportCSV } = await import("./utils");
const { getByText, getByRole } = renderWithProviders(<EntityUsageExportModal {...baseProps} />);
renderWithProviders(<EntityUsageExportModal {...baseProps} />);
// Choose the alternate export type - click the label to trigger radio
const dailyModelLabel = getByText(/Day-by-day by tag and model/i);
const dailyModelLabel = screen.getByText(/Day-by-day by tag and model/i);
await user.click(dailyModelLabel);
// Export with default CSV format
const exportBtn = getByRole("button", { name: /Export CSV/i });
const exportBtn = screen.getByRole("button", { name: /Export CSV/i });
await user.click(exportBtn);
// Ensure the selected scope flowed through

View file

@ -9,6 +9,7 @@ import BaseSSOSettingsForm, {
submitMountedSSOValues,
useSSOSettingsForm,
} from "./BaseSSOSettingsForm";
import { expectControlBesideLabel } from "../../../../../../tests/fieldOrientation";
const user = () => userEvent.setup({ pointerEventsCheck: 0 });
@ -233,6 +234,38 @@ describe("BaseSSOSettingsForm", () => {
expect(screen.queryByText("Use Team Mappings")).not.toBeInTheDocument();
});
it("lays a provider checkbox field out beside its label", async () => {
const TestWrapper = () => {
const form = useSSOSettingsForm("sso-settings");
return <BaseSSOSettingsForm form={form} onFormSubmit={vi.fn()} />;
};
renderWithProviders(<TestWrapper />);
await openProviderDropdown();
await user().click(await screen.findByText(/saml sso/i));
expectControlBesideLabel(
await screen.findByRole("checkbox", { name: "Allow IdP-initiated (unsolicited) responses" }),
);
});
it.each(["Use Role Mappings", "Use Team Mappings"])("lays the %s toggle out beside its label", async (label) => {
const TestWrapper = () => {
const form = useSSOSettingsForm("sso-settings");
return <BaseSSOSettingsForm form={form} onFormSubmit={vi.fn()} />;
};
renderWithProviders(<TestWrapper />);
await openProviderDropdown();
await user().click(await screen.findByText(/okta/i));
expectControlBesideLabel(await screen.findByRole("checkbox", { name: label }));
});
});
describe("renderProviderFields", () => {

View file

@ -303,7 +303,7 @@ const SSOProviderField = ({ field }: { field: SSOProviderConfig["fields"][number
if (field.type === "checkbox") {
return (
<FormField control={control} name={field.name} label={field.label}>
<FormField control={control} name={field.name} label={field.label} orientation="horizontal">
{({ value, onChange, onBlur, id, ...rest }) => (
<Checkbox
id={id}
@ -413,7 +413,7 @@ export const MappingToggleField = ({
const { control } = useFormContext<SSOSettingsFormValues>();
return (
<FormField control={control} name={name} label={label}>
<FormField control={control} name={name} label={label} orientation="horizontal">
{({ value, onChange, onBlur, id, ...rest }) => (
<Checkbox
id={id}

View file

@ -1,4 +1,4 @@
import { act, fireEvent, render, waitFor } from "@testing-library/react";
import { act, fireEvent, render, waitFor, screen } from "@testing-library/react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { MountedFormHost } from "../../../tests/mounted-form-host";
import AdvancedSettings from "./advanced_settings";
@ -35,51 +35,51 @@ describe("AdvancedSettings", () => {
});
it("should render tags list", async () => {
const { getByText } = renderAdvancedSettings();
fireEvent.click(getByText("Advanced Settings"));
renderAdvancedSettings();
fireEvent.click(screen.getByText("Advanced Settings"));
await waitFor(() => {
expect(getByText("Tags")).toBeInTheDocument();
expect(screen.getByText("Tags")).toBeInTheDocument();
});
});
it("should render the litellm params", async () => {
const { getByText } = renderAdvancedSettings();
renderAdvancedSettings();
act(() => {
fireEvent.click(getByText("Advanced Settings"));
fireEvent.click(screen.getByText("Advanced Settings"));
});
await waitFor(() => {
expect(getByText("LiteLLM Params")).toBeInTheDocument();
expect(screen.getByText("LiteLLM Params")).toBeInTheDocument();
});
});
it("hides every PTU field when PTU cost attribution is disabled", async () => {
const { getByText, queryByText } = renderAdvancedSettings();
renderAdvancedSettings();
act(() => {
fireEvent.click(getByText("Advanced Settings"));
fireEvent.click(screen.getByText("Advanced Settings"));
});
await waitFor(() => {
expect(getByText("Tags")).toBeInTheDocument();
expect(screen.getByText("Tags")).toBeInTheDocument();
});
for (const label of PTU_LABELS) {
expect(queryByText(label)).not.toBeInTheDocument();
expect(screen.queryByText(label)).not.toBeInTheDocument();
}
expect(queryByText("PTU Effective To (UTC)")).not.toBeInTheDocument();
expect(screen.queryByText("PTU Effective To (UTC)")).not.toBeInTheDocument();
});
it("shows every PTU field when PTU cost attribution is enabled", async () => {
mockUsePtuCostAttributionEnabled.mockReturnValue(true);
const { getByText } = renderAdvancedSettings();
renderAdvancedSettings();
act(() => {
fireEvent.click(getByText("Advanced Settings"));
fireEvent.click(screen.getByText("Advanced Settings"));
});
await waitFor(() => {
expect(getByText("PTU Count")).toBeInTheDocument();
expect(screen.getByText("PTU Count")).toBeInTheDocument();
});
for (const label of PTU_LABELS) {
expect(getByText(label)).toBeInTheDocument();
expect(screen.getByText(label)).toBeInTheDocument();
}
expect(getByText("PTU Effective To (UTC)")).toBeInTheDocument();
expect(screen.getByText("PTU Effective To (UTC)")).toBeInTheDocument();
});
});

View file

@ -1,4 +1,4 @@
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import { describe, expect, it } from "vitest";
import { getPlaceholder, Providers } from "../provider_info_helpers";
import { MountedFormHost } from "../../../tests/mounted-form-host";
@ -6,7 +6,7 @@ import LiteLLMModelNameField from "./litellm_model_name";
describe("LitellmModelNameField", () => {
it("should render", () => {
const { getByText } = render(
render(
<MountedFormHost>
<LiteLLMModelNameField
selectedProvider={Providers.OpenAI}
@ -15,16 +15,16 @@ describe("LitellmModelNameField", () => {
/>
</MountedFormHost>,
);
expect(getByText("LiteLLM Model Name(s)")).toBeInTheDocument();
expect(screen.getByText("LiteLLM Model Name(s)")).toBeInTheDocument();
});
it("should show Azure placeholder as 'my-deployment'", () => {
const { getByPlaceholderText, queryByPlaceholderText } = render(
render(
<MountedFormHost>
<LiteLLMModelNameField selectedProvider={Providers.Azure} providerModels={[]} getPlaceholder={getPlaceholder} />
</MountedFormHost>,
);
expect(getByPlaceholderText("my-deployment")).toBeInTheDocument();
expect(queryByPlaceholderText("gpt-3.5-turbo")).not.toBeInTheDocument();
expect(screen.getByPlaceholderText("my-deployment")).toBeInTheDocument();
expect(screen.queryByPlaceholderText("gpt-3.5-turbo")).not.toBeInTheDocument();
});
});

View file

@ -26,8 +26,8 @@ const openUploadStep = async () => {
describe("BulkCreateUsersButton", () => {
it("should render", () => {
const { getByText } = render(<BulkCreateUsersButton accessToken="test-token" teams={[]} possibleUIRoles={null} />);
expect(getByText("+ Bulk Invite Users")).toBeInTheDocument();
render(<BulkCreateUsersButton accessToken="test-token" teams={[]} possibleUIRoles={null} />);
expect(screen.getByText("+ Bulk Invite Users")).toBeInTheDocument();
});
it("parses a CSV chosen through the file input", async () => {

View file

@ -1,4 +1,4 @@
import { fireEvent, render } from "@testing-library/react";
import { fireEvent, render, screen } from "@testing-library/react";
import { beforeEach, describe, expect, it } from "vitest";
import CostOptimizationFeedbackBanner from "./cost_optimization_feedback_banner";
@ -10,24 +10,24 @@ describe("CostOptimizationFeedbackBanner", () => {
});
it("renders with a link to the feedback discussion", () => {
const { getByText } = render(<CostOptimizationFeedbackBanner />);
const link = getByText("Share Feedback").closest("a");
render(<CostOptimizationFeedbackBanner />);
const link = screen.getByText("Share Feedback").closest("a");
expect(link).toHaveAttribute("href", "https://github.com/BerriAI/litellm/discussions/32172");
});
it("hides itself and persists the dismissal when the dismiss button is clicked", () => {
const { getByText, queryByText, getByLabelText } = render(<CostOptimizationFeedbackBanner />);
expect(getByText("Help shape cost optimization")).toBeInTheDocument();
render(<CostOptimizationFeedbackBanner />);
expect(screen.getByText("Help shape cost optimization")).toBeInTheDocument();
fireEvent.click(getByLabelText("Dismiss banner"));
fireEvent.click(screen.getByLabelText("Dismiss banner"));
expect(queryByText("Help shape cost optimization")).not.toBeInTheDocument();
expect(screen.queryByText("Help shape cost optimization")).not.toBeInTheDocument();
expect(localStorage.getItem(STORAGE_KEY)).toBe("true");
});
it("stays dismissed on remount once persisted", () => {
localStorage.setItem(STORAGE_KEY, "true");
const { queryByText } = render(<CostOptimizationFeedbackBanner />);
expect(queryByText("Help shape cost optimization")).not.toBeInTheDocument();
render(<CostOptimizationFeedbackBanner />);
expect(screen.queryByText("Help shape cost optimization")).not.toBeInTheDocument();
});
});

View file

@ -108,7 +108,7 @@ beforeEach(() => {
test("renders organization view after loading data", async () => {
mockUseOrganization.mockReturnValue({ data: mockOrg, isLoading: false } as any);
const { findAllByText } = renderWithProviders(
renderWithProviders(
<OrganizationInfoView
organizationId="org_123"
onClose={() => {}}
@ -120,7 +120,7 @@ test("renders organization view after loading data", async () => {
/>,
);
const [orgName] = await findAllByText("Acme Corp");
const [orgName] = await screen.findAllByText("Acme Corp");
expect(orgName).toBeInTheDocument();
});

View file

@ -77,21 +77,21 @@ describe("Settings", () => {
});
it("should render the logging callbacks tab when access token is provided", async () => {
const { getByText } = render(<Settings {...defaultProps} />);
render(<Settings {...defaultProps} />);
await waitFor(() => {
expect(getByText("Active Logging Callbacks")).toBeInTheDocument();
expect(screen.getByText("Active Logging Callbacks")).toBeInTheDocument();
});
});
it("should display additional settings tabs", async () => {
const { getByText } = render(<Settings {...defaultProps} />);
render(<Settings {...defaultProps} />);
await waitFor(() => {
expect(getByText("CloudZero Cost Tracking")).toBeInTheDocument();
expect(getByText("Alerting Types")).toBeInTheDocument();
expect(getByText("Alerting Settings")).toBeInTheDocument();
expect(getByText("Email Alerts")).toBeInTheDocument();
expect(screen.getByText("CloudZero Cost Tracking")).toBeInTheDocument();
expect(screen.getByText("Alerting Types")).toBeInTheDocument();
expect(screen.getByText("Alerting Settings")).toBeInTheDocument();
expect(screen.getByText("Email Alerts")).toBeInTheDocument();
});
});
@ -279,13 +279,13 @@ describe("Settings", () => {
});
it("should display CloudZero Cost Tracking tab", async () => {
const { getByText } = render(<Settings {...defaultProps} />);
render(<Settings {...defaultProps} />);
await waitFor(() => {
expect(getByText("Active Logging Callbacks")).toBeInTheDocument();
expect(screen.getByText("Active Logging Callbacks")).toBeInTheDocument();
});
expect(getByText("CloudZero Cost Tracking")).toBeInTheDocument();
expect(screen.getByText("CloudZero Cost Tracking")).toBeInTheDocument();
});
});

View file

@ -1,5 +1,5 @@
import type { ColumnDef, ExpandedState } from "@tanstack/react-table";
import { render, screen, waitFor } from "@testing-library/react";
import { render, screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { useState } from "react";
import { describe, expect, it, vi } from "vitest";
@ -21,6 +21,12 @@ function person(id: string, name: string, flagged = false): Person {
const names = (): (string | null)[] => screen.getAllByTestId("name-cell").map((el) => el.textContent);
const heightClassesOf = (el: HTMLElement | undefined): string[] =>
(el?.className ?? "")
.split(/\s+/)
.filter((cls) => cls.startsWith("h-"))
.sort();
const nameCellColumns: ColumnDef<Person, unknown>[] = [
{
accessorKey: "name",
@ -229,12 +235,10 @@ describe("DataTable sorting", () => {
describe("DataTable layout", () => {
it("stretches the table to fill the container when resizing is on, so hidden columns leave no right-side gap", () => {
const { container } = render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} enableColumnResizing />);
render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} enableColumnResizing />);
const table = container.querySelector("table");
expect(table).not.toBeNull();
// width pins the natural column total (horizontal scroll on overflow); minWidth:100% fills the gap on underflow.
expect(table?.style.minWidth).toBe("100%");
expect(screen.getByRole("table")).toHaveStyle({ minWidth: "100%" });
});
});
@ -354,17 +358,21 @@ describe("DataTable loading", () => {
const { rerender } = render(
<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} size="compact" isLoading />,
);
const skeletonRow = screen.getAllByTestId("skeleton-row").at(0);
const loadedRowHeight = "h-8";
expect(skeletonRow?.className).toContain(loadedRowHeight);
const skeletonHeight = heightClassesOf(screen.getAllByRole("row").at(-1));
rerender(<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} size="compact" />);
expect(document.querySelector("[data-row-id]")?.className).toContain(loadedRowHeight);
const loadedHeight = heightClassesOf(screen.getByRole("row", { name: /Charlie/ }));
expect(loadedHeight).not.toEqual([]);
expect(skeletonHeight).toEqual(loadedHeight);
});
it("does not force the compact height on default-size skeleton rows", () => {
render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} isLoading />);
expect(screen.getAllByTestId("skeleton-row").at(0)?.className).not.toContain("h-8");
const { rerender } = render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} isLoading />);
const skeletonHeight = heightClassesOf(screen.getAllByRole("row").at(-1));
rerender(<DataTable data={CHARLIE_ALICE_BOB} columns={nameCellColumns} size="compact" isLoading />);
expect(heightClassesOf(screen.getAllByRole("row").at(-1))).not.toEqual(skeletonHeight);
});
it("varies skeleton shape and width per column instead of one fixed bar", () => {
@ -420,7 +428,7 @@ describe("DataTable loading", () => {
describe("DataTable column visibility", () => {
it("hides a column when toggled off in the view-options menu", async () => {
const user = userEvent.setup();
const { container } = render(
render(
<DataTable
data={CHARLIE_ALICE_BOB}
columns={nameEmailColumns}
@ -428,13 +436,13 @@ describe("DataTable column visibility", () => {
/>,
);
expect(container.querySelector('th[data-header-id="email"]')).not.toBeNull();
expect(screen.getByRole("columnheader", { name: "Email" })).toBeInTheDocument();
await user.click(screen.getByTestId("view-options-trigger"));
await user.click(await screen.findByTestId("view-option-email"));
await waitFor(() => expect(container.querySelector('th[data-header-id="email"]')).toBeNull());
await waitFor(() => expect(screen.queryByRole("columnheader", { name: "Email" })).not.toBeInTheDocument());
await user.click(screen.getByTestId("view-option-email"));
await waitFor(() => expect(container.querySelector('th[data-header-id="email"]')).not.toBeNull());
expect(await screen.findByRole("columnheader", { name: "Email" })).toBeInTheDocument();
});
it("omits columns that opt out of hiding from the menu", async () => {
@ -468,14 +476,10 @@ describe("DataTable column visibility", () => {
describe("DataTable pinned columns", () => {
it("applies sticky positioning to a pinned column only", () => {
const { container } = render(<DataTable data={CHARLIE_ALICE_BOB} columns={pinnedColumns} />);
render(<DataTable data={CHARLIE_ALICE_BOB} columns={pinnedColumns} />);
const pinnedHead = container.querySelector<HTMLElement>('th[data-header-id="name"]');
const normalHead = container.querySelector<HTMLElement>('th[data-header-id="email"]');
expect(pinnedHead?.style.position).toBe("sticky");
expect(pinnedHead?.style.left).toBe("0px");
expect(normalHead?.style.position).toBe("");
expect(screen.getByRole("columnheader", { name: "Name" })).toHaveStyle({ position: "sticky", left: "0px" });
expect(screen.getByRole("columnheader", { name: "Email" })).not.toHaveStyle({ position: "sticky" });
});
});
@ -570,7 +574,7 @@ describe("DataTable expansion", () => {
describe("DataTable row styling and footer", () => {
it("applies rowClassName to the matching row only", () => {
const data = [person("a", "Alice", true), person("b", "Bob", false)];
const { container } = render(
render(
<DataTable
data={data}
columns={nameCellColumns}
@ -579,8 +583,8 @@ describe("DataTable row styling and footer", () => {
/>,
);
expect(container.querySelector('tr[data-row-id="a"]')?.className).toContain("flagged-row");
expect(container.querySelector('tr[data-row-id="b"]')?.className).not.toContain("flagged-row");
expect(screen.getByRole("row", { name: /Alice/ })).toHaveClass("flagged-row");
expect(screen.getByRole("row", { name: /Bob/ })).not.toHaveClass("flagged-row");
});
it("renders the footer slot inside a tfoot element", () => {
@ -596,63 +600,57 @@ describe("DataTable row styling and footer", () => {
/>,
);
expect(screen.getByTestId("footer-row").closest("tfoot")).not.toBeNull();
const rowGroups = screen.getAllByRole("rowgroup");
expect(within(rowGroups.at(-1) as HTMLElement).getByText("Total: 3")).toBeInTheDocument();
});
});
describe("DataTable layout", () => {
it("exposes resize handles with stable selectors only when resizing is enabled", () => {
const { container, rerender } = render(
<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} enableColumnResizing />,
);
expect(container.querySelectorAll("[data-resizer][data-header-id]").length).toBe(2);
const { rerender } = render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} enableColumnResizing />);
expect(screen.getByTestId("column-resizer-name")).toBeInTheDocument();
expect(screen.getByTestId("column-resizer-email")).toBeInTheDocument();
rerender(<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} />);
expect(container.querySelectorAll("[data-resizer]").length).toBe(0);
expect(screen.queryByTestId("column-resizer-name")).not.toBeInTheDocument();
});
it("makes the header sticky and constrains body height when maxBodyHeight is set", () => {
const { container } = render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} maxBodyHeight={240} />);
expect(container.querySelector("thead")?.className).toContain("sticky");
const scroller = container.querySelector('[data-slot="table-container"]')?.parentElement as HTMLElement;
expect(scroller).toHaveStyle({ maxHeight: "240px" });
render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} maxBodyHeight={240} />);
expect(screen.getByTestId("data-table-head")).toHaveClass("sticky");
expect(screen.getByTestId("data-table-scroller")).toHaveStyle({ maxHeight: "240px" });
});
it("caps fillHeight at the parent's height instead of stretching to it, so a short table stays short", () => {
const { container } = render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} fillHeight />);
const scroller = container.querySelector('[data-slot="table-container"]')?.parentElement as HTMLElement;
const frame = scroller.parentElement as HTMLElement;
const outer = frame.parentElement as HTMLElement;
render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} fillHeight />);
const outer = screen.getByTestId("data-table-root");
const frame = screen.getByTestId("data-table-frame");
const scroller = screen.getByTestId("data-table-scroller");
// A ceiling, not a stretch: flex-1 here would hold the footer at the bottom on a two-row table.
expect(outer.className).toContain("max-h-full");
expect(outer.className).not.toContain("flex-1");
expect(frame.className).not.toContain("flex-1");
expect(scroller.className).not.toContain("flex-1");
expect(outer).toHaveClass("max-h-full", "flex-col");
expect(outer).not.toHaveClass("flex-1");
expect(frame).toHaveClass("flex-col");
expect(frame).not.toHaveClass("flex-1");
expect(scroller).not.toHaveClass("flex-1");
expect(outer.className).toContain("flex-col");
expect(frame.className).toContain("flex-col");
expect(scroller.className).toContain("min-h-0");
expect(scroller.className).toContain("overflow-auto");
expect(scroller).toHaveClass("min-h-0", "overflow-auto");
expect(scroller).toHaveStyle({ maxHeight: "" });
// Without this the Table primitive's own overflow container captures the sticky header.
expect(scroller.className).toContain("[&_[data-slot=table-container]]:overflow-visible");
expect(scroller).toHaveClass("[&_[data-slot=table-container]]:overflow-visible");
const thead = container.querySelector("thead") as HTMLElement;
expect(thead.className).toContain("sticky");
// Rows pass under the header, so the semi-transparent row tint alone would let them show through.
expect(thead.className).toContain("bg-background");
expect(screen.getByTestId("data-table-head")).toHaveClass("sticky", "bg-background");
});
it("leaves the default layout untouched when neither height mode is set", () => {
const { container } = render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} />);
const scroller = container.querySelector('[data-slot="table-container"]')?.parentElement as HTMLElement;
render(<DataTable data={CHARLIE_ALICE_BOB} columns={nameEmailColumns} />);
const scroller = screen.getByTestId("data-table-scroller");
expect(scroller.className).toContain("overflow-x-auto");
expect(scroller.className).not.toContain("min-h-0");
expect(scroller).toHaveClass("overflow-x-auto");
expect(scroller).not.toHaveClass("min-h-0");
expect(scroller).toHaveStyle({ maxHeight: "" });
expect((scroller.parentElement as HTMLElement).className).not.toContain("flex-col");
expect(container.querySelector("thead")?.className).not.toContain("sticky");
expect(container.querySelector("thead")?.className).not.toContain("bg-background");
expect(screen.getByTestId("data-table-frame")).not.toHaveClass("flex-col");
expect(screen.getByTestId("data-table-head")).not.toHaveClass("sticky", "bg-background");
});
});

View file

@ -195,8 +195,7 @@ function DataTableHeadCell<TData>({ header, size, stickyHeader, enableColumnResi
)}
{canResize && (
<div
data-resizer
data-header-id={header.id}
data-testid={`column-resizer-${header.id}`}
onMouseDown={header.getResizeHandler()}
onTouchStart={header.getResizeHandler()}
onDoubleClick={() => column.resetSize()}
@ -589,15 +588,19 @@ export function DataTable<TData extends RowData, TValue>(props: DataTableProps<T
const paginationNode = renderPagination();
return (
<div className={cn("w-full", fill.outer)}>
<div className={cn("overflow-hidden rounded-lg border border-border", fill.frame)}>
<div data-testid="data-table-root" className={cn("w-full", fill.outer)}>
<div data-testid="data-table-frame" className={cn("overflow-hidden rounded-lg border border-border", fill.frame)}>
{toolbar !== undefined && <div className="shrink-0 border-b border-border px-4 py-3">{toolbar(table)}</div>}
<div
data-testid="data-table-scroller"
className={cn(stickyHeader ? "overflow-auto" : "overflow-x-auto", fill.body)}
style={maxBodyHeight !== undefined ? { maxHeight: maxBodyHeight } : undefined}
>
<TableRoot className={enableColumnResizing ? "table-fixed" : ""} style={tableStyle}>
<TableHeader className={cn(stickyHeader ? "sticky top-0 z-sticky" : "", fill.header)}>
<TableHeader
data-testid="data-table-head"
className={cn(stickyHeader ? "sticky top-0 z-sticky" : "", fill.header)}
>
{table.getHeaderGroups().map((headerGroup) => (
<TableRow key={headerGroup.id} className="bg-muted/50">
{headerGroup.headers.map((header) => (

View file

@ -1,4 +1,4 @@
import { fireEvent, render } from "@testing-library/react";
import { fireEvent, render, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import PaginationStatusAlerts from "./PaginationStatusAlerts";
@ -6,7 +6,7 @@ import PaginationStatusAlerts from "./PaginationStatusAlerts";
describe("PaginationStatusAlerts", () => {
it("shows page progress and wires the Stop button while fetching", () => {
const cancel = vi.fn();
const { getByRole, getByText } = render(
render(
<PaginationStatusAlerts
isFetchingMore={true}
cancelled={false}
@ -15,13 +15,13 @@ describe("PaginationStatusAlerts", () => {
/>,
);
expect(getByText(/Currently fetching spend data: fetched 7 \/ 42 pages/)).toBeInTheDocument();
fireEvent.click(getByRole("button", { name: "Stop" }));
expect(screen.getByText(/Currently fetching spend data: fetched 7 \/ 42 pages/)).toBeInTheDocument();
fireEvent.click(screen.getByRole("button", { name: "Stop" }));
expect(cancel).toHaveBeenCalledTimes(1);
});
it("shows the partial-data notice after a cancel, frozen at the last fetched page", () => {
const { getByText } = render(
render(
<PaginationStatusAlerts
isFetchingMore={false}
cancelled={true}
@ -30,11 +30,11 @@ describe("PaginationStatusAlerts", () => {
/>,
);
expect(getByText("Showing partial spend data (7/42 pages loaded)")).toBeInTheDocument();
expect(screen.getByText("Showing partial spend data (7/42 pages loaded)")).toBeInTheDocument();
});
it("names the subject it is fetching", () => {
const { getByText } = render(
render(
<PaginationStatusAlerts
isFetchingMore={true}
cancelled={false}
@ -44,7 +44,7 @@ describe("PaginationStatusAlerts", () => {
/>,
);
expect(getByText(/Currently fetching agent data: fetched 1 \/ 3 pages/)).toBeInTheDocument();
expect(screen.getByText(/Currently fetching agent data: fetched 1 \/ 3 pages/)).toBeInTheDocument();
});
it("renders nothing when idle", () => {

View file

@ -1,4 +1,4 @@
import { render } from "@testing-library/react";
import { render, screen } from "@testing-library/react";
import React from "react";
import { describe, expect, it } from "vitest";
import { AreaChart } from "./area_chart";
@ -21,9 +21,9 @@ describe("AreaChart", () => {
});
it("renders the No data placeholder instead of a chart when data is empty", () => {
const { container, getByText } = render(<AreaChart data={[]} index="date" categories={["tokens"]} />);
const { container } = render(<AreaChart data={[]} index="date" categories={["tokens"]} />);
expect(getByText("No data")).toBeInTheDocument();
expect(screen.getByText("No data")).toBeInTheDocument();
expect(container.querySelector('[data-slot="chart"]')).toBeNull();
});

View file

@ -21,9 +21,9 @@ describe("BarChart", () => {
});
it("renders the No data placeholder instead of a chart when data is empty", () => {
const { container, getByText } = render(<BarChart data={[]} index="date" categories={["passed"]} />);
const { container } = render(<BarChart data={[]} index="date" categories={["passed"]} />);
expect(getByText("No data")).toBeInTheDocument();
expect(screen.getByText("No data")).toBeInTheDocument();
expect(container.querySelector('[data-slot="chart"]')).toBeNull();
});

View file

@ -445,7 +445,7 @@ describe("KeyInfoView handleKeyUpdate budget_duration", () => {
);
fireEvent.click(screen.getByText("Settings"));
expect(screen.getByText("Budget Reset").parentElement?.textContent).toContain("Every 30d");
expect(screen.getByTestId("budget-reset-value")).toHaveTextContent("Every 30d");
fireEvent.click(screen.getByText("Edit Settings"));
(globalThis as any).__TEST_FORM_VALUES = {
@ -456,7 +456,7 @@ describe("KeyInfoView handleKeyUpdate budget_duration", () => {
fireEvent.click(screen.getByText("Mock Submit"));
await waitFor(() => {
expect(screen.getByText("Budget Reset").parentElement?.textContent).toBe("Budget ResetNever");
expect(screen.getByTestId("budget-reset-value")).toHaveTextContent("Never");
});
});
});

View file

@ -1,7 +1,7 @@
import { fireEvent, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { beforeEach, describe, expect, it, vi } from "vitest";
import { renderWithProviders } from "../../../tests/test-utils";
import { chooseSelectOption, renderWithProviders } from "../../../tests/test-utils";
import { KeyResponse } from "../key_team_helpers/key_list";
import { MODEL_MAX_BUDGET_PREMIUM_HINT } from "../key_team_helpers/ModelMaxBudgetEditor";
import {
@ -300,7 +300,7 @@ describe("KeyEditView", () => {
});
it("should render", async () => {
const { getByText } = renderWithProviders(
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => {}}
@ -313,12 +313,12 @@ describe("KeyEditView", () => {
);
await waitFor(() => {
expect(getByText("Save Changes")).toBeInTheDocument();
expect(screen.getByText("Save Changes")).toBeInTheDocument();
});
});
it("should render tags", async () => {
const { getByText } = renderWithProviders(
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => {}}
@ -331,12 +331,12 @@ describe("KeyEditView", () => {
);
await waitFor(() => {
expect(getByText("test-tag")).toBeInTheDocument();
expect(screen.getByText("test-tag")).toBeInTheDocument();
});
});
it("should not render tags in metadata textarea", async () => {
const { getByLabelText } = renderWithProviders(
renderWithProviders(
<KeyEditView
keyData={MOCK_KEY_DATA}
onCancel={() => {}}
@ -348,7 +348,7 @@ describe("KeyEditView", () => {
/>,
);
const metadataTextarea = getByLabelText("Metadata") as HTMLTextAreaElement;
const metadataTextarea = screen.getByLabelText("Metadata") as HTMLTextAreaElement;
await waitFor(() => {
expect(metadataTextarea).toHaveValue("{}");
});
@ -963,10 +963,7 @@ describe("KeyEditView", () => {
/>,
);
await userEvent.click(await screen.findByLabelText("Reset Budget"));
const weeklyOption = await screen.findByText("weekly");
await userEvent.click(weeklyOption);
await chooseSelectOption(userEvent, await screen.findByLabelText("Reset Budget"), "weekly");
const submitButton = screen.getByRole("button", { name: /save changes/i });
await userEvent.click(submitButton);
@ -1042,8 +1039,7 @@ describe("KeyEditView", () => {
);
const resetBudget = await screen.findByLabelText("Reset Budget");
await userEvent.click(resetBudget);
await userEvent.click(await screen.findByText("Never resets"));
await chooseSelectOption(userEvent, resetBudget, "Never resets");
await waitFor(() => {
expect(resetBudget).toHaveTextContent("Never resets");
@ -1074,8 +1070,7 @@ describe("KeyEditView", () => {
/>,
);
await userEvent.click(await screen.findByLabelText("Reset Budget"));
await userEvent.click(await screen.findByText("Never resets"));
await chooseSelectOption(userEvent, await screen.findByLabelText("Reset Budget"), "Never resets");
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
@ -1946,8 +1941,7 @@ describe("KeyEditView", () => {
await userEvent.clear(duration);
await userEvent.type(duration, "45d");
await userEvent.click(screen.getByLabelText(/TPM Rate Limit Type/));
await userEvent.click(await screen.findByTitle("Guaranteed throughput"));
await chooseSelectOption(userEvent, screen.getByLabelText(/TPM Rate Limit Type/), /^Guaranteed throughput/);
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
@ -2103,8 +2097,7 @@ describe("KeyEditView", () => {
renderForPayload(onSubmitMock);
await screen.findByRole("button", { name: /save changes/i });
await userEvent.click(screen.getByLabelText(/RPM Rate Limit Type/));
await userEvent.click(await screen.findByTitle("Guaranteed throughput"));
await chooseSelectOption(userEvent, screen.getByLabelText(/RPM Rate Limit Type/), /^Guaranteed throughput/);
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));

View file

@ -381,6 +381,6 @@ describe("KeyInfoView budget reset visibility", () => {
await waitFor(() => {
expect(screen.getByText("Budget Reset")).toBeInTheDocument();
});
expect(screen.getByText("Budget Reset").parentElement).toHaveTextContent("Never");
expect(screen.getByTestId("budget-reset-value")).toHaveTextContent("Never");
});
});

View file

@ -895,7 +895,7 @@ export default function KeyInfoView({
<div>
<p className="text-sm font-medium">Budget Reset</p>
<p className="text-sm">
<p data-testid="budget-reset-value" className="text-sm">
{currentKeyData.budget_reset_at
? `${currentKeyData.budget_duration ? `Every ${currentKeyData.budget_duration}, next ` : ""}${formatTimestamp(currentKeyData.budget_reset_at)}`
: "Never"}

View file

@ -14,6 +14,7 @@ import {
FieldSet,
FieldTitle,
} from "./field";
import { ROW_LAYOUT_CLASSES, STRETCH_CHILDREN_CLASS } from "../../../tests/fieldOrientation";
describe("FieldError", () => {
it("renders nothing when there are no errors and no children", () => {
@ -85,6 +86,20 @@ describe("Field", () => {
expect(screen.getByRole("group")).toHaveAttribute("data-orientation", "horizontal");
});
it("stretches every child when vertical, which is what inputs, selects and textareas want", () => {
render(<Field />);
expect(screen.getByRole("group")).toHaveClass("flex-col", STRETCH_CHILDREN_CLASS);
});
it("lays children in a row at their own width when horizontal, so a checkbox stays square", () => {
render(<Field orientation="horizontal" />);
const field = screen.getByRole("group");
expect(field).toHaveClass(...ROW_LAYOUT_CLASSES);
expect(field).not.toHaveClass(STRETCH_CHILDREN_CLASS);
});
});
describe("field primitives forward refs to their DOM node", () => {

View file

@ -0,0 +1,18 @@
import { expect } from "vitest";
export const ROW_LAYOUT_CLASSES = ["flex-row", "items-center"] as const;
export const STRETCH_CHILDREN_CLASS = "*:w-full";
/**
* Asserts a control sits beside its label at its own width instead of being stretched across the
* field. Reaches for the resolved classes because the defect is purely visual: nothing accessible
* distinguishes a square checkbox from a full-width bar.
*/
export const expectControlBesideLabel = (control: HTMLElement): void => {
const field = control.closest('[data-slot="field"]');
if (field === null) throw new Error("control is not rendered inside a form field");
expect(field).toHaveAttribute("data-orientation", "horizontal");
expect(field).toHaveClass(...ROW_LAYOUT_CLASSES);
expect(field).not.toHaveClass(STRETCH_CHILDREN_CLASS);
};

View file

@ -52,7 +52,7 @@ const pointerBlocked = (element: HTMLElement): boolean => {
* the option text alone is a race that React 19's flush timing loses.
*/
export const chooseSelectOption = async (
user: ReturnType<typeof userEvent.setup>,
user: Pick<ReturnType<typeof userEvent.setup>, "click">,
trigger: HTMLElement,
optionName: string | RegExp,
) => {

8
uv.lock generated
View file

@ -10,7 +10,7 @@ resolution-markers = [
]
[options]
exclude-newer = "2026-08-26T18:33:25.773031Z"
exclude-newer = "2026-08-29T17:58:57.633306Z"
exclude-newer-span = "P3D"
[manifest]
@ -4266,7 +4266,7 @@ wheels = [
[[package]]
name = "litellm"
version = "1.100.0"
version = "1.101.0"
source = { editable = "." }
dependencies = [
{ name = "aiohttp" },
@ -4669,12 +4669,12 @@ proxy-dev = [
[[package]]
name = "litellm-enterprise"
version = "0.1.62"
version = "0.1.63"
source = { editable = "enterprise" }
[[package]]
name = "litellm-proxy-extras"
version = "0.4.91"
version = "0.4.92"
source = { editable = "litellm-proxy-extras" }
[[package]]