Merge branch 'main' into litellm_oss_staging_03_18_2026

This commit is contained in:
Krish Dholakia 2026-03-19 17:57:55 -07:00 • committed by GitHub
commit 8d92d8637d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
44 changed files with 2387 additions and 686 deletions

View file

@ -42,7 +42,7 @@ commands:
"pydantic==2.11.0" "mcp==1.25.0" "requests-mock>=1.12.1" \
"responses==0.25.7" "pytest-xdist==3.6.1" "pytest-timeout==2.2.0" \
"pytest-cov==5.0.0" "semantic_router==0.1.10" "fastapi-offline==1.7.3" \
"a2a"
"a2a" "parameterized>=0.9.0"
- setup_litellm_enterprise_pip
- save_cache:
paths:
@ -1115,7 +1115,7 @@ jobs:
for dir in "${IGNORE_DIRS[@]}"; do
IGNORE_ARGS="$IGNORE_ARGS --ignore=$dir"
done
python -m pytest -v tests/llm_translation $IGNORE_ARGS --junitxml=test-results/junit.xml --durations=20 -n 8 --timeout=120 --timeout_method=thread
python -m pytest -v tests/llm_translation $IGNORE_ARGS --junitxml=test-results/junit.xml --durations=20 -n 8 --timeout=120 --timeout_method=thread --retries 2 --retry-delay 5
no_output_timeout: 15m
# Store test results
@ -1331,7 +1331,7 @@ jobs:
command: |
pwd
ls
python -m pytest -vv tests/unified_google_tests --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit.xml --durations=5
python -m pytest -vv tests/unified_google_tests --cov=litellm --cov-report=xml -x -s -v --junitxml=test-results/junit.xml --durations=5 --retries 3 --retry-delay 5
no_output_timeout: 15m
- run:
name: Rename the coverage files

View file

@ -28,9 +28,12 @@ jobs:
find . -type d -name "__pycache__" -exec rm -rf {} + || true
find . -name "*.pyc" -delete || true
- name: Check poetry.lock is up to date
run: |
poetry check --lock || (echo "❌ poetry.lock is out of sync with pyproject.toml. Run 'poetry lock' locally and commit the result." && exit 1)
- name: Install dependencies
run: |
poetry lock
poetry install --with dev
- name: Check Black formatting

View file

@ -163,6 +163,9 @@ run_grype_scans() {
"CVE-2026-25639" # axios - full fix requires 1.x major version bump; pinned to >=0.30.2 to clear other axios CVEs, upgrade to 1.x in follow-up
"CVE-2026-2297" # Python 3.13 SourcelessFileLoader audit hook bypass - no fix available in base image
"GHSA-qffp-2rhf-9h96" # tar hardlink path traversal - from nodejs_wheel bundled npm, not used in application runtime code
"CVE-2026-2673" # OpenSSL 3.6.1 TLS 1.3 key exchange group negotiation issue - no fix available yet
"CVE-2026-3644" # Python 3.13 vulnerability - no fix available in base image
"CVE-2026-4224" # Python 3.13 Expat parser stack overflow in ElementDeclHandler - no fix available in base image
)
# Build JSON array of allowlisted CVE IDs for jq

View file

@ -0,0 +1,78 @@
---
slug: guardrail-logging-secret-exposure-incident
title: "Incident Report: Guardrail logging exposed secret headers in spend logs and traces"
date: 2026-03-18T10:00:00
authors:
- litellm
tags: [incident-report, security, guardrails]
hide_table_of_contents: false
---
**Date:** March 18, 2026
**Duration:** Unknown
**Severity:** High
**Status:** Resolved
## Summary
When a custom guardrail returned the full LiteLLM request/data dictionary, the guardrail response logged by LiteLLM could include `secret_fields.raw_headers`, including plaintext `Authorization` headers containing API keys or other credentials.
This information could then propagate to logging and observability surfaces that consume guardrail metadata, including:
- **Spend logs in the LiteLLM UI:** visible to admins with access to spend-log data
- **OpenTelemetry traces:** visible to anyone with access to the relevant telemetry backend
LLM calls, proxy routing, and provider execution were not blocked by this bug. The impact was exposure of sensitive request headers in observability and logging paths.
{/* truncate */}
---
## Background
LiteLLM keeps internal request data (including request headers) for use during the call. That data is not meant to be written to logs or telemetry.
When custom guardrails run, their outcomes are logged so they can appear in spend logs, OpenTelemetry traces, and other observability backends. If a guardrail returned the full request payload instead of a minimal result, that internal request data could be included in what was logged. Before the fix, the guardrail logging path did not strip that data before sending it to those systems.
```mermaid
flowchart TD
inboundRequest["1. Incoming proxy request"] --> storeSecrets["2. Store internal request data"]
storeSecrets --> guardrailRuns["3. Custom guardrail runs"]
guardrailRuns --> fullDataReturn["4. Guardrail returns full request payload"]
fullDataReturn --> loggingBuild["5. Build guardrail log payload"]
loggingBuild --> spendLogs["6a. Persist to spend logs / UI"]
loggingBuild --> otelTraces["6b. Attach to OTEL guardrail spans"]
```
---
## Root Cause
The root cause was incomplete sanitization in the guardrail logging path. When building the payload that gets sent to spend logs and traces, LiteLLM prepared guardrail responses for logging but did not strip internal request data (such as headers) from them. If a guardrail returned a response that included that data, it was passed through to the logging and observability systems unchanged.
---
## Impact
This issue required all of the following:
1. A custom guardrail returned the full LiteLLM request/data dictionary, or another response object containing `secret_fields`.
2. LiteLLM logged that guardrail response through the standard guardrail logging path.
3. An operator, admin, or telemetry consumer had access to the resulting logs or traces.
When those conditions were met, sensitive values could become visible through:
- **Spend logs / UI responses:** guardrail metadata could be included in spend-log payloads rendered in the admin UI.
- **OpenTelemetry traces:** `guardrail_response` could be written as a span attribute on guardrail spans.
- **Other downstream observability backends:** any integration consuming the same guardrail metadata could receive the leaked values.
This was a logging and telemetry exposure bug. It did not let callers bypass auth, access other tenants directly, or change model behavior, but it could expose plaintext credentials to people with access to those observability systems.
---
## Guidance For Users
- Upgrade to LiteLLM 1.82.3+.
- If you operated custom guardrails that return the full request/data dict, review whether spend logs or telemetry traces were retained during the affected period.
- Rotate any credentials that may have appeared in `Authorization` or other forwarded request headers in those systems.
- Apply least-privilege access controls to spend-log views and telemetry backends that may contain request-derived metadata.

View file

@ -902,6 +902,7 @@ router_settings:
| OTEL_SERVICE_NAME | Service name identifier for OpenTelemetry
| OTEL_TRACER_NAME | Tracer name for OpenTelemetry tracing
| OTEL_LOGS_EXPORTER | Exporter type for OpenTelemetry logs (e.g., console)
| OTEL_IGNORE_CONTEXT_PROPAGATION | When true, ignore parent span context propagation in OpenTelemetry callbacks
| PAGERDUTY_API_KEY | API key for PagerDuty Alerting
| PANW_PRISMA_AIRS_API_KEY | API key for PANW Prisma AIRS service
| PANW_PRISMA_AIRS_API_BASE | Base URL for PANW Prisma AIRS service

View file

@ -602,6 +602,22 @@ Since you shouldn't use 12.5, round down to **10** to leave a safety buffer. Thi
- Total maximum connections: 8 workers × 10 connections = 80 connections
- This stays safely under your database's 100 connection limit
## LiteLLM License Key (Enterprise)
To enable [LiteLLM Enterprise features](https://docs.litellm.ai/docs/proxy/enterprise), set your license key as an environment variable:
```bash
export LITELLM_LICENSE="eyJ..."
```
The license key is a JWT token provided when you purchase a LiteLLM Enterprise license. Once set, LiteLLM will automatically detect and activate enterprise features.
You can also add it to your `.env` file:
```env
LITELLM_LICENSE="eyJ..."
```
## Extras

View file

@ -48,6 +48,20 @@ const sidebars = {
slug: "/guardrail_providers"
},
items: [
{
type: "category",
label: "Contributing to Guardrails",
items: [
"adding_provider/generic_guardrail_api",
"adding_provider/simple_guardrail_tutorial",
"adding_provider/adding_guardrail_support",
]
},
{
type: "doc",
id: "proxy/guardrails/team_based_guardrails",
label: "Team Bring-Your-Own Guardrails",
},
...[
"proxy/guardrails/qualifire",
"proxy/guardrails/aim_security",

View file

@ -757,7 +757,7 @@ def _map_traffic_type_to_service_tier(traffic_type: Optional[str]) -> Optional[s
"""
if traffic_type is None:
return None
service_tier = _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER.get(traffic_type.upper())
service_tier = _GEMINI_TRAFFIC_TYPE_TO_SERVICE_TIER.get(str(traffic_type).upper())
return service_tier

View file

@ -291,7 +291,7 @@ class DataDogLogger(
dd_payload = DatadogPayload(
ddsource=get_datadog_source(),
ddtags=get_datadog_tags(),
ddtags=",".join(get_datadog_tags()),
hostname=get_datadog_hostname(),
message=safe_dumps(message_payload),
service=get_datadog_service(),
@ -442,7 +442,9 @@ class DataDogLogger(
verbose_logger.debug("Datadog: Logger - Logging payload = %s", json_payload)
dd_payload = DatadogPayload(
ddsource=get_datadog_source(),
ddtags=get_datadog_tags(standard_logging_object=standard_logging_object),
ddtags=",".join(
get_datadog_tags(standard_logging_object=standard_logging_object)
),
hostname=get_datadog_hostname(),
message=json_payload,
service=get_datadog_service(),
@ -545,7 +547,7 @@ class DataDogLogger(
_dd_message_str = safe_dumps(_payload_dict)
_dd_payload = DatadogPayload(
ddsource=get_datadog_source(),
ddtags=get_datadog_tags(),
ddtags=",".join(get_datadog_tags()),
hostname=get_datadog_hostname(),
message=_dd_message_str,
service=get_datadog_service(),
@ -587,7 +589,7 @@ class DataDogLogger(
_dd_message_str = safe_dumps(_payload_dict)
_dd_payload = DatadogPayload(
ddsource=get_datadog_source(),
ddtags=get_datadog_tags(),
ddtags=",".join(get_datadog_tags()),
hostname=get_datadog_hostname(),
message=_dd_message_str,
service=get_datadog_service(),
@ -678,7 +680,7 @@ class DataDogLogger(
dd_payload = DatadogPayload(
ddsource=get_datadog_source(),
ddtags=get_datadog_tags(),
ddtags=",".join(get_datadog_tags()),
hostname=get_datadog_hostname(),
message=json_payload,
service=get_datadog_service(),

View file

@ -38,8 +38,13 @@ def get_datadog_pod_name() -> str:
def get_datadog_tags(
standard_logging_object: Optional[StandardLoggingPayload] = None,
) -> str:
"""Build Datadog tags string used by multiple integrations."""
) -> List[str]:
"""Build Datadog tags as a list of individual tag strings.
Returns a list of "key:value" strings suitable for Datadog LLM Observability
(which expects tags as an array). For Datadog Logs API (ddtags), join with
comma: ",".join(get_datadog_tags(...)).
"""
base_tags = {
"env": get_datadog_env(),
@ -66,4 +71,4 @@ def get_datadog_tags(
if team_tag:
tags.append(f"team:{team_tag}")
return ",".join(tags)
return tags

View file

@ -203,7 +203,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
type="span",
attributes=DDSpanAttributes(
ml_app=get_datadog_service(),
tags=[get_datadog_tags()],
tags=get_datadog_tags(),
spans=self.log_queue,
),
),
@ -315,7 +315,7 @@ class DataDogLLMObsLogger(CustomBatchLogger):
duration=int((end_time - start_time).total_seconds() * 1e9),
metrics=metrics,
status="error" if error_info else "ok",
tags=[get_datadog_tags(standard_logging_object=standard_logging_payload)],
tags=get_datadog_tags(standard_logging_object=standard_logging_payload),
)
apm_trace_id = self._get_apm_trace_id()

View file

@ -2983,8 +2983,7 @@ class Logging(LiteLLMLoggingBaseClass):
if (
isinstance(callback, CustomLogger)
and is_sync_request
and self.call_type
!= CallTypes.pass_through.value
and self.call_type != CallTypes.pass_through.value
): # custom logger class
callback.log_failure_event(
start_time=start_time,

View file

@ -1901,15 +1901,19 @@ class CustomStreamWrapper:
"usage",
getattr(complete_streaming_response, "usage"),
)
try:
_cache_copy = complete_streaming_response.model_copy(deep=True)
_log_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:
_cache_copy = complete_streaming_response.model_copy()
_log_copy = complete_streaming_response.model_copy()
self.cache_streaming_response(
processed_chunk=complete_streaming_response.model_copy(
deep=True
),
processed_chunk=_cache_copy,
cache_hit=cache_hit,
)
executor.submit(
self.logging_obj.success_handler,
complete_streaming_response.model_copy(deep=True),
_log_copy,
None,
None,
cache_hit,
@ -2121,11 +2125,13 @@ class CustomStreamWrapper:
"usage",
getattr(complete_streaming_response, "usage"),
)
try:
_copy = complete_streaming_response.model_copy(deep=True)
except RuntimeError:
_copy = complete_streaming_response.model_copy()
asyncio.create_task(
self.async_cache_streaming_response(
processed_chunk=complete_streaming_response.model_copy(
deep=True
),
processed_chunk=_copy,
cache_hit=cache_hit,
)
)

View file

@ -4462,6 +4462,78 @@
"supports_vision": true,
"supports_web_search": true
},
"azure/gpt-5.4-mini": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 7.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
},
"azure/gpt-5.4-nano": {
"cache_read_input_token_cost": 2e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
},
"azure/gpt-image-1": {
"cache_read_input_image_token_cost": 2.5e-06,
"cache_read_input_token_cost": 1.25e-06,

View file

@ -208,26 +208,52 @@ async def exchange_token_with_server(
client_id: str,
client_secret: Optional[str],
code_verifier: Optional[str],
refresh_token: Optional[str] = None,
scope: Optional[str] = None,
):
if grant_type != "authorization_code":
if grant_type not in ("authorization_code", "refresh_token"):
raise HTTPException(status_code=400, detail="Unsupported grant_type")
if mcp_server.token_url is None:
raise HTTPException(status_code=400, detail="MCP server token url is not set")
proxy_base_url = get_request_base_url(request)
token_data = {
"grant_type": "authorization_code",
"client_id": mcp_server.client_id if mcp_server.client_id else client_id,
"client_secret": mcp_server.client_secret
if mcp_server.client_secret
else client_secret,
"code": code,
"redirect_uri": f"{proxy_base_url}/callback",
}
resolved_client_id = mcp_server.client_id if mcp_server.client_id else client_id
resolved_client_secret = (
mcp_server.client_secret if mcp_server.client_secret else client_secret
)
if code_verifier:
token_data["code_verifier"] = code_verifier
if grant_type == "refresh_token":
if not refresh_token:
raise HTTPException(
status_code=400,
detail="refresh_token is required for refresh_token grant",
)
token_data: dict = {
"grant_type": "refresh_token",
"refresh_token": refresh_token,
"client_id": resolved_client_id,
}
if resolved_client_secret is not None:
token_data["client_secret"] = resolved_client_secret
if scope:
token_data["scope"] = scope
else:
if not code:
raise HTTPException(
status_code=400,
detail="code is required for authorization_code grant",
)
proxy_base_url = get_request_base_url(request)
token_data = {
"grant_type": "authorization_code",
"client_id": resolved_client_id,
"code": code,
"redirect_uri": f"{proxy_base_url}/callback",
}
if resolved_client_secret is not None:
token_data["client_secret"] = resolved_client_secret
if code_verifier:
token_data["code_verifier"] = code_verifier
async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
response = await async_client.post(
@ -375,6 +401,8 @@ async def token_endpoint(
client_id: str = Form(...),
client_secret: Optional[str] = Form(None),
code_verifier: str = Form(None),
refresh_token: Optional[str] = Form(None),
scope: Optional[str] = Form(None),
mcp_server_name: Optional[str] = None,
):
"""
@ -408,6 +436,8 @@ async def token_endpoint(
client_id=client_id,
client_secret=client_secret,
code_verifier=code_verifier,
refresh_token=refresh_token,
scope=scope,
)

View file

@ -2955,7 +2955,9 @@ class LiteLLM_ErrorLogs(LiteLLMPydanticObjectBase):
endTime: Union[str, datetime, None]
AUDIT_ACTIONS = Literal["created", "updated", "deleted", "blocked", "rotated"]
AUDIT_ACTIONS = Literal[
"created", "updated", "deleted", "blocked", "unblocked", "rotated"
]
class LiteLLM_AuditLogs(LiteLLMPydanticObjectBase):

View file

@ -29,6 +29,7 @@ from litellm.constants import (
DEFAULT_MAX_RECURSE_DEPTH,
EMAIL_BUDGET_ALERT_MAX_SPEND_ALERT_PERCENTAGE,
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
from litellm.proxy._types import (
RBAC_ROLES,
@ -407,18 +408,21 @@ async def common_checks( # noqa: PLR0915
# 2. If team can call model
if _model and team_object:
if not await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases if valid_token else None,
):
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
type=ProxyErrorTypes.team_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
)
with tracer.trace("litellm.proxy.auth.common_checks.can_team_access_model"):
if not await can_team_access_model(
model=_model,
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases
if valid_token
else None,
):
raise ProxyException(
message=f"Team not allowed to access model. Team={team_object.team_id}, Model={_model}. Allowed team models = {team_object.models}",
type=ProxyErrorTypes.team_model_access_denied,
param="model",
code=status.HTTP_401_UNAUTHORIZED,
)
# Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent
if valid_token is not None and valid_token.agent_id:
@ -443,54 +447,62 @@ async def common_checks( # noqa: PLR0915
## 2.1 If user can call model (if personal key)
if _model and team_object is None and user_object is not None:
await can_user_call_model(
model=_model,
llm_router=llm_router,
user_object=user_object,
)
with tracer.trace("litellm.proxy.auth.common_checks.can_user_call_model"):
await can_user_call_model(
model=_model,
llm_router=llm_router,
user_object=user_object,
)
# 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget)
await _run_project_checks(
project_object=project_object,
_model=_model,
llm_router=llm_router,
skip_budget_checks=skip_budget_checks,
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
with tracer.trace("litellm.proxy.auth.common_checks.run_project_checks"):
await _run_project_checks(
project_object=project_object,
_model=_model,
llm_router=llm_router,
skip_budget_checks=skip_budget_checks,
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
)
# If this is a free model, skip all budget checks
if not skip_budget_checks:
# 3. If team is in budget
await _team_max_budget_check(
team_object=team_object,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
with tracer.trace("litellm.proxy.auth.common_checks.team_max_budget_check"):
await _team_max_budget_check(
team_object=team_object,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
# 3.0.5. If team is over soft budget (alert only, doesn't block)
await _team_soft_budget_check(
team_object=team_object,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
with tracer.trace("litellm.proxy.auth.common_checks.team_soft_budget_check"):
await _team_soft_budget_check(
team_object=team_object,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
# 3.1. If organization is in budget
await _organization_max_budget_check(
valid_token=valid_token,
team_object=team_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
with tracer.trace(
"litellm.proxy.auth.common_checks.organization_max_budget_check"
):
await _organization_max_budget_check(
valid_token=valid_token,
team_object=team_object,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
await _tag_max_budget_check(
request_body=request_body,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
with tracer.trace("litellm.proxy.auth.common_checks.tag_max_budget_check"):
await _tag_max_budget_check(
request_body=request_body,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
)
# 4. If user is in budget
## 4.1 check personal budget, if personal key
@ -508,14 +520,15 @@ async def common_checks( # noqa: PLR0915
)
## 4.2 check team member budget, if team key
await _check_team_member_budget(
team_object=team_object,
user_object=user_object,
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
with tracer.trace("litellm.proxy.auth.common_checks.check_team_member_budget"):
await _check_team_member_budget(
team_object=team_object,
user_object=user_object,
valid_token=valid_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# 5. If end_user ('user' passed to /chat/completions, /embeddings endpoint) is in budget
if (
@ -554,19 +567,21 @@ async def common_checks( # noqa: PLR0915
)
# 11. [OPTIONAL] Vector store checks - is the object allowed to access the vector store
await vector_store_access_check(
request_body=request_body,
team_object=team_object,
valid_token=valid_token,
)
with tracer.trace("litellm.proxy.auth.common_checks.vector_store_access_check"):
await vector_store_access_check(
request_body=request_body,
team_object=team_object,
valid_token=valid_token,
)
# 12. [OPTIONAL] Tool allowlist - key/team allowed_tools (no DB in hot path)
await check_tools_allowlist(
request_body=request_body,
valid_token=valid_token,
team_object=team_object,
route=route,
)
with tracer.trace("litellm.proxy.auth.common_checks.check_tools_allowlist"):
await check_tools_allowlist(
request_body=request_body,
valid_token=valid_token,
team_object=team_object,
route=route,
)
return True

View file

@ -548,13 +548,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
custom_auth_api_key: bool = False
try:
# get the request body
await pre_db_read_auth_checks(
request_data=request_data,
request=request,
route=route,
)
with tracer.trace("litellm.proxy.auth.pre_db_read_auth_checks"):
await pre_db_read_auth_checks(
request_data=request_data,
request=request,
route=route,
)
pass_through_endpoints: Optional[List[dict]] = general_settings.get(
"pass_through_endpoints", None
)
@ -588,9 +587,10 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
### USER-DEFINED AUTH FUNCTION ###
if enterprise_custom_auth is not None:
response = await enterprise_custom_auth(
request=request, api_key=api_key, user_custom_auth=user_custom_auth
)
with tracer.trace("litellm.proxy.auth.enterprise_custom_auth"):
response = await enterprise_custom_auth(
request=request, api_key=api_key, user_custom_auth=user_custom_auth
)
if response is not None and isinstance(response, UserAPIKeyAuth):
validated = UserAPIKeyAuth.model_validate(response)
validated = await _run_post_custom_auth_checks(
@ -706,18 +706,19 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Fall through to virtual key checks
if do_standard_jwt_auth:
result = await JWTAuthManager.auth_builder(
request_data=request_data,
general_settings=general_settings,
api_key=api_key,
jwt_handler=jwt_handler,
route=route,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
parent_otel_span=parent_otel_span,
request_headers=_safe_get_request_headers(request),
)
with tracer.trace("litellm.proxy.auth.jwt_auth_builder"):
result = await JWTAuthManager.auth_builder(
request_data=request_data,
general_settings=general_settings,
api_key=api_key,
jwt_handler=jwt_handler,
route=route,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
parent_otel_span=parent_otel_span,
request_headers=_safe_get_request_headers(request),
)
is_proxy_admin = result["is_proxy_admin"]
team_id = result["team_id"]
@ -909,15 +910,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
try:
end_user_params["end_user_id"] = end_user_id
# get end-user object
_end_user_object = await get_end_user_object(
end_user_id=end_user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
with tracer.trace("litellm.proxy.auth.get_end_user_object"):
_end_user_object = await get_end_user_object(
end_user_id=end_user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
route=route,
)
if _end_user_object is not None:
end_user_params[
"allowed_model_region"
@ -960,14 +961,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
if valid_token is None:
## Check CACHE
try:
valid_token = await get_key_object(
hashed_token=hash_token(api_key),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_cache_only=True,
)
with tracer.trace("litellm.proxy.auth.get_key_object_check_cache"):
valid_token = await get_key_object(
hashed_token=hash_token(api_key),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
check_cache_only=True,
)
except Exception:
verbose_logger.debug("api key not found in cache.")
valid_token = None
@ -1139,13 +1141,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
api_key = hash_token(token=api_key)
try:
valid_token = await get_key_object(
hashed_token=api_key,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
with tracer.trace("litellm.proxy.auth.get_key_object_from_db"):
valid_token = await get_key_object(
hashed_token=api_key,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except ProxyException as e:
if e.code == 401 or e.code == "401":
e.message = "Authentication Error, Invalid proxy server token passed. Received API Key = {}, Key Hash (Token) ={}. Unable to find token in cache or `LiteLLM_VerificationTokenTable`".format(
@ -1233,14 +1236,15 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
# Check 2. If user_id for this token is in budget - done in common_checks()
if valid_token.user_id is not None:
try:
user_obj = await get_user_object(
user_id=valid_token.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
with tracer.trace("litellm.proxy.auth.get_user_object"):
user_obj = await get_user_object(
user_id=valid_token.user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except Exception as e:
verbose_logger.debug(
"litellm.proxy.auth.user_api_key_auth.py::user_api_key_auth() - Unable to get user from db/cache. Setting user_obj to None. Exception received - {}".format(
@ -1329,71 +1333,73 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
)
if not skip_budget_checks:
# Check 4. Token Spend is under budget
if RouteChecks.is_llm_api_route(route=route):
await _virtual_key_max_budget_check(
with tracer.trace("litellm.proxy.auth.budget_checks"):
# Check 4. Token Spend is under budget
if RouteChecks.is_llm_api_route(route=route):
await _virtual_key_max_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 5. Max Budget Alert Check
await _virtual_key_max_budget_alert_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 5. Max Budget Alert Check
await _virtual_key_max_budget_alert_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 6. Soft Budget Check
await _virtual_key_soft_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 5. Token Model Spend is under Model budget
max_budget_per_model = valid_token.model_max_budget
current_model = request_data.get("model", None)
if (
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and prisma_client is not None
and current_model is not None
and valid_token.token is not None
):
## GET THE SPEND FOR THIS MODEL
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
# Check 6. Soft Budget Check
await _virtual_key_soft_budget_check(
valid_token=valid_token,
proxy_logging_obj=proxy_logging_obj,
user_obj=user_obj,
)
# Check 5b. End-user model max budget
end_user_mmb = valid_token.end_user_model_max_budget
if (
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
# Check 5. Token Model Spend is under Model budget
max_budget_per_model = valid_token.model_max_budget
current_model = request_data.get("model", None)
if (
max_budget_per_model is not None
and isinstance(max_budget_per_model, dict)
and len(max_budget_per_model) > 0
and prisma_client is not None
and current_model is not None
and valid_token.token is not None
):
## GET THE SPEND FOR THIS MODEL
await model_max_budget_limiter.is_key_within_model_budget(
user_api_key_dict=valid_token,
model=current_model,
)
# Check 5b. End-user model max budget
end_user_mmb = valid_token.end_user_model_max_budget
if (
end_user_mmb is not None
and isinstance(end_user_mmb, dict)
and len(end_user_mmb) > 0
and current_model is not None
and valid_token.end_user_id is not None
):
await model_max_budget_limiter.is_end_user_within_model_budget(
end_user_id=valid_token.end_user_id,
end_user_model_max_budget=end_user_mmb,
model=current_model,
)
# Check 6: Additional Common Checks across jwt + key auth
if valid_token.team_id is not None:
try:
_team_obj = await get_team_object(
team_id=valid_token.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
with tracer.trace("litellm.proxy.auth.get_team_object"):
_team_obj = await get_team_object(
team_id=valid_token.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
except HTTPException:
_team_obj = LiteLLM_TeamTableCachedObj(
team_id=valid_token.team_id,
@ -1431,11 +1437,14 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
litellm.max_budget > 0 and prisma_client is not None
): # user set proxy max budget
cache_key = "{}:spend".format(litellm_proxy_admin_name)
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
with tracer.trace("litellm.proxy.auth.get_global_proxy_spend"):
global_proxy_spend = (
await _fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
)
if global_proxy_spend is not None:
call_info = CallInfo(
@ -1452,21 +1461,22 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
user_info=call_info,
)
)
_ = await common_checks(
request=request,
request_body=request_data,
team_object=_team_obj,
user_object=user_obj,
end_user_object=_end_user_object,
general_settings=general_settings,
global_proxy_spend=global_proxy_spend,
route=route,
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
skip_budget_checks=skip_budget_checks,
project_object=_project_obj,
)
with tracer.trace("litellm.proxy.auth.common_checks"):
_ = await common_checks(
request=request,
request_body=request_data,
team_object=_team_obj,
user_object=user_obj,
end_user_object=_end_user_object,
general_settings=general_settings,
global_proxy_spend=global_proxy_spend,
route=route,
llm_router=llm_router,
proxy_logging_obj=proxy_logging_obj,
valid_token=valid_token,
skip_budget_checks=skip_budget_checks,
project_object=_project_obj,
)
# Token passed all checks
if valid_token is None:
raise HTTPException(401, detail="Invalid API key")

View file

@ -1260,7 +1260,9 @@ class ProxyBaseLLMRequestProcessing:
custom_headers = ProxyBaseLLMRequestProcessing.get_custom_headers(
user_api_key_dict=user_api_key_dict,
call_id=(
_litellm_logging_obj.litellm_call_id if _litellm_logging_obj else None
_litellm_logging_obj.litellm_call_id
if _litellm_logging_obj
else self.data.get("litellm_call_id")
),
model_id=model_id,
version=version,

View file

@ -41,10 +41,8 @@ from litellm.proxy._experimental.mcp_server.db import (
from litellm.proxy._types import *
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.proxy.auth.auth_checks import (
_cache_key_object,
_delete_cache_key_object,
can_team_access_model,
get_key_object,
get_org_object,
get_project_object,
get_team_object,
@ -1656,7 +1654,7 @@ async def _get_and_validate_existing_key(
LiteLLM_VerificationToken: The existing key row
Raises:
HTTPException: If key is not found
ProxyException: 404 if key is not found
"""
if prisma_client is None:
raise HTTPException(
@ -1664,16 +1662,18 @@ async def _get_and_validate_existing_key(
detail={"error": "Database not connected"},
)
existing_key_row = await prisma_client.get_data(
token=token,
table_name="key",
query_type="find_unique",
hashed_token = _hash_token_if_needed(token=token)
existing_key_row = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
)
if existing_key_row is None:
raise HTTPException(
status_code=404,
detail={"error": f"Key not found: {token}"},
raise ProxyException(
message="Key not found.",
type=ProxyErrorTypes.not_found_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
return existing_key_row
@ -2111,19 +2111,11 @@ async def update_key_fn(
key = data_json.pop("key")
# get the row from db
if prisma_client is None:
raise Exception("Not connected to DB!")
existing_key_row = await prisma_client.get_data(
token=data.key, table_name="key", query_type="find_unique"
existing_key_row = await _get_and_validate_existing_key(
token=data.key,
prisma_client=prisma_client,
)
if existing_key_row is None:
raise HTTPException(
status_code=404,
detail={"error": f"Team not found, passed team_id={data.team_id}"},
)
await _validate_update_key_data(
data=data,
existing_key_row=existing_key_row,
@ -2158,6 +2150,8 @@ async def update_key_fn(
)
_data = {**non_default_values, "token": key}
if prisma_client is None:
raise Exception("Not connected to DB!")
response = await prisma_client.update_data(token=key, data=_data)
# Delete - key from cache, since it's been updated!
@ -2330,6 +2324,8 @@ async def bulk_update_keys(
error_message = error_detail.get("error", str(e))
else:
error_message = str(error_detail)
elif isinstance(e, ProxyException):
error_message = e.message
else:
error_message = str(e)
@ -4945,18 +4941,19 @@ async def block_key(
route="/key/block",
)
if litellm.store_audit_logs is True:
# make an audit log for key update
record = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
# Check if the key exists before trying to block it
existing_record = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
)
if existing_record is None:
raise ProxyException(
message="Key not found.",
type=ProxyErrorTypes.not_found_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
if record is None:
raise ProxyException(
message=f"Key {data.key} not found",
type=ProxyErrorTypes.bad_request_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
if litellm.store_audit_logs is True:
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
@ -4970,7 +4967,7 @@ async def block_key(
object_id=hashed_token,
action="blocked",
updated_values="{}",
before_value=record.model_dump_json(),
before_value=existing_record.model_dump_json(),
)
)
)
@ -4979,24 +4976,9 @@ async def block_key(
where={"token": hashed_token}, data={"blocked": True} # type: ignore
)
## UPDATE KEY CACHE
### get cached object ###
key_object = await get_key_object(
## UPDATE KEY CACHE - invalidate so next read re-fetches from DB
await _delete_cache_key_object(
hashed_token=hashed_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
)
### update cached object ###
key_object.blocked = True
### store cached object ###
await _cache_key_object(
hashed_token=hashed_token,
user_api_key_obj=key_object,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
@ -5068,18 +5050,19 @@ async def unblock_key(
route="/key/unblock",
)
if litellm.store_audit_logs is True:
# make an audit log for key update
record = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
# Check if the key exists before trying to unblock it
existing_record = await prisma_client.db.litellm_verificationtoken.find_unique(
where={"token": hashed_token}
)
if existing_record is None:
raise ProxyException(
message="Key not found.",
type=ProxyErrorTypes.not_found_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
if record is None:
raise ProxyException(
message=f"Key {data.key} not found",
type=ProxyErrorTypes.bad_request_error,
param="key",
code=status.HTTP_404_NOT_FOUND,
)
if litellm.store_audit_logs is True:
asyncio.create_task(
create_audit_log_for_update(
request_data=LiteLLM_AuditLogs(
@ -5091,9 +5074,9 @@ async def unblock_key(
changed_by_api_key=user_api_key_dict.api_key,
table_name=LitellmTableNames.KEY_TABLE_NAME,
object_id=hashed_token,
action="blocked",
action="unblocked",
updated_values="{}",
before_value=record.model_dump_json(),
before_value=existing_record.model_dump_json(),
)
)
)
@ -5102,24 +5085,9 @@ async def unblock_key(
where={"token": hashed_token}, data={"blocked": False} # type: ignore
)
## UPDATE KEY CACHE
### get cached object ###
key_object = await get_key_object(
## UPDATE KEY CACHE - invalidate so next read re-fetches from DB
await _delete_cache_key_object(
hashed_token=hashed_token,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=None,
proxy_logging_obj=proxy_logging_obj,
)
### update cached object ###
key_object.blocked = False
### store cached object ###
await _cache_key_object(
hashed_token=hashed_token,
user_api_key_obj=key_object,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)

View file

@ -1399,6 +1399,8 @@ if MCP_AVAILABLE:
client_id: Optional[str] = Form(None),
client_secret: Optional[str] = Form(None),
code_verifier: Optional[str] = Form(None),
refresh_token: Optional[str] = Form(None),
scope: Optional[str] = Form(None),
):
mcp_server = _get_cached_temporary_mcp_server_or_404(server_id)
resolved_client_id = mcp_server.client_id or client_id or ""
@ -1422,6 +1424,8 @@ if MCP_AVAILABLE:
client_id=resolved_client_id,
client_secret=client_secret,
code_verifier=code_verifier,
refresh_token=refresh_token,
scope=scope,
)
@router.post(

View file

@ -3334,7 +3334,9 @@ def _convert_teams_to_response_models(
use_deleted_table: bool,
) -> List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]]:
"""Convert raw Prisma team rows to response models."""
team_list: List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]] = []
team_list: List[
Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]
] = []
for team in teams:
try:
team_dict = team.model_dump()

View file

@ -12,6 +12,7 @@ import click
import httpx
from dotenv import load_dotenv
import litellm
from litellm.constants import DEFAULT_NUM_WORKERS_LITELLM_PROXY
from litellm.secret_managers.main import get_secret_bool
@ -387,7 +388,7 @@ class ProxyInitializationHelpers:
@click.option("--api_base", default=None, help="API base URL.")
@click.option(
"--api_version",
default="2024-07-01-preview",
default=litellm.AZURE_DEFAULT_API_VERSION,
help="For azure - pass in the api version.",
)
@click.option(

View file

@ -4462,6 +4462,78 @@
"supports_vision": true,
"supports_web_search": true
},
"azure/gpt-5.4-mini": {
"cache_read_input_token_cost": 7.5e-08,
"input_cost_per_token": 7.5e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
},
"azure/gpt-5.4-nano": {
"cache_read_input_token_cost": 2e-08,
"input_cost_per_token": 2e-07,
"litellm_provider": "azure",
"max_input_tokens": 1050000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
},
"azure/gpt-image-1": {
"cache_read_input_image_token_cost": 2.5e-06,
"cache_read_input_token_cost": 1.25e-06,
@ -37032,5 +37104,157 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true
},
"volcengine/doubao-seed-2-0-pro-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"source": "https://www.volcengine.com/docs/82379/1330310",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": false,
"supports_vision": true,
"tiered_pricing": [
{
"input_cost_per_token": 4.6e-07,
"output_cost_per_token": 2.3e-06,
"range": [
0,
32000.0
]
},
{
"input_cost_per_token": 7e-07,
"output_cost_per_token": 3.5e-06,
"range": [
32000.0,
128000.0
]
},
{
"input_cost_per_token": 1.4e-06,
"output_cost_per_token": 7e-06,
"range": [
128000.0,
256000.0
]
}
]
},
"volcengine/doubao-seed-2-0-lite-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"source": "https://www.volcengine.com/docs/82379/1330310",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": false,
"supports_vision": true,
"tiered_pricing": [
{
"input_cost_per_token": 8.7e-08,
"output_cost_per_token": 5.2e-07,
"range": [
0,
32000.0
]
},
{
"input_cost_per_token": 1.3e-07,
"output_cost_per_token": 7.8e-07,
"range": [
32000.0,
128000.0
]
},
{
"input_cost_per_token": 2.6e-07,
"output_cost_per_token": 1.6e-06,
"range": [
128000.0,
256000.0
]
}
]
},
"volcengine/doubao-seed-2-0-mini-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"source": "https://www.volcengine.com/docs/82379/1330310",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": false,
"supports_vision": true,
"tiered_pricing": [
{
"input_cost_per_token": 2.9e-08,
"output_cost_per_token": 2.9e-07,
"range": [
0,
32000.0
]
},
{
"input_cost_per_token": 5.8e-08,
"output_cost_per_token": 5.8e-07,
"range": [
32000.0,
128000.0
]
},
{
"input_cost_per_token": 1.2e-07,
"output_cost_per_token": 1.2e-06,
"range": [
128000.0,
256000.0
]
}
]
},
"volcengine/doubao-seed-2-0-code-preview-260215": {
"litellm_provider": "volcengine",
"max_input_tokens": 256000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"source": "https://www.volcengine.com/docs/82379/1330310",
"supports_function_calling": true,
"supports_reasoning": true,
"supports_tool_choice": false,
"supports_vision": true,
"tiered_pricing": [
{
"input_cost_per_token": 4.6e-07,
"output_cost_per_token": 2.3e-06,
"range": [
0,
32000.0
]
},
{
"input_cost_per_token": 7e-07,
"output_cost_per_token": 3.5e-06,
"range": [
32000.0,
128000.0
]
},
{
"input_cost_per_token": 1.4e-06,
"output_cost_per_token": 7e-06,
"range": [
128000.0,
256000.0
]
}
]
}
}

2
poetry.lock generated
View file

@ -8018,4 +8018,4 @@ utils = ["numpydoc"]
[metadata]
lock-version = "2.1"
python-versions = ">=3.9,<4.0"
content-hash = "eda34dfd8b35474beffee18893d6782c7b3d0d3d2c610f66237eb97176f43527"
content-hash = "2cf958f1a04fd5f1ab0e5cfc33bdbf441b518ed6c82d0f2546bf64cd3d2f89be"

View file

@ -593,7 +593,7 @@ def test_datadog_static_methods():
# Test tags format with default values
assert (
"env:unknown,service:litellm-server,version:unknown,HOSTNAME:"
in get_datadog_tags()
in ",".join(get_datadog_tags())
)
# Test with custom environment variables
@ -631,7 +631,7 @@ def test_datadog_static_methods():
# Test tags format with custom values
expected_custom_tags = "env:production,service:custom-service,version:1.0.0,HOSTNAME:test-host,POD_NAME:pod-123"
print("DataDogLogger._get_datadog_tags()", get_datadog_tags())
assert get_datadog_tags() == expected_custom_tags
assert ",".join(get_datadog_tags()) == expected_custom_tags
@pytest.mark.asyncio
@ -672,11 +672,11 @@ def test_get_datadog_tags():
"""Test the _get_datadog_tags static method with various inputs"""
# Test with no standard_logging_object and default env vars
base_tags = get_datadog_tags()
assert "env:" in base_tags
assert "service:" in base_tags
assert "version:" in base_tags
assert "POD_NAME:" in base_tags
assert "HOSTNAME:" in base_tags
assert any("env:" in t for t in base_tags)
assert any("service:" in t for t in base_tags)
assert any("version:" in t for t in base_tags)
assert any("POD_NAME:" in t for t in base_tags)
assert any("HOSTNAME:" in t for t in base_tags)
# Test with custom env vars
test_env = {
@ -705,12 +705,12 @@ def test_get_datadog_tags():
# Test with empty request_tags
standard_logging_obj["request_tags"] = []
tags_empty_request = get_datadog_tags(standard_logging_obj)
assert "request_tag:" not in tags_empty_request
assert not any(t.startswith("request_tag:") for t in tags_empty_request)
# Test with None request_tags
standard_logging_obj["request_tags"] = None
tags_none_request = get_datadog_tags(standard_logging_obj)
assert "request_tag:" not in tags_none_request
assert not any(t.startswith("request_tag:") for t in tags_none_request)
@pytest.mark.asyncio

View file

@ -44,7 +44,7 @@ class TestDatadogTagsRegression:
assert "env:test-env" in tags_legacy
assert "service:test-service" in tags_legacy
# Verify NO team tag (should not invent one)
assert "team:" not in tags_legacy
assert not any(t.startswith("team:") for t in tags_legacy)
# Case 2: New feature (team info provided)
payload_with_team = StandardLoggingPayload(

View file

@ -1666,3 +1666,141 @@ async def test_oauth_authorize_prefers_request_scope_over_server_config():
redirect_url = response.headers["location"]
assert "scope=custom_scope1+custom_scope2" in redirect_url or "scope=custom_scope1%20custom_scope2" in redirect_url
assert "default_scope" not in redirect_url
@pytest.mark.asyncio
async def test_token_endpoint_refresh_token_grant():
"""Test that token endpoint supports refresh_token grant type."""
try:
from fastapi import Request
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
token_endpoint,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
# Clear registry
global_mcp_server_manager.registry.clear()
# Create mock OAuth2 server
oauth2_server = MCPServer(
server_id="google_mcp",
name="google_mcp",
server_name="google_mcp",
alias="google_mcp",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id="test_client_id",
client_secret="test_secret",
authorization_url="https://accounts.google.com/o/oauth2/v2/auth",
token_url="https://oauth2.googleapis.com/token",
scopes=["openid", "email"],
)
global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server
mock_request = MagicMock(spec=Request)
mock_request.base_url = "https://proxy.litellm.example/"
mock_request.headers = {}
# Mock httpx client response with new tokens
mock_response = MagicMock()
mock_response.json.return_value = {
"access_token": "new_access_token",
"token_type": "Bearer",
"expires_in": 3599,
"refresh_token": "new_refresh_token",
}
mock_response.raise_for_status = MagicMock()
mock_async_client = MagicMock()
mock_async_client.post = AsyncMock(return_value=mock_response)
with patch(
"litellm.proxy._experimental.mcp_server.discoverable_endpoints.get_async_httpx_client"
) as mock_get_client:
mock_get_client.return_value = mock_async_client
response = await token_endpoint(
request=mock_request,
grant_type="refresh_token",
code=None,
redirect_uri=None,
client_id="test_client_id",
mcp_server_name="google_mcp",
client_secret="test_secret",
refresh_token="rt-test",
scope="openid email",
)
# Verify the POST was called with refresh_token grant data
mock_async_client.post.assert_called_once()
call_args = mock_async_client.post.call_args
assert call_args[1]["data"]["grant_type"] == "refresh_token"
assert call_args[1]["data"]["refresh_token"] == "rt-test"
assert call_args[1]["data"]["client_id"] == "test_client_id"
assert call_args[1]["data"]["client_secret"] == "test_secret"
assert call_args[1]["data"]["scope"] == "openid email"
# Verify response contains the new tokens
import json
token_data = json.loads(response.body)
assert token_data["access_token"] == "new_access_token"
assert token_data["refresh_token"] == "new_refresh_token"
@pytest.mark.asyncio
async def test_token_endpoint_authorization_code_missing_code():
"""Test that authorization_code grant rejects missing code param."""
try:
from litellm.proxy._experimental.mcp_server.discoverable_endpoints import (
exchange_token_with_server,
)
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import MCPTransport
from litellm.types.mcp import MCPAuth
from litellm.types.mcp_server.mcp_server_manager import MCPServer
except ImportError:
pytest.skip("MCP discoverable endpoints not available")
global_mcp_server_manager.registry.clear()
server = MCPServer(
server_id="test_server",
name="test_server",
server_name="test_server",
alias="test_server",
transport=MCPTransport.http,
auth_type=MCPAuth.oauth2,
client_id="cid",
token_url="https://example.com/token",
)
global_mcp_server_manager.registry[server.server_id] = server
mock_request = MagicMock()
mock_request.base_url = "https://proxy.example/"
mock_request.headers = {}
with pytest.raises(HTTPException) as exc_info:
await exchange_token_with_server(
request=mock_request,
mcp_server=server,
grant_type="authorization_code",
code=None,
redirect_uri="https://example.com/cb",
client_id="cid",
client_secret=None,
code_verifier=None,
)
assert exc_info.value.status_code == 400
assert "code is required" in str(exc_info.value.detail)

View file

@ -1458,10 +1458,6 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
return_value=mock_key_record
)
# Mock get_key_object and _cache_key_object functions
mock_key_object = MagicMock()
mock_key_object.blocked = True # Initially blocked
# Mock hash_token function
def mock_hash_token(token):
if token == "sk-test123456789":
@ -1482,19 +1478,12 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
) # Disable audit logs for simpler test
# Mock get_key_object and _cache_key_object
async def mock_get_key_object(**kwargs):
return mock_key_object
async def mock_cache_key_object(**kwargs):
async def mock_delete_cache_key_object(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_key_object",
mock_get_key_object,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
mock_cache_key_object,
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
mock_delete_cache_key_object,
)
# Create mock request and user auth
@ -1519,11 +1508,9 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
)
assert result == mock_key_record
assert mock_key_object.blocked == False # Should be updated to unblocked
# Reset mocks for second test
mock_prisma_client.db.litellm_verificationtoken.update.reset_mock()
mock_key_object.blocked = True # Reset to blocked state
# Test Case 2: Using already hashed token
hashed_token_request = BlockKeyRequest(key=test_hashed_token)
@ -1541,7 +1528,6 @@ async def test_unblock_key_supports_both_sk_and_hashed_tokens(monkeypatch):
)
assert result == mock_key_record
assert mock_key_object.blocked == False # Should be updated to unblocked
@pytest.mark.asyncio
@ -1579,6 +1565,249 @@ async def test_unblock_key_invalid_key_format(monkeypatch):
assert "Invalid key format" in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_block_key_nonexistent_key_returns_404(monkeypatch):
"""
Test that block_key returns 404 (not misleading 401) when the key
doesn't exist in the database, even when the caller is authenticated
as a proxy admin.
Previously, block_key would call get_key_object() for cache refresh,
which raised a 401 ProxyException with 'Authentication Error' — making
it look like an auth failure when it was really a missing-key error.
"""
from litellm.proxy._types import BlockKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import block_key
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
# find_unique returns None → key does not exist
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
def mock_hash_token(token):
return "abcd1234" * 8 # 64-char hex
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
monkeypatch.setattr("litellm.store_audit_logs", False)
mock_request = MagicMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
)
data = BlockKeyRequest(key="sk-does-not-exist-key")
with pytest.raises(ProxyException) as exc_info:
await block_key(
data=data,
http_request=mock_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
assert exc_info.value.code == "404"
assert "not found" in str(exc_info.value.message).lower()
# Must NOT contain "Authentication Error"
assert "Authentication Error" not in str(exc_info.value.message)
# update should never be called since the key doesn't exist
mock_prisma_client.db.litellm_verificationtoken.update.assert_not_called()
@pytest.mark.asyncio
async def test_unblock_key_nonexistent_key_returns_404(monkeypatch):
"""
Test that unblock_key returns 404 (not misleading 401) when the key
doesn't exist in the database.
"""
from litellm.proxy._types import BlockKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import (
unblock_key,
)
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
# find_unique returns None → key does not exist
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
def mock_hash_token(token):
return "abcd1234" * 8
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
monkeypatch.setattr("litellm.store_audit_logs", False)
mock_request = MagicMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
)
data = BlockKeyRequest(key="sk-does-not-exist-key")
with pytest.raises(ProxyException) as exc_info:
await unblock_key(
data=data,
http_request=mock_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
assert exc_info.value.code == "404"
assert "not found" in str(exc_info.value.message).lower()
assert "Authentication Error" not in str(exc_info.value.message)
mock_prisma_client.db.litellm_verificationtoken.update.assert_not_called()
@pytest.mark.asyncio
async def test_update_key_nonexistent_key_returns_404(monkeypatch):
"""
Test that update_key_fn returns 404 (not misleading 401) when the body
key doesn't exist in the database, even when the caller is authenticated
as a proxy admin via the Authorization header.
"""
from litellm.proxy.management_endpoints.key_management_endpoints import (
update_key_fn,
)
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
# find_unique returns None → key does not exist
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None)
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
mock_request = MagicMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
)
data = UpdateKeyRequest(key="sk-does-not-exist-key")
with pytest.raises(ProxyException) as exc_info:
await update_key_fn(
request=mock_request,
data=data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
assert exc_info.value.code == "404"
assert "not found" in str(exc_info.value.message).lower()
assert "Authentication Error" not in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_block_key_existing_key_succeeds(monkeypatch):
"""
Test that block_key successfully blocks an existing key and
invalidates the cache entry.
"""
from litellm.proxy._types import BlockKeyRequest
from litellm.proxy.management_endpoints.key_management_endpoints import block_key
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
test_hashed_token = "a1b2c3d4e5f6789012345678901234567890123456789012345678901234abcd"
mock_key_record = MagicMock()
mock_key_record.token = test_hashed_token
mock_key_record.blocked = False
mock_key_record.model_dump_json.return_value = (
f'{{"token": "{test_hashed_token}", "blocked": false}}'
)
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=mock_key_record
)
mock_updated_record = MagicMock()
mock_updated_record.token = test_hashed_token
mock_updated_record.blocked = True
mock_prisma_client.db.litellm_verificationtoken.update = AsyncMock(
return_value=mock_updated_record
)
def mock_hash_token(token):
if token.startswith("sk-"):
return test_hashed_token
return token
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
monkeypatch.setattr("litellm.store_audit_logs", False)
# Mock _delete_cache_key_object
async def mock_delete_cache_key_object(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
mock_delete_cache_key_object,
)
mock_request = MagicMock()
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN, api_key="sk-admin", user_id="admin_user"
)
data = BlockKeyRequest(key="sk-test123456789")
result = await block_key(
data=data,
http_request=mock_request,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify the key was found and updated
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once_with(
where={"token": test_hashed_token}
)
mock_prisma_client.db.litellm_verificationtoken.update.assert_called_once_with(
where={"token": test_hashed_token}, data={"blocked": True}
)
assert result == mock_updated_record
@pytest.mark.asyncio
async def test_validate_key_team_change_with_member_permissions():
"""
@ -4871,14 +5100,16 @@ async def test_validate_max_budget():
async def test_get_and_validate_existing_key():
"""
Test _get_and_validate_existing_key helper function.
Tests:
1. Successfully retrieve existing key
2. Key not found raises HTTPException
2. Key not found raises ProxyException
3. Database not connected raises HTTPException
"""
from fastapi import HTTPException
from litellm.proxy._types import ProxyException
# Test Case 1: Successfully retrieve existing key
mock_prisma_client = AsyncMock()
mock_key = LiteLLM_VerificationToken(
@ -4887,39 +5118,49 @@ async def test_get_and_validate_existing_key():
models=["gpt-4"],
team_id=None,
)
mock_prisma_client.get_data = AsyncMock(return_value=mock_key)
result = await _get_and_validate_existing_key(
token="test-key-123",
prisma_client=mock_prisma_client,
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=mock_key
)
assert result == mock_key
mock_prisma_client.get_data.assert_called_once_with(
token="test-key-123",
table_name="key",
query_type="find_unique",
)
# Test Case 2: Key not found raises HTTPException
mock_prisma_client.get_data = AsyncMock(return_value=None)
with pytest.raises(HTTPException) as exc_info:
await _get_and_validate_existing_key(
token="non-existent-key",
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
return_value="hashed-test-key-123",
):
result = await _get_and_validate_existing_key(
token="test-key-123",
prisma_client=mock_prisma_client,
)
assert exc_info.value.status_code == 404
assert "Key not found" in str(exc_info.value.detail)
assert result == mock_key
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once_with(
where={"token": "hashed-test-key-123"}
)
# Test Case 2: Key not found raises ProxyException
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=None
)
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
return_value="hashed-non-existent-key",
):
with pytest.raises(ProxyException) as exc_info:
await _get_and_validate_existing_key(
token="non-existent-key",
prisma_client=mock_prisma_client,
)
assert str(exc_info.value.code) == "404"
assert "Key not found" in exc_info.value.message
# Test Case 3: Database not connected raises HTTPException
with pytest.raises(HTTPException) as exc_info:
await _get_and_validate_existing_key(
token="test-key-123",
prisma_client=None,
)
assert exc_info.value.status_code == 500
assert "Database not connected" in str(exc_info.value.detail)
@ -4960,75 +5201,82 @@ async def test_process_single_key_update():
"tags": ["production"],
}
mock_prisma_client.get_data = AsyncMock(return_value=existing_key)
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
return_value=existing_key
)
mock_updated_key_obj = MagicMock()
mock_updated_key_obj.model_dump.return_value = updated_key_data
mock_prisma_client.update_data = AsyncMock(
return_value={"data": mock_updated_key_obj}
)
# Mock prepare_key_update_data
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
) as mock_prepare:
mock_prepare.return_value = {"max_budget": 100.0, "tags": ["production"]}
# Mock TeamMemberPermissionChecks
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
) as mock_permission_check:
mock_permission_check.return_value = None
# Mock _delete_cache_key_object
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
) as mock_delete_cache:
mock_delete_cache.return_value = None
# Mock hash_token (imported from litellm.proxy._types)
with patch(
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.return_value = "hashed-test-key-123"
# Mock KeyManagementEventHooks
# Mock _hash_token_if_needed
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
return_value="hashed-test-key-123",
):
# Create update request
key_update_item = BulkUpdateKeyRequestItem(
key="test-key-123",
max_budget=100.0,
tags=["production"],
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call the function
result = await _process_single_key_update(
key_update_item=key_update_item,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
llm_router=mock_llm_router,
)
# Verify results
assert result is not None
assert "token" not in result # Token should be removed
assert result.get("max_budget") == 100.0
assert result.get("tags") == ["production"]
# Verify mocks were called
mock_prisma_client.get_data.assert_called_once()
mock_prisma_client.update_data.assert_called_once()
mock_delete_cache.assert_called_once()
# Mock KeyManagementEventHooks
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create update request
key_update_item = BulkUpdateKeyRequestItem(
key="test-key-123",
max_budget=100.0,
tags=["production"],
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call the function
result = await _process_single_key_update(
key_update_item=key_update_item,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
llm_router=mock_llm_router,
)
# Verify results
assert result is not None
assert "token" not in result # Token should be removed
assert result.get("max_budget") == 100.0
assert result.get("tags") == ["production"]
# Verify mocks were called
mock_prisma_client.db.litellm_verificationtoken.find_unique.assert_called_once()
mock_prisma_client.update_data.assert_called_once()
mock_delete_cache.assert_called_once()
@pytest.mark.asyncio
@ -5090,7 +5338,7 @@ async def test_bulk_update_keys_success(monkeypatch):
"tags": ["staging"],
}
mock_prisma_client.get_data = AsyncMock(
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
side_effect=[existing_key_1, existing_key_2]
)
mock_updated_key_1_obj = MagicMock()
@ -5103,7 +5351,7 @@ async def test_bulk_update_keys_success(monkeypatch):
{"data": mock_updated_key_2_obj},
]
)
# Patch dependencies
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
@ -5115,7 +5363,7 @@ async def test_bulk_update_keys_success(monkeypatch):
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_llm_router)
# Mock helper functions
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
@ -5124,7 +5372,7 @@ async def test_bulk_update_keys_success(monkeypatch):
{"max_budget": 100.0, "tags": ["production"]},
{"max_budget": 200.0, "tags": ["staging"]},
]
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
):
@ -5135,45 +5383,49 @@ async def test_bulk_update_keys_success(monkeypatch):
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.side_effect = ["hashed-key-1", "hashed-key-2"]
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
side_effect=["hashed-key-1", "hashed-key-2"],
):
# Create request
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="test-key-2",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 2
assert len(response.failed_updates) == 0
assert response.successful_updates[0].key == "test-key-1"
assert response.successful_updates[1].key == "test-key-2"
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create request
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="test-key-2",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 2
assert len(response.failed_updates) == 0
assert response.successful_updates[0].key == "test-key-1"
assert response.successful_updates[1].key == "test-key-2"
@pytest.mark.asyncio
@ -5218,7 +5470,7 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
}
# First key exists, second key doesn't exist
mock_prisma_client.get_data = AsyncMock(
mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock(
side_effect=[existing_key_1, None] # Second key not found
)
mock_updated_key_1_obj = MagicMock()
@ -5226,7 +5478,9 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
mock_prisma_client.update_data = AsyncMock(
return_value={"data": mock_updated_key_1_obj}
)
# Mock get_data for the error handler path (used to fetch key_info on failure)
mock_prisma_client.get_data = AsyncMock(return_value=None)
# Patch dependencies
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
@ -5238,13 +5492,13 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_llm_router)
# Mock helper functions
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
) as mock_prepare:
mock_prepare.return_value = {"max_budget": 100.0, "tags": ["production"]}
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
):
@ -5255,46 +5509,50 @@ async def test_bulk_update_keys_partial_failures(monkeypatch):
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.return_value = "hashed-key-1"
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
"litellm.proxy.management_endpoints.key_management_endpoints._hash_token_if_needed",
side_effect=["hashed-key-1", "hashed-non-existent-key"],
):
# Create request with one valid and one invalid key
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="non-existent-key",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 1
assert len(response.failed_updates) == 1
assert response.successful_updates[0].key == "test-key-1"
assert response.failed_updates[0].key == "non-existent-key"
assert "Key not found" in response.failed_updates[0].failed_reason
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create request with one valid and one invalid key
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="non-existent-key",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 1
assert len(response.failed_updates) == 1
assert response.successful_updates[0].key == "test-key-1"
assert response.failed_updates[0].key == "non-existent-key"
assert "Key not found" in response.failed_updates[0].failed_reason
@pytest.mark.parametrize(
@ -7379,19 +7637,12 @@ def _setup_block_unblock_mocks(monkeypatch, mock_key_team_id=None):
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
monkeypatch.setattr("litellm.store_audit_logs", False)
async def mock_get_key_object(**kwargs):
return mock_key_object
async def mock_cache_key_object(**kwargs):
async def mock_delete_cache_key_object(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints.get_key_object",
mock_get_key_object,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
mock_cache_key_object,
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
mock_delete_cache_key_object,
)
return mock_prisma_client, test_hashed_token
@ -7638,16 +7889,9 @@ async def test_update_key_non_budget_fields_allowed_for_internal_user(monkeypatc
monkeypatch.setattr("litellm.proxy.proxy_server.hash_token", mock_hash_token)
async def mock_cache_key_object(**kwargs):
pass
async def mock_delete_cache_key_object(**kwargs):
pass
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._cache_key_object",
mock_cache_key_object,
)
monkeypatch.setattr(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object",
mock_delete_cache_key_object,

View file

@ -1519,6 +1519,8 @@ class TestTemporaryMCPSessionEndpoints:
client_id="client",
client_secret="secret",
code_verifier="verifier",
refresh_token=None,
scope=None,
)
assert result is exchange_response
@ -1532,6 +1534,56 @@ class TestTemporaryMCPSessionEndpoints:
client_id="client",
client_secret="secret",
code_verifier="verifier",
refresh_token=None,
scope=None,
)
@pytest.mark.asyncio
async def test_mcp_token_proxies_refresh_token_grant(self):
from litellm.proxy.management_endpoints.mcp_management_endpoints import (
mcp_token,
)
request = MagicMock()
server = generate_mock_mcp_server_config_record(server_id="server-1")
exchange_response = {"access_token": "new-token", "refresh_token": "new-rt"}
with (
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints._get_cached_temporary_mcp_server_or_404",
return_value=server,
) as get_server,
patch(
"litellm.proxy.management_endpoints.mcp_management_endpoints.exchange_token_with_server",
AsyncMock(return_value=exchange_response),
) as exchange_mock,
):
result = await mcp_token(
request=request,
server_id="server-1",
grant_type="refresh_token",
code=None,
redirect_uri=None,
client_id="client",
client_secret="secret",
code_verifier=None,
refresh_token="rt-123",
scope=None,
)
assert result is exchange_response
get_server.assert_called_once_with("server-1")
exchange_mock.assert_awaited_once_with(
request=request,
mcp_server=server,
grant_type="refresh_token",
code=None,
redirect_uri=None,
client_id="client",
client_secret="secret",
code_verifier=None,
refresh_token="rt-123",
scope=None,
)
@pytest.mark.asyncio

View file

@ -280,6 +280,47 @@ class TestProxyInitializationHelpers:
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
mock_uvicorn_run.assert_called_once()
@patch("uvicorn.run")
@patch("atexit.register")
@patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database")
@patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False)
def test_proxy_default_api_version_uses_azure_default(
self, mock_should_update, mock_setup_db, mock_atexit_register, mock_uvicorn_run
):
"""Proxy default api_version should match litellm.AZURE_DEFAULT_API_VERSION for consistency."""
from click.testing import CliRunner
import litellm
from litellm.proxy.proxy_cli import run_server
runner = CliRunner()
mock_proxy_module = MagicMock(
app=MagicMock(),
ProxyConfig=MagicMock(),
KeyManagementSettings=MagicMock(),
save_worker_config=MagicMock(),
)
clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")}
with patch.dict(os.environ, clean_env, clear=True), patch.dict(
"sys.modules",
{
"proxy_server": mock_proxy_module,
"litellm.proxy.proxy_server": mock_proxy_module,
},
), patch(
"litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args"
) as mock_get_args:
mock_get_args.return_value = {
"app": "litellm.proxy.proxy_server:app",
"host": "localhost",
"port": 8000,
}
result = runner.invoke(run_server, ["--local", "--skip_server_startup"])
assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}"
mock_proxy_module.save_worker_config.assert_called_once()
call_kwargs = mock_proxy_module.save_worker_config.call_args[1]
assert call_kwargs["api_version"] == litellm.AZURE_DEFAULT_API_VERSION
@patch("uvicorn.run")
@patch("builtins.print")
def test_keepalive_timeout_flag(self, mock_print, mock_uvicorn_run):

View file

@ -0,0 +1,146 @@
import { render, screen } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { vi } from "vitest";
import { flexRender, getCoreRowModel, useReactTable } from "@tanstack/react-table";
import { getAgentHubTableColumns, AgentHubData } from "./AgentHubTableColumns";
const mockAgent: AgentHubData = {
agent_id: "agent-1",
protocolVersion: "1.0",
name: "Test Agent",
description: "A test agent for unit testing",
url: "https://agent.example.com",
version: "2.0",
capabilities: { streaming: true, caching: false },
defaultInputModes: ["text"],
defaultOutputModes: ["text", "image"],
skills: [
{ id: "s1", name: "Skill One", description: "First skill" },
{ id: "s2", name: "Skill Two", description: "Second skill" },
{ id: "s3", name: "Skill Three", description: "Third skill" },
],
is_public: true,
};
function TestTable({
data,
publicPage = false,
showModal = vi.fn(),
copyToClipboard = vi.fn(),
}: {
data: AgentHubData[];
publicPage?: boolean;
showModal?: ReturnType<typeof vi.fn>;
copyToClipboard?: ReturnType<typeof vi.fn>;
}) {
const columns = getAgentHubTableColumns(showModal, copyToClipboard, publicPage);
const table = useReactTable({ data, columns, getCoreRowModel: getCoreRowModel() });
return (
<table>
<thead>
{table.getHeaderGroups().map((hg) => (
<tr key={hg.id}>
{hg.headers.map((h) => (
<th key={h.id}>{flexRender(h.column.columnDef.header, h.getContext())}</th>
))}
</tr>
))}
</thead>
<tbody>
{table.getRowModel().rows.map((row) => (
<tr key={row.id}>
{row.getVisibleCells().map((cell) => (
<td key={cell.id}>{flexRender(cell.column.columnDef.cell, cell.getContext())}</td>
))}
</tr>
))}
</tbody>
</table>
);
}
describe("AgentHubTableColumns", () => {
it("should render", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("Test Agent")).toBeInTheDocument();
});
it("should display the agent description", () => {
render(<TestTable data={[mockAgent]} />);
// Description appears in both the description column and the mobile view within agent name column
expect(screen.getAllByText("A test agent for unit testing").length).toBeGreaterThanOrEqual(1);
});
it("should display the version with a 'v' prefix", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("v2.0")).toBeInTheDocument();
});
it("should display the protocol version", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("1.0")).toBeInTheDocument();
});
it("should show skill count with correct pluralization", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("3 skills")).toBeInTheDocument();
});
it("should show first two skills and '+1' for overflow", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("Skill One")).toBeInTheDocument();
expect(screen.getByText("Skill Two")).toBeInTheDocument();
expect(screen.getByText("+1")).toBeInTheDocument();
});
it("should show only true capabilities as badges", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("streaming")).toBeInTheDocument();
expect(screen.queryByText("caching")).not.toBeInTheDocument();
});
it("should display I/O modes", () => {
render(<TestTable data={[mockAgent]} />);
// "In:" and "Out:" are in <span> children; getByText with exact:false
// matches against the element's full textContent across child nodes
expect(screen.getByText((_, el) =>
el?.tagName === "P" && el.textContent === "In: text"
)).toBeInTheDocument();
expect(screen.getByText((_, el) =>
el?.tagName === "P" && el.textContent === "Out: text, image"
)).toBeInTheDocument();
});
it("should display 'Yes' badge for public agents", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByText("Yes")).toBeInTheDocument();
});
it("should display 'No' badge for non-public agents", () => {
const privateAgent = { ...mockAgent, is_public: false };
render(<TestTable data={[privateAgent]} />);
expect(screen.getByText("No")).toBeInTheDocument();
});
it("should display a Details button", () => {
render(<TestTable data={[mockAgent]} />);
expect(screen.getByRole("button", { name: /details|info/i })).toBeInTheDocument();
});
it("should show '-' when agent has no capabilities", () => {
const noCapAgent = { ...mockAgent, capabilities: {} };
render(<TestTable data={[noCapAgent]} />);
// The dash is rendered in the capabilities column
expect(screen.getByText("-")).toBeInTheDocument();
});
it("should show singular 'skill' for one skill", () => {
const oneSkillAgent = {
...mockAgent,
skills: [{ id: "s1", name: "Only Skill", description: "One" }],
};
render(<TestTable data={[oneSkillAgent]} />);
expect(screen.getByText("1 skill")).toBeInTheDocument();
});
});

View file

@ -194,7 +194,6 @@ export const getAgentHubTableColumns = (
return publicA - publicB;
},
cell: ({ row }) => {
console.log(`CHECKPOINT 1: ${JSON.stringify(row.original)}`);
const agent = row.original;
return agent.is_public === true ? (

View file

@ -0,0 +1,73 @@
import { renderWithProviders, screen } from "../../../tests/test-utils";
import userEvent from "@testing-library/user-event";
import { vi } from "vitest";
import UsageExportHeader from "./UsageExportHeader";
import type { EntitySpendData } from "./types";
vi.mock("./EntityUsageExportModal", () => ({
default: ({ isOpen, onClose }: { isOpen: boolean; onClose: () => void }) =>
isOpen ? (
<div data-testid="export-modal">
<button onClick={onClose}>Close</button>
</div>
) : null,
}));
const defaultProps = {
dateValue: { from: new Date("2025-01-01"), to: new Date("2025-01-31") },
entityType: "team" as const,
spendData: {
results: [],
metadata: {
total_spend: 0,
total_api_requests: 0,
total_successful_requests: 0,
total_failed_requests: 0,
total_tokens: 0,
},
} satisfies EntitySpendData,
};
describe("UsageExportHeader", () => {
it("should render", () => {
renderWithProviders(<UsageExportHeader {...defaultProps} />);
expect(screen.getByRole("button", { name: /export data/i })).toBeInTheDocument();
});
it("should open the export modal when the export button is clicked", async () => {
const user = userEvent.setup();
renderWithProviders(<UsageExportHeader {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /export data/i }));
expect(screen.getByTestId("export-modal")).toBeInTheDocument();
});
it("should close the export modal when onClose is called", async () => {
const user = userEvent.setup();
renderWithProviders(<UsageExportHeader {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /export data/i }));
await user.click(screen.getByRole("button", { name: /close/i }));
expect(screen.queryByTestId("export-modal")).not.toBeInTheDocument();
});
it("should not show filter dropdown when showFilters is false", () => {
renderWithProviders(<UsageExportHeader {...defaultProps} showFilters={false} />);
expect(screen.queryByText(/filter/i)).not.toBeInTheDocument();
});
it("should show filter dropdown when showFilters is true and options provided", () => {
renderWithProviders(
<UsageExportHeader
{...defaultProps}
showFilters
filterLabel="Team"
filterPlaceholder="Select teams"
filterOptions={[
{ label: "Team A", value: "team-a" },
{ label: "Team B", value: "team-b" },
]}
onFiltersChange={vi.fn()}
/>,
);
expect(screen.getByText("Team")).toBeInTheDocument();
});
});

View file

@ -0,0 +1,98 @@
import { render, screen, act } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { vi } from "vitest";
import { GuardrailConfig } from "./GuardrailConfig";
describe("GuardrailConfig", () => {
const defaultProps = {
guardrailName: "Content Safety",
guardrailType: "Content Safety",
provider: "bedrock",
};
afterEach(() => {
vi.useRealTimers();
});
it("should render", () => {
render(<GuardrailConfig {...defaultProps} />);
expect(screen.getByText("Parameters")).toBeInTheDocument();
});
it("should display the guardrail name in the parameters description", () => {
render(<GuardrailConfig {...defaultProps} />);
expect(screen.getByText(/Configure Content Safety behavior/)).toBeInTheDocument();
});
// Note: Version history entries are hardcoded placeholders in the component.
// These assertions will need updating when wired to real API data.
it("should show version history when 'View history' is clicked", async () => {
const user = userEvent.setup();
render(<GuardrailConfig {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /view history/i }));
expect(screen.getByText("Initial configuration")).toBeInTheDocument();
expect(screen.getByText("Added custom categories list")).toBeInTheDocument();
});
it("should toggle version history text between View/Hide", async () => {
const user = userEvent.setup();
render(<GuardrailConfig {...defaultProps} />);
const button = screen.getByRole("button", { name: /view history/i });
await user.click(button);
expect(screen.getByRole("button", { name: /hide history/i })).toBeInTheDocument();
});
it("should show custom code textarea when custom code override is toggled on", async () => {
const user = userEvent.setup();
render(<GuardrailConfig {...defaultProps} />);
// Walk up from "Custom Code Override" heading to find the enclosing section,
// then locate the switch within it
const heading = screen.getByText("Custom Code Override");
let container = heading.parentElement;
let customCodeSwitch: Element | null = null;
while (container && !customCodeSwitch) {
customCodeSwitch = container.querySelector('[role="switch"]');
container = container.parentElement;
}
if (!customCodeSwitch) {
throw new Error("Could not find the Custom Code Override switch via DOM traversal");
}
await user.click(customCodeSwitch);
expect(screen.getByPlaceholderText(/async def evaluate/)).toBeInTheDocument();
});
it("should hide custom code textarea when custom code override is off", () => {
render(<GuardrailConfig {...defaultProps} />);
// There's an input for categories, but no textarea
expect(screen.queryByPlaceholderText(/async def evaluate/)).not.toBeInTheDocument();
});
it("should show the re-run button in idle state", () => {
render(<GuardrailConfig {...defaultProps} />);
expect(screen.getByRole("button", { name: /re-run on failing logs/i })).toBeInTheDocument();
});
it("should show loading state when re-run is clicked", async () => {
vi.useFakeTimers({ shouldAdvanceTime: true });
const user = userEvent.setup({ advanceTimers: vi.advanceTimersByTime });
render(<GuardrailConfig {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /re-run on failing logs/i }));
expect(screen.getByText(/Running on 10 samples/)).toBeInTheDocument();
});
it("should show success message after re-run completes", async () => {
vi.useFakeTimers({ shouldAdvanceTime: true });
const user = userEvent.setup({ advanceTimers: vi.advanceTimersByTime });
render(<GuardrailConfig {...defaultProps} />);
await user.click(screen.getByRole("button", { name: /re-run on failing logs/i }));
await act(async () => { vi.advanceTimersByTime(2500); });
expect(screen.getByText(/7\/10 would now pass/)).toBeInTheDocument();
});
it("should display the Revert and Save buttons", () => {
render(<GuardrailConfig {...defaultProps} />);
expect(screen.getByRole("button", { name: /revert/i })).toBeInTheDocument();
// The component's hardcoded default version is "v3", so Save shows "v4"
expect(screen.getByRole("button", { name: /save as v\d+/i })).toBeInTheDocument();
});
});

View file

@ -138,4 +138,18 @@ describe("DocsMenu", () => {
await user.click(button);
expect(button).toHaveAttribute("aria-expanded", "true");
});
it("should close menu when clicking outside", async () => {
const user = userEvent.setup();
renderWithProviders(
<div>
<DocsMenu items={items} />
<button>Outside</button>
</div>,
);
await user.click(screen.getByRole("button", { name: /docs/i }));
expect(screen.getByText("Custom pricing")).toBeInTheDocument();
await user.click(screen.getByRole("button", { name: /outside/i }));
expect(screen.queryByText("Custom pricing")).not.toBeInTheDocument();
});
});

View file

@ -26,6 +26,10 @@ const PERMISSION_OPTIONS = [
"/key/unblock",
"/key/bulk_update",
"/key/{key_id}/reset_spend",
"/key/info",
"/key/list",
"/key/aliases",
"/team/daily/activity",
];
interface SettingRowProps {

View file

@ -13,6 +13,7 @@ import {
CreditCardOutlined,
DatabaseOutlined,
ExperimentOutlined,
ExportOutlined,
FileTextOutlined,
FolderOutlined,
KeyOutlined,
@ -400,7 +401,7 @@ const Sidebar: React.FC<SidebarProps> = ({ setPage, defaultSelectedKey, collapse
onClick={(e) => e.stopPropagation()}
style={{ color: "inherit", textDecoration: "none" }}
>
{label}
{label} <ExportOutlined style={{ fontSize: 10, marginLeft: 4 }} />
</a>
);
}

View file

@ -0,0 +1,296 @@
import { render, screen } from "@testing-library/react";
import { describe, it, expect, vi } from "vitest";
import ChatMessageBubble from "./ChatMessageBubble";
import { EndpointType } from "./mode_endpoint_mapping";
import { MessageType } from "./types";
// Mock child components to isolate bubble rendering logic
vi.mock("react-markdown", () => ({
default: ({ children }: { children: string }) => <div data-testid="react-markdown">{children}</div>,
}));
vi.mock("react-syntax-highlighter", () => ({
Prism: ({ children }: { children: string }) => <pre data-testid="syntax-highlighter">{children}</pre>,
}));
vi.mock("react-syntax-highlighter/dist/esm/styles/prism", () => ({
coy: {},
}));
vi.mock("./ReasoningContent", () => ({
default: ({ reasoningContent }: { reasoningContent: string }) => (
<div data-testid="reasoning-content">{reasoningContent}</div>
),
}));
vi.mock("./MCPEventsDisplay", () => ({
default: ({ events }: { events: unknown[] }) => (
<div data-testid="mcp-events-display">{events.length} events</div>
),
}));
vi.mock("./SearchResultsDisplay", () => ({
SearchResultsDisplay: ({ searchResults }: { searchResults: unknown[] }) => (
<div data-testid="search-results-display">{searchResults.length} results</div>
),
}));
vi.mock("./ResponseMetrics", () => ({
default: ({ timeToFirstToken }: { timeToFirstToken?: number }) => (
<div data-testid="response-metrics">TTFT: {timeToFirstToken}</div>
),
}));
vi.mock("./A2AMetrics", () => ({
default: ({ a2aMetadata }: { a2aMetadata: unknown }) => (
<div data-testid="a2a-metrics">A2A</div>
),
}));
vi.mock("./CodeInterpreterOutput", () => ({
default: ({ code }: { code: string }) => <div data-testid="code-interpreter-output">{code}</div>,
}));
vi.mock("./AudioRenderer", () => ({
default: ({ message }: { message: MessageType }) => (
<div data-testid="audio-renderer">{typeof message.content === "string" ? message.content : ""}</div>
),
}));
vi.mock("./ResponsesImageRenderer", () => ({
default: () => <div data-testid="responses-image-renderer" />,
}));
vi.mock("./ChatImageRenderer", () => ({
default: () => <div data-testid="chat-image-renderer" />,
}));
const defaultProps = {
isLastMessage: false,
endpointType: EndpointType.CHAT,
mcpEvents: [],
codeInterpreterResult: null,
accessToken: "test-token",
};
describe("ChatMessageBubble", () => {
it("should render a user message with right-aligned text", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "user", content: "Hello" }}
/>,
);
expect(screen.getByText("user")).toBeInTheDocument();
expect(screen.getByText("Hello")).toBeInTheDocument();
});
it("should render an assistant message with left-aligned text", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "Hi there" }}
/>,
);
expect(screen.getByText("assistant")).toBeInTheDocument();
expect(screen.getByText("Hi there")).toBeInTheDocument();
});
it("should show model badge for assistant messages when model is provided", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "Reply", model: "gpt-4" }}
/>,
);
expect(screen.getByText("gpt-4")).toBeInTheDocument();
});
it("should not show model badge for user messages even when model is set", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "user", content: "Hello", model: "gpt-4" }}
/>,
);
expect(screen.queryByText("gpt-4")).not.toBeInTheDocument();
});
it("should render markdown content via ReactMarkdown", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "**bold text**" }}
/>,
);
expect(screen.getByTestId("react-markdown")).toHaveTextContent("**bold text**");
});
it("should render an image when isImage is true", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "https://example.com/img.png", isImage: true }}
/>,
);
expect(screen.getByAltText("Generated image")).toHaveAttribute("src", "https://example.com/img.png");
});
it("should render AudioRenderer when isAudio is true", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "audio-url", isAudio: true }}
/>,
);
expect(screen.getByTestId("audio-renderer")).toBeInTheDocument();
});
it("should show ReasoningContent when reasoningContent is present", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{ role: "assistant", content: "answer", reasoningContent: "thinking..." }}
/>,
);
expect(screen.getByTestId("reasoning-content")).toHaveTextContent("thinking...");
});
it("should show MCP events on the last assistant message for RESPONSES endpoint", () => {
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
render(
<ChatMessageBubble
{...defaultProps}
isLastMessage={true}
endpointType={EndpointType.RESPONSES}
mcpEvents={mcpEvents as any}
message={{ role: "assistant", content: "response" }}
/>,
);
expect(screen.getByTestId("mcp-events-display")).toHaveTextContent("1 events");
});
it("should show MCP events on the last assistant message for CHAT endpoint", () => {
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
render(
<ChatMessageBubble
{...defaultProps}
isLastMessage={true}
endpointType={EndpointType.CHAT}
mcpEvents={mcpEvents as any}
message={{ role: "assistant", content: "response" }}
/>,
);
expect(screen.getByTestId("mcp-events-display")).toHaveTextContent("1 events");
});
it("should not show MCP events when isLastMessage is false", () => {
const mcpEvents = [{ type: "tool_call", item_id: "1" }];
render(
<ChatMessageBubble
{...defaultProps}
isLastMessage={false}
endpointType={EndpointType.RESPONSES}
mcpEvents={mcpEvents as any}
message={{ role: "assistant", content: "response" }}
/>,
);
expect(screen.queryByTestId("mcp-events-display")).not.toBeInTheDocument();
});
it("should show SearchResultsDisplay when searchResults are present", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{
role: "assistant",
content: "found results",
searchResults: [{ object: "search", search_query: "q", data: [] }],
}}
/>,
);
expect(screen.getByTestId("search-results-display")).toBeInTheDocument();
});
it("should show ResponseMetrics when usage data is present and no a2aMetadata", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{
role: "assistant",
content: "response",
timeToFirstToken: 150,
usage: { completionTokens: 10, promptTokens: 5, totalTokens: 15 },
}}
/>,
);
expect(screen.getByTestId("response-metrics")).toBeInTheDocument();
});
it("should show A2AMetrics when a2aMetadata is present instead of ResponseMetrics", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{
role: "assistant",
content: "agent response",
timeToFirstToken: 100,
a2aMetadata: { taskId: "task-1", status: { state: "completed" } },
}}
/>,
);
expect(screen.getByTestId("a2a-metrics")).toBeInTheDocument();
expect(screen.queryByTestId("response-metrics")).not.toBeInTheDocument();
});
it("should show CodeInterpreterOutput on the last assistant message for RESPONSES endpoint", () => {
render(
<ChatMessageBubble
{...defaultProps}
isLastMessage={true}
endpointType={EndpointType.RESPONSES}
codeInterpreterResult={{
code: "print('hello')",
containerId: "container-1",
annotations: [],
}}
message={{ role: "assistant", content: "result" }}
/>,
);
expect(screen.getByTestId("code-interpreter-output")).toHaveTextContent("print('hello')");
});
it("should render generated image from chat completions via message.image", () => {
render(
<ChatMessageBubble
{...defaultProps}
message={{
role: "assistant",
content: "Here is your image",
image: { url: "https://example.com/generated.png", detail: "auto" },
}}
/>,
);
const images = screen.getAllByAltText("Generated image");
expect(images.some((img) => img.getAttribute("src") === "https://example.com/generated.png")).toBe(true);
});
});

View file

@ -0,0 +1,214 @@
import { RobotOutlined, UserOutlined } from "@ant-design/icons";
import React from "react";
import ReactMarkdown from "react-markdown";
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
import { coy } from "react-syntax-highlighter/dist/esm/styles/prism";
import { CodeInterpreterResult } from "../llm_calls/code_interpreter_handler";
import A2AMetrics from "./A2AMetrics";
import AudioRenderer from "./AudioRenderer";
import ChatImageRenderer from "./ChatImageRenderer";
import CodeInterpreterOutput from "./CodeInterpreterOutput";
import { EndpointType } from "./mode_endpoint_mapping";
import MCPEventsDisplay from "./MCPEventsDisplay";
import type { MCPEvent } from "../../mcp_tools/types";
import ReasoningContent from "./ReasoningContent";
import ResponseMetrics from "./ResponseMetrics";
import ResponsesImageRenderer from "./ResponsesImageRenderer";
import { SearchResultsDisplay } from "./SearchResultsDisplay";
import { MessageType } from "./types";
interface ChatMessageBubbleProps {
message: MessageType;
/** Whether this is the last message in the chat history. */
isLastMessage: boolean;
endpointType: EndpointType;
/** MCP events to display on the last assistant message. */
mcpEvents: MCPEvent[];
/** Code interpreter result to display on the last assistant message. */
codeInterpreterResult: CodeInterpreterResult | null;
/** API key used to fetch code interpreter file downloads. */
accessToken: string;
}
function ChatMessageBubble({
message,
isLastMessage,
endpointType,
mcpEvents,
codeInterpreterResult,
accessToken,
}: ChatMessageBubbleProps) {
const isUser = message.role === "user";
return (
<div className={`mb-4 ${isUser ? "text-right" : "text-left"}`}>
<div
className="inline-block max-w-[80%] rounded-lg shadow-sm p-3.5 px-4"
style={{
backgroundColor: isUser ? "#f0f8ff" : "#ffffff",
border: isUser ? "1px solid #e6f0fa" : "1px solid #f0f0f0",
textAlign: "left",
}}
>
{/* Header: role icon + name + model badge */}
<div className="flex items-center gap-2 mb-1.5">
<div
className="flex items-center justify-center w-6 h-6 rounded-full mr-1"
style={{
backgroundColor: isUser ? "#e6f0fa" : "#f5f5f5",
}}
>
{isUser ? (
<UserOutlined style={{ fontSize: "12px", color: "#2563eb" }} />
) : (
<RobotOutlined style={{ fontSize: "12px", color: "#4b5563" }} />
)}
</div>
<strong className="text-sm capitalize">{message.role}</strong>
{message.role === "assistant" && message.model && (
<span className="text-xs px-2 py-0.5 rounded bg-gray-100 text-gray-600 font-normal">
{message.model}
</span>
)}
</div>
{/* Reasoning content (chain-of-thought) */}
{message.reasoningContent && <ReasoningContent reasoningContent={message.reasoningContent} />}
{/* MCP events at the start of the last assistant message */}
{message.role === "assistant" &&
isLastMessage &&
mcpEvents.length > 0 &&
(endpointType === EndpointType.RESPONSES || endpointType === EndpointType.CHAT) && (
<div className="mb-3">
<MCPEventsDisplay events={mcpEvents} />
</div>
)}
{/* Search results */}
{message.role === "assistant" && message.searchResults && (
<SearchResultsDisplay searchResults={message.searchResults} />
)}
{/* Code Interpreter output for the last assistant message */}
{message.role === "assistant" &&
isLastMessage &&
codeInterpreterResult &&
endpointType === EndpointType.RESPONSES && (
<CodeInterpreterOutput
code={codeInterpreterResult.code}
containerId={codeInterpreterResult.containerId}
annotations={codeInterpreterResult.annotations}
accessToken={accessToken}
/>
)}
{/* Message body */}
<div
className="whitespace-pre-wrap break-words max-w-full message-content"
style={{
wordWrap: "break-word",
overflowWrap: "break-word",
wordBreak: "break-word",
hyphens: "auto",
}}
>
{message.isImage ? (
<img
src={typeof message.content === "string" ? message.content : ""}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: "500px" }}
/>
) : message.isAudio ? (
<AudioRenderer message={message} />
) : (
<>
{/* Attached image for user messages based on endpoint */}
{endpointType === EndpointType.RESPONSES && <ResponsesImageRenderer message={message} />}
{endpointType === EndpointType.CHAT && <ChatImageRenderer message={message} />}
<ReactMarkdown
components={{
code({
node,
inline,
className,
children,
...props
}: React.ComponentPropsWithoutRef<"code"> & {
inline?: boolean;
node?: unknown;
}) {
const match = /language-(\w+)/.exec(className || "");
return !inline && match ? (
<SyntaxHighlighter
style={coy as any}
language={match[1]}
PreTag="div"
className="rounded-md my-2"
wrapLines={true}
wrapLongLines={true}
{...props}
>
{String(children).replace(/\n$/, "")}
</SyntaxHighlighter>
) : (
<code
className={`${className} px-1.5 py-0.5 rounded bg-gray-100 text-sm font-mono`}
style={{ wordBreak: "break-word" }}
{...props}
>
{children}
</code>
);
},
pre: ({ node, ...props }) => (
<pre style={{ overflowX: "auto", maxWidth: "100%" }} {...props} />
),
}}
>
{typeof message.content === "string" ? message.content : ""}
</ReactMarkdown>
{/* Generated image from chat completions */}
{message.image && (
<div className="mt-3">
<img
src={message.image.url}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: "500px" }}
/>
</div>
)}
</>
)}
{/* Response metrics */}
{message.role === "assistant" &&
(message.timeToFirstToken || message.totalLatency || message.usage) &&
!message.a2aMetadata && (
<ResponseMetrics
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
usage={message.usage}
toolName={message.toolName}
/>
)}
{/* A2A Metrics */}
{message.role === "assistant" && message.a2aMetadata && (
<A2AMetrics
a2aMetadata={message.a2aMetadata}
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
/>
)}
</div>
</div>
</div>
);
}
export default ChatMessageBubble;

View file

@ -63,6 +63,7 @@ import EndpointSelector from "./EndpointSelector";
import FilePreviewCard from "./FilePreviewCard";
import MCPEventsDisplay from "./MCPEventsDisplay";
import type { MCPEvent } from "../../mcp_tools/types";
import ChatMessageBubble from "./ChatMessageBubble";
import { EndpointType, getEndpointType } from "./mode_endpoint_mapping";
import ReasoningContent from "./ReasoningContent";
import ResponseMetrics, { TokenUsage } from "./ResponseMetrics";
@ -1932,168 +1933,14 @@ const ChatUI: React.FC<ChatUIProps> = ({
{chatHistory.map((message, index) => (
<div key={index}>
<div className={`mb-4 ${message.role === "user" ? "text-right" : "text-left"}`}>
<div
className="inline-block max-w-[80%] rounded-lg shadow-sm p-3.5 px-4"
style={{
backgroundColor: message.role === "user" ? "#f0f8ff" : "#ffffff",
border: message.role === "user" ? "1px solid #e6f0fa" : "1px solid #f0f0f0",
textAlign: "left",
}}
>
<div className="flex items-center gap-2 mb-1.5">
<div
className="flex items-center justify-center w-6 h-6 rounded-full mr-1"
style={{
backgroundColor: message.role === "user" ? "#e6f0fa" : "#f5f5f5",
}}
>
{message.role === "user" ? (
<UserOutlined style={{ fontSize: "12px", color: "#2563eb" }} />
) : (
<RobotOutlined style={{ fontSize: "12px", color: "#4b5563" }} />
)}
</div>
<strong className="text-sm capitalize">{message.role}</strong>
{message.role === "assistant" && message.model && (
<span className="text-xs px-2 py-0.5 rounded bg-gray-100 text-gray-600 font-normal">
{message.model}
</span>
)}
</div>
{message.reasoningContent && <ReasoningContent reasoningContent={message.reasoningContent} />}
{/* Show MCP events at the start of assistant messages */}
{message.role === "assistant" &&
index === chatHistory.length - 1 &&
mcpEvents.length > 0 &&
(endpointType === EndpointType.RESPONSES || endpointType === EndpointType.CHAT) && (
<div className="mb-3">
<MCPEventsDisplay events={mcpEvents} />
</div>
)}
{/* Show search results at the start of assistant messages */}
{message.role === "assistant" && message.searchResults && (
<SearchResultsDisplay searchResults={message.searchResults} />
)}
{/* Show Code Interpreter output for the last assistant message */}
{message.role === "assistant" &&
index === chatHistory.length - 1 &&
codeInterpreter.result &&
endpointType === EndpointType.RESPONSES && (
<CodeInterpreterOutput
code={codeInterpreter.result.code}
containerId={codeInterpreter.result.containerId}
annotations={codeInterpreter.result.annotations}
accessToken={apiKeySource === "session" ? accessToken || "" : apiKey}
/>
)}
<div
className="whitespace-pre-wrap break-words max-w-full message-content"
style={{
wordWrap: "break-word",
overflowWrap: "break-word",
wordBreak: "break-word",
hyphens: "auto",
}}
>
{message.isImage ? (
<img
src={typeof message.content === "string" ? message.content : ""}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: "500px" }}
/>
) : message.isAudio ? (
<AudioRenderer message={message} />
) : (
<>
{/* Show attached image for user messages based on current endpoint */}
{endpointType === EndpointType.RESPONSES && <ResponsesImageRenderer message={message} />}
{endpointType === EndpointType.CHAT && <ChatImageRenderer message={message} />}
<ReactMarkdown
components={{
code({
node,
inline,
className,
children,
...props
}: React.ComponentPropsWithoutRef<"code"> & {
inline?: boolean;
node?: any;
}) {
const match = /language-(\w+)/.exec(className || "");
return !inline && match ? (
<SyntaxHighlighter
style={coy as any}
language={match[1]}
PreTag="div"
className="rounded-md my-2"
wrapLines={true}
wrapLongLines={true}
{...props}
>
{String(children).replace(/\n$/, "")}
</SyntaxHighlighter>
) : (
<code
className={`${className} px-1.5 py-0.5 rounded bg-gray-100 text-sm font-mono`}
style={{ wordBreak: "break-word" }}
{...props}
>
{children}
</code>
);
},
pre: ({ node, ...props }) => (
<pre style={{ overflowX: "auto", maxWidth: "100%" }} {...props} />
),
}}
>
{typeof message.content === "string" ? message.content : ""}
</ReactMarkdown>
{/* Show generated image from chat completions */}
{message.image && (
<div className="mt-3">
<img
src={message.image.url}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: "500px" }}
/>
</div>
)}
</>
)}
{message.role === "assistant" &&
(message.timeToFirstToken || message.totalLatency || message.usage) &&
!message.a2aMetadata && (
<ResponseMetrics
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
usage={message.usage}
toolName={message.toolName}
/>
)}
{/* A2A Metrics - show for A2A agent responses */}
{message.role === "assistant" && message.a2aMetadata && (
<A2AMetrics
a2aMetadata={message.a2aMetadata}
timeToFirstToken={message.timeToFirstToken}
totalLatency={message.totalLatency}
/>
)}
</div>
</div>
</div>
<ChatMessageBubble
message={message}
isLastMessage={index === chatHistory.length - 1}
endpointType={endpointType as EndpointType}
mcpEvents={mcpEvents}
codeInterpreterResult={codeInterpreter.result}
accessToken={apiKeySource === "session" ? accessToken || "" : apiKey}
/>
</div>
))}

View file

@ -151,6 +151,39 @@ describe("GuardrailViewer", () => {
expect(screen.queryByText(/Raw Bedrock Guardrail Response/)).not.toBeInTheDocument();
});
it("renders without crashing when guardrail_mode is null", () => {
const data = makeGuardrailInformation({ guardrail_mode: null });
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
// Null mode should display as dash
expect(screen.getByText("—")).toBeInTheDocument();
});
it("renders without crashing when guardrail_mode is an object", () => {
const data = makeGuardrailInformation({
guardrail_mode: { default: "pre_call", tags: {} },
});
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
expect(screen.getByText("PRE-CALL")).toBeInTheDocument();
});
it("renders without crashing when guardrail_mode is an array and shows in both timeline buckets", () => {
const data = makeGuardrailInformation({
guardrail_mode: ["pre_call", "post_call"],
});
renderWithProviders(<GuardrailViewer data={data} />);
expect(screen.getByText("Guardrails & Policy Compliance")).toBeInTheDocument();
// Mode badge shows first element formatted
expect(screen.getByText("PRE-CALL")).toBeInTheDocument();
// Entry should appear in both pre-call and post-call timeline sections
expect(screen.getByText(/Pre-call guardrail:/)).toBeInTheDocument();
expect(screen.getByText(/Post-call guardrail:/)).toBeInTheDocument();
});
it("integration: renders with real Bedrock details without mocks", async () => {
const user = userEvent.setup();
const data = makeGuardrailInformation({

View file

@ -40,7 +40,7 @@ interface GuardrailInformation {
duration: number;
end_time: number;
start_time: number;
guardrail_mode: string;
guardrail_mode: string | string[] | Record<string, unknown> | null;
guardrail_name: string;
guardrail_status: string;
guardrail_response: GuardrailEntity[] | BedrockGuardrailResponse | any;
@ -77,9 +77,50 @@ const PROVIDERS_WITH_CUSTOM_RENDERERS = new Set([
"litellm_content_filter",
]);
const formatMode = (mode: unknown): string => {
if (mode == null || mode === "") return "—";
const s = typeof mode === "string" ? mode : String(mode);
/**
* Extracts a plain string from guardrail_mode for display purposes.
* Returns the first mode when multiple are present.
*/
const resolveMode = (mode: GuardrailInformation["guardrail_mode"]): string | null => {
if (mode == null) return null;
if (typeof mode === "string") return mode;
if (Array.isArray(mode)) {
const first = mode[0];
return typeof first === "string" ? first : null;
}
if (typeof mode === "object" && "default" in mode) {
const def = mode.default;
if (typeof def === "string") return def;
if (Array.isArray(def)) {
const first = def[0];
return typeof first === "string" ? first : null;
}
}
return null;
};
/**
* Checks whether guardrail_mode includes the given target stage.
* Handles arrays (multi-stage guardrails) by checking all elements.
*/
const modeMatches = (
mode: GuardrailInformation["guardrail_mode"],
target: string,
): boolean => {
if (mode == null) return false;
if (typeof mode === "string") return mode === target;
if (Array.isArray(mode)) return mode.includes(target);
if (typeof mode === "object" && "default" in mode) {
const def = mode.default;
if (typeof def === "string") return def === target;
if (Array.isArray(def)) return def.some((x) => typeof x === "string" && x === target);
}
return false;
};
const formatMode = (mode: GuardrailInformation["guardrail_mode"]): string => {
const s = resolveMode(mode);
if (s == null || s === "") return "—";
return s.replace(/_/g, "-").toUpperCase();
};
@ -301,10 +342,13 @@ const RequestLifecycle = ({ entries }: { entries: GuardrailInformation[] }) => {
// Request received
items.push({ type: "request", label: "Request received", offsetMs: 0 });
// Pre-call guardrails
const preCalls = sorted.filter((e) => e.guardrail_mode === "pre_call");
const postCalls = sorted.filter((e) => e.guardrail_mode === "post_call" || e.guardrail_mode === "logging_only");
const duringCalls = sorted.filter((e) => e.guardrail_mode === "during_call");
// Pre-call guardrails — use modeMatches so array modes (e.g. ["pre_call", "post_call"])
// place the entry in every matching bucket.
const preCalls = sorted.filter((e) => modeMatches(e.guardrail_mode, "pre_call"));
const postCalls = sorted.filter(
(e) => modeMatches(e.guardrail_mode, "post_call") || modeMatches(e.guardrail_mode, "logging_only"),
);
const duringCalls = sorted.filter((e) => modeMatches(e.guardrail_mode, "during_call"));
for (const e of preCalls) {
const offsetMs = Math.round((e.end_time - baseTime) * 1000);

View file

@ -23,7 +23,7 @@ export interface GuardrailInformation {
duration: number;
end_time: number;
start_time: number;
guardrail_mode: string;
guardrail_mode: string | string[] | Record<string, unknown> | null;
guardrail_name: string;
guardrail_status: string;
guardrail_response: GuardrailEntity[] | BedrockGuardrailResponse;