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