mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-16 23:41:43 +00:00
feat(rust): add Bedrock Anthropic invoke support to /v1/messages
Adds a Bedrock Messages provider (invoke + invoke-with-response-stream) with SigV4 signing and AWS event-stream to Anthropic SSE transcoding on the Axum gateway. Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
8177230a29
commit
9fa4a70f31
24 changed files with 992 additions and 202 deletions
25
litellm-rust/Cargo.lock
generated
25
litellm-rust/Cargo.lock
generated
|
|
@ -180,6 +180,17 @@ dependencies = [
|
|||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "aws-smithy-eventstream"
|
||||
version = "0.61.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5a9381123ab62d20c13082b151f30f962a3b112b727345394536dfa39a482944"
|
||||
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"
|
||||
|
|
@ -1218,6 +1238,7 @@ version = "0.1.0"
|
|||
dependencies = [
|
||||
"axum",
|
||||
"base64",
|
||||
"bytes",
|
||||
"futures-channel",
|
||||
"futures-util",
|
||||
"litellm-core",
|
||||
|
|
@ -1240,8 +1261,12 @@ dependencies = [
|
|||
"aws-credential-types",
|
||||
"aws-sdk-sts",
|
||||
"aws-sigv4",
|
||||
"aws-smithy-eventstream",
|
||||
"aws-smithy-runtime-api",
|
||||
"aws-smithy-types",
|
||||
"aws-types",
|
||||
"base64",
|
||||
"bytes",
|
||||
"rand 0.8.7",
|
||||
"reqwest",
|
||||
"serde",
|
||||
|
|
|
|||
|
|
@ -29,3 +29,4 @@ 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"
|
||||
aws-smithy-eventstream = "0.61.1"
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ reqwest.workspace = true
|
|||
tokio = { workspace = true, features = ["rt-multi-thread", "macros", "net", "time", "sync"] }
|
||||
tokio-tungstenite.workspace = true
|
||||
futures-util.workspace = true
|
||||
bytes = "1"
|
||||
serde_json.workspace = true
|
||||
base64.workspace = true
|
||||
axum = { workspace = true, features = ["ws"], optional = true }
|
||||
|
|
|
|||
|
|
@ -3,8 +3,8 @@ use std::time::SystemTime;
|
|||
use litellm_core::CoreResult;
|
||||
use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth;
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::providers::bedrock::audio_transcription::aws_auth_config;
|
||||
use litellm_core::providers::bedrock::aws_base::{resolve_credentials, sign_bedrock_post};
|
||||
use litellm_core::providers::bedrock::common_utils::aws_auth_config;
|
||||
use serde_json::Value;
|
||||
|
||||
use super::common_utils::truncate_error_body;
|
||||
|
|
@ -61,13 +61,21 @@ pub(crate) async fn sign_request(
|
|||
})?;
|
||||
let mut headers = super::common_utils::string_headers(None)?;
|
||||
headers.insert("Content-Type".to_string(), "application/json".to_string());
|
||||
let host = reqwest::Url::parse(&request.url)
|
||||
.map_err(|error| CoreError::InvalidRequest(format!("invalid Bedrock URL: {error}")))?
|
||||
.host_str()
|
||||
.ok_or_else(|| CoreError::InvalidRequest("Bedrock URL has no host".to_string()))?
|
||||
.to_string();
|
||||
headers.insert("Host".to_string(), host);
|
||||
headers.extend(request.upstream_headers.iter().cloned());
|
||||
match auth {
|
||||
AudioTranscriptionAuth::Bearer => {}
|
||||
AudioTranscriptionAuth::AwsSigV4 { region, .. } => {
|
||||
let credentials =
|
||||
resolve_credentials(aws_auth_config(optional_params, &env_lookup), &env_lookup)
|
||||
.await?;
|
||||
let credentials = resolve_credentials(
|
||||
aws_auth_config(optional_params, &env_lookup, Some(®ion)),
|
||||
&env_lookup,
|
||||
)
|
||||
.await?;
|
||||
headers.extend(sign_bedrock_post(
|
||||
&request.url,
|
||||
&body,
|
||||
|
|
|
|||
|
|
@ -3,8 +3,10 @@ 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_ANTHROPIC_MESSAGES_CONFIG;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS;
|
||||
|
||||
pub(super) fn truncate_error_body(body: &str) -> String {
|
||||
|
|
@ -19,8 +21,9 @@ pub(super) fn messages_provider_config(
|
|||
provider: &str,
|
||||
) -> Option<&'static dyn AnthropicMessagesProviderConfig> {
|
||||
match provider {
|
||||
"anthropic" => Some(&ANTHROPIC_MESSAGES_CONFIG),
|
||||
ANTHROPIC_MESSAGES_PROVIDER => Some(&ANTHROPIC_MESSAGES_CONFIG),
|
||||
"azure_ai" => Some(&AZURE_ANTHROPIC_MESSAGES_CONFIG),
|
||||
"bedrock" => Some(&BEDROCK_ANTHROPIC_MESSAGES_CONFIG),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,17 +1,27 @@
|
|||
use bytes::Bytes;
|
||||
use futures_util::StreamExt;
|
||||
use futures_util::stream::{self, BoxStream};
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::error::CoreError;
|
||||
use litellm_core::messages::transformation::{MessagesAuthKind, MessagesStreaming};
|
||||
use litellm_core::providers::bedrock::aws_base::{resolve_credentials, sign_bedrock_post};
|
||||
use litellm_core::providers::bedrock::common_utils::aws_auth_config;
|
||||
use serde_json::Value;
|
||||
use std::collections::BTreeMap;
|
||||
use std::time::SystemTime;
|
||||
|
||||
use super::client::http_client;
|
||||
use super::common_utils::truncate_error_body;
|
||||
use super::types::ProviderMessagesRequest;
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
|
||||
pub(super) async fn execute_messages_provider_call(
|
||||
request: ProviderMessagesRequest,
|
||||
) -> CoreResult<Value> {
|
||||
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_headers(&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 {
|
||||
|
|
@ -49,15 +59,19 @@ 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 {
|
||||
) -> CoreResult<MessagesStream> {
|
||||
if matches!(request.streaming, MessagesStreaming::Unsupported) {
|
||||
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_headers(&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 {
|
||||
|
|
@ -79,5 +93,122 @@ pub(super) async fn execute_messages_provider_stream(
|
|||
body: truncate_error_body(&text),
|
||||
});
|
||||
}
|
||||
Ok(response)
|
||||
let content_type = response
|
||||
.headers()
|
||||
.get(reqwest::header::CONTENT_TYPE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.unwrap_or("text/event-stream")
|
||||
.to_string();
|
||||
let cache_control = response
|
||||
.headers()
|
||||
.get(reqwest::header::CACHE_CONTROL)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::to_string);
|
||||
let body = match request.streaming {
|
||||
MessagesStreaming::SsePassthrough | MessagesStreaming::Unsupported => response
|
||||
.bytes_stream()
|
||||
.map(|result| result.map_err(|error| CoreError::Network(error.to_string())))
|
||||
.boxed(),
|
||||
MessagesStreaming::BedrockEventStream => bedrock_stream(response),
|
||||
};
|
||||
Ok(MessagesStream {
|
||||
content_type,
|
||||
cache_control,
|
||||
body,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg_attr(not(feature = "server"), allow(dead_code))]
|
||||
pub(crate) struct MessagesStream {
|
||||
pub(crate) content_type: String,
|
||||
pub(crate) cache_control: Option<String>,
|
||||
pub(crate) body: BoxStream<'static, CoreResult<Bytes>>,
|
||||
}
|
||||
|
||||
fn bedrock_stream(response: reqwest::Response) -> BoxStream<'static, CoreResult<Bytes>> {
|
||||
use std::collections::VecDeque;
|
||||
|
||||
use litellm_core::providers::bedrock::messages::streaming::{
|
||||
BedrockEventStreamDecoder, serialize_sse,
|
||||
};
|
||||
|
||||
let upstream = response.bytes_stream().boxed();
|
||||
stream::unfold(
|
||||
(
|
||||
upstream,
|
||||
BedrockEventStreamDecoder::new(),
|
||||
VecDeque::<Bytes>::new(),
|
||||
),
|
||||
|(mut upstream, mut decoder, mut pending)| async move {
|
||||
loop {
|
||||
if let Some(bytes) = pending.pop_front() {
|
||||
return Some((Ok(bytes), (upstream, decoder, pending)));
|
||||
}
|
||||
let chunk = upstream.next().await?;
|
||||
let chunk = match chunk {
|
||||
Ok(chunk) => chunk,
|
||||
Err(error) => {
|
||||
return Some((
|
||||
Err(CoreError::Network(error.to_string())),
|
||||
(upstream, decoder, pending),
|
||||
));
|
||||
}
|
||||
};
|
||||
match decoder.push(&chunk) {
|
||||
Ok(events) => {
|
||||
for event in events {
|
||||
match serialize_sse(&event) {
|
||||
Ok(bytes) => pending.push_back(Bytes::from(bytes)),
|
||||
Err(error) => {
|
||||
return Some((Err(error), (upstream, decoder, pending)));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
Err(error) => {
|
||||
return Some((Err(error), (upstream, decoder, pending)));
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
.boxed()
|
||||
}
|
||||
|
||||
async fn signed_headers(
|
||||
request: &ProviderMessagesRequest,
|
||||
body: &[u8],
|
||||
) -> CoreResult<Vec<(String, String)>> {
|
||||
if let MessagesAuthKind::AwsSigV4 { region } = &request.auth_kind {
|
||||
let env_lookup = environment_lookup;
|
||||
let credentials = resolve_credentials(
|
||||
aws_auth_config(&serde_json::Map::new(), &env_lookup, Some(region)),
|
||||
&env_lookup,
|
||||
)
|
||||
.await?;
|
||||
let host = reqwest::Url::parse(&request.url)
|
||||
.map_err(|error| CoreError::InvalidRequest(format!("invalid Bedrock URL: {error}")))?
|
||||
.host_str()
|
||||
.ok_or_else(|| CoreError::InvalidRequest("Bedrock URL has no host".to_string()))?
|
||||
.to_string();
|
||||
let mut headers = BTreeMap::from([
|
||||
("content-type".to_string(), "application/json".to_string()),
|
||||
("host".to_string(), host),
|
||||
]);
|
||||
let signed = sign_bedrock_post(
|
||||
&request.url,
|
||||
body,
|
||||
&headers,
|
||||
region,
|
||||
&credentials,
|
||||
SystemTime::now(),
|
||||
)?;
|
||||
headers.extend(signed);
|
||||
return Ok(headers.into_iter().collect());
|
||||
}
|
||||
Ok(request.upstream_headers.clone())
|
||||
}
|
||||
|
||||
fn environment_lookup(key: &str) -> Option<String> {
|
||||
std::env::var(key).ok()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ mod types;
|
|||
|
||||
pub use types::MessagesRequest;
|
||||
|
||||
pub(crate) use handler::MessagesStream;
|
||||
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
|
||||
use prepare::prepare_messages_call;
|
||||
|
||||
|
|
@ -16,6 +17,7 @@ pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
|
|||
match execute_messages(request, false).await? {
|
||||
MessagesResponse::Json(body) => Ok(body),
|
||||
MessagesResponse::Stream(response) => {
|
||||
let _ = response;
|
||||
drop(response);
|
||||
Err(litellm_core::CoreError::InvalidResponse(
|
||||
"non-streaming messages execution returned a stream".to_string(),
|
||||
|
|
@ -26,7 +28,7 @@ pub async fn messages(request: MessagesRequest<'_>) -> CoreResult<Value> {
|
|||
|
||||
pub(crate) enum MessagesResponse {
|
||||
Json(Value),
|
||||
Stream(reqwest::Response),
|
||||
Stream(MessagesStream),
|
||||
}
|
||||
|
||||
pub(crate) async fn execute_messages(
|
||||
|
|
|
|||
|
|
@ -1,7 +1,9 @@
|
|||
use litellm_core::CoreError;
|
||||
use litellm_core::CoreResult;
|
||||
use litellm_core::messages::transformation::MessagesAuthKind;
|
||||
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};
|
||||
|
|
@ -30,43 +32,46 @@ pub(super) fn prepare_messages_call(
|
|||
.ok_or_else(|| CoreError::InvalidProvider(provider.to_string()))?;
|
||||
let env_lookup = |key: &str| std::env::var(key).ok();
|
||||
|
||||
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 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}"))
|
||||
let stream = request.body.get("stream").and_then(Value::as_bool) == Some(true);
|
||||
let auth_kind = config.auth_kind(&model, &env_lookup)?;
|
||||
let headers = match &auth_kind {
|
||||
MessagesAuthKind::AwsSigV4 { .. } => Vec::new(),
|
||||
MessagesAuthKind::ApiKey {
|
||||
strategy,
|
||||
accepts_bearer,
|
||||
} => {
|
||||
let mut headers = string_headers(request.extra_headers)?;
|
||||
let already_authorized = has_header(&headers, strategy.header_name())
|
||||
|| (*accepts_bearer && has_bearer_auth(&headers));
|
||||
if !already_authorized {
|
||||
let api_key = config.resolve_api_key(request.api_key, &env_lookup)?;
|
||||
let auth_header = match strategy {
|
||||
MessagesAuthStrategy::Bearer => {
|
||||
("authorization".to_string(), format!("Bearer {api_key}"))
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
headers.push(auth_header);
|
||||
}
|
||||
MessagesAuthStrategy::Header(name) => (name.to_string(), api_key),
|
||||
};
|
||||
headers.push(auth_header);
|
||||
}
|
||||
|
||||
for (name, value) in config.default_headers() {
|
||||
if !has_header(&headers, name) {
|
||||
headers.push((name.to_string(), value.to_string()));
|
||||
for (name, value) in config.default_headers() {
|
||||
if !has_header(&headers, name) {
|
||||
headers.push((name.to_string(), value.to_string()));
|
||||
}
|
||||
}
|
||||
headers
|
||||
}
|
||||
}
|
||||
|
||||
let url = config.complete_url(request.api_base, &model, &env_lookup)?;
|
||||
};
|
||||
let url = config.complete_url(request.api_base, &model, stream, &env_lookup)?;
|
||||
let typed_request = serde_json::from_value(request.body).map_err(|err| {
|
||||
CoreError::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
|
||||
})?;
|
||||
let transformed = config.transform_request(typed_request)?;
|
||||
let body = serde_json::to_value(transformed).map_err(|err| {
|
||||
CoreError::InvalidRequest(format!(
|
||||
"failed to serialize Anthropic messages request: {err}"
|
||||
))
|
||||
})?;
|
||||
let body = config.upstream_body(typed_request)?;
|
||||
|
||||
Ok(ProviderMessagesRequest {
|
||||
provider: provider.to_string(),
|
||||
model,
|
||||
config,
|
||||
auth_kind,
|
||||
streaming: config.streaming(),
|
||||
url,
|
||||
body,
|
||||
upstream_headers: headers,
|
||||
|
|
|
|||
|
|
@ -55,6 +55,7 @@ fn write_response(body: &str) -> String {
|
|||
fn provider_config_resolves_anthropic_and_azure_ai() {
|
||||
assert!(messages_provider_config("anthropic").is_some());
|
||||
assert!(messages_provider_config("azure_ai").is_some());
|
||||
assert!(messages_provider_config("bedrock").is_some());
|
||||
assert!(messages_provider_config("openai").is_none());
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,8 @@
|
|||
use std::time::Duration;
|
||||
|
||||
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
|
||||
use litellm_core::messages::transformation::{
|
||||
AnthropicMessagesProviderConfig, MessagesAuthKind, MessagesStreaming,
|
||||
};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
pub struct MessagesRequest<'a> {
|
||||
|
|
@ -14,9 +16,10 @@ pub struct MessagesRequest<'a> {
|
|||
}
|
||||
|
||||
pub(crate) struct ProviderMessagesRequest {
|
||||
pub(crate) provider: String,
|
||||
pub(crate) model: String,
|
||||
pub(crate) config: &'static dyn AnthropicMessagesProviderConfig,
|
||||
pub(crate) auth_kind: MessagesAuthKind,
|
||||
pub(crate) streaming: MessagesStreaming,
|
||||
pub(crate) url: String,
|
||||
pub(crate) body: Value,
|
||||
pub(crate) upstream_headers: Vec<(String, String)>,
|
||||
|
|
|
|||
|
|
@ -106,16 +106,16 @@ impl RealTimeStreaming {
|
|||
/// `litellm_call_id`, replacing the gateway-generated fallback.
|
||||
fn on_session(&mut self, event: &RealtimeEvent) {
|
||||
let session = event.data.get("session").and_then(Value::as_object);
|
||||
if let Some(id) = session.and_then(|s| s.get("id")).and_then(Value::as_str) {
|
||||
if !id.is_empty() {
|
||||
self.id = id.to_string();
|
||||
self.litellm_call_id = id.to_string();
|
||||
}
|
||||
if let Some(id) = session.and_then(|s| s.get("id")).and_then(Value::as_str)
|
||||
&& !id.is_empty()
|
||||
{
|
||||
self.id = id.to_string();
|
||||
self.litellm_call_id = id.to_string();
|
||||
}
|
||||
if let Some(model) = session.and_then(|s| s.get("model")).and_then(Value::as_str) {
|
||||
if !model.is_empty() {
|
||||
self.model = model.to_string();
|
||||
}
|
||||
if let Some(model) = session.and_then(|s| s.get("model")).and_then(Value::as_str)
|
||||
&& !model.is_empty()
|
||||
{
|
||||
self.model = model.to_string();
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -33,35 +33,32 @@ 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(stream) => stream_response(stream),
|
||||
}
|
||||
}
|
||||
|
||||
fn stream_response(upstream: reqwest::Response) -> Result<Response, MessagesRouteError> {
|
||||
let content_type = 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| {
|
||||
MessagesRouteError(CoreError::InvalidResponse(format!(
|
||||
"invalid upstream response status: {error}"
|
||||
)))
|
||||
})?,
|
||||
)
|
||||
.header(CONTENT_TYPE, content_type);
|
||||
if let Some(value) = upstream.headers().get(CACHE_CONTROL) {
|
||||
response = response.header(CACHE_CONTROL, value);
|
||||
}
|
||||
response
|
||||
.body(Body::from_stream(upstream.bytes_stream()))
|
||||
.map_err(|error| {
|
||||
fn stream_response(
|
||||
stream: crate::messages::MessagesStream,
|
||||
) -> Result<Response, MessagesRouteError> {
|
||||
let response = Response::builder().status(StatusCode::OK).header(
|
||||
CONTENT_TYPE,
|
||||
HeaderValue::try_from(stream.content_type).map_err(|error| {
|
||||
MessagesRouteError(CoreError::InvalidResponse(format!(
|
||||
"failed to build streaming response: {error}"
|
||||
"invalid upstream content type: {error}"
|
||||
)))
|
||||
})
|
||||
})?,
|
||||
);
|
||||
let response = if let Some(value) = stream.cache_control.as_deref() {
|
||||
response.header(CACHE_CONTROL, value)
|
||||
} else {
|
||||
response
|
||||
};
|
||||
let body = Body::from_stream(stream.body);
|
||||
response.body(body).map_err(|error| {
|
||||
MessagesRouteError(CoreError::InvalidResponse(format!(
|
||||
"failed to build streaming response: {error}"
|
||||
)))
|
||||
})
|
||||
}
|
||||
|
||||
fn forwarded_headers(headers: &HeaderMap) -> Result<Option<Map<String, Value>>, CoreError> {
|
||||
|
|
|
|||
|
|
@ -5,11 +5,12 @@ use litellm_core::{CoreError, CoreResult};
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::messages::MessagesStream;
|
||||
use crate::messages::{MessagesRequest, execute_messages};
|
||||
|
||||
pub(crate) enum MessagesResponse {
|
||||
Json(Value),
|
||||
Stream(reqwest::Response),
|
||||
Stream(MessagesStream),
|
||||
}
|
||||
|
||||
pub async fn run(
|
||||
|
|
|
|||
|
|
@ -51,18 +51,17 @@ where
|
|||
provider_model,
|
||||
params.api_key.as_deref(),
|
||||
params.api_base.as_deref(),
|
||||
) {
|
||||
if let Some(handoff) = pool.take(&key) {
|
||||
return crate::io::realtime::realtime_warm(
|
||||
provider_model,
|
||||
handoff,
|
||||
idle_timeout,
|
||||
observe,
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
) && let Some(handoff) = pool.take(&key)
|
||||
{
|
||||
return crate::io::realtime::realtime_warm(
|
||||
provider_model,
|
||||
handoff,
|
||||
idle_timeout,
|
||||
observe,
|
||||
client_in,
|
||||
client_out,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
// Cold path: fresh dial (the original behavior).
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ repository.workspace = true
|
|||
rand.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
base64.workspace = true
|
||||
bytes = "1"
|
||||
thiserror.workspace = true
|
||||
sha2.workspace = true
|
||||
aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true }
|
||||
|
|
@ -17,6 +19,8 @@ aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rust
|
|||
aws-sigv4 = { version = "1.5.1", optional = true }
|
||||
aws-types = { version = "1.4.0", optional = true }
|
||||
aws-smithy-runtime-api = { version = "1.13.0", optional = true }
|
||||
aws-smithy-eventstream = { workspace = true, optional = true }
|
||||
aws-smithy-types = { version = "1.6.1", optional = true }
|
||||
|
||||
[features]
|
||||
default = []
|
||||
|
|
@ -27,6 +31,8 @@ bedrock-auth = [
|
|||
"dep:aws-sigv4",
|
||||
"dep:aws-types",
|
||||
"dep:aws-smithy-runtime-api",
|
||||
"dep:aws-smithy-eventstream",
|
||||
"dep:aws-smithy-types",
|
||||
]
|
||||
|
||||
[dev-dependencies]
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
use crate::error::CoreResult;
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::types::{AnthropicMessagesRequest, AnthropicMessagesResponse};
|
||||
|
||||
|
|
@ -8,6 +9,24 @@ pub enum MessagesAuthStrategy {
|
|||
Header(&'static str),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum MessagesAuthKind {
|
||||
ApiKey {
|
||||
strategy: MessagesAuthStrategy,
|
||||
accepts_bearer: bool,
|
||||
},
|
||||
AwsSigV4 {
|
||||
region: String,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub enum MessagesStreaming {
|
||||
Unsupported,
|
||||
SsePassthrough,
|
||||
BedrockEventStream,
|
||||
}
|
||||
|
||||
impl MessagesAuthStrategy {
|
||||
pub fn header_name(self) -> &'static str {
|
||||
match self {
|
||||
|
|
@ -22,6 +41,7 @@ pub trait AnthropicMessagesProviderConfig: Sync {
|
|||
&self,
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String>;
|
||||
|
||||
|
|
@ -29,14 +49,26 @@ pub trait AnthropicMessagesProviderConfig: Sync {
|
|||
&self,
|
||||
api_key: Option<&str>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String>;
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
MessagesAuthStrategy::Header("x-api-key")
|
||||
) -> CoreResult<String> {
|
||||
let _ = (api_key, env_lookup);
|
||||
Err(crate::error::CoreError::Auth(
|
||||
"provider does not use API key authentication".to_string(),
|
||||
))
|
||||
}
|
||||
|
||||
fn accepts_bearer_auth(&self) -> bool {
|
||||
false
|
||||
fn auth_kind(
|
||||
&self,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<MessagesAuthKind> {
|
||||
Ok(MessagesAuthKind::ApiKey {
|
||||
strategy: MessagesAuthStrategy::Header("x-api-key"),
|
||||
accepts_bearer: false,
|
||||
})
|
||||
}
|
||||
|
||||
fn streaming(&self) -> MessagesStreaming {
|
||||
MessagesStreaming::Unsupported
|
||||
}
|
||||
|
||||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
|
|
@ -53,6 +85,12 @@ pub trait AnthropicMessagesProviderConfig: Sync {
|
|||
Ok(request)
|
||||
}
|
||||
|
||||
fn upstream_body(&self, request: AnthropicMessagesRequest) -> CoreResult<Value> {
|
||||
serde_json::to_value(self.transform_request(request)?).map_err(|error| {
|
||||
CoreError::InvalidRequest(format!("failed to serialize messages request: {error}"))
|
||||
})
|
||||
}
|
||||
|
||||
fn transform_response(
|
||||
&self,
|
||||
_model: &str,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy};
|
||||
use crate::messages::transformation::{
|
||||
AnthropicMessagesProviderConfig, MessagesAuthKind, MessagesAuthStrategy, MessagesStreaming,
|
||||
};
|
||||
|
||||
const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY";
|
||||
const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE";
|
||||
|
|
@ -51,6 +53,7 @@ impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig {
|
|||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
Ok(complete_anthropic_url(api_base, env_lookup))
|
||||
|
|
@ -64,8 +67,19 @@ impl AnthropicMessagesProviderConfig for AnthropicMessagesConfig {
|
|||
resolve_anthropic_api_key(api_key, env_lookup)
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
MessagesAuthStrategy::Header("x-api-key")
|
||||
fn auth_kind(
|
||||
&self,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<MessagesAuthKind> {
|
||||
Ok(MessagesAuthKind::ApiKey {
|
||||
strategy: MessagesAuthStrategy::Header("x-api-key"),
|
||||
accepts_bearer: false,
|
||||
})
|
||||
}
|
||||
|
||||
fn streaming(&self) -> MessagesStreaming {
|
||||
MessagesStreaming::SsePassthrough
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -127,9 +141,14 @@ mod tests {
|
|||
|
||||
#[test]
|
||||
fn auth_strategy_and_default_headers_match_anthropic() {
|
||||
assert_eq!(
|
||||
ANTHROPIC_MESSAGES_CONFIG.auth_strategy().header_name(),
|
||||
"x-api-key"
|
||||
assert!(
|
||||
ANTHROPIC_MESSAGES_CONFIG
|
||||
.auth_kind("model", &|_| None)
|
||||
.expect("auth")
|
||||
.eq(&MessagesAuthKind::ApiKey {
|
||||
strategy: MessagesAuthStrategy::Header("x-api-key"),
|
||||
accepts_bearer: false,
|
||||
})
|
||||
);
|
||||
assert_eq!(
|
||||
ANTHROPIC_MESSAGES_CONFIG.default_headers(),
|
||||
|
|
|
|||
|
|
@ -1,5 +1,7 @@
|
|||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy};
|
||||
use crate::messages::transformation::{
|
||||
AnthropicMessagesProviderConfig, MessagesAuthKind, MessagesStreaming,
|
||||
};
|
||||
use crate::messages::types::{
|
||||
AnthropicMessage, AnthropicMessagesRequest, AnthropicMessagesResponse, ContentBlock,
|
||||
MessageContent, SystemPrompt,
|
||||
|
|
@ -146,6 +148,7 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
|
|||
&self,
|
||||
api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
complete_azure_anthropic_url(api_base, env_lookup)
|
||||
|
|
@ -159,12 +162,20 @@ impl AnthropicMessagesProviderConfig for AzureAnthropicMessagesConfig {
|
|||
resolve_azure_api_key(api_key, env_lookup)
|
||||
}
|
||||
|
||||
fn auth_strategy(&self) -> MessagesAuthStrategy {
|
||||
self.anthropic.auth_strategy()
|
||||
fn auth_kind(
|
||||
&self,
|
||||
model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<MessagesAuthKind> {
|
||||
let _ = (model, env_lookup);
|
||||
Ok(MessagesAuthKind::ApiKey {
|
||||
strategy: crate::messages::transformation::MessagesAuthStrategy::Header("x-api-key"),
|
||||
accepts_bearer: true,
|
||||
})
|
||||
}
|
||||
|
||||
fn accepts_bearer_auth(&self) -> bool {
|
||||
true
|
||||
fn streaming(&self) -> MessagesStreaming {
|
||||
MessagesStreaming::Unsupported
|
||||
}
|
||||
|
||||
fn default_headers(&self) -> &'static [(&'static str, &'static str)] {
|
||||
|
|
@ -292,15 +303,28 @@ mod tests {
|
|||
fn auth_strategy_is_x_api_key() {
|
||||
assert_eq!(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.auth_strategy()
|
||||
.header_name(),
|
||||
"x-api-key"
|
||||
.auth_kind("model", &|_| None)
|
||||
.expect("auth"),
|
||||
MessagesAuthKind::ApiKey {
|
||||
strategy: crate::messages::transformation::MessagesAuthStrategy::Header(
|
||||
"x-api-key"
|
||||
),
|
||||
accepts_bearer: true,
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_bearer_auth_for_entra_id() {
|
||||
assert!(AZURE_ANTHROPIC_MESSAGES_CONFIG.accepts_bearer_auth());
|
||||
assert!(matches!(
|
||||
AZURE_ANTHROPIC_MESSAGES_CONFIG
|
||||
.auth_kind("model", &|_| None)
|
||||
.expect("auth"),
|
||||
MessagesAuthKind::ApiKey {
|
||||
accepts_bearer: true,
|
||||
..
|
||||
}
|
||||
));
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
|
|||
|
|
@ -8,11 +8,8 @@ use crate::audio_transcription::types::{
|
|||
};
|
||||
use crate::error::{CoreError, CoreResult, json_type_name};
|
||||
|
||||
use super::aws_base::AwsAuthConfig;
|
||||
use super::constants::{
|
||||
AWS_REGION, AWS_REGION_NAME, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE,
|
||||
DEFAULT_BEDROCK_REGION,
|
||||
};
|
||||
use super::common_utils::{bedrock_model_id_and_region, resolve_bedrock_region};
|
||||
use super::constants::{BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE};
|
||||
|
||||
const SUPPORTED_PARAMS: &[&str] = &["language", "prompt", "temperature", "response_format"];
|
||||
|
||||
|
|
@ -21,64 +18,6 @@ pub static BEDROCK_AUDIO_TRANSCRIPTION_CONFIG: BedrockAudioTranscriptionConfig =
|
|||
|
||||
pub struct BedrockAudioTranscriptionConfig;
|
||||
|
||||
pub fn bedrock_model_id_and_region(model: &str) -> (String, Option<String>) {
|
||||
let mut stripped = model;
|
||||
for prefix in ["bedrock/converse/", "bedrock/", "converse/"] {
|
||||
if let Some(value) = stripped.strip_prefix(prefix) {
|
||||
stripped = value;
|
||||
break;
|
||||
}
|
||||
}
|
||||
let mut region = None;
|
||||
if let Some((candidate, remainder)) = stripped.split_once('/')
|
||||
&& is_bedrock_region(candidate)
|
||||
{
|
||||
region = Some(candidate.to_string());
|
||||
stripped = remainder;
|
||||
}
|
||||
for prefix in ["nova-2/", "nova/"] {
|
||||
if let Some(value) = stripped.strip_prefix(prefix) {
|
||||
stripped = value;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if region.is_none() {
|
||||
region = stripped
|
||||
.strip_prefix("arn:")
|
||||
.and_then(|value| value.split(':').nth(3))
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string);
|
||||
}
|
||||
(stripped.to_string(), region)
|
||||
}
|
||||
|
||||
fn is_bedrock_region(value: &str) -> bool {
|
||||
value.len() > 3
|
||||
&& value.contains('-')
|
||||
&& value
|
||||
.chars()
|
||||
.all(|char| char.is_ascii_alphanumeric() || char == '-')
|
||||
}
|
||||
|
||||
pub fn resolve_bedrock_region(
|
||||
model_region: Option<&str>,
|
||||
optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
if let Some(region) = optional_params
|
||||
.get("aws_region_name")
|
||||
.and_then(Value::as_str)
|
||||
{
|
||||
return region.to_string();
|
||||
}
|
||||
if let Some(region) = model_region {
|
||||
return region.to_string();
|
||||
}
|
||||
env_lookup(AWS_REGION_NAME)
|
||||
.or_else(|| env_lookup(AWS_REGION))
|
||||
.unwrap_or_else(|| DEFAULT_BEDROCK_REGION.to_string())
|
||||
}
|
||||
|
||||
fn audio_fields(audio: Value) -> CoreResult<(String, String)> {
|
||||
let object = audio.as_object().ok_or_else(|| CoreError::InvalidType {
|
||||
expected: "object",
|
||||
|
|
@ -203,32 +142,6 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
|
|||
}
|
||||
}
|
||||
|
||||
pub fn aws_auth_config(
|
||||
optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> AwsAuthConfig {
|
||||
let value = |key: &str| {
|
||||
optional_params
|
||||
.get(key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::to_string)
|
||||
};
|
||||
let env = |key: &str| env_lookup(key);
|
||||
AwsAuthConfig {
|
||||
access_key_id: value("aws_access_key_id").or_else(|| env("AWS_ACCESS_KEY_ID")),
|
||||
secret_access_key: value("aws_secret_access_key").or_else(|| env("AWS_SECRET_ACCESS_KEY")),
|
||||
session_token: value("aws_session_token").or_else(|| env("AWS_SESSION_TOKEN")),
|
||||
region_name: value("aws_region_name").or_else(|| env(AWS_REGION_NAME)),
|
||||
session_name: value("aws_session_name").or_else(|| env("AWS_SESSION_NAME")),
|
||||
profile_name: value("aws_profile_name").or_else(|| env("AWS_PROFILE_NAME")),
|
||||
role_name: value("aws_role_name").or_else(|| env("AWS_ROLE_NAME")),
|
||||
web_identity_token: value("aws_web_identity_token")
|
||||
.or_else(|| env("AWS_WEB_IDENTITY_TOKEN")),
|
||||
sts_endpoint: value("aws_sts_endpoint").or_else(|| env("AWS_STS_ENDPOINT")),
|
||||
external_id: value("aws_external_id").or_else(|| env("AWS_EXTERNAL_ID")),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
|
|
|||
|
|
@ -0,0 +1,94 @@
|
|||
use serde_json::{Map, Value};
|
||||
|
||||
use super::aws_base::AwsAuthConfig;
|
||||
use super::constants::{AWS_REGION, AWS_REGION_NAME, DEFAULT_BEDROCK_REGION};
|
||||
|
||||
pub fn bedrock_model_id_and_region(model: &str) -> (String, Option<String>) {
|
||||
let mut stripped = model;
|
||||
for prefix in [
|
||||
"bedrock/converse/",
|
||||
"bedrock/messages/",
|
||||
"bedrock/",
|
||||
"converse/",
|
||||
] {
|
||||
if let Some(value) = stripped.strip_prefix(prefix) {
|
||||
stripped = value;
|
||||
break;
|
||||
}
|
||||
}
|
||||
let mut region = None;
|
||||
if let Some((candidate, remainder)) = stripped.split_once('/')
|
||||
&& is_bedrock_region(candidate)
|
||||
{
|
||||
region = Some(candidate.to_string());
|
||||
stripped = remainder;
|
||||
}
|
||||
for prefix in ["nova-2/", "nova/"] {
|
||||
if let Some(value) = stripped.strip_prefix(prefix) {
|
||||
stripped = value;
|
||||
break;
|
||||
}
|
||||
}
|
||||
if region.is_none() {
|
||||
region = stripped
|
||||
.strip_prefix("arn:")
|
||||
.and_then(|value| value.split(':').nth(2))
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string);
|
||||
}
|
||||
(stripped.to_string(), region)
|
||||
}
|
||||
|
||||
pub fn is_bedrock_region(value: &str) -> bool {
|
||||
value.len() > 3
|
||||
&& value.contains('-')
|
||||
&& value
|
||||
.chars()
|
||||
.all(|character| character.is_ascii_alphanumeric() || character == '-')
|
||||
}
|
||||
|
||||
pub fn resolve_bedrock_region(
|
||||
model_region: Option<&str>,
|
||||
optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
optional_params
|
||||
.get("aws_region_name")
|
||||
.and_then(Value::as_str)
|
||||
.filter(|region| !region.trim().is_empty())
|
||||
.map(str::to_string)
|
||||
.or_else(|| model_region.map(str::to_string))
|
||||
.or_else(|| env_lookup(AWS_REGION_NAME))
|
||||
.or_else(|| env_lookup(AWS_REGION))
|
||||
.unwrap_or_else(|| DEFAULT_BEDROCK_REGION.to_string())
|
||||
}
|
||||
|
||||
pub fn aws_auth_config(
|
||||
optional_params: &Map<String, Value>,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
region: Option<&str>,
|
||||
) -> AwsAuthConfig {
|
||||
let value = |key: &str| {
|
||||
optional_params
|
||||
.get(key)
|
||||
.and_then(Value::as_str)
|
||||
.map(str::to_string)
|
||||
};
|
||||
let env = |key: &str| env_lookup(key);
|
||||
AwsAuthConfig {
|
||||
access_key_id: value("aws_access_key_id").or_else(|| env("AWS_ACCESS_KEY_ID")),
|
||||
secret_access_key: value("aws_secret_access_key").or_else(|| env("AWS_SECRET_ACCESS_KEY")),
|
||||
session_token: value("aws_session_token").or_else(|| env("AWS_SESSION_TOKEN")),
|
||||
region_name: region
|
||||
.map(str::to_string)
|
||||
.or_else(|| value("aws_region_name"))
|
||||
.or_else(|| env(AWS_REGION_NAME)),
|
||||
session_name: value("aws_session_name").or_else(|| env("AWS_SESSION_NAME")),
|
||||
profile_name: value("aws_profile_name").or_else(|| env("AWS_PROFILE_NAME")),
|
||||
role_name: value("aws_role_name").or_else(|| env("AWS_ROLE_NAME")),
|
||||
web_identity_token: value("aws_web_identity_token")
|
||||
.or_else(|| env("AWS_WEB_IDENTITY_TOKEN")),
|
||||
sts_endpoint: value("aws_sts_endpoint").or_else(|| env("AWS_STS_ENDPOINT")),
|
||||
external_id: value("aws_external_id").or_else(|| env("AWS_EXTERNAL_ID")),
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,2 @@
|
|||
pub mod streaming;
|
||||
pub mod transformation;
|
||||
|
|
@ -0,0 +1,223 @@
|
|||
use aws_smithy_eventstream::frame::{DecodedFrame, MessageFrameDecoder};
|
||||
use aws_smithy_types::event_stream::HeaderValue;
|
||||
use base64::Engine;
|
||||
use bytes::BytesMut;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct BedrockMessageEvent {
|
||||
pub event_type: String,
|
||||
pub chunk: Value,
|
||||
}
|
||||
|
||||
pub struct BedrockEventStreamDecoder {
|
||||
buffer: BytesMut,
|
||||
decoder: MessageFrameDecoder,
|
||||
}
|
||||
|
||||
impl BedrockEventStreamDecoder {
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
buffer: BytesMut::new(),
|
||||
decoder: MessageFrameDecoder::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn push(&mut self, bytes: &[u8]) -> CoreResult<Vec<BedrockMessageEvent>> {
|
||||
self.buffer.extend_from_slice(bytes);
|
||||
let mut events = Vec::new();
|
||||
loop {
|
||||
let frame = self
|
||||
.decoder
|
||||
.decode_frame(&mut self.buffer)
|
||||
.map_err(|error| {
|
||||
CoreError::InvalidResponse(format!(
|
||||
"invalid Bedrock event stream frame: {error}"
|
||||
))
|
||||
})?;
|
||||
let DecodedFrame::Complete(message) = frame else {
|
||||
break;
|
||||
};
|
||||
let message_type = message.headers().iter().find_map(|header| {
|
||||
(header.name().as_str() == ":message-type").then(|| match header.value() {
|
||||
HeaderValue::String(value) => value.as_str().to_string(),
|
||||
_ => String::new(),
|
||||
})
|
||||
});
|
||||
if matches!(message_type.as_deref(), Some("error" | "exception")) {
|
||||
let exception_type = message.headers().iter().find_map(|header| {
|
||||
(header.name().as_str() == ":exception-type").then(|| match header.value() {
|
||||
HeaderValue::String(value) => value.as_str().to_string(),
|
||||
_ => "unknown".to_string(),
|
||||
})
|
||||
});
|
||||
let payload = String::from_utf8_lossy(message.payload());
|
||||
let payload = payload.chars().take(512).collect::<String>();
|
||||
return Err(CoreError::InvalidResponse(format!(
|
||||
"Bedrock event stream {message_type:?} ({}){}",
|
||||
exception_type.as_deref().unwrap_or("unknown"),
|
||||
if payload.is_empty() {
|
||||
String::new()
|
||||
} else {
|
||||
format!(": {payload}")
|
||||
}
|
||||
)));
|
||||
}
|
||||
let envelope: Value = serde_json::from_slice(message.payload()).map_err(|error| {
|
||||
CoreError::InvalidResponse(format!("invalid Bedrock event payload: {error}"))
|
||||
})?;
|
||||
let encoded = envelope
|
||||
.get("bytes")
|
||||
.and_then(Value::as_str)
|
||||
.ok_or_else(|| {
|
||||
CoreError::InvalidResponse("Bedrock chunk has no bytes".to_string())
|
||||
})?;
|
||||
let decoded = base64::engine::general_purpose::STANDARD
|
||||
.decode(encoded)
|
||||
.map_err(|error| {
|
||||
CoreError::InvalidResponse(format!("invalid Bedrock chunk bytes: {error}"))
|
||||
})?;
|
||||
let mut chunk: Value = serde_json::from_slice(&decoded).map_err(|error| {
|
||||
CoreError::InvalidResponse(format!("invalid Anthropic stream chunk: {error}"))
|
||||
})?;
|
||||
let event_type = chunk
|
||||
.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.ok_or_else(|| {
|
||||
CoreError::InvalidResponse(
|
||||
"Bedrock Anthropic chunk has no event type".to_string(),
|
||||
)
|
||||
})?
|
||||
.to_string();
|
||||
if let Some(metrics) = chunk
|
||||
.as_object_mut()
|
||||
.and_then(|object| object.remove("amazon-bedrock-invocationMetrics"))
|
||||
{
|
||||
let object = chunk.as_object_mut().ok_or_else(|| {
|
||||
CoreError::InvalidResponse(
|
||||
"Bedrock Anthropic chunk must be an object".to_string(),
|
||||
)
|
||||
})?;
|
||||
let usage = object.entry("usage").or_insert_with(|| json!({}));
|
||||
if !usage.is_object() {
|
||||
*usage = json!({});
|
||||
}
|
||||
let usage = usage.as_object_mut().ok_or_else(|| {
|
||||
CoreError::InvalidResponse("Anthropic usage must be an object".to_string())
|
||||
})?;
|
||||
if let Some(input) = metrics.get("inputTokenCount") {
|
||||
usage.insert("input_tokens".to_string(), input.clone());
|
||||
}
|
||||
if let Some(output) = metrics.get("outputTokenCount") {
|
||||
usage.insert("output_tokens".to_string(), output.clone());
|
||||
}
|
||||
}
|
||||
events.push(BedrockMessageEvent { event_type, chunk });
|
||||
}
|
||||
Ok(events)
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for BedrockEventStreamDecoder {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
pub fn serialize_sse(event: &BedrockMessageEvent) -> CoreResult<Vec<u8>> {
|
||||
let data = serde_json::to_string(&event.chunk)
|
||||
.map_err(|error| CoreError::InvalidResponse(format!("invalid stream chunk: {error}")))?;
|
||||
Ok(format!("event: {}\ndata: {data}\n\n", event.event_type).into_bytes())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use aws_smithy_eventstream::frame::write_message_to;
|
||||
use aws_smithy_types::event_stream::{Header, Message};
|
||||
use bytes::BytesMut;
|
||||
|
||||
fn frame(payload: Value, error: bool) -> Vec<u8> {
|
||||
let mut message = Message::new(serde_json::to_vec(&payload).expect("payload"));
|
||||
message = message.add_header(Header::new(
|
||||
":event-type",
|
||||
HeaderValue::String("chunk".into()),
|
||||
));
|
||||
if error {
|
||||
message = message.add_header(Header::new(
|
||||
":message-type",
|
||||
HeaderValue::String("error".into()),
|
||||
));
|
||||
}
|
||||
let mut bytes = BytesMut::new();
|
||||
write_message_to(&message, &mut bytes).expect("frame");
|
||||
bytes.to_vec()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decoder_handles_every_split_boundary_and_metrics() {
|
||||
let chunk = serde_json::json!({
|
||||
"type": "content_block_delta",
|
||||
"usage": {"cache_read_input_tokens": 3},
|
||||
"amazon-bedrock-invocationMetrics": {
|
||||
"outputTokenCount": 7
|
||||
}
|
||||
});
|
||||
let encoded = base64::engine::general_purpose::STANDARD
|
||||
.encode(serde_json::to_vec(&chunk).expect("chunk"));
|
||||
let bytes = frame(serde_json::json!({"bytes": encoded}), false);
|
||||
for split in 1..bytes.len() {
|
||||
let mut decoder = BedrockEventStreamDecoder::new();
|
||||
let mut events = Vec::new();
|
||||
for part in bytes[..split].chunks(1) {
|
||||
events.extend(decoder.push(part).expect("partial frame"));
|
||||
}
|
||||
events.extend(decoder.push(&bytes[split..]).expect("final frame"));
|
||||
assert_eq!(events.len(), 1);
|
||||
assert_eq!(events[0].event_type, "content_block_delta");
|
||||
assert!(
|
||||
events[0]
|
||||
.chunk
|
||||
.get("amazon-bedrock-invocationMetrics")
|
||||
.is_none()
|
||||
);
|
||||
assert_eq!(events[0].chunk["usage"]["cache_read_input_tokens"], 3);
|
||||
assert!(events[0].chunk["usage"].get("input_tokens").is_none());
|
||||
assert_eq!(events[0].chunk["usage"]["output_tokens"], 7);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decoder_surfaces_error_frames() {
|
||||
let error = frame(serde_json::json!({"message": "bad"}), true);
|
||||
let result = BedrockEventStreamDecoder::new().push(&error);
|
||||
assert!(matches!(result, Err(CoreError::InvalidResponse(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn decoder_surfaces_exception_details() {
|
||||
let mut message =
|
||||
Message::new(serde_json::to_vec(&json!({"message": "throttled"})).expect("payload"));
|
||||
message = message
|
||||
.add_header(Header::new(
|
||||
":message-type",
|
||||
HeaderValue::String("exception".into()),
|
||||
))
|
||||
.add_header(Header::new(
|
||||
":exception-type",
|
||||
HeaderValue::String("ThrottlingException".into()),
|
||||
));
|
||||
let mut bytes = BytesMut::new();
|
||||
write_message_to(&message, &mut bytes).expect("frame");
|
||||
let result = BedrockEventStreamDecoder::new().push(&bytes);
|
||||
match result {
|
||||
Err(CoreError::InvalidResponse(message)) => {
|
||||
assert!(message.contains("ThrottlingException"));
|
||||
assert!(message.contains("throttled"));
|
||||
}
|
||||
other => panic!("unexpected result: {other:?}"),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -0,0 +1,291 @@
|
|||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::messages::transformation::{
|
||||
AnthropicMessagesProviderConfig, MessagesAuthKind, MessagesStreaming,
|
||||
};
|
||||
use crate::messages::types::AnthropicMessagesRequest;
|
||||
use crate::providers::bedrock::common_utils::{
|
||||
bedrock_model_id_and_region, resolve_bedrock_region,
|
||||
};
|
||||
use crate::providers::bedrock::constants::BEDROCK_RUNTIME_ENDPOINT_TEMPLATE;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
const BEDROCK_ANTHROPIC_VERSION: &str = "bedrock-2023-05-31";
|
||||
const BEDROCK_MESSAGES_SUFFIX: &str = "/invoke";
|
||||
const BEDROCK_STREAM_SUFFIX: &str = "/invoke-with-response-stream";
|
||||
const CACHE_TTL_5M: &str = "5m";
|
||||
const CACHE_TTL_1H: &str = "1h";
|
||||
const CONTEXT_EDIT_COMPACT: &str = "compact_20260112";
|
||||
const CONTEXT_EDIT_CLEAR_TOOLS: &str = "clear_tool_uses_20250919";
|
||||
const BETA_COMPACT: &str = "compact-2026-01-12";
|
||||
const BETA_CONTEXT: &str = "context-management-2025-06-27";
|
||||
const ALLOWED_FIELDS: &[&str] = &[
|
||||
"anthropic_version",
|
||||
"max_tokens",
|
||||
"messages",
|
||||
"anthropic_beta",
|
||||
"system",
|
||||
"stop_sequences",
|
||||
"temperature",
|
||||
"top_p",
|
||||
"top_k",
|
||||
"tools",
|
||||
"tool_choice",
|
||||
"thinking",
|
||||
"metadata",
|
||||
"output_config",
|
||||
"context_management",
|
||||
];
|
||||
|
||||
pub struct BedrockAnthropicMessagesConfig;
|
||||
|
||||
pub const BEDROCK_ANTHROPIC_MESSAGES_CONFIG: BedrockAnthropicMessagesConfig =
|
||||
BedrockAnthropicMessagesConfig;
|
||||
|
||||
pub fn complete_bedrock_url(
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> String {
|
||||
let (model_id, model_region) = bedrock_model_id_and_region(model);
|
||||
let region = resolve_bedrock_region(model_region.as_deref(), &Map::new(), env_lookup);
|
||||
let endpoint = api_base
|
||||
.map(str::trim)
|
||||
.filter(|value| !value.is_empty())
|
||||
.map(str::to_string)
|
||||
.unwrap_or_else(|| BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", ®ion));
|
||||
let suffix = if stream {
|
||||
BEDROCK_STREAM_SUFFIX
|
||||
} else {
|
||||
BEDROCK_MESSAGES_SUFFIX
|
||||
};
|
||||
format!(
|
||||
"{}/model/{model_id}{suffix}",
|
||||
endpoint.trim_end_matches('/')
|
||||
)
|
||||
}
|
||||
|
||||
fn is_claude_4_5(model: &str) -> bool {
|
||||
let model = model.to_ascii_lowercase();
|
||||
["sonnet", "haiku", "opus"].iter().any(|family| {
|
||||
let prefix = format!("claude-{family}");
|
||||
model.find(&prefix).is_some_and(|offset| {
|
||||
let suffix = &model[offset + prefix.len()..];
|
||||
matches!(suffix, "-4-5" | "-4.5" | ".4-5" | ".4.5")
|
||||
|| suffix.contains("-4-5-")
|
||||
|| suffix.contains("-4.5-")
|
||||
|| suffix.contains(".4-5-")
|
||||
|| suffix.contains(".4.5-")
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
fn sanitize_blocks(value: &mut Value, keep_ttl: bool) {
|
||||
match value {
|
||||
Value::Array(values) => values
|
||||
.iter_mut()
|
||||
.for_each(|value| sanitize_blocks(value, keep_ttl)),
|
||||
Value::Object(object) => {
|
||||
if let Some(cache_control) = object
|
||||
.get_mut("cache_control")
|
||||
.and_then(Value::as_object_mut)
|
||||
{
|
||||
cache_control.remove("scope");
|
||||
if !keep_ttl
|
||||
|| !matches!(
|
||||
cache_control.get("ttl").and_then(Value::as_str),
|
||||
Some(CACHE_TTL_5M | CACHE_TTL_1H)
|
||||
)
|
||||
{
|
||||
cache_control.remove("ttl");
|
||||
}
|
||||
}
|
||||
object
|
||||
.values_mut()
|
||||
.for_each(|value| sanitize_blocks(value, keep_ttl));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn sanitize_tools(value: &mut Value) {
|
||||
let Some(tools) = value.as_array_mut() else {
|
||||
return;
|
||||
};
|
||||
tools.iter_mut().for_each(|tool| {
|
||||
if let Some(object) = tool.as_object_mut() {
|
||||
object.remove("custom");
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
fn filter_context_management(request: &mut Map<String, Value>) {
|
||||
let Some(edits) = request
|
||||
.get_mut("context_management")
|
||||
.and_then(Value::as_object_mut)
|
||||
.and_then(|context| context.get_mut("edits"))
|
||||
.and_then(Value::as_array_mut)
|
||||
else {
|
||||
return;
|
||||
};
|
||||
edits.retain(|edit| {
|
||||
edit.as_object()
|
||||
.and_then(|object| object.get("type"))
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|kind| matches!(kind, CONTEXT_EDIT_COMPACT | CONTEXT_EDIT_CLEAR_TOOLS))
|
||||
});
|
||||
if edits.is_empty() {
|
||||
request.remove("context_management");
|
||||
return;
|
||||
}
|
||||
let allowed_edits = edits.clone();
|
||||
let betas = request
|
||||
.entry("anthropic_beta")
|
||||
.or_insert_with(|| Value::Array(Vec::new()));
|
||||
let Some(betas) = betas.as_array_mut() else {
|
||||
return;
|
||||
};
|
||||
let has_compact = allowed_edits.iter().any(|edit| {
|
||||
edit.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|kind| kind == CONTEXT_EDIT_COMPACT)
|
||||
});
|
||||
let has_context = allowed_edits.iter().any(|edit| {
|
||||
edit.get("type")
|
||||
.and_then(Value::as_str)
|
||||
.is_some_and(|kind| kind == CONTEXT_EDIT_CLEAR_TOOLS)
|
||||
});
|
||||
if has_compact && !betas.iter().any(|beta| beta.as_str() == Some(BETA_COMPACT)) {
|
||||
betas.push(Value::String(BETA_COMPACT.to_string()));
|
||||
}
|
||||
if has_context && !betas.iter().any(|beta| beta.as_str() == Some(BETA_CONTEXT)) {
|
||||
betas.push(Value::String(BETA_CONTEXT.to_string()));
|
||||
}
|
||||
}
|
||||
|
||||
pub fn transform_bedrock_request(
|
||||
model: &str,
|
||||
request: AnthropicMessagesRequest,
|
||||
) -> CoreResult<Value> {
|
||||
let mut value = serde_json::to_value(request)
|
||||
.map_err(|error| CoreError::InvalidRequest(format!("invalid messages request: {error}")))?;
|
||||
let object = value.as_object_mut().ok_or_else(|| {
|
||||
CoreError::InvalidRequest("messages request must be an object".to_string())
|
||||
})?;
|
||||
object.remove("model");
|
||||
object.remove("stream");
|
||||
object
|
||||
.entry("anthropic_version")
|
||||
.or_insert_with(|| Value::String(BEDROCK_ANTHROPIC_VERSION.to_string()));
|
||||
let keep_ttl = is_claude_4_5(model);
|
||||
for key in ["system", "messages", "tools"] {
|
||||
if let Some(value) = object.get_mut(key) {
|
||||
sanitize_blocks(value, keep_ttl);
|
||||
}
|
||||
}
|
||||
if let Some(tools) = object.get_mut("tools") {
|
||||
sanitize_tools(tools);
|
||||
}
|
||||
filter_context_management(object);
|
||||
object.retain(|key, _| ALLOWED_FIELDS.contains(&key.as_str()));
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
impl AnthropicMessagesProviderConfig for BedrockAnthropicMessagesConfig {
|
||||
fn complete_url(
|
||||
&self,
|
||||
api_base: Option<&str>,
|
||||
model: &str,
|
||||
stream: bool,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
Ok(complete_bedrock_url(api_base, model, stream, env_lookup))
|
||||
}
|
||||
|
||||
fn auth_kind(
|
||||
&self,
|
||||
model: &str,
|
||||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<MessagesAuthKind> {
|
||||
let (_, model_region) = bedrock_model_id_and_region(model);
|
||||
Ok(MessagesAuthKind::AwsSigV4 {
|
||||
region: resolve_bedrock_region(model_region.as_deref(), &Map::new(), env_lookup),
|
||||
})
|
||||
}
|
||||
|
||||
fn streaming(&self) -> MessagesStreaming {
|
||||
MessagesStreaming::BedrockEventStream
|
||||
}
|
||||
|
||||
fn upstream_body(&self, request: AnthropicMessagesRequest) -> CoreResult<Value> {
|
||||
let model = request.model.clone();
|
||||
transform_bedrock_request(&model, request)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
fn request(value: Value) -> AnthropicMessagesRequest {
|
||||
serde_json::from_value(value).expect("request")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn url_selects_invoke_endpoint_and_region_precedence() {
|
||||
let env = |key: &str| match key {
|
||||
"AWS_REGION_NAME" => Some("eu-west-1".to_string()),
|
||||
"AWS_REGION" => Some("ap-southeast-1".to_string()),
|
||||
_ => None,
|
||||
};
|
||||
assert_eq!(
|
||||
complete_bedrock_url(None, "bedrock/us-west-2/claude-test", false, &env),
|
||||
"https://bedrock-runtime.us-west-2.amazonaws.com/model/claude-test/invoke"
|
||||
);
|
||||
assert_eq!(
|
||||
complete_bedrock_url(Some("http://localhost:9000/"), "claude-test", true, &env),
|
||||
"http://localhost:9000/model/claude-test/invoke-with-response-stream"
|
||||
);
|
||||
assert_eq!(
|
||||
complete_bedrock_url(
|
||||
None,
|
||||
"arn:aws:bedrock:ap-south-1:123:model/foo",
|
||||
false,
|
||||
&env
|
||||
),
|
||||
"https://bedrock-runtime.ap-south-1.amazonaws.com/model/arn:aws:bedrock:ap-south-1:123:model/foo/invoke"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn transform_filters_and_sanitizes_bedrock_request() {
|
||||
let input = request(json!({
|
||||
"model": "claude-sonnet-4-5-20250929",
|
||||
"stream": true,
|
||||
"messages": [{"role":"user","content":[{"type":"text","text":"hi","cache_control":{"type":"ephemeral","scope":"request","ttl":"1h"}}]}],
|
||||
"tools": [{"name":"lookup","custom":{"defer_loading":true}}],
|
||||
"context_management": {"edits":[{"type":"unsupported"},{"type":"compact_20260112"}]},
|
||||
"service_tier": "auto",
|
||||
"unknown": true
|
||||
}));
|
||||
let transformed =
|
||||
transform_bedrock_request("claude-sonnet-4-5-20250929", input).expect("transform");
|
||||
let output = serde_json::to_value(transformed).expect("json");
|
||||
assert_eq!(output["anthropic_version"], BEDROCK_ANTHROPIC_VERSION);
|
||||
assert_eq!(
|
||||
output["messages"][0]["content"][0]["cache_control"]["ttl"],
|
||||
"1h"
|
||||
);
|
||||
assert!(output["tools"][0].get("custom").is_none());
|
||||
assert_eq!(
|
||||
output["context_management"]["edits"][0]["type"],
|
||||
CONTEXT_EDIT_COMPACT
|
||||
);
|
||||
assert_eq!(output["anthropic_beta"][0], BETA_COMPACT);
|
||||
assert!(output.get("service_tier").is_none());
|
||||
assert!(output.get("unknown").is_none());
|
||||
assert!(output.get("model").is_none());
|
||||
assert!(output.get("stream").is_none());
|
||||
}
|
||||
}
|
||||
|
|
@ -5,4 +5,7 @@
|
|||
#[cfg(feature = "bedrock-auth")]
|
||||
pub mod audio_transcription;
|
||||
pub mod aws_base;
|
||||
pub mod common_utils;
|
||||
mod constants;
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
pub mod messages;
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue