diff --git a/backend/Cargo.lock b/backend/Cargo.lock index e14885ef67..138083ac5d 100644 --- a/backend/Cargo.lock +++ b/backend/Cargo.lock @@ -6781,6 +6781,7 @@ version = "0.12.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "c7d6d2a27d57148378eb5e111173f4276ad26340ecc5c49a4a2152167a2d6a37" dependencies = [ + "async-compression 0.4.12", "base64 0.22.1", "bytes", "encoding_rs", diff --git a/backend/Cargo.toml b/backend/Cargo.toml index b48a7f7f01..4145ec96b9 100644 --- a/backend/Cargo.toml +++ b/backend/Cargo.toml @@ -156,7 +156,7 @@ mail-send = { version = "0.4.0", features = ["builder"], default-features=false urlencoding = "^2" url = "^2" async-oauth2 = "^0" -reqwest = { version = "^0.12", features = ["json", "stream"] } +reqwest = { version = "^0.12", features = ["json", "stream", "gzip"] } time = "0.3.16" serde_urlencoded = "^0" tokio-tar = "^0" diff --git a/backend/windmill-worker/src/snowflake_executor.rs b/backend/windmill-worker/src/snowflake_executor.rs index 349a2dcd73..ad0c77c0fb 100644 --- a/backend/windmill-worker/src/snowflake_executor.rs +++ b/backend/windmill-worker/src/snowflake_executor.rs @@ -4,6 +4,7 @@ use core::fmt::Write; use futures::future::BoxFuture; use futures::{FutureExt, TryFutureExt}; use jsonwebtoken::{encode, Algorithm, EncodingKey, Header}; +use reqwest::Response; use serde_json::{json, value::RawValue, Value}; use sha2::{Digest, Sha256}; use std::collections::HashMap; @@ -44,6 +45,12 @@ struct SnowflakeDatabase { struct SnowflakeResponse { data: Vec>, resultSetMetaData: SnowflakeResultSetMetaData, + statementHandle: String, +} + +#[derive(Deserialize, Debug)] +struct SnowflakeDataOnlyResponse { + data: Vec>, } #[derive(Deserialize, Debug)] @@ -51,6 +58,7 @@ struct SnowflakeResponse { struct SnowflakeResultSetMetaData { numRows: i64, rowType: Vec, + partitionInfo: Vec, } #[derive(Deserialize, Debug)] @@ -65,6 +73,38 @@ struct SnowflakeError { message: String, } +trait SnowflakeResponseExt { + async fn get_snowflake_response Deserialize<'a>>( + self, + ) -> windmill_common::error::Result; +} + +impl SnowflakeResponseExt for Result { + async fn get_snowflake_response Deserialize<'a>>( + self, + ) -> windmill_common::error::Result { + match self { + Ok(response) => match response.error_for_status_ref() { + Ok(_) => response + .json::() + .await + .map_err(|e| Error::ExecutionErr(e.to_string())), + Err(e) => { + let resp = response.text().await.unwrap_or("".to_string()); + match serde_json::from_str::(&resp) { + Ok(sf_err) => return Err(Error::ExecutionErr(sf_err.message)), + Err(_) => return Err(Error::ExecutionErr(e.to_string())), + } + } + }, + Err(e) => Err(Error::ExecutionErr(format!( + "Could not send request: {:?}", + e + ))), + } + } +} + fn do_snowflake_inner<'a>( query: &'a str, job_args: &HashMap, @@ -105,59 +145,66 @@ fn do_snowflake_inner<'a>( .json(&body) .send() .await - .map_err(|e| Error::ExecutionErr(e.to_string()))?; + .get_snowflake_response::() + .await?; - match response.error_for_status_ref() { - Ok(_) => { - let result = response - .json::() - .await - .map_err(|e| Error::ExecutionErr(e.to_string()))?; + if response.resultSetMetaData.numRows > 10000 { + return Err(Error::ExecutionErr( + "More than 10000 rows were requested, use LIMIT 10000 to limit the number of rows" + .to_string(), + )); + } + if let Some(column_order) = column_order { + *column_order = Some( + response + .resultSetMetaData + .rowType + .iter() + .map(|x| x.name.clone()) + .collect::>(), + ); + } - if result.resultSetMetaData.numRows > 10000 { - return Err(Error::ExecutionErr( - "More than 10000 rows were requested, use LIMIT 10000 to limit the number of rows".to_string(), - )); - } - if let Some(column_order) = column_order { - *column_order = Some( - result - .resultSetMetaData - .rowType - .iter() - .map(|x| x.name.clone()) - .collect::>(), - ); - } - let rows = to_raw_value( - &result - .data - .iter() - .map(|row| { - let mut row_map = serde_json::Map::new(); - row.iter() - .zip(result.resultSetMetaData.rowType.iter()) - .for_each(|(val, row_type)| { - row_map.insert( - row_type.name.clone(), - parse_val(&val, &row_type.r#type), - ); - }); - row_map - }) - .collect::>(), + let mut rows = response.data; + + if response.resultSetMetaData.partitionInfo.len() > 1 { + for idx in 1..response.resultSetMetaData.partitionInfo.len() { + let url = format!( + "https://{}.snowflakecomputing.com/api/v2/statements/{}", + account_identifier.to_uppercase(), + response.statementHandle ); + let response = HTTP_CLIENT + .get(url) + .bearer_auth(token) + .header("X-Snowflake-Authorization-Token-Type", "KEYPAIR_JWT") + .query(&[("partition", idx.to_string())]) + .send() + .await + .get_snowflake_response::() + .await?; - Ok(rows) - } - Err(e) => { - let resp = response.text().await.unwrap_or("".to_string()); - match serde_json::from_str::(&resp) { - Ok(sf_err) => Err(Error::ExecutionErr(sf_err.message)), - Err(_) => Err(Error::ExecutionErr(e.to_string())), - } + rows.extend(response.data); } } + + let rows = to_raw_value( + &rows + .iter() + .map(|row| { + let mut row_map = serde_json::Map::new(); + row.iter() + .zip(response.resultSetMetaData.rowType.iter()) + .for_each(|(val, row_type)| { + row_map + .insert(row_type.name.clone(), parse_val(&val, &row_type.r#type)); + }); + row_map + }) + .collect::>(), + ); + + Ok(rows) }; Ok(result_f.boxed())