Merge branch 'litellm_internal_staging' into litellm_context_management

Bring the context management PR up to date with the latest changes on
litellm_internal_staging.
This commit is contained in:
mateo-berri 2026-05-25 20:55:27 +00:00
commit 335afa6251
No known key found for this signature in database
62 changed files with 1374 additions and 376 deletions

View file

@ -146,6 +146,37 @@ class SensitiveDataMasker:
return masked_data
_default_masker = SensitiveDataMasker()
def mask_sensitive_keys(
data: Dict[str, Any], sensitive_fields: Set[str]
) -> Dict[str, Any]:
"""Return a new dict with values masked for keys listed in ``sensitive_fields``.
Unlike :meth:`SensitiveDataMasker.mask_dict`, this does exact key-name
matching (not segment matching), so callers explicitly enumerate which
fields to mask. Non-string and None values are passed through unchanged.
Values shorter than ``visible_prefix + visible_suffix`` (8 by default)
fall outside :meth:`SensitiveDataMasker._mask_value`'s partial-reveal
range and are replaced with a fixed-length all-mask string, so a short
credential is never returned verbatim.
"""
masked: Dict[str, Any] = {}
mask_char = _default_masker.mask_char
min_visible = _default_masker.visible_prefix + _default_masker.visible_suffix
for key, value in data.items():
if value is not None and key in sensitive_fields and isinstance(value, str):
if len(value) < min_visible:
masked[key] = mask_char * len(value) if value else value
else:
masked[key] = _default_masker._mask_value(value)
else:
masked[key] = value
return masked
# Usage example:
"""
masker = SensitiveDataMasker()

View file

@ -157,8 +157,8 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
def _get_agent_runtime_arn(self, model: str) -> str:
"""
Extract ARN from model string
model = "agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
returns: "arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
model = "agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp"
returns: "arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp"
"""
parts = model.split("/", 1)
if len(parts) != 2 or parts[0] != "agentcore":
@ -170,7 +170,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
def _extract_region_from_arn(self, arn: str) -> str:
"""
Extract region from ARN
arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC
arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp
returns: us-west-2
"""
parts = arn.split(":")

View file

@ -71,6 +71,11 @@ class BasePassthroughUtils:
request_headers.pop("content-length", None)
request_headers.pop("host", None)
custom_header_names = {header_name.lower() for header_name in headers}
for header_name in list(request_headers.keys()):
if header_name.lower() in custom_header_names:
request_headers.pop(header_name, None)
# Combine request headers with custom headers
headers = {**request_headers, **headers}

View file

@ -118,15 +118,19 @@ class MCPRequestHandler:
return b"{}"
request.body = mock_body # type: ignore
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
request_route = get_request_route(request)
# Only OAuth metadata routes registered under /.well-known/ are public.
# Match on request.url.path (path-only, exact prefix) so the substring
# cannot be smuggled via query string, hostname, or a deeper URL segment.
if request.url.path.startswith("/.well-known/"):
if request_route.startswith("/.well-known/"):
validated_user_api_key_auth = UserAPIKeyAuth()
elif (
not litellm_api_key
and MCPRequestHandler._target_servers_delegate_auth_to_upstream( # noqa: E501
path=request.url.path, mcp_servers=mcp_servers
path=request_route, mcp_servers=mcp_servers
)
):
# Operator opted this oauth2 server into upstream-delegated auth
@ -174,7 +178,7 @@ class MCPRequestHandler:
"401",
"403",
) and MCPRequestHandler._target_servers_use_oauth2(
path=request.url.path, mcp_servers=mcp_servers
path=request_route, mcp_servers=mcp_servers
):
verbose_logger.debug(
"MCP OAuth2: target server is OAuth2-mode, treating "

View file

@ -498,9 +498,18 @@ def route_in_additonal_public_routes(current_route: str):
def get_request_route(request: Request) -> str:
"""
Helper to get the route from the request
Resolve the request route from the ASGI scope, with ``root_path`` stripped.
remove base url from path if set e.g. `/genai/chat/completions` -> `/chat/completions
Prefer this over ``request.url.path`` for any auth, ACL, routing, or
audit-log decision: Starlette reconstructs ``url.path`` by interpolating
the Host header into a URL string and re-parsing with ``urlsplit``, so a
malformed Host (e.g. ``localhost/?x=1``) collapses ``url.path`` to ``"/"``
while FastAPI continues to dispatch on ``scope["path"]``. ``scope["path"]``
is uvicorn's parse of the HTTP request line and matches the actual
handler, so it's the authoritative route.
Also normalizes sub-path deployments by stripping ``scope["root_path"]``
e.g. ``/genai/chat/completions`` -> ``/chat/completions``.
"""
try:
scope = request.scope

View file

@ -627,7 +627,11 @@ class RouteChecks:
Returns:
bool: True if `thread` or `assistant` is in the request path, False otherwise
"""
if "thread" in request.url.path or "assistant" in request.url.path:
# Inline import — auth_utils participates in a proxy import cycle.
from .auth_utils import get_request_route # noqa: PLC0415
route = get_request_route(request)
if "thread" in route or "assistant" in route:
return True
return False

View file

@ -546,7 +546,10 @@ def _add_vector_store_id_from_path(request_data: dict, request: Request) -> None
request_data: The request data dictionary to populate
request: The FastAPI Request object
"""
path = request.url.path
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
path = get_request_route(request)
vector_store_match = re.search(r"/vector_stores/([^/]+)/", path)
if vector_store_match:
vector_store_id = vector_store_match.group(1)

View file

@ -23,11 +23,11 @@ model_list:
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
#########################################################
########## batch specific params ########################
s3_bucket_name: litellm-proxy
s3_bucket_name: litellm-proxy-941277531214
s3_region_name: us-west-2
s3_access_key_id: os.environ/AWS_ACCESS_KEY_ID
s3_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_batch_role_arn: arn:aws:iam::888602223428:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV
aws_batch_role_arn: arn:aws:iam::941277531214:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV
model_info:
mode: batch

View file

@ -55,7 +55,7 @@ guardrails:
litellm_params:
guardrail: bedrock # supported values: "bedrock", "lakera"
mode: "during_call"
guardrailIdentifier: ff6ujrregl1q
guardrailIdentifier: 4w3d1di3snt5
guardrailVersion: "DRAFT"
- guardrail_name: "custom-pre-guard"
litellm_params:

View file

@ -151,7 +151,10 @@ async def test_endpoint(request: Request):
dict: A dictionary containing the route of the request URL.
"""
# ping the proxy server to check if its healthy
return {"route": request.url.path}
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
return {"route": get_request_route(request)}
@router.get(

View file

@ -333,8 +333,10 @@ def _get_metadata_variable_name(request: Request) -> str:
For ALL other endpoints we call this "metadata"
"""
path = request.url.path
# Inline imports — auth_utils/route_checks participate in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
path = get_request_route(request)
if "thread" in path or "assistant" in path:
return "litellm_metadata"

View file

@ -19,6 +19,7 @@ from pydantic import BaseModel, Field
import litellm
from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
from litellm.proxy._types import (
AUDIT_ACTIONS,
LiteLLM_AuditLogs,
@ -34,6 +35,10 @@ from litellm.types.management_endpoints import (
router = APIRouter()
# Cache fields holding credentials. Masked on read so plaintext Redis /
# Sentinel passwords never leave the server in a GET response.
_CACHE_SENSITIVE_FIELDS: set = {"password", "sentinel_password"}
_REDACTED_VALUE = "***REDACTED***"
@ -295,7 +300,11 @@ async def get_cache_settings(
else:
decrypted_settings["redis_type"] = "node"
current_values = decrypted_settings
# Mask credential fields so the GET response never carries
# plaintext Redis / Sentinel passwords off the server.
current_values = mask_sensitive_keys(
decrypted_settings, _CACHE_SENSITIVE_FIELDS
)
# Update field values with current values
for field in cache_fields:

View file

@ -1568,6 +1568,9 @@ if MCP_AVAILABLE:
from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( # noqa: PLC0415
global_mcp_server_manager,
)
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
server_id = request.path_params.get("server_id", "")
if server_id:
@ -1584,7 +1587,7 @@ if MCP_AVAILABLE:
):
# For /token, require PKCE authorization_code; refresh_token
# grants must NOT bypass auth (see comment above).
path_lower = (request.url.path or "").rstrip("/").lower()
path_lower = get_request_route(request).rstrip("/").lower()
if path_lower.endswith("/token"):
body_data = await _read_request_body(request=request)
grant_type = (body_data or {}).get("grant_type", "")

View file

@ -2,7 +2,7 @@
## Helper utils for the management endpoints (keys/users/teams)
from datetime import datetime
from functools import wraps
from typing import List, Optional, Tuple
from typing import Any, Callable, List, Optional, Tuple
from fastapi import HTTPException, Request
@ -435,6 +435,63 @@ async def send_management_endpoint_alert(
)
async def _emit_management_endpoint_otel_span(
func: Callable,
kwargs: dict,
parent_otel_span: Any,
start_time: datetime,
end_time: datetime,
result: Any = None,
exception: Optional[Exception] = None,
) -> None:
"""Stamp + end the parent OTEL SERVER span for a management endpoint.
Routes the request/response (or exception) through the OTEL success/failure
hook. Falls back to ``func.__name__`` for the route when the handler has no
``http_request`` param — endpoints like ``/key/generate`` never receive one,
and gating the hook on it leaked their SERVER span (created in auth, never
ended → never exported). Always emitting keeps both success and failure
paths consistent.
"""
from litellm.proxy.proxy_server import open_telemetry_logger
if open_telemetry_logger is None:
return
http_request: Optional[Request] = kwargs.get("http_request")
if http_request is not None:
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
route = get_request_route(http_request)
request_body: dict = await _read_request_body(request=http_request)
else:
route = func.__name__
request_body = {}
logging_payload = ManagementEndpointLoggingPayload(
route=route,
request_data=request_body,
response=None,
start_time=start_time,
end_time=end_time,
exception=exception,
)
if exception is None:
await open_telemetry_logger.async_management_endpoint_success_hook(
logging_payload=logging_payload,
parent_otel_span=parent_otel_span,
)
else:
await open_telemetry_logger.async_management_endpoint_failure_hook(
logging_payload=logging_payload,
parent_otel_span=parent_otel_span,
)
def management_endpoint_wrapper(func):
"""
This wrapper does the following:
@ -446,13 +503,10 @@ def management_endpoint_wrapper(func):
@wraps(func)
async def wrapper(*args, **kwargs):
start_time = datetime.now()
_http_request: Optional[Request] = None
try:
result = await func(*args, **kwargs)
end_time = datetime.now()
try:
if kwargs is None:
kwargs = {}
user_api_key_dict: UserAPIKeyAuth = (
kwargs.get("user_api_key_dict") or UserAPIKeyAuth()
)
@ -462,31 +516,16 @@ def management_endpoint_wrapper(func):
user_api_key_dict=user_api_key_dict,
function_name=func.__name__,
)
_http_request = kwargs.get("http_request", None)
parent_otel_span = getattr(user_api_key_dict, "parent_otel_span", None)
if parent_otel_span is not None:
from litellm.proxy.proxy_server import open_telemetry_logger
if open_telemetry_logger is not None:
if _http_request:
_route = _http_request.url.path
_request_body: dict = await _read_request_body(
request=_http_request
)
_response = dict(result) if result is not None else None
logging_payload = ManagementEndpointLoggingPayload(
route=_route,
request_data=_request_body,
response=_response,
start_time=start_time,
end_time=end_time,
)
await open_telemetry_logger.async_management_endpoint_success_hook( # type: ignore
logging_payload=logging_payload,
parent_otel_span=parent_otel_span,
)
await _emit_management_endpoint_otel_span(
func=func,
kwargs=kwargs,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
result=result,
)
# Delete updated/deleted info from cache
_delete_api_key_from_cache(kwargs=kwargs)
@ -502,39 +541,19 @@ def management_endpoint_wrapper(func):
except Exception as e:
end_time = datetime.now()
if kwargs is None:
kwargs = {}
user_api_key_dict: UserAPIKeyAuth = (
kwargs.get("user_api_key_dict") or UserAPIKeyAuth()
)
parent_otel_span = getattr(user_api_key_dict, "parent_otel_span", None)
if parent_otel_span is not None:
from litellm.proxy.proxy_server import open_telemetry_logger
if open_telemetry_logger is not None:
_http_request = kwargs.get("http_request")
if _http_request:
_route = _http_request.url.path
_request_body: dict = await _read_request_body(
request=_http_request
)
else:
_route = func.__name__
_request_body = {}
logging_payload = ManagementEndpointLoggingPayload(
route=_route,
request_data=_request_body,
response=None,
start_time=start_time,
end_time=end_time,
exception=e,
)
await open_telemetry_logger.async_management_endpoint_failure_hook( # type: ignore
logging_payload=logging_payload,
parent_otel_span=parent_otel_span,
)
await _emit_management_endpoint_otel_span(
func=func,
kwargs=kwargs,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
exception=e,
)
raise e

View file

@ -44,6 +44,7 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
from litellm.proxy.utils import is_known_model
from litellm.proxy.vector_store_endpoints.utils import (
@ -1123,6 +1124,9 @@ async def bedrock_proxy_route(
_forward_headers=True,
) # dynamically construct pass-through endpoint based on incoming path
setattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY, data)
# SigV4 signs an exact payload; pass-through must send prepped.body, not json.dumps
# of a dict that hooks may mutate (logging_obj, metadata, etc.).
setattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, prepped.body)
received_value = await endpoint_func(
request,
fastapi_response,

View file

@ -6,7 +6,7 @@ import posixpath
import traceback
from base64 import b64encode
from datetime import datetime
from typing import Any, Dict, List, Optional, Tuple, Union, cast
from typing import Any, Dict, List, Mapping, Optional, Tuple, Union, cast
from urllib.parse import urlencode, urlparse
import httpx
@ -62,6 +62,7 @@ from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
EndpointType,
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
PassthroughStandardLoggingPayload,
)
@ -735,6 +736,22 @@ async def pass_through_request( # noqa: PLR0915
str(url)
)
# SigV4-signed callers (e.g. Bedrock) attach the exact bytes that were
# signed via request.state; we must send those instead of re-encoding the
# parsed dict (hooks mutate it, breaking the signature / Content-Length).
# Tolerate request objects without `state` (test fixtures) and only honor
# values httpx accepts for `content=`.
_request_state = getattr(request, "state", None)
state_raw_body: Optional[Union[str, bytes]] = (
getattr(_request_state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, None)
if _request_state is not None
else None
)
if state_raw_body is not None and not isinstance(
state_raw_body, (str, bytes, bytearray)
):
state_raw_body = None
# Skip body parsing for multipart requests - make_multipart_http_request will handle it
# But if custom_body is provided (e.g., JSON parsed despite multipart content-type), use it
is_multipart = (
@ -883,12 +900,19 @@ async def pass_through_request( # noqa: PLR0915
)
)
else:
# SigV4-signed callers (Bedrock) supply the exact pre-signed bytes;
# otherwise httpx encodes the parsed JSON dict as before.
body_kwargs: Dict[str, Any] = (
{"content": state_raw_body}
if state_raw_body is not None
else {"json": _parsed_body}
)
req = async_client.build_request(
"POST",
url,
json=_parsed_body,
params=requested_query_params,
headers=headers,
**body_kwargs,
)
response = await async_client.send(req, stream=stream)
@ -917,17 +941,28 @@ async def pass_through_request( # noqa: PLR0915
status_code=response.status_code,
)
response = (
await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler(
request=request,
async_client=async_client,
if state_raw_body is not None:
# SigV4-signed callers (Bedrock) require the exact pre-signed bytes
# to be forwarded so the signature/Content-Length stay valid.
response = await async_client.request(
method=request.method,
url=url,
headers=headers,
requested_query_params=requested_query_params,
_parsed_body=_parsed_body,
forward_multipart=is_multipart,
params=requested_query_params,
content=state_raw_body,
)
else:
response = (
await HttpPassThroughEndpointHelpers.non_streaming_http_request_handler(
request=request,
async_client=async_client,
url=url,
headers=headers,
requested_query_params=requested_query_params,
_parsed_body=_parsed_body,
forward_multipart=is_multipart,
)
)
)
verbose_proxy_logger.debug("response.headers= %s", response.headers)
if _is_streaming_response(response) is True:
@ -1225,7 +1260,7 @@ async def _parse_request_data_by_content_type(
def create_pass_through_route(
endpoint,
target: str,
custom_headers: Optional[dict] = None,
custom_headers: Optional[Mapping[str, Any]] = None,
_forward_headers: Optional[bool] = False,
_merge_query_params: Optional[bool] = False,
dependencies: Optional[List] = None,
@ -1272,11 +1307,14 @@ def create_pass_through_route(
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
subpath: str = "", # captures sub-paths when include_subpath=True
):
from litellm.proxy.auth.auth_utils import ( # noqa: PLC0415
get_request_route,
)
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
)
path = request.url.path
path = get_request_route(request)
# Parse request data based on content type
(
@ -1335,9 +1373,12 @@ def create_pass_through_route(
)
)
# Ensure custom_headers is a dict
# Ensure custom_headers is a dict. Botocore returns a HeadersDict
# for SigV4-prepared requests, which is a Mapping but not a dict.
headers_dict = (
param_custom_headers if isinstance(param_custom_headers, dict) else {}
dict(param_custom_headers)
if isinstance(param_custom_headers, Mapping)
else {}
)
# Ensure query_params and custom_body are dicts or None
@ -1380,6 +1421,8 @@ def create_pass_through_route(
finally:
if hasattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY):
delattr(request.state, LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY)
if hasattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY):
delattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY)
return endpoint_func

View file

@ -241,7 +241,10 @@ from litellm.litellm_core_utils.core_helpers import (
)
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.litellm_core_utils.sensitive_data_masker import (
SensitiveDataMasker,
mask_sensitive_keys,
)
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.proxy._types import *
@ -990,6 +993,15 @@ _OPENAPI_HTTP_METHODS = {
}
# Credentials surfaced by `/get/config/callbacks` in the alerting block: the
# full Slack incoming-webhook URL is itself a credential, and the SMTP
# password is a service password. Masked on read so plaintext never reaches
# the UI. Kept here at module scope to match the analogous
# `_SSO_SENSITIVE_FIELDS` / `_CACHE_SENSITIVE_FIELDS` constants in the SSO
# and cache endpoint files.
_ALERTING_SENSITIVE_VARS: Set[str] = {"SLACK_WEBHOOK_URL", "SMTP_PASSWORD"}
def _strip_operation_id_method_suffix(operation_id: str) -> str:
base, separator, suffix = operation_id.rpartition("_")
if separator and suffix in _OPENAPI_HTTP_METHODS:
@ -14708,6 +14720,9 @@ async def get_config(): # noqa: PLR0915
value=env_variable, key=_var
)
_slack_env_vars[_var] = _decrypted_value
_slack_env_vars = mask_sensitive_keys(
_slack_env_vars, _ALERTING_SENSITIVE_VARS
)
_alerting_types = proxy_logging_obj.slack_alerting_instance.alert_types
_all_alert_types = (
@ -14744,6 +14759,7 @@ async def get_config(): # noqa: PLR0915
# decode + decrypt the value
_decrypted_value = decrypt_value_helper(value=env_variable, key=_var)
_email_env_vars[_var] = _decrypted_value
_email_env_vars = mask_sensitive_keys(_email_env_vars, _ALERTING_SENSITIVE_VARS)
alerting_data.append(
{

View file

@ -1817,7 +1817,10 @@ async def ui_view_spend_logs( # noqa: PLR0915
)
try:
is_v2 = "/spend/logs/v2" in request.url.path
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
is_v2 = "/spend/logs/v2" in get_request_route(request)
formats = ["%Y-%m-%d %H:%M:%S", "%Y-%m-%d"] if is_v2 else ["%Y-%m-%d %H:%M:%S"]
def parse_date(date_str: str) -> datetime:

View file

@ -9,6 +9,7 @@ from pydantic.fields import FieldInfo
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.sensitive_data_masker import mask_sensitive_keys
from litellm.proxy._types import *
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.types.proxy.management_endpoints.ui_sso import (
@ -19,6 +20,16 @@ from litellm.types.proxy.management_endpoints.ui_sso import (
router = APIRouter()
# SSO secret fields returned by /get/sso_settings. These are masked on read so
# the UI can show "(set)" without ever transporting the plaintext OAuth secret
# off the server, matching the write-once + masked-on-read contract used for
# the HashiCorp Vault config override.
_SSO_SENSITIVE_FIELDS: Set[str] = {
"google_client_secret",
"microsoft_client_secret",
"generic_client_secret",
}
class IPAddress(BaseModel):
ip: str
@ -728,8 +739,9 @@ async def get_sso_settings():
schema = TypeAdapter(SSOConfig).json_schema(by_alias=True)
# Convert to dict for response
sso_dict = sso_config.model_dump()
# Convert to dict for response, masking OAuth client secrets so plaintext
# is never sent to the UI.
sso_dict = mask_sensitive_keys(sso_config.model_dump(), _SSO_SENSITIVE_FIELDS)
# Add descriptions to the response
result = {

View file

@ -330,11 +330,16 @@ def is_allowed_to_call_vector_store_endpoint(
provider_config.get_vector_store_endpoints_by_type()
)
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
request_route = get_request_route(request)
# Determine the permission type based on the request
permission_type = None
for endpoint in provider_vector_store_endpoints["read"]:
if request.method == endpoint[0] and _does_endpoint_match(
endpoint[1], request.url.path
endpoint[1], request_route
):
permission_type = "read"
break
@ -342,7 +347,7 @@ def is_allowed_to_call_vector_store_endpoint(
if permission_type is None:
for endpoint in provider_vector_store_endpoints["write"]:
if request.method == endpoint[0] and _does_endpoint_match(
endpoint[1], request.url.path
endpoint[1], request_route
):
permission_type = "write"
break
@ -392,10 +397,15 @@ def is_allowed_to_call_vector_store_files_endpoint(
provider_config.get_vector_store_file_endpoints_by_type()
)
# Inline import — auth_utils participates in a proxy import cycle.
from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415
request_route = get_request_route(request)
permission_type: Optional[str] = None
for endpoint in provider_vector_store_endpoints.get("read", ()):
if request.method == endpoint[0] and _does_endpoint_match(
endpoint[1], request.url.path
endpoint[1], request_route
):
permission_type = "read"
break
@ -403,7 +413,7 @@ def is_allowed_to_call_vector_store_files_endpoint(
if permission_type is None:
for endpoint in provider_vector_store_endpoints.get("write", ()):
if request.method == endpoint[0] and _does_endpoint_match(
endpoint[1], request.url.path
endpoint[1], request_route
):
permission_type = "write"
break

View file

@ -7,6 +7,10 @@ from typing_extensions import TypedDict
# JSON without a FastAPI `custom_body` parameter (which would consume the HTTP body).
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY = "litellm_pass_through_custom_body"
# Request.state key for programmatic pass-through callers that must preserve an
# exact byte/string body, such as AWS SigV4-signed requests.
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY = "litellm_pass_through_raw_body"
class EndpointType(str, Enum):
VERTEX_AI = "vertex-ai"

View file

@ -168,7 +168,7 @@ async def test_a2a_completion_bridge_bedrock_agentcore():
litellm._turn_on_debug()
# Bedrock AgentCore ARN (streaming-capable runtime)
agentcore_arn = "arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
agentcore_arn = "arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp"
send_message_payload = {
"message": {

View file

@ -38,7 +38,7 @@ async def test_async_create_file():
file=open(file_path, "rb"),
purpose="batch",
custom_llm_provider="bedrock",
s3_bucket_name="litellm-proxy",
s3_bucket_name="litellm-proxy-941277531214",
)
@ -55,7 +55,7 @@ async def test_async_file_and_batch():
file=open(file_path, "rb"),
purpose="batch",
custom_llm_provider="bedrock",
s3_bucket_name="litellm-proxy",
s3_bucket_name="litellm-proxy-941277531214",
)
print("CREATED FILE RESPONSE=", file_obj)
@ -70,7 +70,7 @@ async def test_async_file_and_batch():
# bedrock specific params
#########################################################
model="us.anthropic.claude-haiku-4-5-20251001-v1:0",
aws_batch_role_arn="arn:aws:iam::888602223428:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV",
aws_batch_role_arn="arn:aws:iam::941277531214:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV",
)
print("CREATED BATCH RESPONSE=", create_batch_response)
@ -129,7 +129,7 @@ async def test_mock_bedrock_file_url_mapping():
),
purpose="batch",
custom_llm_provider="bedrock",
s3_bucket_name="litellm-proxy",
s3_bucket_name="litellm-proxy-941277531214",
)
print(f"PUT URL: {captured_put_url}")

View file

@ -20,7 +20,7 @@ async def test_bedrock_guardrails_pii_masking():
mock_user_api_key_dict = UserAPIKeyAuth()
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailIdentifier="zgkmukebruil",
guardrailVersion="DRAFT",
)
@ -60,7 +60,7 @@ async def test_bedrock_guardrails_pii_masking_content_list():
mock_user_api_key_dict = UserAPIKeyAuth()
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailIdentifier="zgkmukebruil",
guardrailVersion="DRAFT",
)
@ -115,7 +115,7 @@ async def test_bedrock_guardrails_block_messages_api():
mock_user_api_key_dict = UserAPIKeyAuth()
guardrail = BedrockGuardrail(
guardrailIdentifier="ff6ujrregl1q",
guardrailIdentifier="4w3d1di3snt5",
guardrailVersion="DRAFT",
)
@ -166,7 +166,7 @@ async def test_bedrock_guardrails_block_responses_api():
mock_user_api_key_dict = UserAPIKeyAuth()
guardrail = BedrockGuardrail(
guardrailIdentifier="ff6ujrregl1q",
guardrailIdentifier="4w3d1di3snt5",
guardrailVersion="DRAFT",
)
@ -211,7 +211,7 @@ async def test_bedrock_guardrails_with_streaming():
)
guardrail = BedrockGuardrail(
guardrailIdentifier="ff6ujrregl1q",
guardrailIdentifier="4w3d1di3snt5",
guardrailVersion="DRAFT",
supported_event_hooks=[GuardrailEventHooks.post_call],
guardrail_name="bedrock-post-guard",
@ -255,7 +255,7 @@ async def test_bedrock_guardrails_with_streaming_no_violation():
)
guardrail = BedrockGuardrail(
guardrailIdentifier="ff6ujrregl1q",
guardrailIdentifier="4w3d1di3snt5",
guardrailVersion="DRAFT",
supported_event_hooks=[GuardrailEventHooks.post_call],
guardrail_name="bedrock-post-guard",
@ -299,7 +299,7 @@ async def test_bedrock_guardrails_streaming_request_body_mock():
# Create the guardrail
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailIdentifier="zgkmukebruil",
guardrailVersion="DRAFT",
supported_event_hooks=[GuardrailEventHooks.post_call],
guardrail_name="bedrock-post-guard",
@ -382,7 +382,7 @@ async def test_bedrock_guardrail_aws_param_persistence():
from litellm.types.guardrails import GuardrailEventHooks
guardrail = BedrockGuardrail(
guardrailIdentifier="wf0hkdb5x07f",
guardrailIdentifier="zgkmukebruil",
guardrailVersion="DRAFT",
aws_access_key_id="test-access-key",
aws_secret_access_key="test-secret-key",

View file

@ -1,3 +1,4 @@
import json
import logging
import os
import sys
@ -44,6 +45,9 @@ from litellm.llms.bedrock.image_generation.image_handler import (
)
from litellm.llms.bedrock.common_utils import BedrockError
# Base64 placeholder used for mocked Bedrock image responses (a 1x1 PNG).
_MOCK_BEDROCK_IMAGE_B64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
@pytest.mark.parametrize(
"model,expected",
@ -528,17 +532,34 @@ def test_backward_compatibility_regular_nova_model():
def test_amazon_titan_image_gen():
"""Test Amazon Titan image generation with cost tracking."""
from litellm import image_generation
"""Test Amazon Titan image generation with cost tracking.
The Bedrock CI account is not entitled to amazon.titan-image-generator, so
the network call is mocked and only the transform + cost-tracking path is
exercised.
"""
from litellm.llms.custom_httpx.http_handler import HTTPHandler
# Use v2 as v1 has reached end of life
model_id = "bedrock/amazon.titan-image-generator-v2:0"
response = litellm.image_generation(
model=model_id,
prompt="A serene mountain landscape at sunset with a lake reflection",
aws_region_name="us-east-1",
)
mock_payload = {"images": [_MOCK_BEDROCK_IMAGE_B64]}
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = mock_payload
mock_response.text = json.dumps(mock_payload)
mock_response.headers = {}
client = HTTPHandler()
with patch.object(client, "post", return_value=mock_response):
response = litellm.image_generation(
model=model_id,
prompt="A serene mountain landscape at sunset with a lake reflection",
aws_region_name="us-east-1",
aws_access_key_id="fake-access-key-id",
aws_secret_access_key="fake-secret-access-key",
client=client,
)
print(f"response cost: {response._hidden_params['response_cost']}")

View file

@ -7,7 +7,6 @@ import sys
import traceback
from unittest.mock import AsyncMock, MagicMock, patch
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
@ -136,6 +135,51 @@ class TestVertexAIGeminiImageGeneration(BaseImageGenTest):
}
# Base64 placeholder used for mocked Bedrock image responses (a 1x1 PNG).
_MOCK_BEDROCK_IMAGE_B64 = "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg=="
async def _assert_mocked_bedrock_image_generation(call_args: dict) -> None:
"""Run ``aimage_generation`` with the Bedrock HTTP call mocked.
The CI account is not entitled to Nova Canvas, so the network call is
replaced with a canned Bedrock response. This keeps the request transform,
response transform, and cost-tracking path under test without live access.
"""
mock_payload = {"images": [_MOCK_BEDROCK_IMAGE_B64]}
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = mock_payload
mock_response.text = json.dumps(mock_payload)
mock_response.headers = {}
custom_logger = TestCustomLogger()
litellm.logging_callback_manager._reset_all_callbacks()
litellm.callbacks = [custom_logger]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
return_value=mock_response,
):
response = await litellm.aimage_generation(
**call_args,
prompt="A image of a otter",
aws_access_key_id="fake-access-key-id",
aws_secret_access_key="fake-secret-access-key",
)
await asyncio.sleep(1)
assert custom_logger.standard_logging_payload is not None
assert custom_logger.standard_logging_payload["response_cost"] is not None
assert custom_logger.standard_logging_payload["response_cost"] > 0
assert response.data is not None
for d in response.data:
assert isinstance(d, Image)
assert d.b64_json is not None or d.url is not None
class TestBedrockNovaCanvasTextToImage(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
litellm.in_memory_llm_clients_cache = InMemoryCache()
@ -148,6 +192,12 @@ class TestBedrockNovaCanvasTextToImage(BaseImageGenTest):
"aws_region_name": "us-east-1",
}
@pytest.mark.asyncio(scope="module")
async def test_basic_image_generation(self):
await _assert_mocked_bedrock_image_generation(
self.get_base_image_generation_call_args()
)
class TestBedrockNovaCanvasColorGuidedGeneration(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
@ -162,6 +212,12 @@ class TestBedrockNovaCanvasColorGuidedGeneration(BaseImageGenTest):
"aws_region_name": "us-east-1",
}
@pytest.mark.asyncio(scope="module")
async def test_basic_image_generation(self):
await _assert_mocked_bedrock_image_generation(
self.get_base_image_generation_call_args()
)
class TestOpenAIGPTImage1(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:

View file

@ -82,7 +82,7 @@ async def _vertex_ai_mocks():
"bedrock/mistral.mistral-7b-instruct-v0:2",
"openai/gpt-4o",
"openai/self_hosted",
"bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
"bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"vertex_ai/gemini-1.5-flash",
],
)
@ -147,7 +147,7 @@ async def test_litellm_overhead_non_streaming(model):
[
"bedrock/mistral.mistral-7b-instruct-v0:2",
"openai/gpt-4o",
"bedrock/anthropic.claude-3-5-haiku-20241022-v1:0",
"bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
"openai/self_hosted",
],
)

View file

@ -1,7 +1,6 @@
from dataclasses import dataclass, field
from typing import Dict, FrozenSet, List, Optional, Tuple
OMIT = object()
@ -22,6 +21,7 @@ class ModelEntry:
extra_params: Tuple[Tuple[str, str], ...] = field(default_factory=tuple)
required_env: FrozenSet[str] = field(default_factory=frozenset)
caps: FrozenSet[str] = field(default_factory=frozenset)
fail_reason: Optional[str] = None
def params(self) -> Dict[str, str]:
return dict(self.extra_params)
@ -205,6 +205,12 @@ BEDROCK_CONVERSE_MODELS: Tuple[ModelEntry, ...] = (
extra_params=(("aws_region_name", "us-east-1"),),
required_env=_BEDROCK_REQ,
caps=_CAPS_OPUS_4_7,
fail_reason=(
"claude-opus-4-7 is not entitled on the Bedrock CI account "
"941277531214 (model access requires an AWS Sales request, not "
"self-serve); this cell fails on purpose so it stays loud in CI — "
"remove this fail_reason once access is granted"
),
),
ModelEntry(
alias="bedrock-claude-opus-4-6",

View file

@ -15,7 +15,6 @@ from .grid_spec import (
all_cells,
)
_PROMPT_MESSAGES: List[Dict[str, str]] = [
{"role": "user", "content": "Step by step, calculate 47 * 53. Show your work."}
]
@ -168,6 +167,9 @@ async def test_reasoning_effort_grid(
if skip_reason:
pytest.skip(skip_reason)
if model.fail_reason:
pytest.xfail(model.fail_reason)
if route_name == "bedrock_invoke_messages":
status, exc = await _call_messages(model, effort)
else:

View file

@ -19,8 +19,8 @@ import httpx
@pytest.mark.parametrize(
"model",
[
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_13sf6-cALnp38iZD", # non-streaming invocation
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC", # streaming invocation
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_13sf6-4046UzHSwy", # non-streaming invocation
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp", # streaming invocation
],
)
def test_bedrock_agentcore_basic(model):
@ -44,7 +44,7 @@ def test_bedrock_agentcore_basic(model):
@pytest.mark.parametrize(
"model",
[
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_13sf6-cALnp38iZD", # streaming invocation
"bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_13sf6-4046UzHSwy", # streaming invocation
],
)
async def test_bedrock_agentcore_with_streaming(model):
@ -54,7 +54,7 @@ async def test_bedrock_agentcore_with_streaming(model):
print("running streming test for model=", model)
# litellm._turn_on_debug()
response = await litellm.acompletion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -82,7 +82,7 @@ def test_bedrock_agentcore_with_custom_params():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -105,7 +105,7 @@ def test_bedrock_agentcore_with_custom_params():
url = call_kwargs["url"]
print(f"URL: {url}")
assert (
"/runtimes/arn%3Aaws%3Abedrock-agentcore%3Aus-west-2%3A888602223428%3Aruntime%2Fhosted_agent_r9jvp-3ySZuRHjLC/invocations"
"/runtimes/arn%3Aaws%3Abedrock-agentcore%3Aus-west-2%3A941277531214%3Aruntime%2Fhosted_agent_r9jvp-Rq79QFC2fp/invocations"
in url
)
assert "qualifier=DEFAULT" in url
@ -150,7 +150,7 @@ def test_bedrock_agentcore_with_runtime_user_id():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -189,7 +189,7 @@ def test_bedrock_agentcore_with_session_and_user():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -234,7 +234,7 @@ def test_bedrock_agentcore_with_api_key_bearer_token():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -282,7 +282,7 @@ def test_bedrock_agentcore_with_all_parameters():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -350,7 +350,7 @@ def test_bedrock_agentcore_without_api_key_uses_sigv4():
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",
@ -625,7 +625,7 @@ def test_agentcore_synchronous_non_streaming_response():
with patch.object(client, "post", return_value=mock_response) as mock_post:
# Make a synchronous (non-streaming) completion call
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/hosted_agent_r9jvp-Rq79QFC2fp",
messages=[
{
"role": "user",

View file

@ -115,7 +115,7 @@ def test_completion_bedrock_guardrails(streaming):
],
max_tokens=10,
guardrailConfig={
"guardrailIdentifier": "ff6ujrregl1q",
"guardrailIdentifier": "4w3d1di3snt5",
"guardrailVersion": "DRAFT",
"trace": "enabled",
},
@ -144,7 +144,7 @@ def test_completion_bedrock_guardrails(streaming):
stream=True,
max_tokens=10,
guardrailConfig={
"guardrailIdentifier": "ff6ujrregl1q",
"guardrailIdentifier": "4w3d1di3snt5",
"guardrailVersion": "DRAFT",
"trace": "enabled",
},
@ -475,7 +475,7 @@ def test_bedrock_claude_3(image_url):
],
}
response: ModelResponse = completion(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
num_retries=3,
**data,
) # type: ignore
@ -498,7 +498,7 @@ def test_bedrock_claude_3(image_url):
@pytest.mark.parametrize(
"model",
[
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
# "meta.llama3-70b-instruct-v1:0",
# "anthropic.claude-v2",
# "mistral.mixtral-8x7b-instruct-v0:1",
@ -537,7 +537,7 @@ def test_bedrock_stop_value(stop, model):
@pytest.mark.parametrize(
"model",
[
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
"mistral.mixtral-8x7b-instruct-v0:1",
],
)
@ -602,7 +602,7 @@ def test_bedrock_claude_3_tool_calling():
}
]
response: ModelResponse = completion(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
tools=tools,
tool_choice="auto",
@ -630,7 +630,7 @@ def test_bedrock_claude_3_tool_calling():
)
# In the second response, Claude should deduce answer from tool results
second_response = completion(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
tools=tools,
tool_choice="auto",
@ -737,7 +737,7 @@ def test_bedrock_ptu():
from openai.types.chat import ChatCompletion
model_id = (
"arn:aws:bedrock:us-west-2:888602223428:provisioned-model/8fxff74qyhs3"
"arn:aws:bedrock:us-west-2:941277531214:provisioned-model/8fxff74qyhs3"
)
try:
response = litellm.completion(
@ -752,7 +752,7 @@ def test_bedrock_ptu():
assert "url" in mock_client_post.call_args.kwargs
assert (
mock_client_post.call_args.kwargs["url"]
== "https://bedrock-runtime.us-west-2.amazonaws.com/model/arn%3Aaws%3Abedrock%3Aus-west-2%3A888602223428%3Aprovisioned-model%2F8fxff74qyhs3/converse"
== "https://bedrock-runtime.us-west-2.amazonaws.com/model/arn%3Aaws%3Abedrock%3Aus-west-2%3A941277531214%3Aprovisioned-model%2F8fxff74qyhs3/converse"
)
mock_client_post.assert_called_once()
@ -2327,7 +2327,7 @@ def test_bedrock_cross_region_inference(monkeypatch):
def test_bedrock_empty_content_real_call():
completion(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=[
{
"role": "user",

View file

@ -299,7 +299,10 @@ def test_completion_claude_3():
@pytest.mark.parametrize(
"model",
["anthropic/claude-sonnet-4-5-20250929", "anthropic.claude-3-sonnet-20240229-v1:0"],
[
"anthropic/claude-sonnet-4-5-20250929",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
],
)
def test_completion_claude_3_function_call(model):
litellm.set_verbose = True
@ -385,7 +388,7 @@ def test_completion_claude_3_function_call(model):
[
("gpt-3.5-turbo", None, None),
("claude-sonnet-4-5-20250929", None, None),
("anthropic.claude-3-sonnet-20240229-v1:0", None, None),
("us.anthropic.claude-sonnet-4-5-20250929-v1:0", None, None),
# (
# "azure_ai/command-r-plus",
# os.getenv("AZURE_COHERE_API_KEY"),
@ -1578,7 +1581,7 @@ def test_completion_openai():
[
# ("gpt-4o-2024-08-06", None),
# ("azure/gpt-4.1-mini", None),
("bedrock/anthropic.claude-3-sonnet-20240229-v1:0", None),
("bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", None),
# ("azure/gpt-4o-new-test", "2024-08-01-preview"),
],
)
@ -1666,15 +1669,13 @@ def custom_callback(
#################################################
print(
f"""
print(f"""
Model: {model},
Messages: {messages},
User: {user},
Seed: {kwargs["seed"]},
temperature: {kwargs["temperature"]},
"""
)
""")
assert kwargs["user"] == "ishaans app"
assert kwargs["model"] == "gpt-3.5-turbo-1106"
@ -2699,7 +2700,7 @@ def test_bedrock_deepseek_custom_prompt_dict():
def test_bedrock_deepseek_known_tokenizer_config(monkeypatch):
model = (
"deepseek_r1/arn:aws:bedrock:us-west-2:888602223428:imported-model/bnnr6463ejgf"
"deepseek_r1/arn:aws:bedrock:us-west-2:941277531214:imported-model/bnnr6463ejgf"
)
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from unittest.mock import Mock
@ -2914,8 +2915,8 @@ def response_format_tests(response: litellm.ModelResponse):
"model",
[
"bedrock/mistral.mistral-large-2407-v1:0",
"bedrock/cohere.command-r-plus-v1:0",
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-haiku-4-5-20251001-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
"mistral.mistral-7b-instruct-v0:2",
"meta.llama3-8b-instruct-v1:0",
],

View file

@ -142,7 +142,8 @@ def trade(model_name: str) -> List[Trade]: # type: ignore
@pytest.mark.parametrize(
"model", ["claude-haiku-4-5-20251001", "anthropic.claude-3-haiku-20240307-v1:0"]
"model",
["claude-haiku-4-5-20251001", "us.anthropic.claude-haiku-4-5-20251001-v1:0"],
)
@pytest.mark.flaky(retries=6, delay=10)
def test_function_call_parsing(model):

View file

@ -49,7 +49,7 @@ def get_current_weather(location, unit="fahrenheit"):
"mistral/mistral-large-latest",
"claude-haiku-4-5-20251001",
"gemini/gemini-2.5-flash-lite",
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
],
)
@pytest.mark.flaky(retries=3, delay=1)
@ -267,7 +267,6 @@ def test_aaparallel_function_call_with_anthropic_thinking(model):
from litellm.types.utils import ChatCompletionMessageToolCall, Function, Message
_PARALLEL_TOOL_HISTORY_MESSAGES = [
{
"role": "user",
@ -303,7 +302,7 @@ _PARALLEL_TOOL_HISTORY_MESSAGES = [
[
# Bedrock Converse still requires modify_params to inject the dummy tool.
(
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
_PARALLEL_TOOL_HISTORY_MESSAGES,
True,
),
@ -314,7 +313,7 @@ _PARALLEL_TOOL_HISTORY_MESSAGES = [
False,
),
(
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
[
{
"role": "user",
@ -579,7 +578,7 @@ def test_groq_parallel_function_call():
@pytest.mark.parametrize(
"model",
[
"bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
"bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
],
)
def test_passing_tool_result_as_list(model):

View file

@ -57,7 +57,7 @@ async def test_completion_sagemaker(sync_mode):
print("testing sagemaker")
if sync_mode is True:
response = litellm.completion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -67,7 +67,7 @@ async def test_completion_sagemaker(sync_mode):
)
else:
response = await litellm.acompletion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -158,7 +158,7 @@ async def test_completion_sagemaker_messages_api(sync_mode):
"model",
[
# "sagemaker_chat/huggingface-pytorch-tgi-inference-2024-08-23-15-48-59-245",
"sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"sagemaker/litellm-ci-textgen",
],
)
# @pytest.mark.flaky(retries=3, delay=1)
@ -218,7 +218,7 @@ async def test_completion_sagemaker_stream(sync_mode, model):
"model",
[
# "sagemaker_chat/huggingface-pytorch-tgi-inference-2024-08-23-15-48-59-245",
"sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"sagemaker/litellm-ci-textgen",
],
)
async def test_completion_sagemaker_streaming_bad_request(sync_mode, model):
@ -256,7 +256,7 @@ async def test_acompletion_sagemaker_non_stream():
"id": "cmpl-mockid",
"object": "text_completion",
"created": 1629800000,
"model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"model": "sagemaker/litellm-ci-textgen",
"choices": [
{
"text": "This is a mock response from SageMaker.",
@ -282,7 +282,7 @@ async def test_acompletion_sagemaker_non_stream():
) as mock_post:
# Act: Call the litellm.acompletion function
response = await litellm.acompletion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -302,7 +302,7 @@ async def test_acompletion_sagemaker_non_stream():
assert args_to_sagemaker == expected_payload
assert (
kwargs["url"]
== "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/jumpstart-dft-hf-textgeneration1-mp-20240815-185614/invocations"
== "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/litellm-ci-textgen/invocations"
)
@ -316,7 +316,7 @@ async def test_completion_sagemaker_non_stream():
"id": "cmpl-mockid",
"object": "text_completion",
"created": 1629800000,
"model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"model": "sagemaker/litellm-ci-textgen",
"choices": [
{
"text": "This is a mock response from SageMaker.",
@ -342,7 +342,7 @@ async def test_completion_sagemaker_non_stream():
) as mock_post:
# Act: Call the litellm.acompletion function
response = litellm.completion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -362,7 +362,7 @@ async def test_completion_sagemaker_non_stream():
assert args_to_sagemaker == expected_payload
assert (
kwargs["url"]
== "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/jumpstart-dft-hf-textgeneration1-mp-20240815-185614/invocations"
== "https://runtime.sagemaker.us-west-2.amazonaws.com/endpoints/litellm-ci-textgen/invocations"
)
@ -377,7 +377,7 @@ async def test_completion_sagemaker_prompt_template_non_stream():
"id": "cmpl-mockid",
"object": "text_completion",
"created": 1629800000,
"model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"model": "sagemaker/litellm-ci-textgen",
"choices": [
{
"text": "This is a mock response from SageMaker.",
@ -433,7 +433,7 @@ async def test_completion_sagemaker_non_stream_with_aws_params():
"id": "cmpl-mockid",
"object": "text_completion",
"created": 1629800000,
"model": "sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
"model": "sagemaker/litellm-ci-textgen",
"choices": [
{
"text": "This is a mock response from SageMaker.",
@ -459,7 +459,7 @@ async def test_completion_sagemaker_non_stream_with_aws_params():
) as mock_post:
# Act: Call the litellm.acompletion function
response = litellm.completion(
model="sagemaker/jumpstart-dft-hf-textgeneration1-mp-20240815-185614",
model="sagemaker/litellm-ci-textgen",
messages=[
{"role": "user", "content": "hi"},
],
@ -482,5 +482,5 @@ async def test_completion_sagemaker_non_stream_with_aws_params():
assert args_to_sagemaker == expected_payload
assert (
kwargs["url"]
== "https://runtime.sagemaker.us-west-5.amazonaws.com/endpoints/jumpstart-dft-hf-textgeneration1-mp-20240815-185614/invocations"
== "https://runtime.sagemaker.us-west-5.amazonaws.com/endpoints/litellm-ci-textgen/invocations"
)

View file

@ -1174,7 +1174,7 @@ async def test_completion_replicate_llama3_streaming(sync_mode):
[
# ["bedrock/ai21.jamba-instruct-v1:0", "us-east-1"],
# ["bedrock/cohere.command-r-plus-v1:0", None],
["anthropic.claude-3-sonnet-20240229-v1:0", None],
["us.anthropic.claude-sonnet-4-5-20250929-v1:0", None],
# ["mistral.mistral-7b-instruct-v0:2", None],
# ["meta.llama3-8b-instruct-v1:0", None],
],
@ -1246,7 +1246,7 @@ def test_bedrock_claude_3_streaming():
try:
litellm.set_verbose = True
response: ModelResponse = completion( # type: ignore
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
max_tokens=10, # type: ignore
stream=True,
@ -1276,7 +1276,7 @@ def test_bedrock_claude_3_streaming():
"model",
[
"claude-haiku-4-5-20251001",
"cohere.command-r-plus-v1:0", # bedrock
"us.anthropic.claude-haiku-4-5-20251001-v1:0", # bedrock
"gpt-3.5-turbo",
],
)
@ -3500,7 +3500,7 @@ def test_unit_test_perplexity_citations_chunk():
[
"gpt-3.5-turbo",
"claude-sonnet-4-5-20250929",
"anthropic.claude-3-sonnet-20240229-v1:0",
"us.anthropic.claude-sonnet-4-5-20250929-v1:0",
# "vertex_ai/claude-3-5-sonnet@20240620",
],
)

View file

@ -27,7 +27,7 @@ async def test_basic_s3_logging(sync_mode, streaming):
verbose_logger.setLevel(level=logging.DEBUG)
litellm.success_callback = ["s3"]
litellm.s3_callback_params = {
"s3_bucket_name": "load-testing-oct",
"s3_bucket_name": "load-testing-oct-941277531214",
"s3_aws_secret_access_key": "os.environ/AWS_SECRET_ACCESS_KEY",
"s3_aws_access_key_id": "os.environ/AWS_ACCESS_KEY_ID",
"s3_region_name": "us-west-2",
@ -64,14 +64,14 @@ async def test_basic_s3_logging(sync_mode, streaming):
await asyncio.sleep(2)
print(f"response: {response}")
total_objects, all_s3_keys = list_all_s3_objects("load-testing-oct")
total_objects, all_s3_keys = list_all_s3_objects("load-testing-oct-941277531214")
# assert that atlest one key has response.id in it
assert any(response_id in key for key in all_s3_keys)
s3 = boto3.client("s3")
# delete all objects
for key in all_s3_keys:
s3.delete_object(Bucket="load-testing-oct", Key=key)
s3.delete_object(Bucket="load-testing-oct-941277531214", Key=key)
@pytest.mark.asyncio
@ -82,7 +82,7 @@ async def test_basic_s3_v2_logging(streaming):
from litellm.integrations.s3_v2 import S3Logger
litellm.s3_callback_params = {
"s3_bucket_name": "load-testing-oct",
"s3_bucket_name": "load-testing-oct-941277531214",
"s3_aws_secret_access_key": "test-secret",
"s3_aws_access_key_id": "test-key",
"s3_region_name": "us-west-2",

View file

@ -2,7 +2,6 @@ import io
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import asyncio
@ -67,7 +66,7 @@ def setup_vector_store_registry():
litellm.vector_store_registry = VectorStoreRegistry(
vector_stores=[
LiteLLM_ManagedVectorStore(
vector_store_id="T37J8R4WTM", custom_llm_provider="bedrock"
vector_store_id="LCYXFBR2TU", custom_llm_provider="bedrock"
)
]
)
@ -111,7 +110,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_completion(
response = await litellm.acompletion(
model="anthropic/claude-3.5-sonnet",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
client=client,
)
except Exception as e:
@ -152,7 +151,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call(
response = await litellm.acompletion(
model="bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
client=async_client,
)
print("OPENAI RESPONSE:", json.dumps(dict(response), indent=4, default=str))
@ -196,7 +195,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_streaming(
response = await litellm.acompletion(
model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
stream=True,
client=async_client,
)
@ -255,7 +254,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools(
model=f"anthropic/{os.environ.get('CI_CD_DEFAULT_ANTHROPIC_MODEL', 'claude-haiku-4-5-20251001')}",
messages=[{"role": "user", "content": "what is litellm?"}],
max_tokens=10,
tools=[{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]}],
tools=[{"type": "file_search", "vector_store_ids": ["LCYXFBR2TU"]}],
)
assert response is not None
@ -279,7 +278,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_with_llm_api_call_with_tools_
tools=[
{
"type": "file_search",
"vector_store_ids": ["T37J8R4WTM"],
"vector_store_ids": ["LCYXFBR2TU"],
"filters": {
"key": "user_id",
"value": "fake-user-id",
@ -387,7 +386,7 @@ async def test_bedrock_kb_request_body_has_transformed_filters(
tools=[
{
"type": "file_search",
"vector_store_ids": ["T37J8R4WTM"],
"vector_store_ids": ["LCYXFBR2TU"],
"filters": {
"key": "user_id",
"value": "fake-user-id",
@ -461,7 +460,7 @@ async def test_openai_with_knowledge_base_mock_openai(setup_vector_store_registr
await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
client=client,
)
except Exception as e:
@ -537,7 +536,7 @@ async def test_openai_with_vector_store_ids_in_tool_call_mock_openai(
await litellm.acompletion(
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
tools=[{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]}],
tools=[{"type": "file_search", "vector_store_ids": ["LCYXFBR2TU"]}],
client=client,
)
except Exception as e:
@ -611,7 +610,7 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist
model="gpt-5.5",
messages=[{"role": "user", "content": "what is litellm?"}],
tools=[
{"type": "file_search", "vector_store_ids": ["T37J8R4WTM"]},
{"type": "file_search", "vector_store_ids": ["LCYXFBR2TU"]},
{"type": "file_search", "vector_store_ids": ["unknownVS"]},
],
client=client,
@ -645,7 +644,7 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist
# model="gpt-5.5",
# messages=[{"role": "user", "content": "what is litellm?"}],
# vector_store_ids = [
# "T37J8R4WTM"
# "LCYXFBR2TU"
# ],
# )
@ -667,7 +666,7 @@ async def test_openai_with_mixed_tool_call_mock_openai(setup_vector_store_regist
# # expect the vector store request metadata object to have the correct values
# vector_store_request_metadata = standard_logging_vector_store_request_metadata[0]
# assert vector_store_request_metadata.get("vector_store_id") == "T37J8R4WTM"
# assert vector_store_request_metadata.get("vector_store_id") == "LCYXFBR2TU"
# assert vector_store_request_metadata.get("query") == "what is litellm?"
# assert vector_store_request_metadata.get("custom_llm_provider") == "bedrock"
@ -723,7 +722,7 @@ async def test_e2e_bedrock_knowledgebase_retrieval_without_vector_store_registry
response = await litellm.acompletion(
model="anthropic/claude-3.5-sonnet",
messages=[{"role": "user", "content": "what is litellm?"}],
vector_store_ids=["T37J8R4WTM"],
vector_store_ids=["LCYXFBR2TU"],
client=client,
)
except Exception as e:

View file

@ -217,9 +217,120 @@ def _create_request_with_host_header(path: str, host_header: str) -> Request:
],
)
def test_get_request_route_not_bypassed_by_malformed_host(host_header: str):
for protected_path in ["/health", "/user/new", "/key/generate", "/get/internal_user_settings"]:
request = _create_request_with_host_header(path=protected_path, host_header=host_header)
result = get_request_route(request)
assert result == protected_path, (
f"Host: {host_header!r} caused route {protected_path!r} to resolve as {result!r}"
for protected_path in [
"/health",
"/user/new",
"/key/generate",
"/get/internal_user_settings",
]:
request = _create_request_with_host_header(
path=protected_path, host_header=host_header
)
result = get_request_route(request)
assert (
result == protected_path
), f"Host: {host_header!r} caused route {protected_path!r} to resolve as {result!r}"
# ---------------------------------------------------------------------------
# Regression tests for variant call sites that previously read request.url.path
# (Host-derived) instead of the ASGI scope path. Each test sends a Host header
# crafted to collapse url.path to a substring the call site's decision logic
# would match on, while scope["path"] is the real (unmatching) route.
# ---------------------------------------------------------------------------
_BYPASS_HOSTS = [
"localhost/?x=1",
"localhost:4000/?x=1",
"localhost/#test",
"localhost:4000/#test",
]
def _is_assistants(req):
return RouteChecks._is_assistants_api_request(req)
def _metadata_var_name(req):
from litellm.proxy.litellm_pre_call_utils import _get_metadata_variable_name
return _get_metadata_variable_name(req)
def _vector_store_id_in_path(req):
from litellm.proxy.common_utils.http_parsing_utils import (
_add_vector_store_id_from_path,
)
data: dict = {}
_add_vector_store_id_from_path(request_data=data, request=req)
return "vector_store_id" in data
# (label, scope_path, host_suffix_template, predicate, expected) — host_suffix_template
# receives the host_header via %s substitution. The predicate is invoked on a Request
# whose scope["path"] is scope_path and whose Host header is the formatted suffix.
#
# The MCP entries (well_known_mcp_bypass, pkce_token_suffix) call
# get_request_route directly rather than the surrounding production handler
# (MCPRequestHandler.process_mcp_request / _mcp_oauth_user_api_key_auth) —
# those handlers require an ASGI scope plus MCP state to invoke, and the call
# sites do nothing with the path except feed it to this helper. The helper-
# level assertion is the relevant signal.
_CALL_SITES = [
("assistants_classification", "/key/generate", "%s/thread", _is_assistants, False),
(
"metadata_variable_name",
"/chat/completions",
"%s/thread",
_metadata_var_name,
"metadata",
),
(
"vector_store_id_extraction",
"/key/generate",
"%s/vector_stores/x/files",
_vector_store_id_in_path,
False,
),
(
"well_known_mcp_bypass",
"/mcp/tools/call",
"/.well-known/%s",
lambda r: get_request_route(r).startswith("/.well-known/"),
False,
),
(
"pkce_token_suffix",
"/mcp/server-id/token",
"%s",
lambda r: get_request_route(r).rstrip("/").lower().endswith("/token"),
True,
),
(
"spend_logs_v2_classification",
"/spend/logs",
"%s/spend/logs/v2",
lambda r: "/spend/logs/v2" in get_request_route(r),
False,
),
("health_route_echo", "/test", "%s", lambda r: get_request_route(r), "/test"),
]
@pytest.mark.parametrize("host_header", _BYPASS_HOSTS)
@pytest.mark.parametrize(
"label,scope_path,host_suffix_template,predicate,expected",
_CALL_SITES,
ids=[c[0] for c in _CALL_SITES],
)
def test_call_site_uses_scope_path(
label, scope_path, host_suffix_template, predicate, expected, host_header
):
"""Each call site that previously read request.url.path must now make its
decision against scope["path"]. The Host header is crafted so url.path
would resolve to a value that flips the decision under the old code."""
request = _create_request_with_host_header(
path=scope_path, host_header=host_suffix_template % host_header
)
assert predicate(request) == expected

View file

@ -3,6 +3,7 @@ async_management_endpoint_{success,failure}_hook integration points."""
import asyncio
from datetime import datetime
from unittest.mock import MagicMock
import pytest
@ -14,6 +15,7 @@ from litellm.proxy._types import (
from ._helpers import (
HttpStatusException,
assert_server_span_attrs,
get_server_span,
make_fastapi_http_exception,
make_httpx_status_error,
)
@ -28,6 +30,10 @@ def _real_user_api_key_dict(parent_span):
)
async def _noop_alert(*args, **kwargs):
return None
async def _drive_admin_failure(*, otel, exception, parent_span, route):
payload = ManagementEndpointLoggingPayload(
route=route,
@ -180,3 +186,173 @@ def test_admin_endpoint_failure_stamps_server_span(
expected_url_path=path,
where=f"{path} {expected_status}",
)
def test_management_wrapper_success_ends_server_span_without_http_request(
server_span_factory, otel_with_exporter, monkeypatch
):
"""Regression: management endpoints whose handler does not declare an
``http_request`` parameter (``/key/generate``, ``/user/new``, ``/mcp/*``,
...) must still get their parent SERVER span stamped + ended on success.
The success hook itself stamps 200 and ``end()``s the parent, but the
wrapper only invoked it when ``http_request`` was present — so on success
the span (created in auth) was never ended and never exported. This drives
the real wrapper around an ``http_request``-less handler and asserts the
SERVER span reaches the exporter with status 200.
"""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
otel, exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", otel, raising=False)
monkeypatch.setattr(mgmt_utils, "send_management_endpoint_alert", _noop_alert)
server_span = server_span_factory(KEY_GENERATE_PATH)
@mgmt_utils.management_endpoint_wrapper
async def fake_generate_key_fn(data=None, user_api_key_dict=None):
# No ``http_request`` parameter — mirrors generate_key_fn et al.
return {"key": "sk-xyz", "key_name": "k"}
asyncio.run(
fake_generate_key_fn(
data={},
user_api_key_dict=_real_user_api_key_dict(server_span),
)
)
assert_server_span_attrs(
exporter,
expected_status=200,
expected_url_path=KEY_GENERATE_PATH,
where="management wrapper success without http_request",
)
def test_management_wrapper_failure_ends_server_span(
server_span_factory, otel_with_exporter, monkeypatch
):
"""When the handler raises, the wrapper must route through the failure hook
and stamp + end the parent SERVER span with the error status — even for an
``http_request``-less handler (route falls back to ``func.__name__``)."""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
otel, exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", otel, raising=False)
server_span = server_span_factory(KEY_GENERATE_PATH)
@mgmt_utils.management_endpoint_wrapper
async def failing_fn(data=None, user_api_key_dict=None):
raise HttpStatusException(500, "boom")
with pytest.raises(HttpStatusException):
asyncio.run(
failing_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))
)
assert_server_span_attrs(
exporter,
expected_status=500,
expected_url_path=KEY_GENERATE_PATH,
where="management wrapper failure",
)
def test_management_wrapper_success_with_http_request(
server_span_factory, otel_with_exporter, monkeypatch
):
"""Cover the branch where the handler DOES declare ``http_request``: the
route comes from ``http_request.url.path`` and the body is read from it."""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
otel, exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", otel, raising=False)
monkeypatch.setattr(mgmt_utils, "send_management_endpoint_alert", _noop_alert)
async def _fake_body(request=None):
return {"team_alias": "t"}
monkeypatch.setattr(mgmt_utils, "_read_request_body", _fake_body)
server_span = server_span_factory("/team/new")
http_request = MagicMock()
http_request.url.path = "/team/new"
@mgmt_utils.management_endpoint_wrapper
async def fake_new_team(data=None, http_request=None, user_api_key_dict=None):
return {"team_id": "t-1"}
asyncio.run(
fake_new_team(
data={},
http_request=http_request,
user_api_key_dict=_real_user_api_key_dict(server_span),
)
)
assert_server_span_attrs(
exporter,
expected_status=200,
expected_url_path="/team/new",
where="management wrapper success with http_request",
)
def test_management_wrapper_noop_when_otel_logger_absent(
server_span_factory, otel_with_exporter, monkeypatch
):
"""When no OTEL logger is registered, the helper early-returns and no SERVER
span is exported — and the handler result is still returned unchanged."""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
_otel, exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", None, raising=False)
monkeypatch.setattr(mgmt_utils, "send_management_endpoint_alert", _noop_alert)
server_span = server_span_factory(KEY_GENERATE_PATH)
@mgmt_utils.management_endpoint_wrapper
async def fake_fn(data=None, user_api_key_dict=None):
return {"ok": True}
result = asyncio.run(
fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))
)
assert result == {"ok": True}
assert get_server_span(exporter) is None
def test_management_wrapper_swallows_post_success_errors(
server_span_factory, otel_with_exporter, monkeypatch
):
"""A failure in post-success bookkeeping (cache invalidation, alerting) must
not propagate — the handler result is returned regardless (non-blocking)."""
import litellm.proxy.proxy_server as proxy_server
from litellm.proxy.management_helpers import utils as mgmt_utils
otel, _exporter = otel_with_exporter
monkeypatch.setattr(proxy_server, "open_telemetry_logger", otel, raising=False)
monkeypatch.setattr(mgmt_utils, "send_management_endpoint_alert", _noop_alert)
def _boom(*args, **kwargs):
raise RuntimeError("cache backend down")
monkeypatch.setattr(mgmt_utils, "_delete_api_key_from_cache", _boom)
server_span = server_span_factory(KEY_GENERATE_PATH)
@mgmt_utils.management_endpoint_wrapper
async def fake_fn(data=None, user_api_key_dict=None):
return {"ok": True}
result = asyncio.run(
fake_fn(data={}, user_api_key_dict=_real_user_api_key_dict(server_span))
)
assert result == {"ok": True}

View file

@ -76,7 +76,7 @@ class TestAgentCoreAcceptHeader:
with patch.object(client, "post", return_value=MagicMock()) as mock_post:
try:
litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_runtime",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_runtime",
messages=[{"role": "user", "content": "test"}],
api_key="test-jwt-token",
client=client,
@ -281,7 +281,7 @@ class TestAgentCoreStreamingJsonFallback:
with patch.object(client, "post", return_value=mock_response):
response = litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_agent",
messages=[{"role": "user", "content": "test"}],
stream=True,
client=client,
@ -318,7 +318,7 @@ class TestAgentCoreStreamingJsonFallback:
client, "post", new_callable=AsyncMock, return_value=mock_response
):
response = await litellm.acompletion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_agent",
messages=[{"role": "user", "content": "test"}],
stream=True,
client=client,
@ -353,7 +353,7 @@ class TestAgentCoreStreamingJsonFallback:
Exception, match="Failed to read/parse JSON response body"
):
litellm.completion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_agent",
messages=[{"role": "user", "content": "test"}],
stream=True,
client=client,
@ -383,7 +383,7 @@ class TestAgentCoreStreamingJsonFallback:
Exception, match="Failed to read/parse JSON response body"
):
await litellm.acompletion(
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/test_agent",
model="bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:941277531214:runtime/test_agent",
messages=[{"role": "user", "content": "test"}],
stream=True,
client=client,

View file

@ -20,6 +20,9 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
pass_through_request,
)
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
)
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
@ -2153,7 +2156,12 @@ async def test_create_pass_through_route_custom_body_url_target():
endpoint_func = create_pass_through_route(
endpoint=unique_path,
target="https://bedrock-agent-runtime.us-east-1.amazonaws.com",
custom_headers={"Content-Type": "application/json"},
custom_headers=Headers(
{
"Authorization": "AWS4-HMAC-SHA256 signed",
"Content-Type": "application/json",
}
),
_forward_headers=True,
)
@ -2213,6 +2221,147 @@ async def test_create_pass_through_route_custom_body_url_target():
# The critical assertion: custom_body takes precedence over
# the body parsed from the raw request
assert call_kwargs["custom_body"] == bedrock_body
# HeadersDict-like custom_headers (e.g. botocore SigV4) must be coerced
# to a plain dict so signed headers actually reach the upstream.
assert call_kwargs["custom_headers"] == {
"authorization": "AWS4-HMAC-SHA256 signed",
"content-type": "application/json",
}
@pytest.mark.asyncio
async def test_pass_through_request_non_streaming_uses_content_for_state_raw_body():
"""
Bedrock SigV4 path: exact signed bytes live on request.state; upstream must receive
content=... even if pre_call_hook mutates the parsed dict (would change json=).
"""
# Bytes that were signed (simulated); parsed body + hook will diverge on purpose.
raw_signed = b'{"retrievalQuery":{"text":"signed"},"sig":"intact"}'
parsed_from_wire = {"retrievalQuery": {"text": "signed"}, "sig": "intact"}
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = QueryParams({})
mock_request.headers = Headers({"Content-Type": "application/json"})
mock_request.state = SimpleNamespace()
setattr(mock_request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, raw_signed)
mock_request.body = AsyncMock(
return_value=json.dumps(parsed_from_wire).encode("utf-8")
)
mock_user = MagicMock()
mock_user.api_key = "sk-test"
upstream = httpx.Response(
status_code=200,
headers={"content-type": "application/json"},
content=b'{"ok": true}',
request=httpx.Request(
"POST",
"https://bedrock-agent-runtime.us-east-1.amazonaws.com/knowledgebases/KB/retrieve",
),
)
mock_async_client = AsyncMock()
mock_async_client.request = AsyncMock(return_value=upstream)
mock_client_obj = MagicMock()
mock_client_obj.client = mock_async_client
async def _hook_mutates_body(**kwargs):
data = kwargs["data"]
if isinstance(data, dict):
data["hook_mutated"] = True
return data
with (
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client",
return_value=mock_client_obj,
),
patch(
"litellm.proxy.proxy_server.proxy_logging_obj.pre_call_hook",
new=AsyncMock(side_effect=_hook_mutates_body),
),
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler",
new=AsyncMock(),
),
):
await pass_through_request(
request=mock_request,
target="https://bedrock-agent-runtime.us-east-1.amazonaws.com/knowledgebases/KB/retrieve",
custom_headers={"content-type": "application/json"},
user_api_key_dict=mock_user,
stream=False,
)
mock_async_client.request.assert_called_once()
req_kw = mock_async_client.request.call_args[1]
assert req_kw.get("content") == raw_signed
assert "json" not in req_kw
@pytest.mark.asyncio
async def test_pass_through_request_streaming_uses_content_for_state_raw_body():
"""Streaming pass-through with state raw body must use build_request(..., content=...)."""
raw_signed = b'{"model":"m","stream":true}'
parsed_from_wire = {"model": "m", "stream": True}
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.query_params = QueryParams({})
mock_request.headers = Headers({"Content-Type": "application/json"})
mock_request.state = SimpleNamespace()
setattr(mock_request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY, raw_signed)
mock_request.body = AsyncMock(
return_value=json.dumps(parsed_from_wire).encode("utf-8")
)
mock_user = MagicMock()
mock_user.api_key = "sk-test"
mock_built = MagicMock()
mock_async_client = AsyncMock()
mock_async_client.build_request = MagicMock(return_value=mock_built)
stream_resp = httpx.Response(
status_code=200,
headers={"content-type": "text/event-stream"},
content=b"data: {}\n\n",
request=httpx.Request("POST", "https://example.com/v1/messages"),
)
mock_async_client.send = AsyncMock(return_value=stream_resp)
mock_client_obj = MagicMock()
mock_client_obj.client = mock_async_client
with (
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client",
return_value=mock_client_obj,
),
patch(
"litellm.proxy.proxy_server.proxy_logging_obj.pre_call_hook",
new=AsyncMock(side_effect=lambda **kw: kw["data"]),
),
patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler",
new=AsyncMock(),
),
):
response = await pass_through_request(
request=mock_request,
target="https://example.com/v1/messages",
custom_headers={"Authorization": "Bearer x"},
user_api_key_dict=mock_user,
stream=None,
)
from fastapi.responses import StreamingResponse
assert isinstance(response, StreamingResponse)
mock_async_client.build_request.assert_called_once()
br_kw = mock_async_client.build_request.call_args[1]
assert br_kw.get("content") == raw_signed
assert "json" not in br_kw
@pytest.mark.asyncio

View file

@ -538,6 +538,36 @@ def test_forward_headers_from_request_protected_headers_not_overwritten():
assert "Anthropic-Beta" not in result
def test_forward_headers_custom_wins_case_insensitive_over_request_authorization():
"""
When forwarding request headers, provider-signed/custom headers must win
even if the incoming request uses a different case for the same header name.
"""
from litellm.passthrough.utils import BasePassthroughUtils
request_headers = {
"authorization": "Bearer sk-litellm-key",
"content-type": "application/json",
"x-request-id": "req-123",
}
signed_headers = {
"Authorization": "AWS4-HMAC-SHA256 signed",
"Content-Type": "application/json",
}
result = BasePassthroughUtils.forward_headers_from_request(
request_headers=request_headers,
headers=signed_headers.copy(),
forward_headers=True,
)
assert result["Authorization"] == "AWS4-HMAC-SHA256 signed"
assert "authorization" not in result
assert result["Content-Type"] == "application/json"
assert "content-type" not in result
assert result["x-request-id"] == "req-123"
@pytest.mark.asyncio
async def test_vertex_passthrough_custom_model_name_replaced_in_url():
"""

View file

@ -360,11 +360,15 @@ class TestProxySettingEndpoints:
assert "proxy_base_url" in values
assert "user_email" in values
# Verify values match our mock config
# Verify non-secret values match our mock config. OAuth client
# secrets are masked on read so the GET response never carries
# plaintext credentials.
assert values["google_client_id"] == "test_google_client_id"
assert values["google_client_secret"] == "test_google_client_secret"
assert values["google_client_secret"] != "test_google_client_secret"
assert "*" in values["google_client_secret"]
assert values["microsoft_client_id"] == "test_microsoft_client_id"
assert values["microsoft_client_secret"] == "test_microsoft_client_secret"
assert values["microsoft_client_secret"] != "test_microsoft_client_secret"
assert "*" in values["microsoft_client_secret"]
assert values["proxy_base_url"] == "https://example.com"
assert values["user_email"] == "admin@example.com"
@ -1321,10 +1325,12 @@ class TestProxySettingEndpoints:
assert "values" in data
assert "field_schema" in data
# Verify decrypted values are returned
# Verify decrypted values are returned. OAuth client secrets are
# masked on read so plaintext is never sent to the UI.
values = data["values"]
assert values["google_client_id"] == "decrypted_google_id"
assert values["google_client_secret"] == "decrypted_google_secret"
assert values["google_client_secret"] != "decrypted_google_secret"
assert "*" in values["google_client_secret"]
assert values["microsoft_client_id"] == "decrypted_microsoft_id"
assert values["proxy_base_url"] == "https://decrypted.example.com"

View file

@ -446,7 +446,7 @@ async def test_chat_completion_anthropic_structured_output():
client = AsyncOpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
res = await client.beta.chat.completions.parse(
model="bedrock/us.anthropic.claude-3-sonnet-20240229-v1:0",
model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0",
messages=messages,
response_format=EventsList,
timeout=60,

View file

@ -22,7 +22,7 @@ class TestBedrockVectorStore(BaseVectorStoreTest):
def get_base_request_args(self):
return {
"vector_store_id": "T37J8R4WTM",
"vector_store_id": "LCYXFBR2TU",
"custom_llm_provider": "bedrock",
"query": "what happens after we add a model",
}
@ -106,7 +106,7 @@ async def test_bedrock_search_with_router():
_router = Router(model_list=[])
search_response = await _router.avector_store_search(
query="what happens after we add a model",
vector_store_id="T37J8R4WTM",
vector_store_id="LCYXFBR2TU",
custom_llm_provider="bedrock",
)
print(search_response)
@ -150,7 +150,7 @@ async def test_bedrock_search_with_credentials_managed_registry():
# Create vector store with credential reference
vector_store = LiteLLM_ManagedVectorStore(
vector_store_id="T37J8R4WTM",
vector_store_id="LCYXFBR2TU",
custom_llm_provider="bedrock",
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
@ -162,7 +162,7 @@ async def test_bedrock_search_with_credentials_managed_registry():
litellm.vector_store_registry = registry
# Verify credentials can be retrieved from registry
retrieved_credentials = registry.get_credentials_for_vector_store("T37J8R4WTM")
retrieved_credentials = registry.get_credentials_for_vector_store("LCYXFBR2TU")
assert retrieved_credentials, "Should retrieve credentials from registry"
assert retrieved_credentials.get("aws_access_key_id") == "test_access_key"
assert retrieved_credentials.get("aws_secret_access_key") == "test_secret_key"
@ -194,7 +194,7 @@ async def test_bedrock_search_with_credentials_managed_registry():
search_response = await _router.avector_store_search(
query="what happens after we add a model",
vector_store_id="T37J8R4WTM",
vector_store_id="LCYXFBR2TU",
custom_llm_provider="bedrock",
)
@ -203,7 +203,7 @@ async def test_bedrock_search_with_credentials_managed_registry():
call_kwargs = mock_handler.call_args[1]
# Verify that the credential accessor was called with the correct vector store ID
mock_get_creds.assert_called_with("T37J8R4WTM")
mock_get_creds.assert_called_with("LCYXFBR2TU")
# Verify the credentials were injected into the search call
litellm_params = call_kwargs.get("litellm_params", {})
@ -224,7 +224,7 @@ async def test_bedrock_search_with_credentials_managed_registry():
assert search_response["data"][0]["id"] == "test_result"
print(
f"✅ Test passed: Credential accessor was called with vector store ID: T37J8R4WTM"
f"✅ Test passed: Credential accessor was called with vector store ID: LCYXFBR2TU"
)
print(f"✅ Retrieved credentials: {retrieved_credentials}")
print(f"✅ Credentials were injected into search call")

View file

@ -0,0 +1,36 @@
// hooks/useHideAgentPlatformBanner.ts
import { useSyncExternalStore } from "react";
import { getLocalStorageItem, LOCAL_STORAGE_EVENT } from "@/utils/localStorageUtils";
export const HIDE_AGENT_PLATFORM_BANNER_KEY = "litellmHideAgentPlatformBanner";
function subscribe(callback: () => void) {
const onStorage = (e: StorageEvent) => {
if (e.key === HIDE_AGENT_PLATFORM_BANNER_KEY) {
callback();
}
};
const onCustom = (e: Event) => {
const { key } = (e as CustomEvent).detail;
if (key === HIDE_AGENT_PLATFORM_BANNER_KEY) {
callback();
}
};
window.addEventListener("storage", onStorage);
window.addEventListener(LOCAL_STORAGE_EVENT, onCustom);
return () => {
window.removeEventListener("storage", onStorage);
window.removeEventListener(LOCAL_STORAGE_EVENT, onCustom);
};
}
function getSnapshot() {
return getLocalStorageItem(HIDE_AGENT_PLATFORM_BANNER_KEY) === "true";
}
export function useHideAgentPlatformBanner() {
return useSyncExternalStore(subscribe, getSnapshot);
}

View file

@ -37,7 +37,7 @@ const MIGRATED_PAGES: Record<string, string> = {
function LayoutContent({ children }: { children: React.ReactNode }) {
const router = useRouter();
const searchParams = useSearchParams();
const { accessToken, userRole, userId, userEmail, premiumUser } = useAuthorized();
const { accessToken } = useAuthorized();
const [sidebarCollapsed, setSidebarCollapsed] = React.useState(false);
const [page, setPage] = useState(() => {
return searchParams.get("page") || "api-keys";
@ -70,15 +70,9 @@ function LayoutContent({ children }: { children: React.ReactNode }) {
isPublicPage={false}
sidebarCollapsed={sidebarCollapsed}
onToggleSidebar={toggleSidebar}
userID={userId}
userEmail={userEmail}
userRole={userRole}
premiumUser={premiumUser}
proxySettings={undefined}
setProxySettings={() => { }}
accessToken={accessToken}
isDarkMode={false}
toggleDarkMode={() => { }}
/>
<DebugWarningBanner accessToken={accessToken} />
<div className="flex flex-1 overflow-auto">

View file

@ -473,18 +473,12 @@ function CreateKeyPageContent() {
) : (
<div className="flex flex-col min-h-screen">
<Navbar
userID={userID}
userRole={userRole}
premiumUser={premiumUser}
userEmail={userEmail}
setProxySettings={setProxySettings}
proxySettings={proxySettings}
accessToken={accessToken}
isPublicPage={false}
sidebarCollapsed={sidebarCollapsed}
onToggleSidebar={toggleSidebar}
isDarkMode={isDarkMode}
toggleDarkMode={toggleDarkMode}
/>
<div className="flex flex-1">
<div className="mt-2">

View file

@ -1,6 +1,7 @@
import { useDisableBlogPosts } from "@/app/(dashboard)/hooks/useDisableBlogPosts";
import { useBlogPosts, type BlogPost } from "@/app/(dashboard)/hooks/blogPosts/useBlogPosts";
import { LoadingOutlined } from "@ant-design/icons";
import { NAV_PRODUCT_LINK_CLASS } from "@/components/Navbar/navProductLinkClass";
import { DownOutlined, LoadingOutlined } from "@ant-design/icons";
import { Button, Dropdown, Space, Typography } from "antd";
import type { MenuProps } from "antd";
import React from "react";
@ -74,9 +75,13 @@ export const BlogDropdown: React.FC = () => {
];
}
// Blog opens a post list; Docs is a single outbound link — navbar adds a layout-only chevron there for alignment.
return (
<Dropdown menu={{ items }} trigger={["hover"]} placement="bottomRight">
<Button type="text">Blog</Button>
<Button type="text" className={`${NAV_PRODUCT_LINK_CLASS} !border-0 !bg-transparent`}>
Blog
<DownOutlined className="text-[10px] text-gray-500" aria-hidden />
</Button>
</Dropdown>
);
};

View file

@ -29,14 +29,14 @@ describe("CommunityEngagementButtons", () => {
expect(joinSlackLink).toHaveAttribute("rel", "noopener noreferrer");
});
it("should render Star us on GitHub button with correct link", () => {
it("should render GitHub link with correct href", () => {
renderWithProviders(<CommunityEngagementButtons />);
const starOnGithubLink = screen.getByRole("link", { name: /star us on github/i });
expect(starOnGithubLink).toBeInTheDocument();
expect(starOnGithubLink).toHaveAttribute("href", "https://github.com/BerriAI/litellm");
expect(starOnGithubLink).toHaveAttribute("target", "_blank");
expect(starOnGithubLink).toHaveAttribute("rel", "noopener noreferrer");
const githubLink = screen.getByRole("link", { name: /litellm on github/i });
expect(githubLink).toBeInTheDocument();
expect(githubLink).toHaveAttribute("href", "https://github.com/BerriAI/litellm");
expect(githubLink).toHaveAttribute("target", "_blank");
expect(githubLink).toHaveAttribute("rel", "noopener noreferrer");
});
it("should not render buttons when prompts are disabled", () => {
@ -45,6 +45,6 @@ describe("CommunityEngagementButtons", () => {
renderWithProviders(<CommunityEngagementButtons />);
expect(screen.queryByRole("link", { name: /join slack/i })).not.toBeInTheDocument();
expect(screen.queryByRole("link", { name: /star us on github/i })).not.toBeInTheDocument();
expect(screen.queryByRole("link", { name: /litellm on github/i })).not.toBeInTheDocument();
});
});

View file

@ -1,36 +1,45 @@
import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts";
import { GithubOutlined, SlackOutlined } from "@ant-design/icons";
import { Button } from "antd";
import { Tooltip } from "antd";
import React from "react";
const iconBtnClass =
"inline-flex h-9 w-9 shrink-0 items-center justify-center rounded-md border-0 bg-transparent text-gray-500 transition-colors hover:bg-gray-100 hover:text-gray-700 cursor-pointer";
export const CommunityEngagementButtons: React.FC = () => {
const disableShowPrompts = useDisableShowPrompts();
// Hide buttons if prompts are disabled
if (disableShowPrompts) {
return null;
}
return (
<>
<Button
href="https://www.litellm.ai/support"
target="_blank"
rel="noopener noreferrer"
icon={<SlackOutlined />}
className="shadow-md shadow-indigo-500/20 hover:shadow-indigo-500/50 transition-shadow"
>
Join Slack
</Button>
<Button
href="https://github.com/BerriAI/litellm"
target="_blank"
rel="noopener noreferrer"
className="shadow-md shadow-indigo-500/20 hover:shadow-indigo-500/50 transition-shadow"
icon={<GithubOutlined />}
>
Star us on GitHub
</Button>
</>
<div
className="flex items-center gap-0.5 rounded-md border border-gray-200/80 bg-gray-50 px-0.5 py-0"
aria-label="Community links"
>
<Tooltip title="LiteLLM Slack community">
<a
href="https://www.litellm.ai/support"
target="_blank"
rel="noopener noreferrer"
className={iconBtnClass}
aria-label="Join Slack"
>
<SlackOutlined className="text-lg" />
</a>
</Tooltip>
<Tooltip title="LiteLLM on GitHub">
<a
href="https://github.com/BerriAI/litellm"
target="_blank"
rel="noopener noreferrer"
className={iconBtnClass}
aria-label="LiteLLM on GitHub"
>
<GithubOutlined className="text-lg" />
</a>
</Tooltip>
</div>
);
};

View file

@ -0,0 +1,69 @@
import { renderWithProviders, screen } from "../../../../tests/test-utils";
import { NotificationsBell, AGENT_PLATFORM_URL } from "./NotificationsBell";
import React from "react";
import userEvent from "@testing-library/user-event";
describe("NotificationsBell", () => {
beforeEach(() => {
localStorage.clear();
});
it("should open notifications with Agent Platform details and GitHub link", async () => {
const user = userEvent.setup();
renderWithProviders(<NotificationsBell />);
await user.click(screen.getByRole("button", { name: /^notifications$/i }));
expect(screen.getByText(/LiteLLM Agent Platform/i)).toBeInTheDocument();
const githubBtn = screen.getByRole("link", { name: /^GitHub$/i });
expect(githubBtn).toHaveAttribute("href", AGENT_PLATFORM_URL);
expect(githubBtn).toHaveAttribute("target", "_blank");
expect(githubBtn).toHaveAttribute("rel", "noopener noreferrer");
});
it("should offer mark as read when announcement is unread", async () => {
const user = userEvent.setup();
renderWithProviders(<NotificationsBell />);
await user.click(screen.getByRole("button", { name: /^notifications$/i }));
expect(screen.getByRole("button", { name: /^mark as read$/i })).toBeInTheDocument();
});
it("should hide mark as read and persist after marking read", async () => {
const user = userEvent.setup();
renderWithProviders(<NotificationsBell />);
await user.click(screen.getByRole("button", { name: /^notifications$/i }));
await user.click(screen.getByRole("button", { name: /^mark as read$/i }));
expect(localStorage.getItem("litellmHideAgentPlatformBanner")).toBe("true");
await user.click(screen.getByRole("button", { name: /^notifications$/i }));
expect(screen.queryByRole("button", { name: /^mark as read$/i })).not.toBeInTheDocument();
});
it("should not show mark as read when previously dismissed", async () => {
localStorage.setItem("litellmHideAgentPlatformBanner", "true");
const user = userEvent.setup();
renderWithProviders(<NotificationsBell />);
await user.click(screen.getByRole("button", { name: /^notifications$/i }));
expect(screen.queryByRole("button", { name: /^mark as read$/i })).not.toBeInTheDocument();
});
it("should sync sibling instances when one is dismissed", async () => {
const user = userEvent.setup();
renderWithProviders(
<>
<div data-testid="bell-a">
<NotificationsBell />
</div>
<div data-testid="bell-b">
<NotificationsBell />
</div>
</>,
);
// Both bells start unread → both render the "Mark as read" affordance once opened.
const [bellA, bellB] = screen.getAllByRole("button", { name: /^notifications$/i });
await user.click(bellA);
await user.click(screen.getByRole("button", { name: /^mark as read$/i }));
// Dismissing in bell A must also clear bell B without a remount.
await user.click(bellB);
expect(screen.queryByRole("button", { name: /^mark as read$/i })).not.toBeInTheDocument();
});
});

View file

@ -0,0 +1,59 @@
"use client";
import {
HIDE_AGENT_PLATFORM_BANNER_KEY,
useHideAgentPlatformBanner,
} from "@/app/(dashboard)/hooks/useHideAgentPlatformBanner";
import { emitLocalStorageChange, setLocalStorageItem } from "@/utils/localStorageUtils";
import { BellOutlined } from "@ant-design/icons";
import { Badge, Button, Popover, Typography } from "antd";
import React, { useState } from "react";
export const AGENT_PLATFORM_URL = "https://github.com/BerriAI/litellm-agent-platform";
export const NotificationsBell: React.FC = () => {
const hidden = useHideAgentPlatformBanner();
const hasUnread = !hidden;
const [open, setOpen] = useState(false);
const markDismissed = () => {
setLocalStorageItem(HIDE_AGENT_PLATFORM_BANNER_KEY, "true");
emitLocalStorageChange(HIDE_AGENT_PLATFORM_BANNER_KEY);
setOpen(false);
};
const content = (
<div className="max-w-[280px]">
<Typography.Title level={5} className="!mt-0 !mb-2">
LiteLLM Agent Platform
</Typography.Title>
<Typography.Paragraph type="secondary" className="!mb-3 text-sm leading-snug">
Open-source agent infra — sandboxes, durable sessions, and workers on AWS Fargate.
</Typography.Paragraph>
<div className="flex flex-wrap items-center gap-2">
<Button type="primary" size="small" href={AGENT_PLATFORM_URL} target="_blank" rel="noopener noreferrer">
GitHub
</Button>
{hasUnread ? (
<Button type="link" size="small" className="!px-1" onClick={markDismissed}>
Mark as read
</Button>
) : null}
</div>
</div>
);
return (
<Popover content={content} trigger="click" open={open} onOpenChange={setOpen} placement="bottomRight">
<Button
type="text"
className="!flex !h-9 !w-9 items-center justify-center !rounded-md text-gray-600 transition-colors hover:!bg-gray-100 hover:!text-gray-900"
aria-label="Notifications"
>
<Badge dot={hasUnread} color="#1677ff" size="small" offset={[8, 2]}>
<BellOutlined className="text-base" aria-hidden />
</Badge>
</Button>
</Popover>
);
};

View file

@ -37,6 +37,8 @@ vi.mock("@/utils/localStorageUtils", () => ({
describe("UserDropdown", () => {
const mockOnLogout = vi.fn();
const getAccountTrigger = () => screen.getByRole("button", { name: /account menu/i });
beforeEach(() => {
vi.clearAllMocks();
mockUseAuthorizedImpl = () => ({
@ -55,22 +57,23 @@ describe("UserDropdown", () => {
it("should render", () => {
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
expect(screen.getByRole("button")).toBeInTheDocument();
expect(getAccountTrigger()).toBeInTheDocument();
});
it("should display user button with User text", () => {
it("should surface initials and account menu affordance", () => {
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
expect(screen.getByText("User")).toBeInTheDocument();
expect(getAccountTrigger()).toBeInTheDocument();
expect(screen.getByText("TE")).toBeInTheDocument();
});
it("should show user email when dropdown is opened", async () => {
const user = userEvent.setup();
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("test@example.com")).toBeInTheDocument();
expect(screen.getAllByText("test@example.com").length).toBeGreaterThan(0);
});
});
@ -78,7 +81,7 @@ describe("UserDropdown", () => {
const user = userEvent.setup();
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("test-user-id")).toBeInTheDocument();
@ -89,10 +92,10 @@ describe("UserDropdown", () => {
const user = userEvent.setup();
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("Admin")).toBeInTheDocument();
expect(screen.getAllByText("Admin").length).toBeGreaterThan(0);
});
});
@ -100,7 +103,7 @@ describe("UserDropdown", () => {
const user = userEvent.setup();
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("Standard")).toBeInTheDocument();
@ -118,7 +121,7 @@ describe("UserDropdown", () => {
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("Premium")).toBeInTheDocument();
@ -129,10 +132,10 @@ describe("UserDropdown", () => {
const user = userEvent.setup();
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("test@example.com")).toBeInTheDocument();
expect(screen.getAllByText("test@example.com").length).toBeGreaterThan(0);
});
await user.click(screen.getByText("Logout"));
@ -144,10 +147,10 @@ describe("UserDropdown", () => {
const user = userEvent.setup();
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("test@example.com")).toBeInTheDocument();
expect(screen.getAllByText("test@example.com").length).toBeGreaterThan(0);
});
const toggle = screen.getByLabelText("Toggle hide new feature indicators");
@ -169,10 +172,10 @@ describe("UserDropdown", () => {
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("test@example.com")).toBeInTheDocument();
expect(screen.getAllByText("test@example.com").length).toBeGreaterThan(0);
});
const toggle = screen.getByLabelText("Toggle hide new feature indicators");
@ -189,10 +192,10 @@ describe("UserDropdown", () => {
const user = userEvent.setup();
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("test@example.com")).toBeInTheDocument();
expect(screen.getAllByText("test@example.com").length).toBeGreaterThan(0);
});
const toggle = screen.getByLabelText("Toggle hide all prompts");
@ -215,10 +218,10 @@ describe("UserDropdown", () => {
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("test@example.com")).toBeInTheDocument();
expect(screen.getAllByText("test@example.com").length).toBeGreaterThan(0);
});
const toggle = screen.getByLabelText("Toggle hide all prompts");
@ -231,6 +234,17 @@ describe("UserDropdown", () => {
expect(localStorageUtils.emitLocalStorageChange).toHaveBeenCalledWith("disableShowPrompts");
});
it("should show Account in the trigger when user id is the default placeholder", () => {
mockUseAuthorizedImpl = () => ({
userId: "default_user_id",
userEmail: null as any,
userRole: "Admin",
premiumUser: false,
});
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
expect(screen.getByText("Account")).toBeInTheDocument();
});
it("should display dash when user email is not available", async () => {
const user = userEvent.setup();
mockUseAuthorizedImpl = () => ({
@ -242,7 +256,7 @@ describe("UserDropdown", () => {
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("-")).toBeInTheDocument();
@ -260,7 +274,7 @@ describe("UserDropdown", () => {
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
const dashElements = screen.getAllByText("-");
@ -277,10 +291,10 @@ describe("UserDropdown", () => {
renderWithProviders(<UserDropdown onLogout={mockOnLogout} />);
await user.click(screen.getByText("User"));
await user.click(getAccountTrigger());
await waitFor(() => {
expect(screen.getByText("test@example.com")).toBeInTheDocument();
expect(screen.getAllByText("test@example.com").length).toBeGreaterThan(0);
});
const toggle = screen.getByLabelText("Toggle hide new feature indicators");

View file

@ -9,6 +9,7 @@ import {
removeLocalStorageItem,
setLocalStorageItem,
} from "@/utils/localStorageUtils";
import { navAccountDisplayName } from "@/components/Navbar/navDisplayName";
import {
CrownOutlined,
DownOutlined,
@ -23,6 +24,39 @@ import React, { useEffect, useState } from "react";
const { Text } = Typography;
function hueFromString(seed: string): number {
let h = 0;
for (let i = 0; i < seed.length; i += 1) {
h = seed.charCodeAt(i) + ((h << 5) - h);
}
return Math.abs(h) % 360;
}
function initialsFromIdentity(email: string | null, userId: string | null): string {
const local = email?.split("@")[0]?.trim();
if (local) {
const parts = local
.replace(/[^a-zA-Z0-9]+/g, " ")
.trim()
.split(/\s+/)
.filter(Boolean);
if (parts.length >= 2) {
return `${parts[0]!.charAt(0)}${parts[1]!.charAt(0)}`.toUpperCase();
}
if (parts.length === 1) {
const p = parts[0]!;
return p.length >= 2 ? p.slice(0, 2).toUpperCase() : `${p.charAt(0)}`.toUpperCase();
}
}
if (userId && userId.length >= 2) {
return userId.slice(0, 2).toUpperCase();
}
if (userId && userId.length === 1) {
return `${userId.toUpperCase()}•`;
}
return "?";
}
interface UserDropdownProps {
onLogout: () => void;
}
@ -61,19 +95,12 @@ const UserDropdown: React.FC<UserDropdownProps> = ({ onLogout }) => {
<Text type="secondary">{userEmail || "-"}</Text>
</Space>
{premiumUser ? (
<Tag
icon={<CrownOutlined />}
color="gold"
>
<Tag icon={<CrownOutlined />} color="gold">
Premium
</Tag>
) : (
<Tooltip title="Upgrade to Premium for advanced features" placement="left">
<Tag
icon={<CrownOutlined />}
>
Standard
</Tag>
<Tag icon={<CrownOutlined />}>Standard</Tag>
</Tooltip>
)}
</Space>
@ -83,12 +110,7 @@ const UserDropdown: React.FC<UserDropdownProps> = ({ onLogout }) => {
<UserOutlined />
<Text type="secondary">User ID</Text>
</Space>
<Text
copyable
ellipsis
style={{ maxWidth: "150px" }}
title={userId || "-"}
>
<Text copyable ellipsis style={{ maxWidth: "150px" }} title={userId || "-"}>
{userId || "-"}
</Text>
</Space>
@ -189,13 +211,17 @@ const UserDropdown: React.FC<UserDropdownProps> = ({ onLogout }) => {
</Space>
);
const seed = userEmail || userId || "user";
const initials = initialsFromIdentity(userEmail, userId);
const hue = hueFromString(seed);
const displayName = navAccountDisplayName(userEmail, userId);
return (
<Dropdown
trigger={["click"]}
menu={{ items: userItems }}
popupRender={(menu) => (
<div
className="bg-white rounded-lg shadow-lg"
>
<div className="rounded-lg bg-white shadow-lg">
{renderUserInfoSection()}
<Divider style={{ margin: 0 }} />
{React.cloneElement(menu as React.ReactElement, {
@ -204,12 +230,23 @@ const UserDropdown: React.FC<UserDropdownProps> = ({ onLogout }) => {
</div>
)}
>
<Button type="text" >
<Space>
<UserOutlined />
<Text>User</Text>
<DownOutlined />
</Space>
<Button
type="text"
className="!flex max-w-[min(200px,34vw)] items-center gap-2 !rounded-md !py-0.5 !pl-1 !pr-2 transition-colors hover:!bg-gray-100"
aria-label={`Account menu — ${userRole ?? "Unknown role"} — signed in as ${userEmail || userId || "unknown"}`}
aria-haspopup="menu"
>
<span
className="flex h-8 w-8 shrink-0 items-center justify-center rounded-full text-xs font-semibold text-white shadow-inner ring-1 ring-black/5"
style={{ backgroundColor: `hsl(${hue} 46% 38%)` }}
aria-hidden
>
{initials}
</span>
<span className="hidden min-w-0 truncate text-left text-sm font-medium leading-none text-gray-900 md:inline">
{displayName}
</span>
<DownOutlined className="hidden shrink-0 text-[10px] text-gray-400 md:inline" aria-hidden />
</Button>
</Dropdown>
);

View file

@ -0,0 +1,17 @@
import { describe, expect, it } from "vitest";
import { navAccountDisplayName } from "./navDisplayName";
describe("navAccountDisplayName", () => {
it("should prefer email when present", () => {
expect(navAccountDisplayName("x@y.com", "ignored")).toBe("x@y.com");
});
it("should map default_user_id placeholder to Account", () => {
expect(navAccountDisplayName(null, "default_user_id")).toBe("Account");
expect(navAccountDisplayName(null, "DEFAULT_USER_ID")).toBe("Account");
});
it("should show a sensible token when user id is non-placeholder", () => {
expect(navAccountDisplayName(null, "user-uuid-123")).toBe("user-uuid-123");
});
});

View file

@ -0,0 +1,15 @@
/** Primary label for the navbar account control — avoids raw placeholder JWT/user IDs in the UI. */
export function navAccountDisplayName(userEmail: string | null, userId: string | null): string {
const email = userEmail?.trim();
if (email) {
return email;
}
const id = userId?.trim();
if (!id) {
return "Account";
}
if (/^default[_\s-]?user[_\s-]?id$/i.test(id)) {
return "Account";
}
return id;
}

View file

@ -0,0 +1,3 @@
/** Shared styling for Docs / Blog in the top nav (product navigation zone). */
export const NAV_PRODUCT_LINK_CLASS =
"inline-flex h-9 shrink-0 items-center justify-center gap-1 rounded-md px-2 text-sm font-medium leading-none text-gray-800 transition-colors hover:bg-gray-100 hover:text-gray-950";

View file

@ -30,6 +30,7 @@ const mockUserDropdownData = vi.hoisted(() => ({
vi.mock("./Navbar/UserDropdown/UserDropdown", async (importOriginal) => {
const React = await import("react");
const { useState } = React;
const { Button } = await import("antd");
const localStorageUtils = await import("@/utils/localStorageUtils");
return {
default: function MockUserDropdown({ onLogout }: { onLogout: () => void }) {
@ -37,9 +38,9 @@ vi.mock("./Navbar/UserDropdown/UserDropdown", async (importOriginal) => {
const [open, setOpen] = useState(false);
return (
<div>
<button type="button" onClick={() => setOpen(!open)}>
User
</button>
<Button type="text" aria-label="Open account menu" onClick={() => setOpen(!open)}>
Account
</Button>
{open && (
<div data-testid="user-dropdown-content">
<span>{userId}</span>
@ -136,30 +137,25 @@ Object.defineProperty(window, "location", {
describe("Navbar", () => {
const defaultProps = {
userID: "test-user",
userEmail: "test@example.com",
userRole: "Admin",
premiumUser: false,
proxySettings: {},
setProxySettings: vi.fn(),
accessToken: "test-token",
isPublicPage: false,
isDarkMode: false,
toggleDarkMode: vi.fn(),
};
it("should render without crashing", () => {
renderWithProviders(<Navbar {...defaultProps} />);
expect(screen.getByRole("button", { name: /^notifications$/i })).toBeInTheDocument();
expect(screen.getByText("Docs")).toBeInTheDocument();
expect(screen.getByText("User")).toBeInTheDocument();
expect(screen.getByRole("button", { name: /open account menu/i })).toBeInTheDocument();
});
it("should display user information in dropdown", async () => {
const user = userEvent.setup();
renderWithProviders(<Navbar {...defaultProps} />);
await user.click(screen.getByText("User"));
await user.click(screen.getByRole("button", { name: /open account menu/i }));
await waitFor(() => {
expect(screen.getByText("test-user")).toBeInTheDocument();
@ -198,7 +194,7 @@ describe("Navbar", () => {
});
renderWithProviders(<Navbar {...defaultProps} />);
await user.click(screen.getByText("User"));
await user.click(screen.getByRole("button", { name: /open account menu/i }));
await waitFor(() => {
expect(screen.getByText("Premium")).toBeInTheDocument();
@ -247,11 +243,12 @@ describe("Navbar", () => {
mockUseThemeImpl = () => ({ logoUrl: null });
});
it("should hide user dropdown on public pages", () => {
it("should hide user dropdown and notifications on public pages", () => {
const publicPageProps = { ...defaultProps, isPublicPage: true };
renderWithProviders(<Navbar {...publicPageProps} />);
expect(screen.queryByText("User")).not.toBeInTheDocument();
expect(screen.queryByRole("button", { name: /open account menu/i })).not.toBeInTheDocument();
expect(screen.queryByRole("button", { name: /^notifications$/i })).not.toBeInTheDocument();
});
it("should handle hide new features toggle", async () => {
@ -265,7 +262,7 @@ describe("Navbar", () => {
renderWithProviders(<Navbar {...defaultProps} />);
await user.click(screen.getByText("User"));
await user.click(screen.getByRole("button", { name: /open account menu/i }));
await waitFor(() => {
expect(screen.getByText("test-user")).toBeInTheDocument();
@ -290,7 +287,7 @@ describe("Navbar", () => {
renderWithProviders(<Navbar {...defaultProps} />);
await user.click(screen.getByText("User"));
await user.click(screen.getByRole("button", { name: /open account menu/i }));
await waitFor(() => {
expect(screen.getByText("test-user")).toBeInTheDocument();

View file

@ -1,47 +1,39 @@
import { useHealthReadinessDetails } from "@/app/(dashboard)/hooks/healthReadiness/useHealthReadinessDetails";
import { useDisableBouncingIcon } from "@/app/(dashboard)/hooks/useDisableBouncingIcon";
import { useDisableShowPrompts } from "@/app/(dashboard)/hooks/useDisableShowPrompts";
import { useWorker } from "@/hooks/useWorker";
import { getProxyBaseUrl } from "@/components/networking";
import { useTheme } from "@/contexts/ThemeContext";
import { clearTokenCookies } from "@/utils/cookieUtils";
import { clearStoredReturnUrl } from "@/utils/returnUrlUtils";
import { fetchProxySettings } from "@/utils/proxyUtils";
import { MenuFoldOutlined, MenuUnfoldOutlined, MoonOutlined, SunOutlined } from "@ant-design/icons";
import { Button, Switch, Tag } from "antd";
import { DownOutlined, MenuFoldOutlined, MenuUnfoldOutlined } from "@ant-design/icons";
import { Tag } from "antd";
import Link from "next/link";
import React, { useEffect, useState } from "react";
import { BlogDropdown } from "./Navbar/BlogDropdown/BlogDropdown";
import { CommunityEngagementButtons } from "./Navbar/CommunityEngagementButtons/CommunityEngagementButtons";
import { NAV_PRODUCT_LINK_CLASS } from "./Navbar/navProductLinkClass";
import { NotificationsBell } from "./Navbar/NotificationsBell/NotificationsBell";
import UserDropdown from "./Navbar/UserDropdown/UserDropdown";
import WorkerDropdown from "./Navbar/WorkerDropdown/WorkerDropdown";
interface NavbarProps {
userID: string | null;
userEmail: string | null;
userRole: string | null;
premiumUser: boolean;
proxySettings: any;
setProxySettings: React.Dispatch<React.SetStateAction<any>>;
accessToken: string | null;
isPublicPage: boolean;
sidebarCollapsed?: boolean;
onToggleSidebar?: () => void;
isDarkMode: boolean;
toggleDarkMode: () => void;
}
const Navbar: React.FC<NavbarProps> = ({
userID,
userEmail,
userRole,
premiumUser,
proxySettings,
setProxySettings,
accessToken,
isPublicPage = false,
sidebarCollapsed = false,
onToggleSidebar,
isDarkMode,
toggleDarkMode,
}) => {
const baseUrl = getProxyBaseUrl();
const [logoutUrl, setLogoutUrl] = useState("");
@ -49,8 +41,10 @@ const Navbar: React.FC<NavbarProps> = ({
const { data: healthData } = useHealthReadinessDetails(accessToken);
const version = healthData?.litellm_version;
const disableBouncingIcon = useDisableBouncingIcon();
const hideCommunityLinks = useDisableShowPrompts();
const { isControlPlane, selectedWorker } = useWorker();
const showWorkerSwitch = isControlPlane && selectedWorker !== null;
// Simple logo URL: use custom logo if available, otherwise default
const imageUrl = logoUrl || `${baseUrl}/get_image`;
useEffect(() => {
@ -87,14 +81,14 @@ const Navbar: React.FC<NavbarProps> = ({
};
return (
<nav className="bg-white border-b border-gray-200 sticky top-0 z-10">
<nav className="sticky top-0 z-10 border-b border-gray-200 bg-white">
<div className="w-full">
<div className="flex items-center h-14 px-4">
<div className="flex items-center flex-shrink-0">
<div className="flex h-14 items-center px-4">
<div className="flex flex-shrink-0 items-center">
{onToggleSidebar && (
<button
onClick={onToggleSidebar}
className="flex items-center justify-center w-10 h-10 mr-2 text-gray-600 hover:text-gray-900 hover:bg-gray-100 rounded transition-colors"
className="mr-2 flex h-9 w-9 items-center justify-center rounded-md text-gray-600 transition-colors hover:bg-gray-100 hover:text-gray-900"
title={sidebarCollapsed ? "Expand sidebar" : "Collapse sidebar"}
>
<span className="text-lg">{sidebarCollapsed ? <MenuUnfoldOutlined /> : <MenuFoldOutlined />}</span>
@ -104,11 +98,11 @@ const Navbar: React.FC<NavbarProps> = ({
<div className="flex items-center gap-2">
<Link href={baseUrl ? baseUrl : "/"} className="flex items-center">
<div className="relative">
<div className="h-10 max-w-48 flex items-center justify-center overflow-hidden">
<div className="flex h-10 max-w-48 items-center justify-center overflow-hidden">
<img
src={imageUrl}
alt="LiteLLM Brand"
className="max-w-full max-h-full w-auto h-auto object-contain"
className="h-auto max-h-full w-auto max-w-full object-contain"
/>
</div>
</div>
@ -117,14 +111,14 @@ const Navbar: React.FC<NavbarProps> = ({
<div className="relative">
{!disableBouncingIcon && (
<span
className="absolute -top-1 -left-2 text-lg animate-bounce"
className="absolute -left-2 -top-1 animate-bounce text-lg"
style={{ animationDuration: "2s" }}
title="Thanks for using LiteLLM!"
>
🌑
</span>
)}
<Tag className="relative text-xs font-medium cursor-pointer z-10">
<Tag className="relative z-10 cursor-pointer text-xs font-medium">
<a
href="https://docs.litellm.ai/release_notes"
target="_blank"
@ -138,28 +132,48 @@ const Navbar: React.FC<NavbarProps> = ({
)}
</div>
</div>
{/* Right side nav items */}
<div className="flex items-center space-x-5 ml-auto">
<WorkerDropdown onWorkerSwitch={handleWorkerSwitch} />
<CommunityEngagementButtons />
{/* Dark mode is currently a work in progress. To test, you can change 'false' to 'true' below.
Do not set this to true by default until all components are confirmed to support dark mode styles. */}
{false && (
<Switch
data-testid="dark-mode-toggle"
checked={isDarkMode}
onChange={toggleDarkMode}
checkedChildren={<MoonOutlined />}
unCheckedChildren={<SunOutlined />}
/>
)}
<Button type="text" href="https://docs.litellm.ai/docs/" target="_blank" rel="noopener noreferrer">
Docs
</Button>
<BlogDropdown />
{!isPublicPage && <UserDropdown onLogout={handleLogout} />}
<div className="ml-auto flex min-w-0 flex-1 items-center justify-end gap-4">
{showWorkerSwitch && (
<div className="flex shrink-0 items-center">
<WorkerDropdown onWorkerSwitch={handleWorkerSwitch} />
</div>
)}
<nav
aria-label="Product documentation"
className={`flex min-w-0 items-center gap-2 ${showWorkerSwitch ? "border-l border-gray-200 pl-4" : ""}`}
>
<a
href="https://docs.litellm.ai/docs/"
target="_blank"
rel="noopener noreferrer"
className={NAV_PRODUCT_LINK_CLASS}
>
Docs
{/* Layout parity with Blog chevron — intentional single-level link */}
<DownOutlined className="pointer-events-none text-[10px] opacity-0" aria-hidden />
</a>
<BlogDropdown />
</nav>
{!hideCommunityLinks && (
<div className="flex shrink-0 items-center border-l border-gray-200 pl-4">
<CommunityEngagementButtons />
</div>
)}
{!isPublicPage && (
<div className="flex shrink-0 items-center border-l border-gray-200 pl-4">
<div className="flex items-center gap-0.5 rounded-lg bg-gray-50 px-1 py-0 transition-colors hover:bg-gray-100">
<NotificationsBell />
<span className="mx-0.5 h-6 w-px shrink-0 bg-gray-200" aria-hidden />
<UserDropdown onLogout={handleLogout} />
</div>
</div>
)}
</div>
{/* Dark mode toggle: keep disabled until the dashboard supports dark styles end-to-end. */}
</div>
</div>
</nav>

View file

@ -983,16 +983,10 @@ const PublicModelHub: React.FC<PublicModelHubProps> = ({ accessToken, isEmbedded
{/* Navigation - only show when not embedded */}
{!isEmbedded && (
<Navbar
userID={null}
userEmail={null}
userRole={null}
premiumUser={false}
setProxySettings={setProxySettings}
proxySettings={proxySettings}
accessToken={accessToken || null}
isPublicPage={true}
isDarkMode={false}
toggleDarkMode={() => {}}
/>
)}