diff --git a/backend/windmill-worker/src/snowflake_executor.rs b/backend/windmill-worker/src/snowflake_executor.rs index dc0b46e094..e1dad165c0 100644 --- a/backend/windmill-worker/src/snowflake_executor.rs +++ b/backend/windmill-worker/src/snowflake_executor.rs @@ -82,10 +82,76 @@ struct SnowflakeError { message: String, } -trait SnowflakeResponseExt { - async fn parse_snowflake_response Deserialize<'a>>( - self, - ) -> windmill_common::error::Result; +#[derive(Deserialize, Debug)] +#[serde(rename_all = "camelCase")] +struct SnowflakeAsyncResponse { + statement_handle: String, +} + +async fn poll_snowflake_async_query( + http_client: &Client, + account_identifier: &str, + statement_handle: &str, + token: &str, + token_is_keypair: bool, + deadline: std::time::Instant, +) -> windmill_common::error::Result { + let url = format!( + "https://{}.snowflakecomputing.com/api/v2/statements/{}", + account_identifier.to_uppercase(), + statement_handle + ); + + loop { + if std::time::Instant::now() > deadline { + return Err(Error::ExecutionErr( + "Snowflake query timed out while polling for results".to_string(), + )); + } + + let mut request = http_client.get(&url).bearer_auth(token); + if token_is_keypair { + request = request.header("X-Snowflake-Authorization-Token-Type", "KEYPAIR_JWT"); + } + + let response = request.send().await.map_err(|e| { + Error::ExecutionErr(format!("Could not poll Snowflake status: {:?}", e)) + })?; + + let status = response.status(); + let body = response.text().await.map_err(|e| { + Error::ExecutionErr(format!("error reading poll response body: {}", e)) + })?; + + tracing::debug!("Snowflake poll response status: {}, body: {}", status, &body[..body.len().min(500)]); + + if status == reqwest::StatusCode::ACCEPTED { + // Still running, wait and poll again + tracing::info!("Snowflake query still running, polling again in 1s..."); + tokio::time::sleep(std::time::Duration::from_secs(1)).await; + continue; + } + + if !status.is_success() { + return Err(Error::ExecutionErr(format!( + "Snowflake poll returned error status {}: {}", + status, + &body[..body.len().min(500)] + ))); + } + + // Query completed, parse the response + let response: SnowflakeResponse = serde_json::from_str(&body).map_err(|e| { + Error::ExecutionErr(format!( + "error decoding poll response: {}. Status: {}. Body preview: {}", + e, + status, + &body[..body.len().min(500)] + )) + })?; + + return Ok(response); + } } async fn handle_snowflake_result( @@ -109,18 +175,6 @@ async fn handle_snowflake_result( } } -impl SnowflakeResponseExt for Result { - async fn parse_snowflake_response Deserialize<'a>>( - self, - ) -> windmill_common::error::Result { - let response = handle_snowflake_result(self).await?; - response - .json::() - .await - .map_err(|e| Error::ExecutionErr(e.to_string())) - } -} - fn do_snowflake_inner<'a>( query: &'a str, job_args: &HashMap, @@ -134,6 +188,7 @@ fn do_snowflake_inner<'a>( http_client: &'a Client, s3: Option, reserved_variables: &HashMap, + deadline: std::time::Instant, ) -> windmill_common::error::Result>>>> { let sig = parse_snowflake_sig(&query) @@ -180,12 +235,85 @@ fn do_snowflake_inner<'a>( let result = request.send().await; if skip_collect { - handle_snowflake_result(result).await?; + // Still need to handle async (202) responses even when not collecting results + let raw_response = handle_snowflake_result(result).await?; + let status = raw_response.status(); + + if status == reqwest::StatusCode::ACCEPTED { + let body = raw_response.text().await.map_err(|e| { + Error::ExecutionErr(format!("error reading response body: {}", e)) + })?; + let async_resp: SnowflakeAsyncResponse = serde_json::from_str(&body).map_err(|e| { + Error::ExecutionErr(format!( + "error decoding async response: {}. Body preview: {}", + e, + &body[..body.len().min(500)] + )) + })?; + + tracing::info!( + "Snowflake statement running asynchronously, polling for completion (handle: {})", + async_resp.statement_handle + ); + + // Poll until complete, but discard the results + poll_snowflake_async_query( + http_client, + account_identifier, + &async_resp.statement_handle, + token, + token_is_keypair, + deadline, + ) + .await?; + } + Ok(vec![]) } else { - let response = result - .parse_snowflake_response::() - .await?; + // Handle both sync (200) and async (202) responses + let raw_response = handle_snowflake_result(result).await?; + let status = raw_response.status(); + let body = raw_response.text().await.map_err(|e| { + Error::ExecutionErr(format!("error reading response body: {}", e)) + })?; + + tracing::debug!("Snowflake response status: {}, body: {}", status, &body[..body.len().min(1000)]); + + let response = if status == reqwest::StatusCode::ACCEPTED { + // Async execution - need to poll for results + let async_resp: SnowflakeAsyncResponse = serde_json::from_str(&body).map_err(|e| { + Error::ExecutionErr(format!( + "error decoding async response: {}. Body preview: {}", + e, + &body[..body.len().min(500)] + )) + })?; + + tracing::info!( + "Snowflake query running asynchronously, polling for results (handle: {})", + async_resp.statement_handle + ); + + poll_snowflake_async_query( + http_client, + account_identifier, + &async_resp.statement_handle, + token, + token_is_keypair, + deadline, + ) + .await? + } else { + // Sync execution - parse directly + serde_json::from_str::(&body).map_err(|e| { + Error::ExecutionErr(format!( + "error decoding response body: {}. Status: {}. Body preview: {}", + e, + status, + &body[..body.len().min(500)] + )) + })? + }; if s3.is_none() && response.resultSetMetaData.numRows > 10000 { return Err(Error::ExecutionErr( @@ -233,13 +361,113 @@ fn do_snowflake_inner<'a>( request.header("X-Snowflake-Authorization-Token-Type", "KEYPAIR_JWT"); } - let response = request - .send() - .await - .parse_snowflake_response::() - .await?; + let result = request.send().await; + let raw_response = match handle_snowflake_result(result).await { + Ok(r) => r, + Err(e) => { + yield Err(e); + return; + } + }; + let status = raw_response.status(); + let body = match raw_response.text().await { + Ok(b) => b, + Err(e) => { + yield Err(Error::ExecutionErr(format!("error reading partition response: {}", e))); + return; + } + }; - for row in response.data { + // Handle async (202) response for partition fetch + let partition_data: SnowflakeDataOnlyResponse = if status == reqwest::StatusCode::ACCEPTED { + // Poll until complete - partition fetches should be fast, but handle async just in case + let mut poll_body = body; + loop { + if std::time::Instant::now() > deadline { + yield Err(Error::ExecutionErr( + "Snowflake partition fetch timed out while polling".to_string(), + )); + return; + } + + let async_resp: SnowflakeAsyncResponse = match serde_json::from_str(&poll_body) { + Ok(r) => r, + Err(e) => { + yield Err(Error::ExecutionErr(format!( + "error decoding async partition response: {}", + e + ))); + return; + } + }; + + tokio::time::sleep(std::time::Duration::from_secs(1)).await; + + let poll_url = format!( + "https://{}.snowflakecomputing.com/api/v2/statements/{}", + cloned_account_identifier.to_uppercase(), + async_resp.statement_handle + ); + let mut poll_request = HTTP_CLIENT.get(&poll_url).bearer_auth(cloned_token.as_str()); + if token_is_keypair { + poll_request = poll_request.header("X-Snowflake-Authorization-Token-Type", "KEYPAIR_JWT"); + } + + let poll_response = match poll_request.send().await { + Ok(r) => r, + Err(e) => { + yield Err(Error::ExecutionErr(format!("partition poll error: {:?}", e))); + return; + } + }; + + let poll_status = poll_response.status(); + poll_body = match poll_response.text().await { + Ok(b) => b, + Err(e) => { + yield Err(Error::ExecutionErr(format!("error reading partition poll response: {}", e))); + return; + } + }; + + if poll_status == reqwest::StatusCode::ACCEPTED { + continue; + } + + if !poll_status.is_success() { + yield Err(Error::ExecutionErr(format!( + "partition poll returned error: {}", + &poll_body[..poll_body.len().min(500)] + ))); + return; + } + + match serde_json::from_str(&poll_body) { + Ok(r) => break r, + Err(e) => { + yield Err(Error::ExecutionErr(format!( + "error decoding partition poll response: {}", + e + ))); + return; + } + } + } + } else { + match serde_json::from_str(&body) { + Ok(r) => r, + Err(e) => { + yield Err(Error::ExecutionErr(format!( + "error decoding partition response: {}. Body: {}", + e, + &body[..body.len().min(500)] + ))); + return; + } + } + }; + + for row in partition_data.data { yield Ok(row); } } @@ -416,6 +644,8 @@ pub async fn do_snowflake( let http_client = build_http_client(timeout_duration)?; + let deadline = std::time::Instant::now() + timeout_duration; + let reserved_variables = get_reserved_variables(job, &client.token, conn, parent_runnable_path).await?; @@ -444,6 +674,7 @@ pub async fn do_snowflake( &http_client, s3.clone(), &reserved_variables, + deadline, )? .await?; results.push(result);