feat(ai): support IAM auth for bedrock provider (#7379)

* support iam for bedrock ai

* lock

* cleaning
This commit is contained in:
centdix
2025-12-16 23:00:14 +01:00
committed by GitHub
parent fa992db75e
commit b80d0e23f1
7 changed files with 206 additions and 22 deletions
+1
View File
@@ -15239,6 +15239,7 @@ dependencies = [
"async-trait",
"async_zip",
"aws-config",
"aws-credential-types",
"aws-sdk-config",
"aws-sdk-sqs",
"aws-sdk-sso",
+1
View File
@@ -149,6 +149,7 @@ rustls = { workspace = true }
aws-sigv4.workspace = true
aws-sdk-config.workspace = true
aws-config = { workspace = true, optional = true }
aws-credential-types.workspace = true
async-trait.workspace = true
google-cloud-pubsub = { workspace = true, optional = true }
google-cloud-googleapis = { workspace = true , optional = true }
+90 -20
View File
@@ -133,6 +133,10 @@ struct AIStandardResource {
api_key: Option<String>,
organization_id: Option<String>,
region: Option<String>,
#[serde(alias = "awsAccessKeyId")]
aws_access_key_id: Option<String>,
#[serde(alias = "awsSecretAccessKey")]
aws_secret_access_key: Option<String>,
}
#[derive(Deserialize, Debug)]
@@ -154,6 +158,9 @@ struct AIRequestConfig {
pub access_token: Option<String>,
pub organization_id: Option<String>,
pub user: Option<String>,
pub region: Option<String>,
pub aws_access_key_id: Option<String>,
pub aws_secret_access_key: Option<String>,
}
impl AIRequestConfig {
@@ -163,8 +170,18 @@ impl AIRequestConfig {
w_id: &str,
resource: AIResource,
) -> Result<Self> {
let (api_key, access_token, organization_id, base_url, user) = match resource {
let (
api_key,
access_token,
organization_id,
base_url,
user,
region,
aws_access_key_id,
aws_secret_access_key,
) = match resource {
AIResource::Standard(resource) => {
let region = resource.region.clone();
let base_url = provider
.get_base_url(resource.base_url, resource.region, db)
.await?;
@@ -178,8 +195,28 @@ impl AIRequestConfig {
} else {
None
};
let aws_access_key_id = if let Some(access_key_id) = resource.aws_access_key_id {
Some(get_variable_or_self(access_key_id, db, w_id).await?)
} else {
None
};
let aws_secret_access_key =
if let Some(secret_access_key) = resource.aws_secret_access_key {
Some(get_variable_or_self(secret_access_key, db, w_id).await?)
} else {
None
};
(api_key, None, organization_id, base_url, None)
(
api_key,
None,
organization_id,
base_url,
None,
region,
aws_access_key_id,
aws_secret_access_key,
)
}
AIResource::OAuth(resource) => {
let user = if let Some(user) = resource.user.clone() {
@@ -190,11 +227,20 @@ impl AIRequestConfig {
let token = Self::get_token_using_oauth(resource, db, w_id).await?;
let base_url = provider.get_base_url(None, None, db).await?;
(None, Some(token), None, base_url, user)
(None, Some(token), None, base_url, user, None, None, None)
}
};
Ok(Self { base_url, organization_id, api_key, access_token, user })
Ok(Self {
base_url,
organization_id,
api_key,
access_token,
user,
region,
aws_access_key_id,
aws_secret_access_key,
})
}
async fn get_token_using_oauth(
@@ -251,6 +297,10 @@ impl AIRequestConfig {
let is_anthropic_sdk = headers.get("X-Anthropic-SDK").is_some();
let is_bedrock = matches!(provider, AIProvider::AWSBedrock);
// Check if using IAM credentials for Bedrock (instead of bearer token)
let use_iam_auth =
is_bedrock && self.aws_access_key_id.is_some() && self.aws_secret_access_key.is_some();
// Handle AWS Bedrock transformation
let (url, body) = if is_bedrock && method != Method::GET {
let (model, transformed_body, is_streaming) =
@@ -282,7 +332,7 @@ impl AIRequestConfig {
tracing::debug!("AI request URL: {}", url);
let mut request = HTTP_CLIENT
.request(method, url)
.request(method.clone(), &url)
.header("content-type", "application/json");
for (header_name, header_value) in headers.iter() {
@@ -291,23 +341,43 @@ impl AIRequestConfig {
}
}
// For Bedrock with IAM credentials, sign the request using SigV4
if use_iam_auth {
let region = self.region.as_deref().ok_or_else(|| {
Error::internal_err("AWS region must be set for IAM authentication with Bedrock")
})?;
let signed_headers = bedrock::sign_bedrock_request(
method.as_str(),
&url,
&body,
self.aws_access_key_id.as_ref().unwrap(),
self.aws_secret_access_key.as_ref().unwrap(),
region,
)?;
for (header_name, header_value) in signed_headers {
request = request.header(header_name, header_value);
}
} else {
// For non-IAM auth, use bearer token or API key
if let Some(api_key) = self.api_key {
if is_azure {
request = request.header("api-key", api_key.clone())
} else {
request = request.header("authorization", format!("Bearer {}", api_key.clone()))
}
if is_anthropic {
request = request.header("X-API-Key", api_key);
}
}
if let Some(access_token) = self.access_token {
request = request.header("authorization", format!("Bearer {}", access_token))
}
}
request = request.body(body);
if let Some(api_key) = self.api_key {
if is_azure {
request = request.header("api-key", api_key.clone())
} else {
request = request.header("authorization", format!("Bearer {}", api_key.clone()))
}
if is_anthropic {
request = request.header("X-API-Key", api_key);
}
}
if let Some(access_token) = self.access_token {
request = request.header("authorization", format!("Bearer {}", access_token))
}
if let Some(org_id) = self.organization_id {
request = request.header("OpenAI-Organization", org_id);
}
+64
View File
@@ -1,9 +1,73 @@
use axum::body::Bytes;
use aws_sigv4::http_request::{sign, SignableBody, SignableRequest, SigningSettings};
use aws_sigv4::sign::v4;
use bytes;
use futures;
use std::time::SystemTime;
use uuid;
use windmill_common::error::{Error, Result};
/// Sign a request for AWS Bedrock using SigV4
///
/// Returns a vector of (header_name, header_value) tuples to add to the request
pub fn sign_bedrock_request(
method: &str,
uri: &str,
body: &[u8],
access_key_id: &str,
secret_access_key: &str,
region: &str,
) -> Result<Vec<(String, String)>> {
let identity = aws_credential_types::Credentials::new(
access_key_id,
secret_access_key,
None, // session token
None, // expiration
"windmill",
)
.into();
let signing_settings = SigningSettings::default();
let signing_params = v4::SigningParams::builder()
.identity(&identity)
.region(region)
.name("bedrock")
.time(SystemTime::now())
.settings(signing_settings)
.build()
.map_err(|e| Error::internal_err(format!("Failed to build signing params: {}", e)))?;
// Parse the URI to extract path and query
let parsed_uri: http::Uri = uri
.parse()
.map_err(|e| Error::internal_err(format!("Failed to parse URI: {}", e)))?;
let path_and_query = parsed_uri
.path_and_query()
.map(|pq| pq.as_str())
.unwrap_or("/");
let signable_request = SignableRequest::new(
method,
path_and_query,
std::iter::once(("host", parsed_uri.host().unwrap_or(""))),
SignableBody::Bytes(body),
)
.map_err(|e| Error::internal_err(format!("Failed to create signable request: {}", e)))?;
let (signing_instructions, _signature) = sign(signable_request, &signing_params.into())
.map_err(|e| Error::internal_err(format!("Failed to sign request: {}", e)))?
.into_parts();
// Collect the headers to add
let mut headers = Vec::new();
for (name, value) in signing_instructions.headers() {
headers.push((name.to_string(), value.to_string()));
}
Ok(headers)
}
/// Transform OpenAI format request to AWS Bedrock Converse format
/// Returns: (model_id, transformed_body, is_streaming)
pub fn transform_openai_to_bedrock(body: &[u8]) -> Result<(String, Bytes, bool)> {
@@ -56,6 +56,28 @@ impl BedrockClient {
Ok(Self { client: BedrockRuntimeClient::from_conf(config) })
}
pub async fn from_credentials(
access_key_id: String,
secret_access_key: String,
region: &str,
) -> Result<Self, Error> {
let credentials = aws_credential_types::Credentials::new(
access_key_id,
secret_access_key,
None, // session token
None, // expiration
"windmill",
);
let config = aws_sdk_bedrockruntime::config::Builder::new()
.region(aws_config::Region::new(region.to_string()))
.behavior_version(BehaviorVersion::latest())
.credentials_provider(credentials)
.build();
Ok(Self { client: BedrockRuntimeClient::from_conf(config) })
}
pub fn client(&self) -> &BedrockRuntimeClient {
&self.client
}
@@ -569,9 +591,21 @@ impl BedrockQueryBuilder {
client: &AuthedClient,
workspace_id: &str,
structured_output_tool_name: Option<&str>,
aws_access_key_id: Option<&str>,
aws_secret_access_key: Option<&str>,
) -> Result<ParsedResponse, Error> {
// Create Bedrock client with bearer token authentication
let bedrock_client = BedrockClient::from_bearer_token(api_key.to_string(), region).await?;
// Create Bedrock client - use IAM credentials if provided, otherwise fall back to bearer token
let bedrock_client = match (aws_access_key_id, aws_secret_access_key) {
(Some(access_key_id), Some(secret_access_key)) => {
BedrockClient::from_credentials(
access_key_id.to_string(),
secret_access_key.to_string(),
region,
)
.await?
}
_ => BedrockClient::from_bearer_token(api_key.to_string(), region).await?,
};
// Prepare messages: convert S3Objects to ImageUrls by downloading from S3
let prepared_messages = prepare_messages_for_api(messages, client, workspace_id).await?;
+12
View File
@@ -128,6 +128,10 @@ pub struct ProviderResource {
#[serde(alias = "baseUrl")]
pub base_url: Option<String>,
pub region: Option<String>,
#[serde(alias = "awsAccessKeyId")]
pub aws_access_key_id: Option<String>,
#[serde(alias = "awsSecretAccessKey")]
pub aws_secret_access_key: Option<String>,
}
#[derive(Deserialize, Debug)]
@@ -159,6 +163,14 @@ impl ProviderWithResource {
pub fn get_region(&self) -> Option<&str> {
self.resource.region.as_deref()
}
pub fn get_aws_access_key_id(&self) -> Option<&str> {
self.resource.aws_access_key_id.as_deref()
}
pub fn get_aws_secret_access_key(&self) -> Option<&str> {
self.resource.aws_secret_access_key.as_deref()
}
}
#[derive(Serialize)]
@@ -581,6 +581,8 @@ pub async fn run_agent(
client,
&job.workspace_id,
structured_output_tool_name.as_deref(),
args.provider.get_aws_access_key_id(),
args.provider.get_aws_secret_access_key(),
)
.await?
} else {