mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
fix(rust): harden Bedrock messages gateway
Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
32a1d55e56
commit
ea7fc3747e
9 changed files with 472 additions and 50 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -1237,6 +1237,7 @@ name = "litellm-ai-gateway"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"aws-smithy-eventstream",
|
||||
"aws-smithy-types",
|
||||
"axum",
|
||||
"base64",
|
||||
"bytes",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 }
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
));
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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() {
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue