From 32a1d55e56e896c1f440253619180451bfd0a725 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 29 Jul 2026 02:35:33 +0000 Subject: [PATCH] feat(rust): support Bedrock Anthropic messages Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm-rust/Cargo.lock | 22 ++ litellm-rust/Cargo.toml | 2 + litellm-rust/bedrock_messages_harness.sh | 35 +++ litellm-rust/crates/ai-gateway/Cargo.toml | 2 + litellm-rust/crates/ai-gateway/src/main.rs | 11 + .../ai-gateway/src/messages/common_utils.rs | 2 + .../crates/ai-gateway/src/messages/handler.rs | 86 ++++++- .../crates/ai-gateway/src/messages/mod.rs | 11 +- .../crates/ai-gateway/src/messages/prepare.rs | 25 +- .../crates/ai-gateway/src/messages/types.rs | 2 + .../ai-gateway/src/routes/messages/mod.rs | 114 ++++++++- .../ai-gateway/src/routes/messages/service.rs | 9 +- .../core/src/messages/transformation.rs | 11 + .../crates/core/src/messages/types.rs | 2 + .../anthropic/messages/transformation.rs | 1 + .../azure_ai/messages/transformation.rs | 1 + .../src/providers/bedrock/messages/mod.rs | 1 + .../bedrock/messages/transformation.rs | 232 ++++++++++++++++++ .../crates/core/src/providers/bedrock/mod.rs | 1 + 19 files changed, 547 insertions(+), 23 deletions(-) create mode 100755 litellm-rust/bedrock_messages_harness.sh create mode 100644 litellm-rust/crates/core/src/providers/bedrock/messages/mod.rs create mode 100644 litellm-rust/crates/core/src/providers/bedrock/messages/transformation.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index ce28f737334..c3317017d55 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -180,6 +180,17 @@ dependencies = [ "tokio", ] +[[package]] +name = "aws-smithy-eventstream" +version = "0.60.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78d8391e65fcea47c586a22e1a41f173b38615b112b2c6b7a44e80cec3e6b706" +dependencies = [ + "aws-smithy-types", + "bytes", + "crc32fast", +] + [[package]] name = "aws-smithy-http" version = "0.64.0" @@ -596,6 +607,15 @@ dependencies = [ "libc", ] +[[package]] +name = "crc32fast" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9481c1c90cbf2ac953f07c8d4a58aa3945c425b7185c9154d67a65e4230da511" +dependencies = [ + "cfg-if", +] + [[package]] name = "crypto-common" version = "0.1.7" @@ -1216,8 +1236,10 @@ checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" name = "litellm-ai-gateway" version = "0.1.0" dependencies = [ + "aws-smithy-eventstream", "axum", "base64", + "bytes", "futures-channel", "futures-util", "litellm-core", diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 6d63be05d00..db7ad8a6919 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -29,3 +29,5 @@ tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"] tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] } futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] } base64 = "0.22" +bytes = "1" +aws-smithy-eventstream = "0.60.3" diff --git a/litellm-rust/bedrock_messages_harness.sh b/litellm-rust/bedrock_messages_harness.sh new file mode 100755 index 00000000000..04d9e2dba5a --- /dev/null +++ b/litellm-rust/bedrock_messages_harness.sh @@ -0,0 +1,35 @@ +#!/usr/bin/env bash +set -euo pipefail + +export AWS_REGION_NAME="${AWS_REGION_NAME:-us-west-2}" +export BEDROCK_MODEL="${BEDROCK_MODEL:-us.anthropic.claude-sonnet-4-5-20250929-v1:0}" +export LITELLM_MASTER_KEY="${LITELLM_MASTER_KEY:-harness-master-key}" +export PORT="${PORT:-4001}" + +if [[ -z "${AWS_BEARER_TOKEN_BEDROCK:-}" ]]; then + echo "AWS_BEARER_TOKEN_BEDROCK is required" >&2 + exit 1 +fi + +cargo run -p litellm-ai-gateway --features server >/tmp/litellm-bedrock-gateway.log 2>&1 & +gateway_pid=$! +trap 'kill "$gateway_pid" 2>/dev/null || true' EXIT +for _ in {1..60}; do + curl -sf "http://127.0.0.1:${PORT}/health/readiness" >/dev/null && break + sleep 1 +done + +headers=(-H "authorization: Bearer ${LITELLM_MASTER_KEY}" -H "content-type: application/json") +url="http://127.0.0.1:${PORT}/v1/messages" + +echo "simple" +curl -sS "${headers[@]}" "$url" -d "{\"model\":\"${BEDROCK_MODEL}\",\"max_tokens\":32,\"messages\":[{\"role\":\"user\",\"content\":\"Reply with one word: hello\"}]}" | jq '{type,id,model,usage}' + +echo "streaming" +curl -sS "${headers[@]}" "$url" -d "{\"model\":\"${BEDROCK_MODEL}\",\"stream\":true,\"max_tokens\":32,\"messages\":[{\"role\":\"user\",\"content\":\"Reply with one word: hello\"}]}" | grep -E '^(event:|data:)' | head -20 + +echo "tool_use" +curl -sS "${headers[@]}" "$url" -d "{\"model\":\"${BEDROCK_MODEL}\",\"max_tokens\":64,\"tools\":[{\"name\":\"get_weather\",\"description\":\"Get weather\",\"input_schema\":{\"type\":\"object\",\"properties\":{\"city\":{\"type\":\"string\"}},\"required\":[\"city\"]}}],\"tool_choice\":{\"type\":\"tool\",\"name\":\"get_weather\"},\"messages\":[{\"role\":\"user\",\"content\":\"What is the weather in Paris?\"}]}" | jq '{type,model,content}' + +echo "bad-model" +curl -sS -o /tmp/litellm-bedrock-bad-model.json -w 'HTTP %{http_code}\n' "${headers[@]}" "$url" -d '{"model":"us.anthropic.invalid-v1:0","max_tokens":8,"messages":[{"role":"user","content":"hello"}]}' diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index 541beabe170..95444ff86a0 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -24,6 +24,8 @@ tokio-tungstenite.workspace = true futures-util.workspace = true serde_json.workspace = true base64.workspace = true +bytes.workspace = true +aws-smithy-eventstream.workspace = true axum = { workspace = true, features = ["ws"], optional = true } serde.workspace = true subtle = { workspace = true, optional = true } diff --git a/litellm-rust/crates/ai-gateway/src/main.rs b/litellm-rust/crates/ai-gateway/src/main.rs index da3a486d4ee..abfceb0e866 100644 --- a/litellm-rust/crates/ai-gateway/src/main.rs +++ b/litellm-rust/crates/ai-gateway/src/main.rs @@ -142,6 +142,17 @@ fn build_router() -> Router { /// A real deployment loads `model_list` from config; this is the minimal stand-in /// so the gateway has one OpenAI deployment to route to. fn build_router_from_env() -> Router { + if let Ok(model) = std::env::var("BEDROCK_MODEL") { + let api_key = std::env::var("AWS_BEARER_TOKEN_BEDROCK").ok(); + return Router::new(vec![Deployment { + model_name: model.clone(), + litellm_params: LiteLLMParams { + model: format!("bedrock/{model}"), + api_key, + api_base: None, + }, + }]); + } let model = std::env::var("OPENAI_REALTIME_MODEL").unwrap_or_else(|_| "gpt-realtime".to_string()); let api_key = std::env::var("OPENAI_API_KEY").ok(); diff --git a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs b/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs index 68ecc3f17c1..52f8c9e9af2 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs @@ -3,6 +3,7 @@ use litellm_core::error::{CoreError, json_type_name}; use litellm_core::messages::transformation::AnthropicMessagesProviderConfig; use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; +use litellm_core::providers::bedrock::messages::transformation::BEDROCK_MESSAGES_CONFIG; use serde_json::{Map, Value}; use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS; @@ -21,6 +22,7 @@ pub(super) fn messages_provider_config( match provider { "anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG), "azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG), + "bedrock" => Some(&BEDROCK_MESSAGES_CONFIG), _ => None, } } diff --git a/litellm-rust/crates/ai-gateway/src/messages/handler.rs b/litellm-rust/crates/ai-gateway/src/messages/handler.rs index 90c12367f50..f06684a5565 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/handler.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/handler.rs @@ -1,5 +1,12 @@ +use std::collections::BTreeMap; +use std::time::SystemTime; + use litellm_core::CoreResult; use litellm_core::error::CoreError; +use litellm_core::messages::transformation::MessagesAuthStrategy; +use litellm_core::providers::bedrock::aws_base::{ + AwsAuthConfig, resolve_credentials, sign_bedrock_post, +}; use serde_json::Value; use super::client::http_client; @@ -7,11 +14,76 @@ use super::common_utils::truncate_error_body; use super::types::ProviderMessagesRequest; use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; +fn environment_lookup(key: &str) -> Option { + std::env::var(key).ok() +} + +async fn signed_request( + request: &ProviderMessagesRequest, + body: &[u8], +) -> CoreResult> { + if !matches!( + request.config.auth_strategy(), + MessagesAuthStrategy::AwsSigV4 + ) { + return Ok(request.upstream_headers.clone()); + } + if let Some(token) = &request.bearer_token { + return Ok(request + .upstream_headers + .iter() + .filter(|(name, _)| { + !matches!( + name.to_ascii_lowercase().as_str(), + "authorization" | "x-api-key" | "anthropic-version" + ) + }) + .cloned() + .chain([ + ("Authorization".to_string(), format!("Bearer {token}")), + ("content-type".to_string(), "application/json".to_string()), + ]) + .collect()); + } + let headers = request + .upstream_headers + .iter() + .filter(|(name, _)| { + !matches!( + name.to_ascii_lowercase().as_str(), + "authorization" | "x-api-key" | "anthropic-version" | "host" | "content-length" + ) + }) + .cloned() + .chain(std::iter::once(( + "content-type".to_string(), + "application/json".to_string(), + ))) + .collect::>(); + let region = request.signing_region.as_deref().ok_or_else(|| { + CoreError::InvalidRequest("Bedrock signing region was not resolved".to_string()) + })?; + let credentials = resolve_credentials(AwsAuthConfig::default(), &environment_lookup).await?; + let signed = sign_bedrock_post( + &request.url, + body, + &headers, + region, + &credentials, + SystemTime::now(), + )?; + Ok(signed.into_iter().collect()) +} + pub(super) async fn execute_messages_provider_call( request: ProviderMessagesRequest, ) -> CoreResult { - let mut request_builder = http_client().post(&request.url).json(&request.body); - for (key, value) in &request.upstream_headers { + let body = serde_json::to_vec(&request.body).map_err(|error| { + CoreError::InvalidRequest(format!("invalid messages request body: {error}")) + })?; + let headers = signed_request(&request, &body).await?; + let mut request_builder = http_client().post(&request.url).body(body); + for (key, value) in &headers { request_builder = request_builder.header(key, value); } if let Some(duration) = request.timeout { @@ -50,14 +122,18 @@ pub(super) async fn execute_messages_provider_call( pub(super) async fn execute_messages_provider_stream( request: ProviderMessagesRequest, ) -> CoreResult { - if request.provider != ANTHROPIC_MESSAGES_PROVIDER { + if request.provider != ANTHROPIC_MESSAGES_PROVIDER && request.signing_region.is_none() { return Err(CoreError::InvalidRequest( "streaming messages is not supported for this provider".to_string(), )); } - let mut request_builder = http_client().post(&request.url).json(&request.body); - for (key, value) in &request.upstream_headers { + let body = serde_json::to_vec(&request.body).map_err(|error| { + CoreError::InvalidRequest(format!("invalid messages request body: {error}")) + })?; + let headers = signed_request(&request, &body).await?; + let mut request_builder = http_client().post(&request.url).body(body); + for (key, value) in &headers { request_builder = request_builder.header(key, value); } if let Some(duration) = request.timeout { diff --git a/litellm-rust/crates/ai-gateway/src/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/messages/mod.rs index fd2dd546941..c8192b1674a 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/mod.rs @@ -15,7 +15,8 @@ use prepare::prepare_messages_call; pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { match execute_messages(request, false).await? { MessagesResponse::Json(body) => Ok(body), - MessagesResponse::Stream(response) => { + MessagesResponse::Stream { response, provider } => { + drop(provider); drop(response); Err(litellm_core::CoreError::InvalidResponse( "non-streaming messages execution returned a stream".to_string(), @@ -26,7 +27,10 @@ pub async fn messages(request: MessagesRequest<'_>) -> CoreResult { pub(crate) enum MessagesResponse { Json(Value), - Stream(reqwest::Response), + Stream { + provider: String, + response: reqwest::Response, + }, } pub(crate) async fn execute_messages( @@ -35,9 +39,10 @@ pub(crate) async fn execute_messages( ) -> CoreResult { let prepared = prepare_messages_call(request)?; if stream { + let provider = prepared.provider.clone(); execute_messages_provider_stream(prepared) .await - .map(MessagesResponse::Stream) + .map(|response| MessagesResponse::Stream { provider, response }) } else { execute_messages_provider_call(prepared) .await diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs index 9a027490eb6..4f8426cb321 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs @@ -2,6 +2,7 @@ use litellm_core::CoreError; use litellm_core::CoreResult; use litellm_core::messages::transformation::MessagesAuthStrategy; use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; +use serde_json::Value; use super::common_utils::{has_bearer_auth, has_header, messages_provider_config, string_headers}; use super::types::{MessagesRequest, ProviderMessagesRequest}; @@ -33,15 +34,27 @@ pub(super) fn prepare_messages_call( let mut headers = string_headers(request.extra_headers)?; let auth_strategy = config.auth_strategy(); - let already_authorized = has_header(&headers, auth_strategy.header_name()) - || (config.accepts_bearer_auth() && has_bearer_auth(&headers)); - if !already_authorized { + let bearer_token = if matches!(auth_strategy, MessagesAuthStrategy::AwsSigV4) { + request + .api_key + .map(str::to_string) + .or_else(|| env_lookup("AWS_BEARER_TOKEN_BEDROCK")) + .filter(|token| !token.trim().is_empty()) + } else { + None + }; + let already_authorized = matches!(auth_strategy, MessagesAuthStrategy::AwsSigV4) + || !matches!(auth_strategy, MessagesAuthStrategy::AwsSigV4) + && (has_header(&headers, auth_strategy.header_name()) + || (config.accepts_bearer_auth() && has_bearer_auth(&headers))); + if !already_authorized && !matches!(auth_strategy, MessagesAuthStrategy::AwsSigV4) { let api_key = config.resolve_api_key(request.api_key, &env_lookup)?; let auth_header = match auth_strategy { MessagesAuthStrategy::Bearer => { ("authorization".to_string(), format!("Bearer {api_key}")) } MessagesAuthStrategy::Header(name) => (name.to_string(), api_key), + MessagesAuthStrategy::AwsSigV4 => unreachable!(), }; headers.push(auth_header); } @@ -52,7 +65,9 @@ pub(super) fn prepare_messages_call( } } - let url = config.complete_url(request.api_base, &model, &env_lookup)?; + let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true); + let url = config.complete_url(request.api_base, &model, stream, &env_lookup)?; + let signing_region = config.signing_region(request.api_base, &env_lookup); let typed_request = serde_json::from_value(request.body).map_err(|err| { CoreError::InvalidRequest(format!("invalid Anthropic messages request: {err}")) })?; @@ -70,6 +85,8 @@ pub(super) fn prepare_messages_call( url, body, upstream_headers: headers, + signing_region, + bearer_token, timeout: request.timeout, }) } diff --git a/litellm-rust/crates/ai-gateway/src/messages/types.rs b/litellm-rust/crates/ai-gateway/src/messages/types.rs index 848fadb4b02..2215d269e62 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/types.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/types.rs @@ -20,5 +20,7 @@ pub(crate) struct ProviderMessagesRequest { pub(crate) url: String, pub(crate) body: Value, pub(crate) upstream_headers: Vec<(String, String)>, + pub(crate) signing_region: Option, + pub(crate) bearer_token: Option, pub(crate) timeout: Option, } 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 a34b2edd7b8..3b6ef313c4d 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -2,6 +2,7 @@ mod service; +use aws_smithy_eventstream::frame::{DecodedFrame, MessageFrameDecoder}; use axum::Router; use axum::body::Body; use axum::extract::{Json, State}; @@ -9,6 +10,9 @@ use axum::http::StatusCode; use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE, HeaderMap, HeaderValue}; use axum::response::{IntoResponse, Response}; use axum::routing::post; +use bytes::Bytes; +use futures_util::StreamExt; +use futures_util::stream::{self, BoxStream}; use litellm_core::CoreError; use serde_json::{Map, Value}; @@ -33,16 +37,26 @@ async fn handle( .map_err(MessagesRouteError::from)? { service::MessagesResponse::Json(body) => Ok(Json(body).into_response()), - service::MessagesResponse::Stream(upstream) => stream_response(upstream), + service::MessagesResponse::Stream { provider, response } => { + stream_response(provider, response) + } } } -fn stream_response(upstream: reqwest::Response) -> Result { - let content_type = upstream - .headers() - .get(CONTENT_TYPE) - .cloned() - .unwrap_or_else(|| HeaderValue::from_static("text/event-stream")); +fn stream_response( + provider: String, + upstream: reqwest::Response, +) -> Result { + let is_bedrock = provider == "bedrock"; + let content_type = if is_bedrock { + HeaderValue::from_static("text/event-stream") + } else { + upstream + .headers() + .get(CONTENT_TYPE) + .cloned() + .unwrap_or_else(|| HeaderValue::from_static("text/event-stream")) + }; let mut response = Response::builder() .status( StatusCode::from_u16(upstream.status().as_u16()).map_err(|error| { @@ -55,8 +69,16 @@ fn stream_response(upstream: reqwest::Response) -> Result Result>, + buffer: bytes::BytesMut, + decoder: MessageFrameDecoder, +} + +fn bedrock_sse_stream( + upstream: BoxStream<'static, Result>, +) -> BoxStream<'static, Result> { + stream::unfold( + EventStreamState { + upstream, + buffer: bytes::BytesMut::new(), + decoder: MessageFrameDecoder::new(), + }, + |mut state| async move { + loop { + if let Ok(DecodedFrame::Complete(message)) = + state.decoder.decode_frame(&mut state.buffer) + { + let bytes = message + .headers() + .iter() + .find(|header| header.name().as_str() == ":message-type") + .and_then(|header| header.value().as_string().ok()) + .map_or_else( + || sse_data(message.payload()), + |message_type| { + if message_type.as_str() == "exception" + || message_type.as_str() == "error" + { + sse_error(message.payload()) + } else { + sse_data(message.payload()) + } + }, + ); + return Some((Ok(Bytes::from(bytes)), state)); + } + match state.upstream.next().await { + Some(Ok(chunk)) => state.buffer.extend_from_slice(&chunk), + Some(Err(error)) => { + return Some((Err(std::io::Error::other(error.to_string())), state)); + } + None => return None, + } + } + }, + ) + .boxed() +} + +fn sse_data(payload: &[u8]) -> String { + let value = serde_json::from_slice::(payload) + .ok() + .and_then(|value| value.get("bytes").and_then(serde_json::Value::as_str).map(str::to_string)) + .and_then(|encoded| { + base64::Engine::decode(&base64::engine::general_purpose::STANDARD, encoded).ok() + }) + .and_then(|payload| serde_json::from_slice::(&payload).ok()) + .unwrap_or_else(|| serde_json::json!({"type": "error", "error": {"type": "invalid_request_error", "message": "invalid Bedrock event"}})); + let event = value + .get("type") + .and_then(serde_json::Value::as_str) + .unwrap_or("message"); + format!("event: {event}\ndata: {value}\n\n") +} + +fn sse_error(payload: &[u8]) -> String { + let message = String::from_utf8_lossy(payload); + format!( + "event: error\ndata: {}\n\n", + serde_json::json!({"type": "error", "error": {"type": "api_error", "message": message}}) + ) +} + fn forwarded_headers(headers: &HeaderMap) -> Result>, CoreError> { let forwarded = headers .iter() 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 75ed26e5be8..30ce3357a63 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -9,7 +9,10 @@ use crate::messages::{MessagesRequest, execute_messages}; pub(crate) enum MessagesResponse { Json(Value), - Stream(reqwest::Response), + Stream { + provider: String, + response: reqwest::Response, + }, } pub async fn run( @@ -57,8 +60,8 @@ pub async fn run( .await .map(|response| match response { crate::messages::MessagesResponse::Json(body) => MessagesResponse::Json(body), - crate::messages::MessagesResponse::Stream(upstream) => { - MessagesResponse::Stream(upstream) + crate::messages::MessagesResponse::Stream { provider, response } => { + MessagesResponse::Stream { provider, response } } }) } diff --git a/litellm-rust/crates/core/src/messages/transformation.rs b/litellm-rust/crates/core/src/messages/transformation.rs index b478e20d24b..e75f54276a3 100644 --- a/litellm-rust/crates/core/src/messages/transformation.rs +++ b/litellm-rust/crates/core/src/messages/transformation.rs @@ -6,6 +6,7 @@ use super::types::{AnthropicMessagesRequest, AnthropicMessagesResponse}; pub enum MessagesAuthStrategy { Bearer, Header(&'static str), + AwsSigV4, } impl MessagesAuthStrategy { @@ -13,6 +14,7 @@ impl MessagesAuthStrategy { match self { Self::Bearer => "authorization", Self::Header(header_name) => header_name, + Self::AwsSigV4 => "", } } } @@ -22,9 +24,18 @@ pub trait AnthropicMessagesProviderConfig: Sync { &self, api_base: Option<&str>, model: &str, + stream: bool, env_lookup: &dyn Fn(&str) -> Option, ) -> CoreResult; + fn signing_region( + &self, + _api_base: Option<&str>, + _env_lookup: &dyn Fn(&str) -> Option, + ) -> Option { + None + } + fn resolve_api_key( &self, api_key: Option<&str>, diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index 11fe17ea40f..fa12ef374da 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -45,6 +45,7 @@ pub struct AnthropicMessage { #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] pub struct AnthropicMessagesRequest { + #[serde(skip_serializing_if = "String::is_empty")] pub model: String, pub messages: Vec, #[serde(skip_serializing_if = "Option::is_none")] @@ -95,6 +96,7 @@ pub struct AnthropicMessagesResponse { #[serde(rename = "type")] pub message_type: String, pub role: String, + #[serde(default)] pub model: String, pub content: Vec, // Anthropic always includes stop_reason / stop_sequence, null until the turn diff --git a/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs b/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs index 829f2260d3c..19f1ad633fc 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs @@ -51,6 +51,7 @@ impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig { &self, api_base: Option<&str>, _model: &str, + _stream: bool, env_lookup: &dyn Fn(&str) -> Option, ) -> CoreResult { Ok(complete_anthropic_url(api_base, env_lookup)) diff --git a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs index 7b958c77ba3..545cae508bc 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs @@ -146,6 +146,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig { &self, api_base: Option<&str>, _model: &str, + _stream: bool, env_lookup: &dyn Fn(&str) -> Option, ) -> CoreResult { complete_azure_anthropic_url(api_base, env_lookup) diff --git a/litellm-rust/crates/core/src/providers/bedrock/messages/mod.rs b/litellm-rust/crates/core/src/providers/bedrock/messages/mod.rs new file mode 100644 index 00000000000..f239b6921fa --- /dev/null +++ b/litellm-rust/crates/core/src/providers/bedrock/messages/mod.rs @@ -0,0 +1 @@ +pub mod transformation; diff --git a/litellm-rust/crates/core/src/providers/bedrock/messages/transformation.rs b/litellm-rust/crates/core/src/providers/bedrock/messages/transformation.rs new file mode 100644 index 00000000000..7ae8e41febd --- /dev/null +++ b/litellm-rust/crates/core/src/providers/bedrock/messages/transformation.rs @@ -0,0 +1,232 @@ +use crate::error::{CoreError, CoreResult}; +use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy}; +use crate::messages::types::{AnthropicMessagesRequest, AnthropicMessagesResponse}; +use crate::providers::anthropic::messages::transformation::non_empty; +use crate::providers::bedrock::constants::{ + AWS_REGION, AWS_REGION_NAME, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, DEFAULT_BEDROCK_REGION, +}; +use serde_json::Value; + +const AWS_DEFAULT_REGION: &str = "AWS_DEFAULT_REGION"; +const API_BASE_SCHEME: &str = "https://"; +const MODEL_PATH_PREFIX: &str = "/model/"; +const INVOKE_PATH: &str = "/invoke"; +const STREAM_PATH: &str = "/invoke-with-response-stream"; +const ANTHROPIC_VERSION_FIELD: &str = "anthropic_version"; +const ANTHROPIC_VERSION: &str = "bedrock-2023-05-31"; +const UNSUPPORTED_FIELDS: &[&str] = &[ + "metadata", + "service_tier", + "container", + "mcp_servers", + "context_management", + "output_format", + "output_config", + "speed", + "inference_geo", +]; + +pub struct BedrockMessagesConfig; + +pub const BEDROCK_MESSAGES_CONFIG: BedrockMessagesConfig = BedrockMessagesConfig; + +fn resolve_region(api_base: Option<&str>, env_lookup: &dyn Fn(&str) -> Option) -> String { + api_base + .and_then(|base| base.split('.').nth(1)) + .filter(|region| !region.is_empty()) + .map(str::to_string) + .or_else(|| env_lookup(AWS_REGION_NAME)) + .or_else(|| env_lookup(AWS_REGION)) + .or_else(|| env_lookup(AWS_DEFAULT_REGION)) + .unwrap_or_else(|| DEFAULT_BEDROCK_REGION.to_string()) +} + +fn encode_path_segment(value: &str) -> String { + value.bytes().fold(String::new(), |mut encoded, byte| { + if byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') { + encoded.push(byte as char); + } else { + encoded.push('%'); + encoded.push_str(&format!("{byte:02X}")); + } + encoded + }) +} + +fn endpoint_base(api_base: Option<&str>, env_lookup: &dyn Fn(&str) -> Option) -> String { + non_empty(api_base) + .map(str::to_string) + .unwrap_or_else(|| { + BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", &resolve_region(None, env_lookup)) + }) + .trim_end_matches('/') + .to_string() +} + +pub fn complete_bedrock_url( + api_base: Option<&str>, + model: &str, + stream: bool, + env_lookup: &dyn Fn(&str) -> Option, +) -> CoreResult { + let model = non_empty(Some(model)) + .ok_or_else(|| CoreError::InvalidRequest("Bedrock model cannot be empty".to_string()))?; + let suffix = if stream { STREAM_PATH } else { INVOKE_PATH }; + let base = endpoint_base(api_base, env_lookup); + let base = if base.starts_with(API_BASE_SCHEME) || base.starts_with("http://") { + base + } else { + format!("{API_BASE_SCHEME}{base}") + }; + Ok(format!( + "{base}{MODEL_PATH_PREFIX}{}{suffix}", + encode_path_segment(model) + )) +} + +impl AnthropicMessagesProviderConfig for BedrockMessagesConfig { + fn complete_url( + &self, + api_base: Option<&str>, + model: &str, + stream: bool, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + complete_bedrock_url(api_base, model, stream, env_lookup) + } + + fn signing_region( + &self, + api_base: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, + ) -> Option { + Some(resolve_region(api_base, env_lookup)) + } + + fn resolve_api_key( + &self, + _api_key: Option<&str>, + _env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + Ok(String::new()) + } + + fn auth_strategy(&self) -> MessagesAuthStrategy { + MessagesAuthStrategy::AwsSigV4 + } + + fn transform_request( + &self, + mut request: AnthropicMessagesRequest, + ) -> CoreResult { + request.model.clear(); + request.stream = None; + request.metadata = None; + request.service_tier = None; + request.container = None; + request.mcp_servers = None; + request.context_management = None; + request.output_format = None; + request.output_config = None; + request.speed = None; + request.inference_geo = None; + request + .extra + .retain(|key, _| !UNSUPPORTED_FIELDS.contains(&key.as_str())); + request.extra.insert( + ANTHROPIC_VERSION_FIELD.to_string(), + Value::String(ANTHROPIC_VERSION.to_string()), + ); + Ok(request) + } + + fn transform_response( + &self, + model: &str, + mut response: AnthropicMessagesResponse, + ) -> CoreResult { + if response.model.trim().is_empty() { + response.model = model.to_string(); + } + Ok(response) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn request(value: Value) -> AnthropicMessagesRequest { + serde_json::from_value(value).expect("valid request") + } + + #[test] + fn builds_default_and_streaming_urls_with_encoded_arn() { + let env = |key: &str| (key == AWS_REGION).then(|| "eu-west-1".to_string()); + let model = "arn:aws:bedrock:us-east-1:123456789012:inference-profile/foo/bar"; + assert_eq!( + complete_bedrock_url(None, model, false, &env).expect("url"), + "https://bedrock-runtime.eu-west-1.amazonaws.com/model/arn%3Aaws%3Abedrock%3Aus-east-1%3A123456789012%3Ainference-profile%2Ffoo%2Fbar/invoke" + ); + assert!( + complete_bedrock_url(None, "claude", true, &env) + .expect("url") + .ends_with("/invoke-with-response-stream") + ); + } + + #[test] + fn api_base_region_wins_over_environment() { + let env = |key: &str| (key == AWS_REGION_NAME).then(|| "us-west-2".to_string()); + assert_eq!( + BEDROCK_MESSAGES_CONFIG.signing_region( + Some("https://bedrock-runtime.ap-south-1.amazonaws.com"), + &env + ), + Some("ap-south-1".to_string()) + ); + } + + #[test] + fn request_removes_path_and_unsupported_fields() { + let transformed = BEDROCK_MESSAGES_CONFIG + .transform_request(request(json!({ + "model": "claude", + "stream": true, + "max_tokens": 10, + "messages": [{"role": "user", "content": "hello"}], + "metadata": {"user_id": "ignored"}, + "tools": [{"name": "search"}] + }))) + .expect("transform"); + let value = serde_json::to_value(transformed).expect("json"); + assert!(value.get("model").is_none()); + assert!(value.get("stream").is_none()); + assert!(value.get("metadata").is_none()); + assert_eq!(value["anthropic_version"], ANTHROPIC_VERSION); + assert!(value.get("tools").is_some()); + } + + #[test] + fn rejects_empty_model_and_restamps_empty_response_model() { + assert!(complete_bedrock_url(None, " ", false, &|_| None).is_err()); + let response: AnthropicMessagesResponse = serde_json::from_value(json!({ + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "", + "content": [], + "stop_reason": null, + "stop_sequence": null + })) + .expect("response"); + assert_eq!( + BEDROCK_MESSAGES_CONFIG + .transform_response("claude", response) + .expect("response") + .model, + "claude" + ); + } +} diff --git a/litellm-rust/crates/core/src/providers/bedrock/mod.rs b/litellm-rust/crates/core/src/providers/bedrock/mod.rs index b09675ad7dd..1139407862f 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/mod.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/mod.rs @@ -6,3 +6,4 @@ pub mod audio_transcription; pub mod aws_base; mod constants; +pub mod messages;