fix(rust): harden Bedrock messages gateway

Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-07-29 02:41:24 +00:00
parent 32a1d55e56
commit ea7fc3747e
9 changed files with 472 additions and 50 deletions

View file

@ -1237,6 +1237,7 @@ name = "litellm-ai-gateway"
version = "0.1.0"
dependencies = [
"aws-smithy-eventstream",
"aws-smithy-types",
"axum",
"base64",
"bytes",

View file

@ -31,3 +31,4 @@ futures-util = { version = "0.3", default-features = false, features = ["sink",
base64 = "0.22"
bytes = "1"
aws-smithy-eventstream = "0.60.3"
aws-smithy-types = "1.6.1"

View file

@ -26,6 +26,7 @@ serde_json.workspace = true
base64.workspace = true
bytes.workspace = true
aws-smithy-eventstream.workspace = true
aws-smithy-types.workspace = true
axum = { workspace = true, features = ["ws"], optional = true }
serde.workspace = true
subtle = { workspace = true, optional = true }

View file

@ -51,6 +51,8 @@ pub(crate) const MESSAGES_ROUTE_PATH: &str = "/v1/messages";
/// Provider name used by the Anthropic Messages route when a deployment's
/// provider model does not carry an explicit provider prefix.
pub(crate) const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic";
pub(crate) const AZURE_ANTHROPIC_MESSAGES_PROVIDER: &str = "azure_ai";
pub(crate) const BEDROCK_MESSAGES_PROVIDER: &str = "bedrock";
/// Request headers owned by the gateway and never forwarded upstream.
#[cfg(feature = "server")]

View file

@ -12,7 +12,9 @@ use serde_json::Value;
use super::client::http_client;
use super::common_utils::truncate_error_body;
use super::types::ProviderMessagesRequest;
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use crate::constants::{
ANTHROPIC_MESSAGES_PROVIDER, AZURE_ANTHROPIC_MESSAGES_PROVIDER, BEDROCK_MESSAGES_PROVIDER,
};
fn environment_lookup(key: &str) -> Option<String> {
std::env::var(key).ok()
@ -122,7 +124,10 @@ pub(super) async fn execute_messages_provider_call(
pub(super) async fn execute_messages_provider_stream(
request: ProviderMessagesRequest,
) -> CoreResult<reqwest::Response> {
if request.provider != ANTHROPIC_MESSAGES_PROVIDER && request.signing_region.is_none() {
if !matches!(
request.provider.as_str(),
ANTHROPIC_MESSAGES_PROVIDER | AZURE_ANTHROPIC_MESSAGES_PROVIDER | BEDROCK_MESSAGES_PROVIDER
) {
return Err(CoreError::InvalidRequest(
"streaming messages is not supported for this provider".to_string(),
));

View file

@ -15,18 +15,15 @@ use prepare::prepare_messages_call;
pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
match execute_messages(request, false).await? {
MessagesResponse::Json(body) => Ok(body),
MessagesResponse::Stream { response, provider } => {
drop(provider);
drop(response);
Err(litellm_core::CoreError::InvalidResponse(
"non-streaming messages execution returned a stream".to_string(),
))
}
MessagesResponse::Stream { .. } => Err(litellm_core::CoreError::InvalidResponse(
"non-streaming messages execution returned a stream".to_string(),
)),
}
}
pub(crate) enum MessagesResponse {
Json(Value),
#[allow(dead_code)]
Stream {
provider: String,
response: reqwest::Response,

View file

@ -43,20 +43,31 @@ pub(super) fn prepare_messages_call(
} 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);
let auth_header = match auth_strategy {
MessagesAuthStrategy::AwsSigV4 => None,
MessagesAuthStrategy::Bearer
if has_header(&headers, "authorization")
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers)) =>
{
None
}
MessagesAuthStrategy::Header(name)
if has_header(&headers, name)
|| (config.accepts_bearer_auth() && has_bearer_auth(&headers)) =>
{
None
}
MessagesAuthStrategy::Bearer => {
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
Some(("authorization".to_string(), format!("Bearer {api_key}")))
}
MessagesAuthStrategy::Header(name) => {
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
Some((name.to_string(), api_key))
}
};
if let Some(header) = auth_header {
headers.push(header);
}
for (name, value) in config.default_headers() {

View file

@ -17,7 +17,9 @@ use litellm_core::CoreError;
use serde_json::{Map, Value};
use crate::auth::RequireMasterKey;
use crate::constants::{MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH};
use crate::constants::{
BEDROCK_MESSAGES_PROVIDER, MESSAGES_HEADERS_NOT_FORWARDED, MESSAGES_ROUTE_PATH,
};
use crate::state::AppState;
/// This route's contribution to the app router.
@ -47,7 +49,7 @@ fn stream_response(
provider: String,
upstream: reqwest::Response,
) -> Result<Response, MessagesRouteError> {
let is_bedrock = provider == "bedrock";
let is_bedrock = provider == BEDROCK_MESSAGES_PROVIDER;
let content_type = if is_bedrock {
HeaderValue::from_static("text/event-stream")
} else {
@ -90,6 +92,7 @@ struct EventStreamState {
upstream: BoxStream<'static, Result<Bytes, reqwest::Error>>,
buffer: bytes::BytesMut,
decoder: MessageFrameDecoder,
terminated: bool,
}
fn bedrock_sse_stream(
@ -100,37 +103,58 @@ fn bedrock_sse_stream(
upstream,
buffer: bytes::BytesMut::new(),
decoder: MessageFrameDecoder::new(),
terminated: false,
},
|mut state| async move {
if state.terminated {
return None;
}
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.decoder.decode_frame(&mut state.buffer) {
Ok(DecodedFrame::Complete(message)) => {
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));
}
Ok(DecodedFrame::Incomplete) => {}
Err(error) => {
state.terminated = true;
return Some((
Ok(Bytes::from(sse_error(error.to_string().as_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,
None if state.buffer.is_empty() => return None,
None => {
state.terminated = true;
return Some((
Ok(Bytes::from(sse_error(
b"incomplete Bedrock event stream frame",
))),
state,
));
}
}
}
},
@ -221,11 +245,16 @@ impl IntoResponse for MessagesRouteError {
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::{Mutex, OnceLock};
use std::time::Duration;
use aws_smithy_eventstream::frame::write_message_to;
use aws_smithy_types::event_stream::{Header, HeaderValue, Message};
use axum::body::Body;
use axum::http::Request;
use axum::http::StatusCode;
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE};
use bytes::BytesMut;
use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter};
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
@ -261,6 +290,113 @@ mod tests {
}
}
fn bedrock_state(
model_alias: &str,
provider_model: &str,
api_base: String,
api_key: Option<&str>,
master_key: Option<&str>,
) -> AppState {
AppState {
router: Arc::new(ModelRouter::new(vec![Deployment {
model_name: model_alias.to_string(),
litellm_params: LiteLLMParams {
model: format!("bedrock/{provider_model}"),
api_key: api_key.map(str::to_string),
api_base: Some(api_base),
},
}])),
master_key: master_key.map(Arc::from),
loggers: Arc::new(Vec::new()),
realtime_pool: RealtimePool::disabled(),
}
}
fn event_frame(message_type: &str, payload: serde_json::Value) -> Vec<u8> {
let mut frame = BytesMut::new();
let message = Message::new(
serde_json::to_vec(&serde_json::json!({
"bytes": base64::Engine::encode(
&base64::engine::general_purpose::STANDARD,
serde_json::to_vec(&payload).expect("payload"),
)
}))
.expect("wrapper"),
)
.add_header(Header::new(
":message-type",
HeaderValue::String(message_type.to_string().into()),
));
write_message_to(&message, &mut frame).expect("event frame");
frame.to_vec()
}
async fn read_request(socket: &mut tokio::net::TcpStream) -> String {
let mut request = Vec::new();
let mut buffer = [0_u8; 4096];
loop {
let read = socket.read(&mut buffer).await.expect("reads request");
request.extend_from_slice(&buffer[..read]);
if request.windows(4).any(|window| window == b"\r\n\r\n") {
break;
}
}
let header_end = request
.windows(4)
.position(|window| window == b"\r\n\r\n")
.expect("header end")
+ 4;
let headers = String::from_utf8_lossy(&request[..header_end]);
let content_length = headers
.lines()
.find_map(|line| {
let (name, value) = line.split_once(':')?;
name.eq_ignore_ascii_case("content-length")
.then(|| value.trim().parse::<usize>().ok())
.flatten()
})
.unwrap_or(0);
while request.len().saturating_sub(header_end) < content_length {
let read = socket.read(&mut buffer).await.expect("reads body");
request.extend_from_slice(&buffer[..read]);
}
String::from_utf8(request).expect("request utf8")
}
async fn bedrock_streaming_upstream(
listener: TcpListener,
frames: Vec<Vec<u8>>,
) -> (String, tokio::task::JoinHandle<String>) {
let address = listener.local_addr().expect("listener has address");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts request");
let request = read_request(&mut socket).await;
let body = frames.into_iter().flatten().collect::<Vec<_>>();
let head = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/vnd.amazon.eventstream\r\ncontent-length: {}\r\nconnection: close\r\n\r\n",
body.len()
);
socket
.write_all(head.as_bytes())
.await
.expect("writes head");
let split = body.len() / 2;
socket
.write_all(&body[..split])
.await
.expect("writes first frame part");
tokio::time::sleep(Duration::from_millis(10)).await;
socket
.write_all(&body[split..])
.await
.expect("writes second frame part");
request
});
(format!("http://{address}"), server)
}
static AWS_ENV_LOCK: OnceLock<Mutex<()>> = OnceLock::new();
async fn upstream(listener: TcpListener) -> (String, tokio::task::JoinHandle<String>) {
let address = listener.local_addr().expect("listener has address");
let server = tokio::spawn(async move {
@ -364,6 +500,7 @@ mod tests {
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.header("x-api-key", "request-upstream-key")
.header("anthropic-beta", "beta-feature")
.header("content-type", "application/json")
@ -544,6 +681,251 @@ mod tests {
server.await.expect("upstream task completes");
}
#[tokio::test]
#[allow(clippy::await_holding_lock)]
async fn route_constructs_bedrock_sigv4_invoke_request() {
let _lock = AWS_ENV_LOCK
.get_or_init(|| Mutex::new(()))
.lock()
.expect("env lock");
unsafe {
std::env::remove_var("AWS_BEARER_TOKEN_BEDROCK");
std::env::set_var("AWS_ACCESS_KEY_ID", "test-access-key");
std::env::set_var("AWS_SECRET_ACCESS_KEY", "test-secret-key");
std::env::set_var("AWS_REGION_NAME", "us-east-1");
}
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let (api_base, server) = upstream(listener).await;
let model = "arn:aws:bedrock:us-east-1:123:model/foo";
let app = app(bedrock_state(
model,
model,
api_base,
None,
Some("master-key"),
));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": model,
"stream": false,
"max_tokens": 16,
"messages": [{"role": "user", "content": "hello"}]
})
.to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::OK);
let request = server.await.expect("server task completes");
let (head, body) = request.split_once("\r\n\r\n").expect("request body");
assert!(head.starts_with(
"POST /model/arn%3Aaws%3Abedrock%3Aus-east-1%3A123%3Amodel%2Ffoo/invoke "
));
let head_lower = head.to_ascii_lowercase();
assert!(head_lower.contains("authorization: aws4-hmac-sha256"));
assert!(head_lower.contains("x-amz-date:"));
let body: serde_json::Value = serde_json::from_str(body).expect("body json");
assert_eq!(body["anthropic_version"], "bedrock-2023-05-31");
assert!(body.get("model").is_none());
assert!(body.get("stream").is_none());
}
#[tokio::test]
async fn route_constructs_bedrock_bearer_request_without_anthropic_headers() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let (api_base, server) = upstream(listener).await;
let app = app(bedrock_state(
"claude-test",
"claude-test",
api_base,
Some("bedrock-token"),
Some("master-key"),
));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("x-api-key", "client-key")
.header("anthropic-version", "2023-06-01")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "claude-test",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hello"}]
})
.to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::OK);
let request = server.await.expect("server task completes");
let (head, _) = request.split_once("\r\n\r\n").expect("request body");
let head = head.to_ascii_lowercase();
assert!(head.contains("authorization: bearer bedrock-token"));
assert!(!head.contains("x-api-key:"));
assert!(!head.contains("anthropic-version:"));
}
#[tokio::test]
async fn route_decodes_bedrock_eventstream_split_frames_and_tool_use() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let frames = vec![
event_frame("event", json!({"type": "message_start"})),
event_frame(
"event",
json!({
"type": "content_block_start",
"content_block": {"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {}},
"index": 0
}),
),
event_frame(
"event",
json!({
"type": "content_block_delta",
"delta": {"type": "input_json_delta", "partial_json": "{\"city\":\"Paris\"}"}
}),
),
event_frame("event", json!({"type": "message_stop"})),
];
let (api_base, server) = bedrock_streaming_upstream(listener, frames).await;
let app = app(bedrock_state(
"claude-test",
"claude-test",
api_base,
Some("bedrock-token"),
Some("master-key"),
));
let response = app
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.body(Body::from(
json!({
"model": "claude-test",
"stream": true,
"max_tokens": 16,
"tools": [{"name": "get_weather", "input_schema": {"type": "object"}}],
"messages": [{"role": "user", "content": "weather"}]
})
.to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::OK);
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body");
let body = String::from_utf8(body.to_vec()).expect("sse");
assert!(body.contains("event: message_start"));
assert!(body.contains("\"type\":\"tool_use\""));
assert!(body.contains("event: content_block_delta"));
assert!(body.contains("input_json_delta"));
assert!(body.contains("event: message_stop"));
let request = server.await.expect("server task completes");
let (_, request_body) = request.split_once("\r\n\r\n").expect("request body");
let request_body: serde_json::Value =
serde_json::from_str(request_body).expect("request json");
assert_eq!(request_body["tools"][0]["name"], "get_weather");
}
#[tokio::test]
async fn route_maps_bedrock_errors_and_eventstream_exceptions() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let address = listener.local_addr().expect("listener address");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accept");
let _ = read_request(&mut socket).await;
socket
.write_all(b"HTTP/1.1 400 Bad Request\r\ncontent-length: 3\r\nconnection: close\r\n\r\nbad")
.await
.expect("write error");
});
let response = app(bedrock_state(
"claude-test",
"claude-test",
format!("http://{address}"),
Some("bedrock-token"),
Some("master-key"),
))
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.body(Body::from(
json!({"model": "claude-test", "max_tokens": 8, "messages": []}).to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
assert_eq!(response.status(), StatusCode::BAD_GATEWAY);
server.await.expect("server task");
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let (api_base, server) = bedrock_streaming_upstream(
listener,
vec![event_frame(
"exception",
json!({
"message": "model unavailable"
}),
)],
)
.await;
let response = app(bedrock_state(
"claude-test",
"claude-test",
api_base,
Some("bedrock-token"),
Some("master-key"),
))
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/messages")
.header("authorization", "Bearer master-key")
.header("content-type", "application/json")
.body(Body::from(
json!({"model": "claude-test", "stream": true, "max_tokens": 8, "messages": []})
.to_string(),
))
.expect("request builds"),
)
.await
.expect("route responds");
let body = axum::body::to_bytes(response.into_body(), usize::MAX)
.await
.expect("response body");
assert!(
String::from_utf8(body.to_vec())
.expect("sse")
.contains("event: error")
);
server.await.expect("server task");
}
#[tokio::test]
async fn route_rejects_missing_master_key() {
let app = app(state(

View file

@ -32,15 +32,28 @@ pub const BEDROCK_MESSAGES_CONFIG: BedrockMessagesConfig = BedrockMessagesConfig
fn resolve_region(api_base: Option<&str>, env_lookup: &dyn Fn(&str) -> Option<String>) -> String {
api_base
.and_then(|base| base.split('.').nth(1))
.filter(|region| !region.is_empty())
.map(str::to_string)
.and_then(bedrock_region_from_api_base)
.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 bedrock_region_from_api_base(api_base: &str) -> Option<String> {
let host = api_base
.trim()
.trim_start_matches("https://")
.trim_start_matches("http://")
.split('/')
.next()?
.split(':')
.next()?;
let region = host
.strip_prefix("bedrock-runtime.")?
.strip_suffix(".amazonaws.com")?;
(!region.is_empty()).then(|| 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'~') {
@ -188,6 +201,15 @@ mod tests {
);
}
#[test]
fn non_bedrock_api_base_does_not_supply_a_region() {
let env = |key: &str| (key == AWS_REGION).then(|| "eu-west-1".to_string());
assert_eq!(
BEDROCK_MESSAGES_CONFIG.signing_region(Some("http://127.0.0.1:8080"), &env),
Some("eu-west-1".to_string())
);
}
#[test]
fn request_removes_path_and_unsupported_fields() {
let transformed = BEDROCK_MESSAGES_CONFIG