mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
commit
335afa6251
62 changed files with 1374 additions and 376 deletions
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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(":")
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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", "")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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}")
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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']}")
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
],
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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>
|
||||
);
|
||||
};
|
||||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
15
ui/litellm-dashboard/src/components/Navbar/navDisplayName.ts
Normal file
15
ui/litellm-dashboard/src/components/Navbar/navDisplayName.ts
Normal 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;
|
||||
}
|
||||
|
|
@ -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";
|
||||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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={() => {}}
|
||||
/>
|
||||
)}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue