diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs index 7dbb0f8b1ff..1230fa74641 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs @@ -29,7 +29,15 @@ pub(super) fn prepare_messages_call( let config = messages_provider_config(provider) .ok_or_else(|| CoreError::InvalidProvider(provider.to_string()))?; - let env_lookup = |key: &str| std::env::var(key).ok(); + let env_lookup = |key: &str| { + if key == "AWS_REGION_NAME" { + return request + .aws_region_name + .map(str::to_string) + .or_else(|| std::env::var(key).ok()); + } + std::env::var(key).ok() + }; let mut headers = string_headers(request.extra_headers)?; diff --git a/litellm-rust/crates/ai-gateway/src/messages/tests.rs b/litellm-rust/crates/ai-gateway/src/messages/tests.rs index 23a53e98045..53a884ba5ad 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/tests.rs @@ -148,6 +148,7 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through() custom_llm_provider: Some("azure_ai"), extra_headers: None, timeout: Some(Duration::from_secs(5)), + aws_region_name: None, }) .await .expect("messages request succeeds"); @@ -204,6 +205,7 @@ async fn messages_round_trip_builds_native_anthropic_request() { custom_llm_provider: Some("anthropic"), extra_headers: None, timeout: Some(Duration::from_secs(5)), + aws_region_name: None, }) .await .expect("messages request succeeds"); @@ -257,6 +259,7 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() { custom_llm_provider: Some("azure_ai"), extra_headers: Some(headers), timeout: Some(Duration::from_secs(5)), + aws_region_name: None, }) .await .expect("messages request succeeds"); @@ -311,6 +314,7 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() { custom_llm_provider: Some("azure_ai"), extra_headers: Some(headers), timeout: Some(Duration::from_secs(5)), + aws_region_name: None, }) .await .expect("entra id request succeeds without api key"); @@ -335,6 +339,7 @@ async fn messages_requires_auth_when_no_key_and_no_header() { custom_llm_provider: Some("azure_ai"), extra_headers: None, timeout: Some(Duration::from_millis(50)), + aws_region_name: None, }) .await .expect_err("missing auth errors"); @@ -373,6 +378,7 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() { custom_llm_provider: Some("azure_ai"), extra_headers: Some(headers), timeout: Some(Duration::from_secs(5)), + aws_region_name: None, }) .await .expect("falls back to api key"); @@ -414,6 +420,7 @@ async fn messages_maps_provider_error_status_to_http_error() { custom_llm_provider: Some("azure_ai"), extra_headers: None, timeout: Some(Duration::from_secs(5)), + aws_region_name: None, }) .await .expect_err("provider error propagates"); @@ -431,6 +438,7 @@ async fn messages_rejects_unsupported_provider() { custom_llm_provider: Some("openai"), extra_headers: None, timeout: Some(Duration::from_millis(50)), + aws_region_name: None, }) .await .expect_err("unsupported provider errors"); diff --git a/litellm-rust/crates/ai-gateway/src/messages/types.rs b/litellm-rust/crates/ai-gateway/src/messages/types.rs index 2215d269e62..18fb393f7a4 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/types.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/types.rs @@ -11,6 +11,7 @@ pub struct MessagesRequest<'a> { pub custom_llm_provider: Option<&'a str>, pub extra_headers: Option>, pub timeout: Option, + pub aws_region_name: Option<&'a str>, } pub(crate) struct ProviderMessagesRequest { diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index 84812596008..f58534a4908 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -282,6 +282,7 @@ mod tests { model: format!("anthropic/{provider_model}"), api_key: Some("upstream-key".to_string()), api_base: Some(api_base), + aws_region_name: None, }, }])), master_key: master_key.map(Arc::from), @@ -304,6 +305,7 @@ mod tests { model: format!("bedrock/{provider_model}"), api_key: api_key.map(str::to_string), api_base: Some(api_base), + aws_region_name: None, }, }])), master_key: master_key.map(Arc::from), diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index 30ce3357a63..824166459df 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -54,6 +54,7 @@ pub async fn run( custom_llm_provider, extra_headers, timeout: None, + aws_region_name: deployment.litellm_params.aws_region_name.as_deref(), }; let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true); execute_messages(request, stream) diff --git a/litellm-rust/crates/core/src/router/deployment.rs b/litellm-rust/crates/core/src/router/deployment.rs index 1ee88e682a3..d9c69f78669 100644 --- a/litellm-rust/crates/core/src/router/deployment.rs +++ b/litellm-rust/crates/core/src/router/deployment.rs @@ -12,6 +12,8 @@ pub struct LiteLLMParams { pub api_key: Option, #[serde(default)] pub api_base: Option, + #[serde(default)] + pub aws_region_name: Option, } /// One entry of the `model_list`, mirroring Python's deployment dict. diff --git a/litellm-rust/crates/core/src/router/mod.rs b/litellm-rust/crates/core/src/router/mod.rs index 96bc91bc6b5..e526217453b 100644 --- a/litellm-rust/crates/core/src/router/mod.rs +++ b/litellm-rust/crates/core/src/router/mod.rs @@ -69,6 +69,7 @@ mod tests { model: model.to_string(), api_key: None, api_base: None, + aws_region_name: None, }, } } diff --git a/litellm-rust/crates/core/src/router/strategy/simple_shuffle.rs b/litellm-rust/crates/core/src/router/strategy/simple_shuffle.rs index 74ce0c21e80..5c46bed0c7a 100644 --- a/litellm-rust/crates/core/src/router/strategy/simple_shuffle.rs +++ b/litellm-rust/crates/core/src/router/strategy/simple_shuffle.rs @@ -22,6 +22,7 @@ mod tests { model: model.to_string(), api_key: None, api_base: None, + aws_region_name: None, }, } } diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index ee9bdd0b81f..f8747a60cf2 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -354,7 +354,7 @@ fn marshal_messages_inputs( } #[pyfunction] -#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] +#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, aws_region_name=None))] #[allow(clippy::too_many_arguments)] fn messages( py: Python<'_>, @@ -365,6 +365,7 @@ fn messages( custom_llm_provider: Option, extra_headers: Option>, timeout_seconds: Option, + aws_region_name: Option, ) -> PyResult> { let (body, extra_headers, timeout) = marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?; @@ -378,6 +379,7 @@ fn messages( custom_llm_provider: custom_llm_provider.as_deref(), extra_headers, timeout, + aws_region_name: aws_region_name.as_deref(), })) }); @@ -388,7 +390,7 @@ fn messages( } #[pyfunction] -#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] +#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None, aws_region_name=None))] #[allow(clippy::too_many_arguments)] fn amessages( py: Python<'_>, @@ -399,6 +401,7 @@ fn amessages( custom_llm_provider: Option, extra_headers: Option>, timeout_seconds: Option, + aws_region_name: Option, ) -> PyResult> { let (body, extra_headers, timeout) = marshal_messages_inputs(py, body, extra_headers, timeout_seconds)?; @@ -412,6 +415,7 @@ fn amessages( custom_llm_provider: custom_llm_provider.as_deref(), extra_headers, timeout, + aws_region_name: aws_region_name.as_deref(), }) .await .map_err(core_error_to_pyerr)?; diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index ec1301e5923..019d9b51372 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -2273,7 +2273,7 @@ class BaseLLMHTTPHandler: request_body: dict, timeout: float | httpx.Timeout | None, ) -> AnthropicMessagesResponse | None: - if custom_llm_provider not in ("azure_ai", "anthropic"): + if custom_llm_provider not in ("azure_ai", "anthropic", "bedrock"): return None if litellm_params.get("rust") is not True and not BaseLLMHTTPHandler._rust_env_enabled(): return None @@ -2283,6 +2283,7 @@ class BaseLLMHTTPHandler: from litellm.rust_bridge import messages as rust_messages_bridge upstream_body = {key: value for key, value in request_body.items() if key != "stream"} + upstream_body.setdefault("model", model) try: rust_response = await rust_messages_bridge.amessages( model=model, @@ -2292,6 +2293,7 @@ class BaseLLMHTTPHandler: custom_llm_provider=custom_llm_provider, extra_headers=headers, timeout=timeout, + aws_region_name=litellm_params.get("aws_region_name"), ) except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path verbose_logger.debug( diff --git a/litellm/rust_bridge/messages.py b/litellm/rust_bridge/messages.py index 5abb21879d3..c87842acc57 100644 --- a/litellm/rust_bridge/messages.py +++ b/litellm/rust_bridge/messages.py @@ -20,6 +20,7 @@ class RustMessages(Protocol): custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout_seconds: float | None, + aws_region_name: str | None = None, ) -> dict[str, object]: raise NotImplementedError @@ -34,6 +35,7 @@ class RustAmessages(Protocol): custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout_seconds: float | None, + aws_region_name: str | None = None, ) -> Awaitable[dict[str, object]]: raise NotImplementedError @@ -96,6 +98,7 @@ def messages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: Union[float, httpx.Timeout] | None, + aws_region_name: str | None = None, ) -> dict[str, object] | None: rust_messages = load_rust_messages() if rust_messages is None: @@ -108,6 +111,7 @@ def messages( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), + **({"aws_region_name": aws_region_name} if aws_region_name is not None else {}), ) @@ -120,6 +124,7 @@ async def amessages( custom_llm_provider: str | None, extra_headers: dict[str, object] | None, timeout: Union[float, httpx.Timeout] | None, + aws_region_name: str | None = None, ) -> dict[str, object] | None: rust_amessages = load_rust_amessages() if rust_amessages is None: @@ -132,4 +137,5 @@ async def amessages( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, timeout_seconds=timeout_to_seconds(timeout), + **({"aws_region_name": aws_region_name} if aws_region_name is not None else {}), ) diff --git a/tests/test_litellm/anthropic_interface/test_bedrock_rust_bridge_e2e.py b/tests/test_litellm/anthropic_interface/test_bedrock_rust_bridge_e2e.py new file mode 100644 index 00000000000..de4b585a2ca --- /dev/null +++ b/tests/test_litellm/anthropic_interface/test_bedrock_rust_bridge_e2e.py @@ -0,0 +1,52 @@ +import os + +import pytest + +import litellm + + +@pytest.fixture(autouse=True) +def isolate_host_aws_config(monkeypatch, isolated_aws_credentials_dir): + monkeypatch.setenv("AWS_SHARED_CREDENTIALS_FILE", isolated_aws_credentials_dir["credentials"]) + monkeypatch.setenv("AWS_CONFIG_FILE", isolated_aws_credentials_dir["config"]) + monkeypatch.setenv("AWS_EC2_METADATA_DISABLED", "true") + monkeypatch.delenv("AWS_PROFILE", raising=False) + monkeypatch.delenv("AWS_DEFAULT_PROFILE", raising=False) + monkeypatch.delenv("AWS_CONTAINER_CREDENTIALS_FULL_URI", raising=False) + monkeypatch.delenv("AWS_CONTAINER_CREDENTIALS_RELATIVE_URI", raising=False) + monkeypatch.delenv("AWS_SESSION_TOKEN", raising=False) + monkeypatch.delenv("AWS_ROLE_ARN", raising=False) + monkeypatch.delenv("AWS_WEB_IDENTITY_TOKEN_FILE", raising=False) + + +@pytest.mark.asyncio +async def test_bedrock_messages_with_rust() -> None: + response = await litellm.anthropic.messages.acreate( + model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", + messages=[{"role": "user", "content": "Say hello"}], + max_tokens=20, + api_key=os.getenv("AWS_BEARER_TOKEN_BEDROCK"), + aws_region_name="us-west-2", + rust=True, + ) + + assert response["content"][0]["text"] + assert response["_hidden_params"]["additional_headers"]["x-litellm-rust"] == "true" + + +@pytest.mark.asyncio +async def test_bedrock_messages_streaming_with_rust() -> None: + response = await litellm.anthropic.messages.acreate( + model="bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0", + messages=[{"role": "user", "content": "Say hello"}], + max_tokens=20, + stream=True, + api_key=os.getenv("AWS_BEARER_TOKEN_BEDROCK"), + aws_region_name="us-west-2", + rust=True, + ) + + chunks = [chunk async for chunk in response] + + assert chunks + assert response._hidden_params["additional_headers"]["x-litellm-rust"] == "true"