feat(rust): route Bedrock messages through Python bridge

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-29 03:51:14 +00:00
parent 89e951e04b
commit f13425bb4f
12 changed files with 92 additions and 4 deletions

View file

@ -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)?;

View file

@ -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");

View file

@ -11,6 +11,7 @@ pub struct MessagesRequest<'a> {
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
pub aws_region_name: Option<&'a str>,
}
pub(crate) struct ProviderMessagesRequest {

View file

@ -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),

View file

@ -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)

View file

@ -12,6 +12,8 @@ pub struct LiteLLMParams {
pub api_key: Option<String>,
#[serde(default)]
pub api_base: Option<String>,
#[serde(default)]
pub aws_region_name: Option<String>,
}
/// One entry of the `model_list`, mirroring Python's deployment dict.

View file

@ -69,6 +69,7 @@ mod tests {
model: model.to_string(),
api_key: None,
api_base: None,
aws_region_name: None,
},
}
}

View file

@ -22,6 +22,7 @@ mod tests {
model: model.to_string(),
api_key: None,
api_base: None,
aws_region_name: None,
},
}
}

View file

@ -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<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
aws_region_name: Option<String>,
) -> PyResult<Py<PyAny>> {
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<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
aws_region_name: Option<String>,
) -> PyResult<Bound<'_, PyAny>> {
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)?;

View file

@ -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(

View file

@ -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 {}),
)

View file

@ -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"