diff --git a/backend/Cargo.lock b/backend/Cargo.lock index 359021af50..943495da13 100644 --- a/backend/Cargo.lock +++ b/backend/Cargo.lock @@ -15239,6 +15239,7 @@ dependencies = [ "async-trait", "async_zip", "aws-config", + "aws-credential-types", "aws-sdk-config", "aws-sdk-sqs", "aws-sdk-sso", diff --git a/backend/windmill-api/Cargo.toml b/backend/windmill-api/Cargo.toml index 7b8888b38e..133a895fcf 100644 --- a/backend/windmill-api/Cargo.toml +++ b/backend/windmill-api/Cargo.toml @@ -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 } diff --git a/backend/windmill-api/src/ai.rs b/backend/windmill-api/src/ai.rs index 6552dc9b4a..866b29404f 100644 --- a/backend/windmill-api/src/ai.rs +++ b/backend/windmill-api/src/ai.rs @@ -133,6 +133,10 @@ struct AIStandardResource { api_key: Option, organization_id: Option, region: Option, + #[serde(alias = "awsAccessKeyId")] + aws_access_key_id: Option, + #[serde(alias = "awsSecretAccessKey")] + aws_secret_access_key: Option, } #[derive(Deserialize, Debug)] @@ -154,6 +158,9 @@ struct AIRequestConfig { pub access_token: Option, pub organization_id: Option, pub user: Option, + pub region: Option, + pub aws_access_key_id: Option, + pub aws_secret_access_key: Option, } impl AIRequestConfig { @@ -163,8 +170,18 @@ impl AIRequestConfig { w_id: &str, resource: AIResource, ) -> Result { - 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); } diff --git a/backend/windmill-api/src/bedrock.rs b/backend/windmill-api/src/bedrock.rs index ea93b5b3ea..81ee7056b6 100644 --- a/backend/windmill-api/src/bedrock.rs +++ b/backend/windmill-api/src/bedrock.rs @@ -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> { + 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)> { diff --git a/backend/windmill-worker/src/ai/providers/bedrock.rs b/backend/windmill-worker/src/ai/providers/bedrock.rs index 57f44a8ab7..e0078ad931 100644 --- a/backend/windmill-worker/src/ai/providers/bedrock.rs +++ b/backend/windmill-worker/src/ai/providers/bedrock.rs @@ -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 { + 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 { - // 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?; diff --git a/backend/windmill-worker/src/ai/types.rs b/backend/windmill-worker/src/ai/types.rs index c68bf16a0b..724ea6bbf6 100644 --- a/backend/windmill-worker/src/ai/types.rs +++ b/backend/windmill-worker/src/ai/types.rs @@ -128,6 +128,10 @@ pub struct ProviderResource { #[serde(alias = "baseUrl")] pub base_url: Option, pub region: Option, + #[serde(alias = "awsAccessKeyId")] + pub aws_access_key_id: Option, + #[serde(alias = "awsSecretAccessKey")] + pub aws_secret_access_key: Option, } #[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)] diff --git a/backend/windmill-worker/src/ai_executor.rs b/backend/windmill-worker/src/ai_executor.rs index ca62324bdb..00a58e2c41 100644 --- a/backend/windmill-worker/src/ai_executor.rs +++ b/backend/windmill-worker/src/ai_executor.rs @@ -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 {