Merge branch 'litellm_internal_staging' into litellm_lit_7022_azure_ai_passthrough_config

This commit is contained in:
mateo-berri 2026-09-09 19:42:28 -07:00
commit 966a58d1c8
34 changed files with 1150 additions and 175 deletions

View file

@ -13,6 +13,7 @@ from litellm._logging import verbose_logger
from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import (
BedrockAgentCoreA2ATransformation,
)
from litellm.llms.bedrock.base_aws_llm import run_aws_signing
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.types.llms.custom_http import httpxSpecialProvider
@ -45,7 +46,8 @@ class BedrockAgentCoreA2AHandler:
Returns:
A2A JSON-RPC response dict from the AgentCore agent
"""
url, headers, body = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
url, headers, body = await run_aws_signing(
BedrockAgentCoreA2ATransformation.get_url_and_signed_request,
request_id=request_id,
params=params,
litellm_params=litellm_params,
@ -91,7 +93,8 @@ class BedrockAgentCoreA2AHandler:
Yields:
A2A streaming response events from the AgentCore agent
"""
url, headers, body = BedrockAgentCoreA2ATransformation.get_url_and_signed_request(
url, headers, body = await run_aws_signing(
BedrockAgentCoreA2ATransformation.get_url_and_signed_request,
request_id=request_id,
params=params,
litellm_params=litellm_params,

View file

@ -582,6 +582,7 @@ LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS: Final = float(
LOGGING_EXECUTOR_MAX_THREADS: Final = get_env_int("LOGGING_EXECUTOR_MAX_THREADS", 100)
LOGGING_EXECUTOR_MAX_PENDING_TASKS: Final = get_env_int("LOGGING_EXECUTOR_MAX_PENDING_TASKS", 10_000)
LOGGING_EXECUTOR_DROPPED_TASK_LOG_INTERVAL_SECONDS: Final = 30.0
AWS_SIGNING_MAX_THREADS: Final = 16
DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE: Final = os.getenv(
"DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE", "streaming.chunk.yield"
)

View file

@ -24,7 +24,7 @@ from litellm.integrations.s3 import (
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing
from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
get_async_httpx_client,
@ -366,7 +366,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Sign the request
aws_request: Final = AWSRequest(method="PUT", url=url, data=json_string, headers=headers)
aws_region_name: Final = self.get_aws_region_name_for_non_llm_api_calls(aws_region_name=self.s3_region_name)
S3SigV4Auth(credentials, "s3", aws_region_name).add_auth(aws_request)
await run_aws_signing(S3SigV4Auth(credentials, "s3", aws_region_name).add_auth, aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())
@ -597,7 +597,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Sign the request
aws_request: Final = AWSRequest(method="GET", url=url, headers=headers)
S3SigV4Auth(credentials, "s3", self.s3_region_name).add_auth(aws_request)
await run_aws_signing(S3SigV4Auth(credentials, "s3", self.s3_region_name).add_auth, aws_request)
# Prepare the signed headers
signed_headers: Final = dict(aws_request.headers.items())

View file

@ -22,7 +22,7 @@ from litellm.constants import (
SQS_SEND_MESSAGE_ACTION,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -295,7 +295,7 @@ class SQSLogger(CustomBatchLogger, BaseAWSLLM):
data=prepped.body,
headers=prepped.headers,
)
SigV4Auth(credentials, "sqs", self.sqs_region_name).add_auth(aws_request)
await run_aws_signing(SigV4Auth(credentials, "sqs", self.sqs_region_name).add_auth, aws_request)
signed_headers: Final = dict(aws_request.headers.items())

View file

@ -1,13 +1,17 @@
import asyncio
import base64
import contextvars
import hashlib
import json
import os
import re
import urllib.parse
from collections.abc import Callable, Mapping
from concurrent.futures import ThreadPoolExecutor
from datetime import datetime
from functools import partial
from threading import Lock
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, cast, get_args, overload
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, ParamSpec, TypeVar, cast, get_args, overload
import httpx
from pydantic import BaseModel, ValidationError
@ -16,6 +20,7 @@ from litellm._logging import verbose_logger
from litellm.caching.caching import DualCache
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
AWS_SIGNING_MAX_THREADS,
BEDROCK_EMBEDDING_PROVIDERS_LITERAL,
BEDROCK_IAM_CACHE_FETCH_LOCK_STRIPES,
BEDROCK_IAM_CACHE_MAX_ENTRIES,
@ -80,7 +85,11 @@ class AwsAuthError(Exception):
super().__init__(self.message) # Call the base class constructor with the parameters it needs
class BaseAWSLLM:
class SignsRequestsWithAWS:
pass
class BaseAWSLLM(SignsRequestsWithAWS):
# Process-wide IAM credential cache (shared across instances — Bedrock passthrough is per-request).
# Storage is in-process memory only: no Redis backend unless attached elsewhere. Entry TTL: static
# access-key + secret + region use ``_get_default_ttl_for_boto3_credentials`` (~59 minutes); ambient
@ -1668,3 +1677,52 @@ class BaseAWSLLM:
request_headers_dict["Authorization"] = incoming_authorization
return request_headers_dict, request.body
def sign_aws_json_post(
get_credentials: Callable[[], Credentials],
service_name: str,
aws_region_name: str | None,
url: str,
body: str,
headers: Mapping[str, str],
) -> AWSPreparedRequest:
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError(f"Missing boto3 to call {service_name}. Run 'pip install boto3'.")
aws_request: Final = AWSRequest(method="POST", url=url, data=body, headers=headers)
SigV4Auth(get_credentials(), service_name, aws_region_name).add_auth(aws_request)
return aws_request.prepare()
_SignParams = ParamSpec("_SignParams")
_SignedRequest = TypeVar("_SignedRequest")
AWS_SIGNING_EXECUTOR: Final = ThreadPoolExecutor(max_workers=AWS_SIGNING_MAX_THREADS, thread_name_prefix="aws-signing")
async def run_aws_signing(
sign: Callable[_SignParams, _SignedRequest],
/,
*args: _SignParams.args,
**kwargs: _SignParams.kwargs, # kwargs-ok: ParamSpec forwarding keeps the wrapped signing signature
) -> _SignedRequest:
context: Final = contextvars.copy_context()
return await asyncio.get_running_loop().run_in_executor(
AWS_SIGNING_EXECUTOR, partial(context.run, sign, *args, **kwargs)
)
async def sign_request_off_loop_if_aws(
provider_config: object,
sign_request: Callable[_SignParams, _SignedRequest],
/,
*args: _SignParams.args,
**kwargs: _SignParams.kwargs, # kwargs-ok: ParamSpec forwarding keeps the wrapped sign_request signature
) -> _SignedRequest:
if isinstance(provider_config, SignsRequestsWithAWS):
return await run_aws_signing(sign_request, *args, **kwargs)
return sign_request(*args, **kwargs)

View file

@ -21,7 +21,7 @@ from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token, run_aws_signing
from ..common_utils import BedrockError, _get_all_bedrock_regions, error_response_text
from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call
@ -136,7 +136,8 @@ class BedrockConverseLLM(BaseAWSLLM):
)
data: Final = json.dumps(request_data)
prepped: Final = self.get_request_headers(
prepped: Final = await run_aws_signing(
self.get_request_headers,
credentials=credentials,
aws_region_name=litellm_params.get("aws_region_name") or "us-west-2",
extra_headers=headers,
@ -206,7 +207,8 @@ class BedrockConverseLLM(BaseAWSLLM):
)
data: Final = json.dumps(request_data)
prepped: Final = self.get_request_headers(
prepped: Final = await run_aws_signing(
self.get_request_headers,
credentials=credentials,
aws_region_name=litellm_params.get("aws_region_name") or "us-west-2",
extra_headers=headers,

View file

@ -10,9 +10,10 @@ import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.llms.bedrock.base_aws_llm import run_aws_signing
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.count_tokens.transformation import BedrockCountTokensConfig
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client
class BedrockCountTokensHandler(BedrockCountTokensConfig):
@ -27,6 +28,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
request_data: dict[str, Any],
litellm_params: dict[str, Any],
resolved_model: str,
client: AsyncHTTPHandler | None = None,
) -> dict[str, Any]:
"""
Handle a CountTokens request using existing LiteLLM patterns.
@ -75,7 +77,8 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
# Extract api_key for bearer token auth if provided
api_key: Final = litellm_params.get("api_key", None)
headers: Final = {"Content-Type": "application/json"}
signed_headers, signed_body = self._sign_request(
signed_headers, signed_body = await run_aws_signing(
self._sign_request,
service_name="bedrock",
headers=headers,
optional_params=litellm_params,
@ -85,7 +88,7 @@ class BedrockCountTokensHandler(BedrockCountTokensConfig):
api_key=api_key,
)
async_client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK)
async_client: Final = client or get_async_httpx_client(llm_provider=litellm.LlmProviders.BEDROCK)
response: Final = await async_client.post(
endpoint_url,

View file

@ -5,7 +5,7 @@ Handles embedding calls to Bedrock's `/invoke` endpoint
import copy
import json
import urllib.parse
from collections.abc import Callable
from collections.abc import Callable, Mapping
from typing import TYPE_CHECKING, Final, get_args, overload
import httpx
@ -26,7 +26,7 @@ from litellm.types.llms.bedrock import (
)
from litellm.types.utils import EmbeddingResponse, LlmProviders
from ..base_aws_llm import BaseAWSLLM, Credentials, bedrock_bearer_token
from ..base_aws_llm import AWSPreparedRequest, BaseAWSLLM, Credentials, bedrock_bearer_token, run_aws_signing
from ..common_utils import BedrockError
from .amazon_nova_transformation import AmazonNovaEmbeddingConfig
from .amazon_titan_g1_transformation import AmazonTitanG1Config
@ -41,6 +41,20 @@ if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
def _sign_get_request(
credentials: Credentials, url: str, headers: Mapping[str, str], aws_region_name: str
) -> AWSPreparedRequest:
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
request: Final = AWSRequest(method="GET", url=url, data=None, headers=headers)
SigV4Auth(credentials, "bedrock", aws_region_name).add_auth(request)
return request.prepare()
class BedrockEmbedding(BaseAWSLLM):
@overload
def _load_credentials(
@ -342,7 +356,8 @@ class BedrockEmbedding(BaseAWSLLM):
if extra_headers is not None:
headers = {"Content-Type": "application/json", **extra_headers}
prepped = self.get_request_headers(
prepped = await run_aws_signing(
self.get_request_headers,
credentials=credentials,
aws_region_name=aws_region_name,
extra_headers=extra_headers,
@ -600,9 +615,6 @@ class BedrockEmbedding(BaseAWSLLM):
dict: Status response from AWS Bedrock
"""
# Get AWS credentials using the same method as other Bedrock methods
credentials, _ = self._load_credentials(kwargs)
# Get the runtime endpoint
endpoint_url, _ = self.get_runtime_endpoint(
api_base=None,
@ -619,27 +631,13 @@ class BedrockEmbedding(BaseAWSLLM):
# Prepare headers for GET request
headers: Final = {"Content-Type": "application/json"}
# Use AWSRequest directly for GET requests (get_request_headers hardcodes POST)
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
def sign_status_request() -> AWSPreparedRequest:
credentials, _ = self._load_credentials(kwargs)
return _sign_get_request(
credentials=credentials, url=status_url, headers=headers, aws_region_name=aws_region_name
)
# Create AWSRequest with GET method and encoded URL
request: Final = AWSRequest(
method="GET",
url=status_url,
data=None, # GET request, no body
headers=headers,
)
# Sign the request - SigV4Auth will create canonical string from request URL
sigv4: Final = SigV4Auth(credentials, "bedrock", aws_region_name)
sigv4.add_auth(request)
# Prepare the request
prepped: Final = request.prepare()
prepped: Final = await run_aws_signing(sign_status_request)
# LOGGING
if logging_obj is not None:

View file

@ -21,7 +21,7 @@ from litellm.litellm_core_utils.realtime_streaming import DefaultLoggedRealTimeE
from litellm.types.llms.openai import OpenAIRealtimeEvents
from litellm.types.realtime import RealtimeResponseTransformInput
from ..base_aws_llm import BaseAWSLLM
from ..base_aws_llm import BaseAWSLLM, run_aws_signing
from ..common_utils import BedrockError
from .transformation import BedrockRealtimeConfig
@ -149,7 +149,8 @@ class BedrockRealtime(BaseAWSLLM):
verbose_proxy_logger.debug("Bedrock Realtime: Connecting to %s with model %s", endpoint_uri, model)
credentials: Final = self.get_credentials(
credentials: Final = await run_aws_signing(
self.get_credentials,
aws_access_key_id=aws_access_key_id,
aws_secret_access_key=aws_secret_access_key,
aws_session_token=aws_session_token,
@ -169,7 +170,7 @@ class BedrockRealtime(BaseAWSLLM):
"or configure credentials in the environment"
),
)
frozen_credentials: Final = credentials.get_frozen_credentials()
frozen_credentials: Final = await run_aws_signing(credentials.get_frozen_credentials)
# Initialize Bedrock client with aws_sdk_bedrock_runtime
config: Final = Config(

View file

@ -23,7 +23,7 @@ from botocore.exceptions import (
ProfileNotFound,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, SignsRequestsWithAWS
from litellm.secret_managers.main import get_secret_str
BEDROCK_MANTLE_DEFAULT_REGION: Final = "us-east-1"
@ -55,7 +55,7 @@ def resolve_mantle_region(params: Mapping[str, object]) -> str:
)
class BedrockMantleAuthMixin:
class BedrockMantleAuthMixin(SignsRequestsWithAWS):
_aws_signer: BaseAWSLLM
@staticmethod

View file

@ -77,6 +77,7 @@ from litellm.llms.base_llm.vector_store_files.transformation import (
BaseVectorStoreFilesConfig,
)
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS, run_aws_signing, sign_request_off_loop_if_aws
from litellm.llms.custom_httpx.container_handler import raise_for_error_status
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
@ -637,7 +638,12 @@ class BaseLLMHTTPHandler:
headers=request_headers,
),
)
return await dispatch_async(*await asyncio.to_thread(sign_and_log, transformed))
signed_request: Final = await (
run_aws_signing(sign_and_log, transformed)
if isinstance(provider_config, SignsRequestsWithAWS)
else asyncio.to_thread(sign_and_log, transformed)
)
return await dispatch_async(*signed_request)
return transform_then_dispatch()
@ -1973,7 +1979,9 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
signed_headers, signed_json_body = provider_config.sign_request(
signed_headers, signed_json_body = await sign_request_off_loop_if_aws(
provider_config,
provider_config.sign_request,
headers=headers,
optional_params=optional_params,
request_data=data,
@ -2074,7 +2082,9 @@ class BaseLLMHTTPHandler:
max_attempts,
)
provider_config.transform_anthropic_messages_request_on_http_error(e=e, request_data=request_body)
headers, signed_json_body = provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
provider_config,
provider_config.sign_request,
headers=headers,
optional_params=optional_params_dict,
request_data=request_body,
@ -2234,7 +2244,9 @@ class BaseLLMHTTPHandler:
stream=stream,
)
headers, signed_json_body = anthropic_messages_provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
anthropic_messages_provider_config,
anthropic_messages_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params), # dynamic aws_* params are passed under litellm_params
request_data=request_body,
@ -2910,7 +2922,9 @@ class BaseLLMHTTPHandler:
fake_stream=fake_stream,
)
headers, signed_body = responses_api_provider_config.sign_request(
headers, signed_body = await sign_request_off_loop_if_aws(
responses_api_provider_config,
responses_api_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
@ -4618,7 +4632,9 @@ class BaseLLMHTTPHandler:
)
data = BaseResponsesAPIConfig.normalize_responses_api_request_dict(data)
headers, signed_body = responses_api_provider_config.sign_request(
headers, signed_body = await sign_request_off_loop_if_aws(
responses_api_provider_config,
responses_api_provider_config.sign_request,
headers=headers,
optional_params=dict(litellm_params),
request_data=data,
@ -9845,7 +9861,9 @@ class BaseLLMHTTPHandler:
)
all_optional_params: Final[dict[str, object]] = dict(litellm_params)
all_optional_params.update(vector_store_search_optional_params or {})
headers, signed_json_body = vector_store_provider_config.sign_request(
headers, signed_json_body = await sign_request_off_loop_if_aws(
vector_store_provider_config,
vector_store_provider_config.sign_request,
headers=headers,
optional_params=all_optional_params,
request_data=request_body,

View file

@ -4,6 +4,7 @@ import subprocess
import sys
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final, TypeAlias
@ -12,6 +13,7 @@ import requests
from pydantic import BaseModel, TypeAdapter, ValidationError
from .auth import CliContextObj, context_secret_vault, get_stored_api_key, login
from .claude_settings import claude_settings_path, lite_api_key_helper_configured
from .cmd_quoting import quote_for_cmd
from .pi import (
LITELLM_PROXY_API_KEY_ENV,
@ -84,6 +86,8 @@ def build_agent_env(
base_url: str,
api_key: str,
profiles: frozenset[str],
*,
export_anthropic_token: bool = True,
) -> dict[str, str]:
"""Return a copy of base_env wired to route the agent through the proxy.
@ -98,12 +102,19 @@ def build_agent_env(
proxy's /v1/models; likewise left alone when already set.
pi ignores both base URL variables and instead resolves $LITELLM_PROXY_API_KEY
from its synced models.json provider entry.
With export_anthropic_token=False the bearer is left out (and any inherited
one dropped) so Claude Code asks its configured apiKeyHelper instead; Claude
Code prefers ANTHROPIC_AUTH_TOKEN over the helper and warns when both are set.
"""
env: Final = dict(base_env)
root: Final = base_url.rstrip("/")
if PROFILE_ANTHROPIC in profiles:
env[ANTHROPIC_BASE_URL_ENV] = root
env[ANTHROPIC_AUTH_TOKEN_ENV] = api_key
if export_anthropic_token:
env[ANTHROPIC_AUTH_TOKEN_ENV] = api_key
else:
env.pop(ANTHROPIC_AUTH_TOKEN_ENV, None)
env.pop(ANTHROPIC_API_KEY_ENV, None)
if ENABLE_TOOL_SEARCH_ENV not in env:
env[ENABLE_TOOL_SEARCH_ENV] = ENABLE_TOOL_SEARCH_VALUE
@ -463,6 +474,7 @@ def run_agent(
launcher: Callable[[str, Sequence[str], Mapping[str, str]], None] = _hand_off,
reattach_terminal: Callable[[], None] | None = None,
preparers: Mapping[str, _Preparer] = MappingProxyType(_PREPARERS),
export_anthropic_token: bool = True,
) -> None:
"""Validate, wire the environment, and hand off to the agent.
@ -494,7 +506,9 @@ def run_agent(
env: Final = MappingProxyType(
{
**build_agent_env(env_before_sync, base_url, api_key, profiles),
**build_agent_env(
env_before_sync, base_url, api_key, profiles, export_anthropic_token=export_anthropic_token
),
**(_NO_EXTRA_ENV if isinstance(synced, ModelSyncSkipped) else synced),
}
)
@ -532,14 +546,26 @@ def resolve_api_key(ctx: click.Context) -> str:
_SKIP_VERIFY_HELP: Final = "Skip the pre-launch key check against the proxy."
def _helper_supplies_token(
ctx_obj: CliContextObj, base_url: str, profiles: frozenset[str], settings_path: Path
) -> bool:
if PROFILE_ANTHROPIC not in profiles or not ctx_obj.get("api_key_from_token_file"):
return False
return lite_api_key_helper_configured(base_url, settings_path)
def _launch(ctx: click.Context, binary: str, args: Sequence[str], *, skip_verify: bool) -> None:
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
started_interactive: Final = _is_interactive()
api_key: Final = resolve_api_key(ctx)
display_name, _ = agent_profile(binary)
display_name, profiles = agent_profile(binary)
settings_path: Final = claude_settings_path(os.environ)
helper_supplies_token: Final = _helper_supplies_token(ctx_obj, base_url, profiles, settings_path)
click.echo(f"litellm: routing {display_name} through proxy at {base_url.rstrip('/')}")
if helper_supplies_token:
click.echo(f"litellm: {display_name} reads its key from the apiKeyHelper in {settings_path}")
try:
run_agent(
@ -548,6 +574,7 @@ def _launch(ctx: click.Context, binary: str, args: Sequence[str], *, skip_verify
[binary, *args],
skip_verify=skip_verify,
reattach_terminal=(_restore_controlling_terminal if started_interactive else None),
export_anthropic_token=not helper_supplies_token,
)
except AgentRunError as e:
raise click.ClickException(str(e))

View file

@ -1,3 +1,4 @@
import os
import sys
import time
import webbrowser
@ -40,16 +41,16 @@ from litellm.litellm_core_utils.cli_token_utils import (
)
from .claude_settings import (
CLAUDE_SETTINGS_PATH,
CONFIGURE_STATE_PATH,
SETTINGS_FILE_OWNERS,
STARTING_MODEL_ROLE,
ApiKeyHelper,
ClaudeSettingsError,
KeepModel,
claude_settings_path,
configure_claude_settings,
configure_state_path,
refuse_while_owned,
resolve_api_key_helper,
settings_file_owners,
)
from .pkce_login import (
Http,
@ -784,19 +785,20 @@ def _render_and_prompt_for_team_selection(teams: list[CliTeam]) -> str | None:
def _configure_claude_code(base_url: str) -> None:
"""Point Claude Code at base_url by patching ~/.claude/settings.json, undoable with `lite unconfigure claude`."""
"""Point Claude Code at base_url by patching the settings.json it reads, undoable with `lite unconfigure claude`."""
settings_path: Final = claude_settings_path(os.environ)
try:
configure_claude_settings(
base_url,
ApiKeyHelper(resolve_api_key_helper(base_url)),
KeepModel(),
CLAUDE_SETTINGS_PATH,
CONFIGURE_STATE_PATH,
SETTINGS_FILE_OWNERS,
settings_path,
configure_state_path(settings_path),
settings_file_owners(settings_path),
)
except ClaudeSettingsError as e:
raise click.ClickException(f"Logged in, but could not configure Claude Code: {e}")
click.echo(f"\nConfigured Claude Code: {CLAUDE_SETTINGS_PATH} now routes through {base_url.rstrip('/')}.")
click.echo(f"\nConfigured Claude Code: {settings_path} now routes through {base_url.rstrip('/')}.")
click.echo(
"Your other Claude Code settings were left untouched. Restart Claude Code to pick this up. "
f"Undo with `lite unconfigure claude`; `lite configure claude --model` sets {STARTING_MODEL_ROLE}."
@ -870,8 +872,9 @@ def login(ctx: click.Context, config_claude: bool, pkce: bool) -> None:
ctx_obj: Final[CliContextObj] = ctx.obj
base_url: Final = ctx_obj["base_url"]
if config_claude:
settings_path: Final = claude_settings_path(os.environ)
try:
refuse_while_owned(CLAUDE_SETTINGS_PATH, SETTINGS_FILE_OWNERS)
refuse_while_owned(settings_path, settings_file_owners(settings_path))
except ClaudeSettingsError as e:
raise click.ClickException(f"Cannot configure Claude Code, so not logging in: {e}")

View file

@ -62,6 +62,7 @@ _BASE_URL_PATH: Final = f"{ENV_KEY}.{ANTHROPIC_BASE_URL_KEY}"
STARTING_MODEL_ROLE: Final = "the /model picker's default row, the model Claude Code starts on"
CLAUDE_SETTINGS_PATH: Final = Path.home() / ".claude" / "settings.json"
CLAUDE_CONFIG_DIR_ENV: Final = "CLAUDE_CONFIG_DIR"
BACKUP_PATH: Final = Path.home() / ".litellm" / "claude_settings_backup.json"
AUTOROUTE_BACKUP_PATH: Final = Path.home() / ".litellm" / "autorouter" / "claude_settings_backup.json"
CONFIGURE_STATE_PATH: Final = Path.home() / ".litellm" / "claude_configure_state.json"
@ -88,6 +89,33 @@ class ClaudeSettingsError(Exception):
"""Raised for any user-actionable failure while reading or writing Claude Code settings."""
def claude_settings_path(environ: Mapping[str, str]) -> Path:
"""The settings.json Claude Code reads: under CLAUDE_CONFIG_DIR when set, else ~/.claude/settings.json."""
config_dir: Final = environ.get(CLAUDE_CONFIG_DIR_ENV, "")
if not config_dir:
return CLAUDE_SETTINGS_PATH
return Path(config_dir).expanduser() / "settings.json"
def _is_default_settings_file(settings_path: Path) -> bool:
return settings_path.resolve() == CLAUDE_SETTINGS_PATH.resolve()
def settings_file_owners(settings_path: Path) -> tuple[SettingsFileOwner, ...]:
"""The commands whose backups guard settings_path: `lite up` and `lite autoroute up` only ever manage the default file."""
return SETTINGS_FILE_OWNERS if _is_default_settings_file(settings_path) else ()
def configure_state_path(settings_path: Path) -> Path:
"""The receipt describing settings_path: the default file keeps CONFIGURE_STATE_PATH, and any other file
(a CLAUDE_CONFIG_DIR) gets its own beside it, keyed by its resolved path, so two settings files never
share one undo record."""
if _is_default_settings_file(settings_path):
return CONFIGURE_STATE_PATH
digest: Final = hashlib.sha256(str(settings_path.resolve()).encode()).hexdigest()
return CONFIGURE_STATE_PATH.parent / CONFIGURE_STATE_PATH.stem / f"{digest}.json"
@dataclass(frozen=True, slots=True)
class StaticToken:
"""A long-lived virtual key, written into env.ANTHROPIC_AUTH_TOKEN."""
@ -325,6 +353,19 @@ def resolve_api_key_helper(base_url: str, platform: str = sys.platform) -> str:
return " ".join(quote(token) for token in (lite_path, "--base-url", base_url, "auth", "print-token"))
def lite_api_key_helper_configured(base_url: str, settings_path: Path) -> bool:
"""Whether settings_path already carries the apiKeyHelper `lite login --config-claude` writes for base_url.
Only an exact match counts: a helper for another proxy, a hand-written one, or
settings that cannot be read leave the caller on the env-token path.
"""
try:
configured_helper: Final = load_json_or_empty(settings_path).get(API_KEY_HELPER_KEY)
return configured_helper == resolve_api_key_helper(base_url.rstrip("/"))
except ClaudeSettingsError:
return False
def _owned(container: Mapping[str, JsonValue], key: str) -> OwnedValue:
return OwnedValue(present=key in container, value=container.get(key))
@ -543,6 +584,7 @@ __all__ = (
"API_KEY_HELPER_KEY",
"AUTOROUTE_BACKUP_PATH",
"BACKUP_PATH",
"CLAUDE_CONFIG_DIR_ENV",
"CLAUDE_SETTINGS_PATH",
"CONFIGURE_STATE_PATH",
"ENABLE_GATEWAY_MODEL_DISCOVERY_KEY",
@ -569,11 +611,15 @@ __all__ = (
"UnconfigureOutcome",
"UnpinModel",
"WithheldCredential",
"claude_settings_path",
"configure_claude_settings",
"configure_state_path",
"lite_api_key_helper_configured",
"load_json_or_empty",
"merge_claude_settings",
"read_configure_receipt",
"refuse_while_owned",
"resolve_api_key_helper",
"settings_file_owners",
"unconfigure_claude_settings",
)

View file

@ -1,5 +1,6 @@
"""`lite configure claude` and `lite unconfigure claude`: persistent Claude Code wiring, undoable."""
import os
import re
import sys
from collections.abc import Callable, Sequence
@ -12,9 +13,6 @@ from InquirerPy.base.control import Choice
from .auth import CliContextObj, context_secret_vault, get_stored_api_key
from .claude_settings import (
CLAUDE_SETTINGS_PATH,
CONFIGURE_STATE_PATH,
SETTINGS_FILE_OWNERS,
STARTING_MODEL_ROLE,
ApiKeyHelper,
ClaudeCredential,
@ -24,9 +22,12 @@ from .claude_settings import (
StaticToken,
UnconfigureOutcome,
UnpinModel,
claude_settings_path,
configure_claude_settings,
configure_state_path,
refuse_while_owned,
resolve_api_key_helper,
settings_file_owners,
unconfigure_claude_settings,
)
from .pi import ListingFailure, PiSyncError, fetch_model_ids
@ -67,8 +68,9 @@ def resolve_credential(ctx: click.Context, api_key: str | None) -> tuple[ClaudeC
def _start(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, tuple[str, ...]]:
"""Every configure path begins the same way: the local ownership check first, so a `lite up`
session is refused before any login prompt or request, then the credential, then the listing."""
settings_path: Final = claude_settings_path(os.environ)
try:
refuse_while_owned(CLAUDE_SETTINGS_PATH, SETTINGS_FILE_OWNERS)
refuse_while_owned(settings_path, settings_file_owners(settings_path))
credential, key = resolve_credential(ctx, api_key)
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
@ -106,14 +108,20 @@ def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listed: Sequ
raise click.ClickException(
f"{model!r} is not served by {base_url} for this key. /v1/models lists: {shown}{more}."
)
settings_path: Final = claude_settings_path(os.environ)
try:
configure_claude_settings(
base_url, credential, _model_choice(model), CLAUDE_SETTINGS_PATH, CONFIGURE_STATE_PATH, SETTINGS_FILE_OWNERS
base_url,
credential,
_model_choice(model),
settings_path,
configure_state_path(settings_path),
settings_file_owners(settings_path),
)
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
in_picker: Final = sum(1 for listed_model in listed if _CLAUDE_CODE_PICKER_FILTER.search(listed_model))
click.echo(f"Configured Claude Code: {CLAUDE_SETTINGS_PATH} now routes through {base_url}.")
click.echo(f"Configured Claude Code: {settings_path} now routes through {base_url}.")
click.echo(
"Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN."
if isinstance(credential, StaticToken)
@ -130,9 +138,9 @@ def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listed: Sequ
"'claude' or 'anthropic')."
)
click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.")
if isinstance(credential, StaticToken) and CLAUDE_SETTINGS_PATH.is_symlink():
if isinstance(credential, StaticToken) and settings_path.is_symlink():
click.echo(
f"Note: {CLAUDE_SETTINGS_PATH} is a symlink to {CLAUDE_SETTINGS_PATH.resolve()}, so your key now lives in "
f"Note: {settings_path} is a symlink to {settings_path.resolve()}, so your key now lives in "
"that file; keep it out of version control.",
err=True,
)
@ -221,11 +229,13 @@ def unconfigure_claude() -> None:
Also undoes `lite login --config-claude`. Only keys still holding what configure wrote are
put back; anything you changed since is left as it is and named in the output.
"""
settings_path: Final = claude_settings_path(os.environ)
state_path: Final = configure_state_path(settings_path)
try:
outcome: Final = unconfigure_claude_settings(CLAUDE_SETTINGS_PATH, CONFIGURE_STATE_PATH, SETTINGS_FILE_OWNERS)
outcome: Final = unconfigure_claude_settings(settings_path, state_path, settings_file_owners(settings_path))
except ClaudeSettingsError as e:
raise click.ClickException(str(e))
_report_unconfigure(CLAUDE_SETTINGS_PATH, CONFIGURE_STATE_PATH, outcome)
_report_unconfigure(settings_path, state_path, outcome)
def _report_unconfigure(settings_path: Path, state_path: Path, outcome: UnconfigureOutcome) -> None:

View file

@ -44,7 +44,7 @@ from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicM
from litellm.llms.base_llm.guardrail_translation.utils import (
effective_scan_only_tool_results_for_guardrail,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, bedrock_bearer_token
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, bedrock_bearer_token, run_aws_signing
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -917,7 +917,9 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
source,
)
return BedrockGuardrailResponse()
credentials, aws_region_name = self._load_credentials(bearer_token=bedrock_bearer_token(api_key))
credentials, aws_region_name = await run_aws_signing(
self._load_credentials, bearer_token=bedrock_bearer_token(api_key)
)
allow_chunking: Final = not self._content_uses_contextual_grounding(content)
completed_chunk_usages: Final[list[BedrockGuardrailUsage]] = [] # mutable-ok: billed-chunk usage accumulator
@ -1178,7 +1180,8 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
**base_request_data,
"content": content,
} # mutable-ok: outbound JSON request body
prepared_request: Final = self._prepare_request(
prepared_request: Final = await run_aws_signing(
self._prepare_request,
credentials=credentials,
data=bedrock_request_data,
optional_params=self.optional_params,
@ -1875,10 +1878,13 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM):
return BedrockGuardrailResponse()
api_key: Final[str | None] = request_data.get("api_key") if request_data else None
credentials, aws_region_name = self._load_credentials(bearer_token=bedrock_bearer_token(api_key))
credentials, aws_region_name = await run_aws_signing(
self._load_credentials, bearer_token=bedrock_bearer_token(api_key)
)
body: Final[dict[str, object]] = {"messages": checks_messages, "checks": self.checks}
prepared_request: Final = self._prepare_request(
prepared_request: Final = await run_aws_signing(
self._prepare_request,
credentials=credentials,
data=body,
optional_params=self.optional_params,

View file

@ -15,6 +15,7 @@ import os
import re
from collections.abc import AsyncGenerator, Callable, Mapping, Sequence
from dataclasses import dataclass
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
@ -1120,13 +1121,6 @@ async def bedrock_proxy_route(
"""
create_request_copy(request)
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
aws_region_name: Final = get_secret_str(secret_name="AWS_REGION_NAME")
if not _is_bedrock_agent_runtime_route(endpoint=endpoint):
return await bedrock_llm_proxy_route(
@ -1157,20 +1151,24 @@ async def bedrock_proxy_route(
)
# Add or update query parameters
from litellm.llms.bedrock.base_aws_llm import run_aws_signing, sign_aws_json_post
from litellm.llms.bedrock.chat import BedrockConverseLLM
bedrock_llm: Final = BedrockConverseLLM()
credentials: Final[Credentials] = bedrock_llm.get_credentials()
sigv4: Final = SigV4Auth(credentials, "bedrock", aws_region_name)
headers: Final = {"Content-Type": "application/json"}
# Assuming the body contains JSON data, parse it
try:
data: Final = await _json_request_body(request)
except Exception as e:
raise HTTPException(status_code=400, detail={"error": e})
_request: Final = AWSRequest(method="POST", url=str(updated_url), data=json.dumps(data), headers=headers)
sigv4.add_auth(_request)
prepped: Final = _request.prepare()
prepped: Final = await run_aws_signing(
sign_aws_json_post,
get_credentials=bedrock_llm.get_credentials,
service_name="bedrock",
aws_region_name=aws_region_name,
url=str(updated_url),
body=json.dumps(data),
headers=MappingProxyType({"Content-Type": "application/json"}),
)
## check for streaming
is_streaming_request = False
@ -1228,13 +1226,6 @@ async def comprehend_medical_proxy_route(
[Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical)
"""
try:
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
from botocore.credentials import Credentials
except ImportError:
raise ImportError("Missing boto3 to call comprehendmedical. Run 'pip install boto3'.")
from .llm_provider_handlers.comprehend_medical_passthrough_logging_handler import (
COMPREHEND_MEDICAL_SUPPORTED_OPERATIONS,
)
@ -1265,20 +1256,23 @@ async def comprehend_medical_proxy_route(
if "stream" in data:
raise HTTPException(status_code=400, detail="'stream' is not a Comprehend Medical request member")
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM, run_aws_signing, sign_aws_json_post
credentials: Final[Credentials] = BaseAWSLLM().get_credentials(aws_region_name=aws_region_name)
sigv4: Final = SigV4Auth(credentials, "comprehendmedical", aws_region_name)
headers: Final = MappingProxyType(
{
"Content-Type": "application/x-amz-json-1.1",
"X-Amz-Target": f"{COMPREHEND_MEDICAL_TARGET_PREFIX}.{operation}",
}
)
target_url: Final = f"https://comprehendmedical.{aws_region_name}.{get_aws_dns_suffix(aws_region_name)}/"
_request: Final = AWSRequest(method="POST", url=target_url, data=json.dumps(data), headers=headers)
sigv4.add_auth(_request)
prepped: Final = _request.prepare()
prepped: Final = await run_aws_signing(
sign_aws_json_post,
get_credentials=partial(BaseAWSLLM().get_credentials, aws_region_name=aws_region_name),
service_name="comprehendmedical",
aws_region_name=aws_region_name,
url=target_url,
body=json.dumps(data),
headers=MappingProxyType(
{
"Content-Type": "application/x-amz-json-1.1",
"X-Amz-Target": f"{COMPREHEND_MEDICAL_TARGET_PREFIX}.{operation}",
}
),
)
endpoint_func: Final = create_pass_through_route(
endpoint=operation,

View file

@ -6,6 +6,7 @@ extension, and AWS credential resolution is stubbed so nothing reaches STS.
from __future__ import annotations
import asyncio
from unittest.mock import MagicMock, patch
import httpx
@ -16,6 +17,7 @@ from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.rust_bridge import chat_completions as bridge
from litellm.types.utils import ModelResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
RUST_RESPONSE = {
"created": 1_700_000_000,
@ -308,7 +310,9 @@ CONVERSE_RESPONSE = {
}
async def _drive_async_completion(*, skip_pre_call_logging: bool, logging_obj):
async def _drive_async_completion(
*, skip_pre_call_logging: bool, logging_obj, credentials: Credentials = RESOLVED_CREDENTIALS
):
"""Run the real `async_completion` with a stubbed transport."""
import httpx as _httpx
@ -335,7 +339,7 @@ async def _drive_async_completion(*, skip_pre_call_logging: bool, logging_obj):
stream=None,
optional_params={"maxTokens": 16},
litellm_params={"aws_region_name": "us-west-2"},
credentials=RESOLVED_CREDENTIALS,
credentials=credentials,
headers={},
client=client,
skip_pre_call_logging=skip_pre_call_logging,
@ -357,6 +361,23 @@ async def test_async_completion_logs_pre_call_by_default():
assert logging_obj.pre_call.call_count == 1
@pytest.mark.asyncio
async def test_async_completion_signs_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: botocore refreshes expiring credentials inside SigV4 signing with a
blocking HTTP call, so `async_completion` must sign on a worker thread to keep the loop serving."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe = EventLoopProbe()
release = asyncio.create_task(probe.release_refresh_from_the_loop())
response = await _drive_async_completion(
skip_pre_call_logging=False, logging_obj=MagicMock(), credentials=probe.credentials()
)
await release
assert response.choices[0].message.content == "hi"
assert probe.served_during_refresh is True
def _sync_client_returning_converse_response():
client = MagicMock()
client.post.side_effect = lambda **_kwargs: httpx.Response(

View file

@ -0,0 +1,54 @@
import asyncio
from unittest.mock import AsyncMock
import httpx
import pytest
from botocore.credentials import RefreshableCredentials
from litellm.llms.bedrock.count_tokens.handler import BedrockCountTokensHandler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
class _ProbedCountTokensHandler(BedrockCountTokensHandler):
def __init__(self, probe: EventLoopProbe) -> None:
super().__init__()
self._probe = probe
def get_credentials(
self,
**kwargs: object, # kwargs-ok: mirrors the base resolver's keyword contract, which the probe ignores
) -> RefreshableCredentials:
return self._probe.credentials()
@pytest.mark.asyncio
async def test_handle_count_tokens_request_signs_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: the count_tokens handler signed on the loop, so botocore's blocking
credential refresh inside SigV4 stalled every other request on the worker."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe = EventLoopProbe()
client = AsyncMock(spec=AsyncHTTPHandler)
client.post = AsyncMock(
return_value=httpx.Response(
200,
json={"inputTokens": 7},
request=httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com/"),
)
)
release = asyncio.create_task(probe.release_refresh_from_the_loop())
result = await _ProbedCountTokensHandler(probe).handle_count_tokens_request(
request_data={
"model": "us.anthropic.claude-haiku-4-5-20251001-v1:0",
"messages": [{"role": "user", "content": "hi"}],
},
litellm_params={"aws_region_name": "us-west-2"},
resolved_model="us.anthropic.claude-haiku-4-5-20251001-v1:0",
client=client,
)
await release
assert result == {"input_tokens": 7}
assert client.post.call_args.kwargs["headers"]["Authorization"].startswith("AWS4-HMAC-SHA256")
assert probe.served_during_refresh is True

View file

@ -1,11 +1,15 @@
import json
import asyncio
from unittest.mock import Mock, patch
import httpx
import pytest
import respx
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.types.llms.base import HiddenParams
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
# Mock async invoke responses
async_invoke_response = {
@ -422,3 +426,34 @@ class TestBedrockAsyncInvokeEmbedding:
async_endpoint
== "https://bedrock-runtime.us-east-1.amazonaws.com/async-invoke"
)
@pytest.mark.asyncio
async def test_async_invoke_status_signs_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: the GetAsyncInvoke poll is a signed GET, and botocore refreshes
expiring credentials inside that signing with a blocking HTTP call, so it must run on a worker
thread to keep the loop serving other requests."""
from litellm.llms.bedrock.embed.embedding import BedrockEmbedding
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
litellm.in_memory_llm_clients_cache.flush_cache()
embedder = BedrockEmbedding()
probe = EventLoopProbe()
with (
patch.object(embedder, "_load_credentials", return_value=(probe.credentials(), "us-east-1")),
respx.mock,
):
route = respx.get(url__regex=r"https://bedrock-runtime\.us-east-1\.amazonaws\.com/async-invoke/.*").mock(
return_value=httpx.Response(200, json=async_invoke_status_response)
)
release = asyncio.create_task(probe.release_refresh_from_the_loop())
status = await embedder._get_async_invoke_status(
invocation_arn=async_invoke_status_response["invocationArn"], aws_region_name="us-east-1"
)
await release
assert status["status"] == "InProgress"
assert "Authorization" in route.calls.last.request.headers
assert probe.served_during_refresh is True

View file

@ -1,12 +1,17 @@
import json
import asyncio
import os
from unittest.mock import Mock, patch
from unittest.mock import AsyncMock, MagicMock
import pytest
import httpx
import litellm
from litellm.llms.bedrock.embed.twelvelabs_marengo_transformation import TwelveLabsMarengoEmbeddingConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.bedrock.embed.embedding import BedrockEmbedding
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
# Mock responses for different embedding models
titan_embedding_response = {"embedding": [0.1, 0.2, 0.3], "inputTextTokenCount": 10}
@ -1062,6 +1067,41 @@ def test_bedrock_embedding_bearer_token_never_runs_the_sigv4_credential_chain(mo
assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"
@pytest.mark.asyncio
async def test_async_single_func_embeddings_signs_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: Titan, Nova, and TwelveLabs embeddings sign one SigV4 request per
input, and botocore refreshes expiring credentials inside that signing with a blocking HTTP call,
so each signing must run on a worker thread to keep the loop serving other requests."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe = EventLoopProbe()
client = MagicMock()
client.__class__ = AsyncHTTPHandler
client.post = AsyncMock(
return_value=httpx.Response(
200,
json=titan_embedding_response,
request=httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com"),
)
)
release = asyncio.create_task(probe.release_refresh_from_the_loop())
response = await BedrockEmbedding()._async_single_func_embeddings(
client=client,
timeout=None,
batch_data=[{"inputText": test_input}],
credentials=probe.credentials(),
extra_headers=None,
endpoint_url="https://bedrock-runtime.us-west-2.amazonaws.com/model/amazon.titan-embed-text-v1/invoke",
aws_region_name="us-west-2",
model="amazon.titan-embed-text-v1",
logging_obj=MagicMock(),
provider="amazon",
)
await release
assert response.data[0]["embedding"] == titan_embedding_response["embedding"]
assert "Authorization" in client.post.call_args.kwargs["headers"]
assert probe.served_during_refresh is True
marengo_3_embedding_response = {"data": [{"embedding": [0.01 * i for i in range(512)]}]}
MARENGO_3_DUCK = "data:image/png;base64,ZHVjaw=="

View file

@ -0,0 +1,57 @@
"""Refreshable credentials whose refresh only completes while the event loop keeps serving."""
from __future__ import annotations
import asyncio
import threading
import time
from datetime import datetime, timedelta, timezone
from typing import Final
from botocore.credentials import RefreshableCredentials
REFRESH_RELEASE_TIMEOUT_SECONDS: Final = 2.0
REFRESH_START_TIMEOUT_SECONDS: Final = 10.0
class EventLoopProbe:
"""Blocks inside botocore's credential refresh until a coroutine on the loop releases it.
Signing on the event loop thread can never be released, so `served_during_refresh` reads False there
and True only when the refresh ran on another thread while the loop stayed responsive.
"""
def __init__(self) -> None:
self.refresh_started: Final = threading.Event()
self.loop_served: Final = threading.Event()
self.served_during_refresh: bool | None = None
def refresh(self) -> dict[str, str | None]:
self.refresh_started.set()
served: Final = self.loop_served.wait(timeout=REFRESH_RELEASE_TIMEOUT_SECONDS)
if self.served_during_refresh is None:
self.served_during_refresh = served
return {
"access_key": "AKIAREFRESHED",
"secret_key": "refreshed-secret",
"token": None,
"expiry_time": (datetime.now(timezone.utc) + timedelta(hours=1)).isoformat(),
}
def credentials(self) -> RefreshableCredentials:
return RefreshableCredentials(
access_key="AKIASTALE",
secret_key="stale-secret",
token=None,
expiry_time=datetime.now(timezone.utc) + timedelta(seconds=60),
refresh_using=self.refresh,
method="event-loop-probe",
)
async def release_refresh_from_the_loop(self) -> None:
deadline: Final = time.monotonic() + REFRESH_START_TIMEOUT_SECONDS
while not self.refresh_started.is_set():
if time.monotonic() > deadline:
raise TimeoutError("signing finished without ever starting a credential refresh")
await asyncio.sleep(0.005)
self.loop_served.set()

View file

@ -1,4 +1,6 @@
import asyncio
import json
from concurrent.futures import ThreadPoolExecutor
import os
import threading
import time
@ -22,7 +24,10 @@ from litellm.llms.bedrock.base_aws_llm import (
AwsAuthError,
BaseAWSLLM,
Boto3CredentialsInfo,
run_aws_signing,
sign_request_off_loop_if_aws,
)
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
# Global variable for the base_aws_llm.py file path
@ -3223,3 +3228,53 @@ class TestGetRequestHeadersResign:
extra_headers={"Authorization": "Bearer foo"},
)
assert prepped.headers["Authorization"] == "Bearer foo"
@pytest.mark.asyncio
async def test_sign_request_off_loop_if_aws_keeps_the_loop_serving_while_credentials_refresh():
"""Regression for issue #40165: an AWS provider's signing (and the botocore credential refresh
inside it) must run off the event loop, so other requests keep being served meanwhile."""
probe = EventLoopProbe()
def sign(headers: dict[str, str]) -> dict[str, str]:
request = AWSRequest(
method="POST", url="https://bedrock-runtime.us-west-2.amazonaws.com/", data="{}", headers=headers
)
SigV4Auth(probe.credentials(), "bedrock", "us-west-2").add_auth(request)
return dict(request.headers)
release = asyncio.create_task(probe.release_refresh_from_the_loop())
signed = await sign_request_off_loop_if_aws(BaseAWSLLM(), sign, headers={"Content-Type": "application/json"})
await release
assert "Authorization" in signed
assert probe.served_during_refresh is True
def test_run_aws_signing_leaves_the_default_executor_free_for_other_providers():
"""A signing parked on botocore's refresh lock must not hold a default-executor thread, since every
other provider's async entry point hops through that same executor. The scenario runs on its own loop
so the one-thread default executor it pins never leaks into the session loop."""
async def scenario() -> tuple[str, str]:
loop = asyncio.get_running_loop()
loop.set_default_executor(ThreadPoolExecutor(max_workers=1))
signing_parked = asyncio.Event()
refresh_done = threading.Event()
def sign() -> str:
loop.call_soon_threadsafe(signing_parked.set)
refresh_done.wait()
return threading.current_thread().name
signing = asyncio.create_task(run_aws_signing(sign))
try:
await asyncio.wait_for(signing_parked.wait(), timeout=5)
other_provider = await asyncio.wait_for(loop.run_in_executor(None, threading.current_thread), timeout=5)
finally:
refresh_done.set()
return other_provider.name, await signing
other_provider, signing_thread = asyncio.run(scenario())
assert other_provider != signing_thread
assert signing_thread.startswith("aws-signing")

View file

@ -6,15 +6,20 @@ API docs: https://docs.aws.amazon.com/bedrock/latest/userguide/bedrock-mantle.ht
"""
import json
import asyncio
from unittest.mock import patch
import httpx
import pytest
from botocore.auth import SigV4Auth
from botocore.awsrequest import AWSRequest
import litellm
from litellm.llms.bedrock_mantle.chat.transformation import BedrockMantleChatConfig
from litellm.llms.bedrock.base_aws_llm import sign_request_off_loop_if_aws
from litellm.types.utils import LlmProviders
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
@pytest.fixture
@ -710,3 +715,26 @@ def test_gemma_4_models_register_under_bedrock_mantle(local_cost_map, model_id):
resolved_model, provider, _, _ = litellm.get_llm_provider(full_model_name)
assert provider == "bedrock_mantle"
assert resolved_model == model_id
@pytest.mark.asyncio
async def test_mantle_signing_runs_off_the_event_loop():
"""Regression for issue #40165: Mantle signs with SigV4 through a composed BaseAWSLLM, so the
off-loop gate must recognise it too, or its credential refresh blocks the loop like Bedrock's did."""
probe = EventLoopProbe()
def sign(headers: dict[str, str]) -> dict[str, str]:
request = AWSRequest(
method="POST", url="https://bedrock-mantle.us-east-1.api.aws/v1/responses", data="{}", headers=headers
)
SigV4Auth(probe.credentials(), "bedrock", "us-east-1").add_auth(request)
return dict(request.headers)
release = asyncio.create_task(probe.release_refresh_from_the_loop())
signed = await sign_request_off_loop_if_aws(
BedrockMantleChatConfig(), sign, headers={"Content-Type": "application/json"}
)
await release
assert "Authorization" in signed
assert probe.served_during_refresh is True

View file

@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, Mock, patch
import httpx
import pytest
from botocore.credentials import RefreshableCredentials
import litellm
from litellm._logging import verbose_logger
@ -19,6 +20,7 @@ from litellm.llms.base_llm.audio_transcription.transformation import (
BaseAudioTranscriptionConfig,
)
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS
from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.llms.custom_httpx.llm_http_handler import (
@ -29,11 +31,15 @@ from litellm.llms.custom_httpx.llm_http_handler import (
_rust_responses_websocket_enabled,
)
from litellm.llms.azure.videos.transformation import AzureVideoConfig
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
from litellm.llms.mistral.ocr.transformation import MistralOCRConfig
from litellm.llms.openai.videos.transformation import OpenAIVideoConfig
from litellm.types.llms.openai import ResponsesAPIResponse
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import ImageObject, ImageResponse, ModelResponse, TranscriptionResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
_ACTIVE_KEY = "_code_interpreter_interception_active"
_SANDBOX_KEY = "_code_interpreter_interception_sandbox_key"
@ -813,6 +819,65 @@ async def test_anthropic_messages_streaming_response_aclose_closes_agentic_upstr
assert tracker.closed is True
class _ProbedBedrockMessagesConfig(AmazonAnthropicClaudeMessagesConfig):
def __init__(self, probe: EventLoopProbe) -> None:
super().__init__()
self._probe = probe
def get_credentials(
self,
**kwargs: object, # kwargs-ok: mirrors the base resolver's keyword contract, which the probe ignores
) -> RefreshableCredentials:
return self._probe.credentials()
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_signs_bedrock_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: /v1/messages on Bedrock signed on the loop, so botocore's blocking
credential refresh inside SigV4 stalled every other request on the worker."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe = EventLoopProbe()
handler = BaseLLMHTTPHandler()
upstream_response = httpx.Response(
200,
json={
"id": "msg_123",
"type": "message",
"role": "assistant",
"content": [{"type": "text", "text": "hi"}],
"model": "claude-haiku-4-5-20251001",
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 1, "output_tokens": 1},
},
request=httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com/"),
)
mock_client = AsyncMock(spec=AsyncHTTPHandler)
mock_client.post = AsyncMock(return_value=upstream_response)
mock_logging_obj = Mock()
mock_logging_obj.model_call_details = {}
mock_logging_obj.dynamic_success_callbacks = None
release = asyncio.create_task(probe.release_refresh_from_the_loop())
await handler.async_anthropic_messages_handler(
model="us.anthropic.claude-haiku-4-5-20251001-v1:0",
messages=[{"role": "user", "content": "hi"}],
anthropic_messages_provider_config=_ProbedBedrockMessagesConfig(probe),
anthropic_messages_optional_request_params={"max_tokens": 16},
custom_llm_provider="bedrock",
litellm_params=GenericLiteLLMParams(aws_region_name="us-west-2"),
logging_obj=mock_logging_obj,
client=mock_client,
stream=False,
kwargs={},
)
await release
sent_headers = mock_client.post.call_args.kwargs["headers"]
assert sent_headers["Authorization"].startswith("AWS4-HMAC-SHA256")
assert probe.served_during_refresh is True
@pytest.mark.asyncio
async def test_async_anthropic_messages_handler_passes_litellm_metadata():
"""Ensure litellm_metadata from kwargs is forwarded via update_from_kwargs.
@ -3367,6 +3432,26 @@ async def test_completion_signs_and_logs_off_the_event_loop_after_the_async_tran
assert captured["body"] == {"transformed_by": "async"}
assert config.sign_threads and all(thread is not loop_thread for thread in config.sign_threads)
assert pre_call_threads and all(thread is not loop_thread for thread in pre_call_threads)
assert not any(thread.name.startswith("aws-signing") for thread in config.sign_threads + pre_call_threads)
class _AWSTransformRecordingConfig(SignsRequestsWithAWS, _TransformRecordingConfig):
pass
async def test_completion_signs_aws_configs_on_the_aws_signing_pool_after_the_async_transform():
config = _AWSTransformRecordingConfig(transform_async=True)
pre_call_threads = []
logging_obj = Mock(dynamic_success_callbacks=None, model_call_details={})
logging_obj.pre_call.side_effect = lambda **kwargs: pre_call_threads.append(threading.current_thread())
pending, captured = _start_async_completion(config, logging_obj)
response = await pending
assert response.choices[0].message.content == "async"
assert captured["body"] == {"transformed_by": "async"}
assert config.sign_threads and all(thread.name.startswith("aws-signing") for thread in config.sign_threads)
assert pre_call_threads and all(thread.name.startswith("aws-signing") for thread in pre_call_threads)
async def test_completion_keeps_sync_transform_request_before_returning_by_default():

View file

@ -0,0 +1,32 @@
import os
from collections.abc import Iterator
from pathlib import Path
from typing import Final
import pytest
REAL_CLAUDE_SETTINGS: Final = Path(os.path.expanduser("~")) / ".claude" / "settings.json"
def _current_bytes() -> bytes | None:
return REAL_CLAUDE_SETTINGS.read_bytes() if REAL_CLAUDE_SETTINGS.exists() else None
@pytest.fixture(autouse=True)
def isolated_claude_home(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Iterator[Path]:
before: Final = _current_bytes()
monkeypatch.setenv("HOME", str(tmp_path))
monkeypatch.setenv("USERPROFILE", str(tmp_path))
monkeypatch.setenv("CLAUDE_CONFIG_DIR", str(tmp_path / ".claude"))
yield tmp_path
after: Final = _current_bytes()
if after == before:
return
if before is None:
REAL_CLAUDE_SETTINGS.unlink()
else:
REAL_CLAUDE_SETTINGS.write_bytes(before)
pytest.fail(
f"this test wrote the developer's real {REAL_CLAUDE_SETTINGS}; the original bytes were restored. "
"Resolve the Claude settings path at call time (never Path.home() at import) and point the test at tmp_path"
)

View file

@ -28,6 +28,7 @@ from litellm.proxy.client.cli.commands.agents import (
)
AGENTS_MODULE = "litellm.proxy.client.cli.commands.agents"
CLAUDE_SETTINGS_MODULE = "litellm.proxy.client.cli.commands.claude_settings"
def _agent_command(name):
@ -163,6 +164,27 @@ class TestBuildAgentEnv:
assert env["PATH"] == "/usr/bin"
assert base == {"PATH": "/usr/bin", "ANTHROPIC_API_KEY": "real-key"}
def test_anthropic_profile_leaves_the_bearer_to_the_api_key_helper(self):
env = build_agent_env(
{"ANTHROPIC_AUTH_TOKEN": "stale-token", "ANTHROPIC_API_KEY": "real-key"},
"http://localhost:4000/",
"sk-key",
frozenset({"anthropic"}),
export_anthropic_token=False,
)
assert "ANTHROPIC_AUTH_TOKEN" not in env
assert "ANTHROPIC_API_KEY" not in env
assert env["ANTHROPIC_BASE_URL"] == "http://localhost:4000"
assert env["ENABLE_TOOL_SEARCH"] == "true"
assert env["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1"
def test_helper_mode_still_exports_the_openai_key(self):
env = build_agent_env(
{}, "http://localhost:4000", "sk-key", frozenset({"anthropic", "openai"}), export_anthropic_token=False
)
assert "ANTHROPIC_AUTH_TOKEN" not in env
assert env["OPENAI_API_KEY"] == "sk-key"
class TestAgentLaunchArgs:
def test_claude_and_opencode_get_no_extra_args(self):
@ -509,6 +531,25 @@ class TestRunAgent:
assert "ANTHROPIC_API_KEY" not in env
assert "OPENAI_BASE_URL" not in env
def test_helper_supplied_token_never_reaches_the_launch_env(self):
calls = {}
verified = []
run_agent(
"http://localhost:4000",
"sk-key",
["claude"],
base_env={"PATH": "/usr/bin", "ANTHROPIC_AUTH_TOKEN": "stale-token"},
which=lambda name: "/usr/local/bin/claude",
verify=lambda base_url, api_key: verified.append(api_key),
launcher=lambda p, a, e: calls.update(env=dict(e)),
export_anthropic_token=False,
)
assert verified == ["sk-key"]
assert "ANTHROPIC_AUTH_TOKEN" not in calls["env"]
assert calls["env"]["ANTHROPIC_BASE_URL"] == "http://localhost:4000"
def test_codex_gets_openai_env(self):
calls = {}
run_agent(
@ -1052,6 +1093,101 @@ class TestAgentCommands:
in result.output
)
def _invoke_claude_with_settings(self, tmp_path, settings, obj, *, default_settings=None):
config_dir = tmp_path / "claude-config"
config_dir.mkdir()
if settings is not None:
(config_dir / "settings.json").write_text(json.dumps(settings))
default_path = tmp_path / "home-claude" / "settings.json"
default_path.parent.mkdir()
if default_settings is not None:
default_path.write_text(json.dumps(default_settings))
captured = {}
with (
patch(f"{CLAUDE_SETTINGS_MODULE}.CLAUDE_SETTINGS_PATH", default_path),
patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value="/usr/local/bin/lite"),
patch(f"{AGENTS_MODULE}.run_agent", side_effect=lambda b, k, c, **kw: captured.update(kw)),
):
result = self.runner.invoke(
_agent_command("claude"), [], obj=obj, env={"CLAUDE_CONFIG_DIR": str(config_dir)}
)
assert result.exit_code == 0, result.output
return captured, result.output
def test_helper_is_read_from_the_config_dir_claude_code_uses(self, tmp_path):
captured, output = self._invoke_claude_with_settings(
tmp_path,
{"apiKeyHelper": "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"},
{"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True},
)
assert captured["export_anthropic_token"] is False
assert str(tmp_path / "claude-config" / "settings.json") in output
def test_helper_only_in_the_default_file_keeps_the_env_token_when_config_dir_points_elsewhere(self, tmp_path):
captured, output = self._invoke_claude_with_settings(
tmp_path,
None,
{"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True},
default_settings={"apiKeyHelper": "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"},
)
assert captured["export_anthropic_token"] is True
assert "apiKeyHelper" not in output
def test_stored_login_with_a_matching_helper_leaves_the_token_to_the_helper(self, tmp_path):
captured, output = self._invoke_claude_with_settings(
tmp_path,
{"apiKeyHelper": "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"},
{"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True},
)
assert captured["export_anthropic_token"] is False
assert "reads its key from the apiKeyHelper" in output
def test_explicit_key_is_exported_even_when_a_helper_matches(self, tmp_path):
captured, output = self._invoke_claude_with_settings(
tmp_path,
{"apiKeyHelper": "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"},
{"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": False},
)
assert captured["export_anthropic_token"] is True
assert "apiKeyHelper" not in output
def test_helper_for_another_proxy_keeps_the_env_token(self, tmp_path):
captured, _ = self._invoke_claude_with_settings(
tmp_path,
{"apiKeyHelper": "/usr/local/bin/lite --base-url https://other.example.com auth print-token"},
{"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True},
)
assert captured["export_anthropic_token"] is True
def test_no_claude_settings_keeps_the_env_token(self, tmp_path):
captured, _ = self._invoke_claude_with_settings(
tmp_path,
None,
{"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True},
)
assert captured["export_anthropic_token"] is True
def test_codex_never_consults_claude_settings(self):
captured = {}
with (
patch(f"{AGENTS_MODULE}.lite_api_key_helper_configured", side_effect=AssertionError("consulted")),
patch(f"{AGENTS_MODULE}.run_agent", side_effect=lambda b, k, c, **kw: captured.update(kw)),
):
result = self.runner.invoke(
_agent_command("codex"),
[],
obj={"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True},
)
assert result.exit_code == 0, result.output
assert captured["export_anthropic_token"] is True
def test_codex_shows_friendly_name(self):
captured = {}
with patch(

View file

@ -6,7 +6,6 @@ from pathlib import Path
from unittest.mock import Mock, patch
import pytest
from click.testing import CliRunner
@ -27,7 +26,7 @@ from litellm.proxy.client.cli.commands.auth import (
print_token,
whoami,
)
from litellm.proxy.client.cli.commands import auth as auth_module
from litellm.proxy.client.cli.commands import claude_settings as claude_settings_module
from litellm.proxy.client.cli.commands.claude_settings import SettingsFileOwner
@ -85,7 +84,7 @@ class TestPollingErrorSurfacing:
}
with patch("requests.get", return_value=mock_response) as mock_get, patch("time.sleep"):
with pytest.raises(ValueError, match='Your litellm CLI is out of date and uses a login flow') as exc_info:
with pytest.raises(ValueError, match="Your litellm CLI is out of date and uses a login flow") as exc_info:
_poll_for_ready_data("http://test/sso/cli/poll/sk-legacy")
assert mock_get.call_count == 1
@ -152,7 +151,7 @@ class TestStartCliSsoFlowErrors:
mock_response.status_code = 404
with patch("requests.post", return_value=mock_response):
with pytest.raises(ValueError, match='Either --base-url is wrong, or the proxy is older than') as exc_info:
with pytest.raises(ValueError, match="Either --base-url is wrong, or the proxy is older than") as exc_info:
_start_cli_sso_flow("https://old-proxy.example.com")
message = str(exc_info.value)
@ -168,7 +167,7 @@ class TestStartCliSsoFlowErrors:
mock_response.json.return_value = {"detail": "Too many CLI login attempts. Try again later."}
with patch("requests.post", return_value=mock_response):
with pytest.raises(ValueError, match='Too many CLI login attempts\\. Try again later\\.') as exc_info:
with pytest.raises(ValueError, match="Too many CLI login attempts\\. Try again later\\.") as exc_info:
_start_cli_sso_flow("https://test.example.com")
assert "HTTP 429" in str(exc_info.value)
@ -184,7 +183,7 @@ class TestStartCliSsoFlowErrors:
mock_response.text = "<html>Sign in to corporate VPN</html>"
with patch("requests.post", return_value=mock_response):
with pytest.raises(ValueError, match='A proxy, load balancer, or auth gateway in front of') as exc_info:
with pytest.raises(ValueError, match="A proxy, load balancer, or auth gateway in front of") as exc_info:
_start_cli_sso_flow("https://test.example.com")
message = str(exc_info.value)
@ -198,7 +197,7 @@ class TestStartCliSsoFlowErrors:
from litellm.proxy.client.cli.commands.auth import _start_cli_sso_flow
with patch("requests.post", side_effect=requests.ConnectionError("Connection refused")):
with pytest.raises(ValueError, match='Connection refused\\. Check that the proxy is running') as exc_info:
with pytest.raises(ValueError, match="Connection refused\\. Check that the proxy is running") as exc_info:
_start_cli_sso_flow("https://unreachable.example.com")
message = str(exc_info.value)
@ -585,13 +584,9 @@ class TestLogoutCommand:
assert "could not be checked" in result.output
assert DISABLE_KEYRING_ENV_VAR in result.output
def test_logout_warns_when_the_keychain_refuses_to_release_the_entry(
self, isolated_home, secret_vault_factory
):
def test_logout_warns_when_the_keychain_refuses_to_release_the_entry(self, isolated_home, secret_vault_factory):
"""A locked keychain leaves a live credential behind that the user believes is gone."""
vault = secret_vault_factory(
blob=_secret_blob("https://test.example.com", "sk-stored"), erasable=False
)
vault = secret_vault_factory(blob=_secret_blob("https://test.example.com", "sk-stored"), erasable=False)
_write_token_file(isolated_home, key=None)
result = self.runner.invoke(logout, obj={"secret_vault": vault})
@ -1211,9 +1206,7 @@ class TestKeychainBackedCommands:
assert str(token_file) in result.output
assert json.loads(token_file.read_text())["key"] == "sk-minted"
def test_login_points_a_user_missing_the_keyring_package_at_the_install(
self, isolated_home, secret_vault_factory
):
def test_login_points_a_user_missing_the_keyring_package_at_the_install(self, isolated_home, secret_vault_factory):
"""`lite` ships with every install, the keyring package only with the cli extra. Telling
that user their machine has no keychain sends them looking for a problem they do not have."""
result = self._login(secret_vault_factory(available=False, failure=KeyringNotInstalled()))
@ -1224,9 +1217,7 @@ class TestKeychainBackedCommands:
assert "No OS keychain available" not in result.output
assert json.loads(token_file.read_text())["key"] == "sk-minted"
def test_login_keeps_the_credential_when_the_backend_keeps_nothing(
self, isolated_home, secret_vault_factory
):
def test_login_keeps_the_credential_when_the_backend_keeps_nothing(self, isolated_home, secret_vault_factory):
"""A backend that accepts writes and stores nothing must not be reported as keychain
storage, because the file is then told to drop the only remaining copy."""
result = self._login(secret_vault_factory(discards=True))
@ -1237,9 +1228,7 @@ class TestKeychainBackedCommands:
assert "keyring --enable" in result.output
assert json.loads(token_file.read_text())["key"] == "sk-minted"
def test_login_names_the_kill_switch_instead_of_blaming_the_machine(
self, isolated_home, secret_vault_factory
):
def test_login_names_the_kill_switch_instead_of_blaming_the_machine(self, isolated_home, secret_vault_factory):
result = self._login(secret_vault_factory(available=False, failure=KeyringDisabled()))
assert result.exit_code == 0
@ -1298,9 +1287,7 @@ class TestKeychainBackedCommands:
assert "could not be read" in result.output
assert "lite login" in result.output
def test_whoami_does_not_call_a_credential_it_cannot_read_authenticated(
self, isolated_home, secret_vault_factory
):
def test_whoami_does_not_call_a_credential_it_cannot_read_authenticated(self, isolated_home, secret_vault_factory):
"""A login whose secret is stuck in an unreachable keychain authenticates nothing. Leading
with "Authenticated" and a token age reads as a working session, and sends the user looking
for the problem somewhere other than the keychain the notice underneath names."""
@ -1317,9 +1304,7 @@ class TestKeychainBackedCommands:
assert "the credential cannot be read" in result.output
assert "could not be read" in result.output
def test_whoami_names_the_kill_switch_rather_than_a_missing_package(
self, isolated_home, secret_vault_factory
):
def test_whoami_names_the_kill_switch_rather_than_a_missing_package(self, isolated_home, secret_vault_factory):
"""Every unreachable keychain used to be described as a locked one needing the keyring
package installed. Someone who set the kill switch has the package and an unlocked keychain,
so that advice sends them to fix two things that were never wrong."""
@ -1331,9 +1316,7 @@ class TestKeychainBackedCommands:
assert DISABLE_KEYRING_ENV_VAR in result.output
assert "pip install" not in result.output
def test_print_token_points_an_install_without_keyring_at_the_package(
self, isolated_home, secret_vault_factory
):
def test_print_token_points_an_install_without_keyring_at_the_package(self, isolated_home, secret_vault_factory):
_write_token_file(isolated_home, key=None)
vault = secret_vault_factory(available=False, failure=KeyringNotInstalled())
obj = {"base_url": "https://test.example.com", "secret_vault": vault}
@ -1399,11 +1382,22 @@ class TestLoginConfigClaude:
def setup_method(self):
self.runner = CliRunner()
def _run_login(self, tmp_path, monkeypatch, args, base_url="https://test.example.com"):
settings_path = tmp_path / "claude" / "settings.json"
monkeypatch.setattr(auth_module, "CLAUDE_SETTINGS_PATH", settings_path)
monkeypatch.setattr(auth_module, "CONFIGURE_STATE_PATH", tmp_path / "claude_configure_state.json")
def _isolate_default_settings(self, tmp_path, monkeypatch):
"""The default file, its `lite up` backup and its configure receipt all live under tmp_path."""
backup_path = tmp_path / "claude_settings_backup.json"
monkeypatch.setattr(
claude_settings_module, "SETTINGS_FILE_OWNERS", (SettingsFileOwner(backup_path, "lite up", "lite down"),)
)
monkeypatch.setattr(
claude_settings_module, "CLAUDE_SETTINGS_PATH", tmp_path / "default-home" / ".claude" / "settings.json"
)
monkeypatch.setattr(claude_settings_module, "CONFIGURE_STATE_PATH", tmp_path / "claude_configure_state.json")
return backup_path
def _run_login(self, tmp_path, monkeypatch, args, base_url="https://test.example.com", *, config_dir_env=None):
settings_path = tmp_path / "claude" / "settings.json"
backup_path = self._isolate_default_settings(tmp_path, monkeypatch)
env = {"CLAUDE_CONFIG_DIR": str(settings_path.parent)} if config_dir_env is None else config_dir_env
poll_response = Mock()
poll_response.status_code = 200
poll_response.json.return_value = {
@ -1419,16 +1413,12 @@ class TestLoginConfigClaude:
patch("requests.get", return_value=poll_response),
patch("litellm.proxy.client.cli.commands.auth.save_cli_token"),
patch("litellm.proxy.client.cli.interface.show_commands"),
patch(
"litellm.proxy.client.cli.commands.auth.SETTINGS_FILE_OWNERS",
(SettingsFileOwner(backup_path, "lite up", "lite down"),),
),
patch(
"litellm.proxy.client.cli.commands.claude_settings.shutil.which",
return_value="/usr/local/bin/lite",
),
):
result = self.runner.invoke(login, args, obj={"base_url": base_url})
result = self.runner.invoke(login, args, obj={"base_url": base_url}, env=env)
return result, settings_path, backup_path
def test_default_login_does_not_touch_claude_settings(self, tmp_path, monkeypatch):
@ -1447,7 +1437,7 @@ class TestLoginConfigClaude:
assert written["env"]["ANTHROPIC_BASE_URL"] == "https://test.example.com"
assert written["env"]["ENABLE_TOOL_SEARCH"] == "true"
assert written["apiKeyHelper"] == "/usr/local/bin/lite --base-url https://test.example.com auth print-token"
assert "Configured Claude Code" in result.output
assert f"Configured Claude Code: {settings_path} now routes through https://test.example.com." in result.output
assert "pins a proxy model for every tier" not in result.output
assert "the model Claude Code starts on" in result.output
@ -1463,21 +1453,55 @@ class TestLoginConfigClaude:
assert written["theme"] == "dark"
assert written["env"]["KEEP"] == "me"
def test_refuses_before_logging_in_while_lite_up_holds_the_settings(self, tmp_path, monkeypatch):
# The local precondition comes first: no browser, no token stored, no "Login successful!".
backup_path = tmp_path / "claude_settings_backup.json"
backup_path.write_text("{}")
monkeypatch.setattr(auth_module, "CLAUDE_SETTINGS_PATH", tmp_path / "claude" / "settings.json")
monkeypatch.setattr(
auth_module, "SETTINGS_FILE_OWNERS", (SettingsFileOwner(backup_path, "lite up", "lite down"),)
)
def _run_login_refused_before_the_sso_flow(self, tmp_path, monkeypatch, config_dir):
self._isolate_default_settings(tmp_path, monkeypatch).write_text("{}")
with patch("requests.post") as post, patch("webbrowser.open") as browser:
result = self.runner.invoke(login, ["--config-claude"], obj={"base_url": "https://test.example.com"})
result = self.runner.invoke(
login,
["--config-claude"],
obj={"base_url": "https://test.example.com"},
env={"CLAUDE_CONFIG_DIR": config_dir},
)
assert result.exit_code != 0
assert "not logging in" in result.output and "lite down" in result.output
assert "`lite up` is currently managing" in result.output
assert "Login successful!" not in result.output
post.assert_not_called()
browser.assert_not_called()
assert not (tmp_path / "default-home" / ".claude" / "settings.json").exists()
def test_refuses_before_logging_in_while_lite_up_holds_the_default_settings_file(self, tmp_path, monkeypatch):
self._run_login_refused_before_the_sso_flow(tmp_path, monkeypatch, config_dir="")
def test_refuses_before_logging_in_while_lite_up_holds_the_default_file_reached_through_a_symlink(
self, tmp_path, monkeypatch
):
default_config_dir = tmp_path / "default-home" / ".claude"
default_config_dir.mkdir(parents=True)
alias = tmp_path / "claude-alias"
alias.symlink_to(default_config_dir, target_is_directory=True)
self._run_login_refused_before_the_sso_flow(tmp_path, monkeypatch, config_dir=str(alias))
def test_flag_writes_an_alternate_config_dir_even_while_lite_up_holds_the_default_file(self, tmp_path, monkeypatch):
(tmp_path / "claude_settings_backup.json").write_text("{}")
result, settings_path, _backup_path = self._run_login(tmp_path, monkeypatch, ["--config-claude"])
assert result.exit_code == 0, result.output
written = json.loads(settings_path.read_text())
assert written["apiKeyHelper"] == "/usr/local/bin/lite --base-url https://test.example.com auth print-token"
assert f"Configured Claude Code: {settings_path} now routes through https://test.example.com." in result.output
def test_flag_keeps_a_config_dir_receipt_apart_from_the_default_file_receipt(self, tmp_path, monkeypatch):
result, settings_path, _backup_path = self._run_login(tmp_path, monkeypatch, ["--config-claude"])
assert result.exit_code == 0, result.output
default_receipt = tmp_path / "claude_configure_state.json"
assert not default_receipt.exists()
receipts = list((tmp_path / "claude_configure_state").glob("*.json"))
assert len(receipts) == 1
assert json.loads(receipts[0].read_text())["file_existed"] is False
def test_settings_failure_is_reported_without_claiming_login_failed(self, tmp_path, monkeypatch):
settings_path = tmp_path / "claude" / "settings.json"
@ -1873,7 +1897,10 @@ class TestPkcePrintToken:
assert result.stdout == ""
assert sum(len(session.posts) for session in _FakeSession.instances) == 1
assert result.output.count("Could not renew the key") == 1
assert "Could not renew the key: token request failed with 400: the refresh token was already used" in result.output
assert (
"Could not renew the key: token request failed with 400: the refresh token was already used"
in result.output
)
assert "Key expired. Run 'lite login --pkce' again." in result.output
save.assert_not_called()

View file

@ -3,6 +3,7 @@ import os
import shlex
import stat
import time
from pathlib import Path
from unittest.mock import patch
import pytest
@ -15,6 +16,8 @@ from litellm.proxy.client.cli.commands.claude_settings import (
ANTHROPIC_DEFAULT_MODEL_ENV_KEYS,
AUTOROUTE_BACKUP_PATH,
BACKUP_PATH,
CLAUDE_SETTINGS_PATH,
CONFIGURE_STATE_PATH,
OWNED_ENV_KEYS,
OWNED_TOP_LEVEL_KEYS,
SETTINGS_FILE_OWNERS,
@ -25,7 +28,10 @@ from litellm.proxy.client.cli.commands.claude_settings import (
StartOn,
StaticToken,
UnpinModel,
claude_settings_path,
configure_claude_settings,
configure_state_path,
lite_api_key_helper_configured,
merge_claude_settings,
resolve_api_key_helper,
unconfigure_claude_settings,
@ -390,6 +396,117 @@ class TestDoesNotDestroyUserOwnedStructure:
assert json.loads(settings_path.read_text())["env"] == "not-an-object"
class TestClaudeSettingsPath:
def test_defaults_to_the_home_settings_file(self):
assert claude_settings_path({}) == CLAUDE_SETTINGS_PATH
assert claude_settings_path({"CLAUDE_CONFIG_DIR": ""}) == CLAUDE_SETTINGS_PATH
def test_follows_claude_config_dir_like_claude_code_does(self, tmp_path):
assert claude_settings_path({"CLAUDE_CONFIG_DIR": str(tmp_path)}) == tmp_path / "settings.json"
def test_expands_a_tilde_in_claude_config_dir(self):
assert claude_settings_path({"CLAUDE_CONFIG_DIR": "~/.claude-work"}) == (
Path.home() / ".claude-work" / "settings.json"
)
class TestConfigureStatePath:
"""Each settings file gets its own undo receipt: the default file keeps the long-standing path, and
a CLAUDE_CONFIG_DIR file gets one keyed by its resolved location, so `lite unconfigure claude`
under one config dir never restores the other file's history."""
@pytest.fixture
def default_paths(self, tmp_path):
default_settings = tmp_path / "home" / ".claude" / "settings.json"
default_state = tmp_path / "home" / ".litellm" / "claude_configure_state.json"
with (
patch(f"{CLAUDE_SETTINGS_MODULE}.CLAUDE_SETTINGS_PATH", default_settings),
patch(f"{CLAUDE_SETTINGS_MODULE}.CONFIGURE_STATE_PATH", default_state),
):
yield default_settings, default_state
def test_the_default_file_keeps_the_default_receipt(self, default_paths):
default_settings, default_state = default_paths
assert configure_state_path(default_settings) == default_state
def test_a_symlink_alias_of_the_default_file_shares_its_receipt(self, default_paths):
default_settings, default_state = default_paths
default_settings.parent.mkdir(parents=True)
alias = default_settings.parent.parent / "claude-alias"
alias.symlink_to(default_settings.parent, target_is_directory=True)
assert configure_state_path(alias / "settings.json") == default_state
def test_another_settings_file_gets_a_receipt_of_its_own_beside_the_default_one(self, default_paths, tmp_path):
_default_settings, default_state = default_paths
work_state = configure_state_path(tmp_path / "work" / "settings.json")
play_state = configure_state_path(tmp_path / "play" / "settings.json")
assert work_state != default_state and play_state != default_state
assert work_state != play_state
assert work_state.parent == play_state.parent == default_state.parent / "claude_configure_state"
assert work_state == configure_state_path(tmp_path / "work" / "settings.json")
def test_configure_and_unconfigure_under_a_config_dir_leave_the_default_receipt_alone(
self, default_paths, tmp_path, lite_on_path
):
_default_settings, default_state = default_paths
work_settings = tmp_path / "work" / "settings.json"
work_state = configure_state_path(work_settings)
configure_claude_settings(
"https://proxy.example.com",
ApiKeyHelper(resolve_api_key_helper("https://proxy.example.com")),
KeepModel(),
work_settings,
work_state,
(),
)
assert work_state.exists() and not default_state.exists()
outcome = unconfigure_claude_settings(work_settings, work_state, ())
assert outcome.file_removed and not work_settings.exists()
assert not work_state.exists()
class TestLiteApiKeyHelperConfigured:
def _settings(self, tmp_path, payload):
settings_path = tmp_path / "settings.json"
settings_path.write_text(payload)
return settings_path
def test_recognises_the_helper_lite_login_wrote_for_this_proxy(self, tmp_path, lite_on_path):
settings_path = tmp_path / "settings.json"
_helper_configure("https://proxy.example.com/", settings_path, (), tmp_path / "state.json")
assert lite_api_key_helper_configured("https://proxy.example.com/", settings_path) is True
assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is True
def test_a_helper_for_another_proxy_does_not_count(self, tmp_path, lite_on_path):
settings_path = tmp_path / "settings.json"
_helper_configure("https://other.example.com", settings_path, (), tmp_path / "state.json")
assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False
def test_a_hand_written_helper_does_not_count(self, tmp_path, lite_on_path):
settings_path = self._settings(tmp_path, json.dumps({"apiKeyHelper": "cat ~/.my-proxy-key"}))
assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False
def test_missing_or_helperless_settings_do_not_count(self, tmp_path, lite_on_path):
assert lite_api_key_helper_configured("https://proxy.example.com", tmp_path / "absent.json") is False
helperless = json.dumps({"env": {"ANTHROPIC_BASE_URL": "https://proxy.example.com"}})
settings_path = self._settings(tmp_path, helperless)
assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False
def test_unreadable_settings_fall_back_to_false(self, tmp_path, lite_on_path):
settings_path = self._settings(tmp_path, "{not json")
assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False
def test_lite_missing_from_path_falls_back_to_false(self, tmp_path):
helper = "/usr/local/bin/lite --base-url https://proxy.example.com auth print-token"
settings_path = self._settings(tmp_path, json.dumps({"apiKeyHelper": helper}))
with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value=None):
assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False
class TestMergeClaudeSettings:
"""One merge for every way Claude Code gets wired: `lite up`, `lite login --config-claude`,
`lite configure claude` and `lite autoroute up`."""

View file

@ -9,6 +9,7 @@ import responses
from click.testing import CliRunner
from litellm.proxy.client.cli import cli
from litellm.proxy.client.cli.commands import claude_settings as claude_settings_module
from litellm.proxy.client.cli.commands import configure as configure_module
from litellm.proxy.client.cli.commands.claude_settings import SettingsFileOwner
from litellm.proxy.client.cli.commands.configure import configure_claude, configure_group, interactive_configure
@ -29,10 +30,12 @@ def _mock_models():
@pytest.fixture
def paths(monkeypatch, tmp_path):
"""The default settings file, reached the way Claude Code reaches it: CLAUDE_CONFIG_DIR names its directory."""
settings_path = tmp_path / "claude" / "settings.json"
state_path = tmp_path / "litellm" / "claude_configure_state.json"
monkeypatch.setattr(configure_module, "CLAUDE_SETTINGS_PATH", settings_path)
monkeypatch.setattr(configure_module, "CONFIGURE_STATE_PATH", state_path)
monkeypatch.setenv("CLAUDE_CONFIG_DIR", str(settings_path.parent))
monkeypatch.setattr(claude_settings_module, "CLAUDE_SETTINGS_PATH", settings_path)
monkeypatch.setattr(claude_settings_module, "CONFIGURE_STATE_PATH", state_path)
return settings_path, state_path
@ -58,7 +61,9 @@ def lite_up_backup(monkeypatch, tmp_path):
"""A `lite up` session holding its backup, the local precondition every settings write refuses on."""
backup = tmp_path / "claude_settings_backup.json"
backup.write_text("{}")
monkeypatch.setattr(configure_module, "SETTINGS_FILE_OWNERS", (SettingsFileOwner(backup, "lite up", "lite down"),))
monkeypatch.setattr(
claude_settings_module, "SETTINGS_FILE_OWNERS", (SettingsFileOwner(backup, "lite up", "lite down"),)
)
return backup
@ -325,6 +330,30 @@ class TestUnconfigureClaude:
result = runner.invoke(cli, ["unconfigure", "claude"])
assert result.exit_code != 0 and "lite down" in result.output
@responses.activate
def test_a_config_dir_is_configured_and_undone_apart_from_the_default_file(
self, runner, paths, monkeypatch, tmp_path, lite_up_backup
):
_mock_models()
default_settings, default_state = paths
work_dir = tmp_path / "claude-work"
work_dir.mkdir()
original = {"theme": "dark"}
(work_dir / "settings.json").write_text(json.dumps(original))
monkeypatch.setenv("CLAUDE_CONFIG_DIR", str(work_dir))
configured = _configure(runner, "--api-key", VALID_KEY, "--model", "claude-auto")
assert configured.exit_code == 0, configured.output
assert f"Configured Claude Code: {work_dir / 'settings.json'}" in configured.output
assert json.loads((work_dir / "settings.json").read_text())["env"]["ANTHROPIC_AUTH_TOKEN"] == VALID_KEY
assert not default_settings.exists() and not default_state.exists()
undone = runner.invoke(cli, ["unconfigure", "claude"])
assert undone.exit_code == 0, undone.output
assert json.loads((work_dir / "settings.json").read_text()) == original
assert not default_settings.exists() and not default_state.exists()
assert runner.invoke(cli, ["unconfigure", "claude"]).exit_code != 0, "the receipt is gone with the undo"
def test_without_a_receipt_it_fails_loudly(self, runner, paths):
result = runner.invoke(cli, ["unconfigure", "claude"])
assert result.exit_code != 0

View file

@ -620,10 +620,7 @@ class TestUpCanInvokeTheRealLoginCommand:
ctx.obj = {"base_url": "http://127.0.0.1:9"}
ctx.invoke(real_login, pkce=False)
with (
patch(f"{AUTH_MODULE}.CLAUDE_SETTINGS_PATH", settings_path),
patch(f"{AUTH_MODULE}._start_cli_sso_flow", side_effect=RuntimeError("stop")),
):
CliRunner().invoke(driver, [], standalone_mode=False)
with patch(f"{AUTH_MODULE}._start_cli_sso_flow", side_effect=RuntimeError("stop")):
CliRunner().invoke(driver, [], standalone_mode=False, env={"CLAUDE_CONFIG_DIR": str(tmp_path)})
assert not settings_path.exists()

View file

@ -3,6 +3,8 @@ Unit tests for Bedrock Guardrails
"""
import json
import asyncio
from datetime import datetime, timezone
import sys
from unittest.mock import AsyncMock, MagicMock, patch
@ -28,6 +30,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockTextContent,
)
from litellm.types.utils import CallTypes, ModelResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
@pytest.mark.asyncio
@ -5842,3 +5845,36 @@ async def test_bearer_token_never_runs_the_sigv4_credential_chain(monkeypatch):
assert response["action"] == "NONE"
assert mock_post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"
@pytest.mark.asyncio
async def test_apply_guardrail_signs_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: the ApplyGuardrail request is signed with SigV4, and botocore
refreshes expiring credentials inside that signing with a blocking HTTP call, so it must run
on a worker thread to keep the loop serving other requests."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
guardrail = BedrockGuardrail(guardrailIdentifier="test-guardrail", guardrailVersion="DRAFT")
probe = EventLoopProbe()
allowed = httpx.Response(
200,
json={"action": "NONE", "outputs": [], "assessments": []},
request=httpx.Request("POST", "https://bedrock-runtime.us-east-1.amazonaws.com"),
)
with patch.object(guardrail.async_handler, "post", new=AsyncMock(return_value=allowed)):
release = asyncio.create_task(probe.release_refresh_from_the_loop())
response = await guardrail._post_apply_guardrail_content(
content=[{"text": {"text": "hello"}}],
base_request_data={"source": "INPUT"},
credentials=probe.credentials(),
aws_region_name="us-east-1",
api_key=None,
request_data={},
event_type=GuardrailEventHooks.pre_call,
start_time=datetime.now(timezone.utc),
completed_chunk_usages=[],
)
await release
assert response["action"] == "NONE"
assert probe.served_during_refresh is True

View file

@ -5,10 +5,12 @@ All Bedrock HTTP calls are mocked; no real AWS calls are made.
"""
import json
import asyncio
import logging
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import httpx
from fastapi import HTTPException
@ -21,6 +23,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.bedrock_guardrails import (
BedrockGuardrailResponse,
)
from litellm.types.utils import Choices, Message, ModelResponse
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
CONTENT_FILTER_CHECKS = {"contentFilter": {"categories": [{"category": "VIOLENCE"}]}}
@ -861,3 +864,33 @@ async def test_checks_bearer_token_never_runs_the_sigv4_credential_chain(monkeyp
{"check": "contentFilter", "category": "VIOLENCE", "severityScore": 0.8}
]
assert post.call_args.kwargs["headers"]["Authorization"] == "Bearer env-bearer-token-12345"
@pytest.mark.asyncio
async def test_invoke_guardrail_checks_signs_off_the_event_loop(monkeypatch):
"""Regression for issue #40165: the checks request is signed with SigV4, and botocore refreshes
expiring credentials inside that signing with a blocking HTTP call, so it must run on a worker
thread to keep the loop serving other requests."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
g = BedrockGuardrail(checks=CONTENT_FILTER_CHECKS, content_filter_threshold=0.5)
probe = EventLoopProbe()
allowed = httpx.Response(
200,
json={"results": {"contentFilter": {"results": [{"category": "VIOLENCE", "severityScore": 0.1}]}}},
request=httpx.Request("POST", "https://bedrock-runtime.us-east-1.amazonaws.com"),
)
with (
patch.object(g, "_load_credentials", return_value=(probe.credentials(), "us-east-1")),
patch.object(g.async_handler, "post", new=AsyncMock(return_value=allowed)),
):
release = asyncio.create_task(probe.release_refresh_from_the_loop())
response = await g.make_bedrock_api_request(
source="INPUT",
messages=[{"role": "user", "content": "hello"}],
request_data={"messages": []},
)
await release
assert response == BedrockGuardrailResponse()
assert probe.served_during_refresh is True

View file

@ -1,3 +1,4 @@
import asyncio
import base64
import contextlib
import json
@ -19,6 +20,7 @@ from starlette.datastructures import FormData
import litellm
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe
from litellm.constants import LITELLM_PROXY_MASTER_KEY_ALIAS
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
BaseOpenAIPassThroughHandler,
@ -1852,11 +1854,11 @@ class TestBedrockAgentRuntimePassthroughToggle:
return request
@contextlib.contextmanager
def _patched_dispatch(self, general_settings: Mapping[str, object]):
def _patched_dispatch(self, general_settings: Mapping[str, object], credentials: object | None = None):
from botocore.credentials import Credentials
bedrock_llm: Final = Mock()
bedrock_llm.get_credentials = Mock(return_value=Credentials("ak", "sk"))
bedrock_llm.get_credentials = Mock(return_value=credentials or Credentials("ak", "sk"))
forwarder: Final = AsyncMock(return_value="forwarded")
with (
@ -1891,6 +1893,27 @@ class TestBedrockAgentRuntimePassthroughToggle:
forwarder.assert_awaited_once()
assert "bedrock-agent-runtime.us-east-1.amazonaws.com" in create_route.call_args.kwargs["target"]
@pytest.mark.asyncio
async def test_agent_runtime_dispatch_signs_off_the_event_loop(self, monkeypatch):
"""Regression for issue #40165: the agent-runtime pass-through signed on the loop, so botocore's
blocking credential refresh inside SigV4 stalled every other request on the worker."""
monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False)
probe: Final = EventLoopProbe()
release: Final = asyncio.create_task(probe.release_refresh_from_the_loop())
with self._patched_dispatch(MappingProxyType({}), credentials=probe.credentials()) as (create_route, forwarder):
result: Final = await bedrock_proxy_route(
endpoint=self.AGENT_RUNTIME_ENDPOINT,
request=self._mock_request(),
fastapi_response=Mock(),
user_api_key_dict=UserAPIKeyAuth(),
)
await release
assert result == "forwarded"
assert create_route.call_args.kwargs["custom_headers"]["Authorization"].startswith("AWS4-HMAC-SHA256")
assert probe.served_during_refresh is True
@pytest.mark.asyncio
@pytest.mark.parametrize("value", (True, "true", "True"))
async def test_agent_runtime_dispatch_rejected_when_disabled(self, value: bool | str):