diff --git a/backend/oauth_connect.json b/backend/oauth_connect.json index 994d9999fd..0608f0d3c4 100644 --- a/backend/oauth_connect.json +++ b/backend/oauth_connect.json @@ -171,5 +171,6 @@ "user-library-modify", "user-library-read" ] - } -} \ No newline at end of file + }, + "snowflake_oauth": {} +} diff --git a/backend/windmill-worker/src/snowflake_executor.rs b/backend/windmill-worker/src/snowflake_executor.rs index f89b29832a..9f8d3e12fd 100644 --- a/backend/windmill-worker/src/snowflake_executor.rs +++ b/backend/windmill-worker/src/snowflake_executor.rs @@ -32,9 +32,9 @@ struct Claims { #[derive(Deserialize)] struct SnowflakeDatabase { account_identifier: String, - public_key: String, - private_key: String, - username: String, + public_key: Option, + private_key: Option, + username: Option, database: Option, schema: Option, warehouse: Option, @@ -119,6 +119,7 @@ fn do_snowflake_inner<'a>( mut body: serde_json::Map, account_identifier: &'a str, token: &'a str, + token_is_keypair: bool, column_order: Option<&'a mut Option>>, skip_collect: bool, ) -> windmill_common::error::Result>>> { @@ -144,16 +145,19 @@ fn do_snowflake_inner<'a>( } let result_f = async move { - let result = HTTP_CLIENT + let mut request = HTTP_CLIENT .post(format!( "https://{}.snowflakecomputing.com/api/v2/statements/", account_identifier.to_uppercase() )) .bearer_auth(token) - .header("X-Snowflake-Authorization-Token-Type", "KEYPAIR_JWT") - .json(&body) - .send() - .await; + .json(&body); + + if token_is_keypair { + request = request.header("X-Snowflake-Authorization-Token-Type", "KEYPAIR_JWT"); + } + + let result = request.send().await; if skip_collect { handle_snowflake_result(result).await?; @@ -189,11 +193,17 @@ fn do_snowflake_inner<'a>( account_identifier.to_uppercase(), response.statementHandle ); - let response = HTTP_CLIENT + let mut request = HTTP_CLIENT .get(url) .bearer_auth(token) - .header("X-Snowflake-Authorization-Token-Type", "KEYPAIR_JWT") - .query(&[("partition", idx.to_string())]) + .query(&[("partition", idx.to_string())]); + + if token_is_keypair { + request = + request.header("X-Snowflake-Authorization-Token-Type", "KEYPAIR_JWT"); + } + + let response = request .send() .await .parse_snowflake_response::() @@ -258,7 +268,7 @@ pub async fn do_snowflake( snowflake_args.get("database").cloned() }; - let database = if let Some(db) = db_arg { + let database = if let Some(ref db) = db_arg { serde_json::from_value::(db.clone()) .map_err(|e| Error::ExecutionErr(e.to_string()))? } else { @@ -267,37 +277,59 @@ pub async fn do_snowflake( let annotations = windmill_common::worker::SqlAnnotations::parse(query); - let qualified_username = format!( - "{}.{}", - database.account_identifier.split('.').next().unwrap_or(""), // get first part of account identifier - database.username - ) - .to_uppercase(); + // Check if the token is present in db_arg and use it if available + let (token, token_is_keypair) = if let Some(token) = db_arg + .as_ref() + .and_then(|db| db.get("token")) + .and_then(|t| t.as_str()) + .filter(|t| !t.is_empty()) + { + tracing::debug!("Using oauth token from db_arg"); + (token.to_string(), false) + } else { + tracing::debug!("Generating new oauth token"); - let public_key = pem::parse(database.public_key.as_bytes()).map_err(|e| { - Error::ExecutionErr(format!("Failed to parse public key: {}", e.to_string())) - })?; - let mut public_key_hash = Sha256::new(); - public_key_hash.update(public_key.contents()); + let qualified_username = format!( + "{}.{}", + database.account_identifier.split('.').next().unwrap_or(""), + database.username.as_deref().unwrap_or("") + ) + .to_uppercase(); - let public_key_fp = engine::general_purpose::STANDARD.encode(public_key_hash.finalize()); + let public_key = match database.public_key.as_deref() { + Some(key) => pem::parse(key.as_bytes()).map_err(|e| { + Error::ExecutionErr(format!("Failed to parse public key: {}", e.to_string())) + })?, + None => return Err(Error::ExecutionErr("Public key is missing".to_string())), + }; + let mut public_key_hash = Sha256::new(); + public_key_hash.update(public_key.contents()); - let iss = format!("{}.SHA256:{}", qualified_username, public_key_fp); + let public_key_fp = engine::general_purpose::STANDARD.encode(public_key_hash.finalize()); - let claims = Claims { - iss: iss, - sub: qualified_username, - iat: chrono::Utc::now().timestamp(), - exp: (chrono::Utc::now() + chrono::Duration::try_hours(1).unwrap()).timestamp(), + let iss = format!("{}.SHA256:{}", qualified_username, public_key_fp); + + let claims = Claims { + iss: iss, + sub: qualified_username, + iat: chrono::Utc::now().timestamp(), + exp: (chrono::Utc::now() + chrono::Duration::try_hours(1).unwrap()).timestamp(), + }; + + let private_key = match database.private_key.as_deref() { + Some(key) => EncodingKey::from_rsa_pem(key.as_bytes()).map_err(|e| { + Error::ExecutionErr(format!("Failed to parse private key: {}", e.to_string())) + })?, + None => return Err(Error::ExecutionErr("Private key is missing".to_string())), + }; + + ( + encode(&Header::new(Algorithm::RS256), &claims, &private_key) + .map_err(|e| Error::ExecutionErr(e.to_string()))?, + true, + ) }; - let private_key = EncodingKey::from_rsa_pem(database.private_key.as_bytes()).map_err(|e| { - Error::ExecutionErr(format!("Failed to parse private key: {}", e.to_string())) - })?; - - let token = encode(&Header::new(Algorithm::RS256), &claims, &private_key) - .map_err(|e| Error::ExecutionErr(e.to_string()))?; - tracing::debug!("Snowflake token: {}", token); let mut body = serde_json::Map::new(); @@ -344,6 +376,7 @@ pub async fn do_snowflake( body.clone(), &database.account_identifier, &token, + token_is_keypair, None, annotations.return_last_result && i < queries.len() - 1, ) @@ -371,6 +404,7 @@ pub async fn do_snowflake( body.clone(), &database.account_identifier, &token, + token_is_keypair, Some(column_order), false, )? diff --git a/frontend/src/lib/components/AppConnectInner.svelte b/frontend/src/lib/components/AppConnectInner.svelte index e568589b49..8ea77cc12d 100644 --- a/frontend/src/lib/components/AppConnectInner.svelte +++ b/frontend/src/lib/components/AppConnectInner.svelte @@ -250,6 +250,14 @@ throw Error(`Resource at path ${path} already exists. Delete it or pick another path`) } + + if (resourceType == 'snowflake_oauth') { + const account_identifier = extra_params.find(([key, _]) => key == 'account_identifier') + if (account_identifier) { + args['account_identifier'] = account_identifier[1] + } + } + let account: number | undefined = undefined if (valueToken?.expires_in != undefined) { account = Number( diff --git a/frontend/src/lib/components/InstanceSettings.svelte b/frontend/src/lib/components/InstanceSettings.svelte index e5f4bb0b4d..9cdd9f165f 100644 --- a/frontend/src/lib/components/InstanceSettings.svelte +++ b/frontend/src/lib/components/InstanceSettings.svelte @@ -108,9 +108,24 @@ loading = false latestKeyRenewalAttempt = await SettingService.getLatestKeyRenewalAttempt() + + // populate snowflake account identifier from db + const account_identifier = + oauths?.snowflake_oauth?.connect_config?.extra_params?.account_identifier + if (account_identifier) { + snowflakeAccountIdentifier = account_identifier + } } export async function saveSettings() { + if ( + oauths?.snowflake_oauth && + oauths?.snowflake_oauth?.connect_config?.extra_params?.account_identifier !== + snowflakeAccountIdentifier + ) { + setupSnowflakeUrls() + } + let shouldReloadPage = false if (values) { const allSettings = Object.values(settings).flatMap((x) => Object.entries(x)) @@ -217,7 +232,8 @@ 'linkedin', 'quickbooks', 'visma', - 'spotify' + 'spotify', + 'snowflake_oauth' ] let oauth_name = undefined @@ -269,6 +285,23 @@ } return true } + + let snowflakeAccountIdentifier = '' + + function setupSnowflakeUrls() { + // strip all whitespaces from account identifier + snowflakeAccountIdentifier = snowflakeAccountIdentifier.replace(/\s/g, '') + + const connect_config = { + scopes: [], + auth_url: `https://${snowflakeAccountIdentifier}.snowflakecomputing.com/oauth/authorize`, + token_url: `https://${snowflakeAccountIdentifier}.snowflakecomputing.com/oauth/token-request`, + req_body_auth: false, + extra_params: { account_identifier: snowflakeAccountIdentifier }, + extra_params_callback: {} + } + oauths['snowflake_oauth'].connect_config = connect_config + }
@@ -513,6 +546,22 @@ {#if !windmillBuiltins.includes(k) && k != 'slack'} {/if} + {#if k == 'snowflake_oauth'} + + {/if}
{/if}