fix: support polling for long duration queries in snowflake (#7511)
This commit is contained in:
@@ -82,10 +82,76 @@ struct SnowflakeError {
|
||||
message: String,
|
||||
}
|
||||
|
||||
trait SnowflakeResponseExt {
|
||||
async fn parse_snowflake_response<T: for<'a> Deserialize<'a>>(
|
||||
self,
|
||||
) -> windmill_common::error::Result<T>;
|
||||
#[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<SnowflakeResponse> {
|
||||
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<Response, reqwest::Error> {
|
||||
async fn parse_snowflake_response<T: for<'a> Deserialize<'a>>(
|
||||
self,
|
||||
) -> windmill_common::error::Result<T> {
|
||||
let response = handle_snowflake_result(self).await?;
|
||||
response
|
||||
.json::<T>()
|
||||
.await
|
||||
.map_err(|e| Error::ExecutionErr(e.to_string()))
|
||||
}
|
||||
}
|
||||
|
||||
fn do_snowflake_inner<'a>(
|
||||
query: &'a str,
|
||||
job_args: &HashMap<String, Value>,
|
||||
@@ -134,6 +188,7 @@ fn do_snowflake_inner<'a>(
|
||||
http_client: &'a Client,
|
||||
s3: Option<S3ModeWorkerData>,
|
||||
reserved_variables: &HashMap<String, String>,
|
||||
deadline: std::time::Instant,
|
||||
) -> windmill_common::error::Result<BoxFuture<'a, windmill_common::error::Result<Vec<Box<RawValue>>>>>
|
||||
{
|
||||
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::<SnowflakeResponse>()
|
||||
.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::<SnowflakeResponse>(&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::<SnowflakeDataOnlyResponse>()
|
||||
.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);
|
||||
|
||||
Reference in New Issue
Block a user