mirror of
https://github.com/windmill-labs/windmill.git
synced 2026-08-22 00:01:34 +00:00
feat(ai): support IAM auth for bedrock provider (#7379)
* support iam for bedrock ai * lock * cleaning
This commit is contained in:
Generated
+1
@@ -15239,6 +15239,7 @@ dependencies = [
|
||||
"async-trait",
|
||||
"async_zip",
|
||||
"aws-config",
|
||||
"aws-credential-types",
|
||||
"aws-sdk-config",
|
||||
"aws-sdk-sqs",
|
||||
"aws-sdk-sso",
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
|
||||
@@ -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?;
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user