diff --git a/backend/Cargo.lock b/backend/Cargo.lock index eab6d6d2f5..97203b8559 100644 --- a/backend/Cargo.lock +++ b/backend/Cargo.lock @@ -14594,6 +14594,7 @@ dependencies = [ "crc", "cron", "croner", + "datafusion", "futures", "futures-core", "gethostname", @@ -14633,6 +14634,8 @@ dependencies = [ "thiserror 2.0.12", "tikv-jemalloc-ctl", "tokio", + "tokio-stream", + "tokio-util", "tonic", "tracing", "tracing-appender", @@ -14641,6 +14644,7 @@ dependencies = [ "tracing-subscriber", "uuid", "windmill-macros", + "windmill-parser-sql", ] [[package]] @@ -14944,6 +14948,7 @@ version = "1.490.0" dependencies = [ "anyhow", "async-recursion", + "async-stream", "backon", "base64 0.22.1", "bit-vec 0.6.3", diff --git a/backend/parsers/windmill-parser-sql/src/lib.rs b/backend/parsers/windmill-parser-sql/src/lib.rs index 5b9e1ddd6f..21ebe45f7b 100644 --- a/backend/parsers/windmill-parser-sql/src/lib.rs +++ b/backend/parsers/windmill-parser-sql/src/lib.rs @@ -120,6 +120,60 @@ pub fn parse_db_resource(code: &str) -> Option { cap.map(|x| x.get(1).map(|x| x.as_str().to_string()).unwrap()) } +#[derive(Clone, Copy, Debug)] +pub enum S3ModeFormat { + Json, + Csv, + Parquet, +} +pub fn s3_mode_extension(format: S3ModeFormat) -> &'static str { + match format { + S3ModeFormat::Json => "json", + S3ModeFormat::Csv => "csv", + S3ModeFormat::Parquet => "parquet", + } +} +pub struct S3ModeArgs { + pub prefix: Option, + pub storage: Option, + pub format: S3ModeFormat, +} +pub fn parse_s3_mode(code: &str) -> anyhow::Result> { + let cap = match RE_S3_MODE.captures(code) { + Some(x) => x, + None => return Ok(None), + }; + let args_str = cap + .get(1) + .map(|x| x.as_str().to_string()) + .unwrap_or_default(); + + let mut prefix = None; + let mut storage = None; + let mut format = S3ModeFormat::Json; + + for kv in args_str.split(' ').map(|kv| kv.trim()) { + if kv.is_empty() { + continue; + } + let mut it = kv.split('='); + let (Some(key), Some(value)) = (it.next(), it.next()) else { + return Err(anyhow!("Invalid S3 mode argument: {}", kv)); + }; + match (key.trim(), value.trim()) { + ("prefix", _) => prefix = Some(value.to_string()), + ("storage", _) => storage = Some(value.to_string()), + ("format", "json") => format = S3ModeFormat::Json, + ("format", "parquet") => format = S3ModeFormat::Parquet, + ("format", "csv") => format = S3ModeFormat::Csv, + ("format", format) => return Err(anyhow!("Invalid S3 mode format: {}", format)), + (_, _) => return Err(anyhow!("Invalid S3 mode argument: {}", kv)), + } + } + + Ok(Some(S3ModeArgs { prefix, storage, format })) +} + pub fn parse_sql_blocks(code: &str) -> Vec<&str> { let mut blocks = vec![]; let mut last_idx = 0; @@ -147,6 +201,7 @@ lazy_static::lazy_static! { static ref RE_NONEMPTY_SQL_BLOCK: Regex = Regex::new(r#"(?m)^\s*[^\s](?:[^-]|$)"#).unwrap(); static ref RE_DB: Regex = Regex::new(r#"(?m)^-- database (\S+) *(?:\r|\n|$)"#).unwrap(); + static ref RE_S3_MODE: Regex = Regex::new(r#"(?m)^-- s3( (.+))? *(?:\r|\n|$)"#).unwrap(); // -- $1 name (type) = default static ref RE_ARG_MYSQL: Regex = Regex::new(r#"(?m)^-- \? (\w+) \((\w+)\)(?: ?\= ?(.+))? *(?:\r|\n|$)"#).unwrap(); diff --git a/backend/windmill-common/Cargo.toml b/backend/windmill-common/Cargo.toml index 76a6f16962..d3f7ed4c00 100644 --- a/backend/windmill-common/Cargo.toml +++ b/backend/windmill-common/Cargo.toml @@ -12,7 +12,7 @@ tantivy = [] prometheus = ["dep:prometheus"] loki = ["dep:tracing-loki"] benchmark = [] -parquet = ["dep:object_store", "dep:aws-config", "dep:aws-sdk-sts"] +parquet = ["dep:object_store", "dep:aws-config", "dep:aws-sdk-sts", "dep:datafusion"] aws_auth = ["dep:aws-sdk-sts", "dep:aws-config"] otel = ["dep:opentelemetry-semantic-conventions", "dep:opentelemetry-otlp", "dep:opentelemetry_sdk", "dep:opentelemetry", "dep:tracing-opentelemetry", "dep:opentelemetry-appender-tracing", "dep:tonic"] @@ -44,6 +44,9 @@ tracing = { workspace = true } axum = { workspace = true } hyper = { workspace = true } tokio = { workspace = true } +tokio-stream.workspace = true +tokio-util.workspace = true +datafusion = { workspace = true, optional = true} reqwest = { workspace = true } tracing-subscriber = { workspace = true } lazy_static.workspace = true @@ -67,6 +70,7 @@ async-stream.workspace = true const_format.workspace = true crc.workspace = true windmill-macros.workspace = true +windmill-parser-sql.workspace = true jsonwebtoken.workspace = true backon.workspace = true diff --git a/backend/windmill-common/src/s3_helpers.rs b/backend/windmill-common/src/s3_helpers.rs index 495f45912e..d49e64865e 100644 --- a/backend/windmill-common/src/s3_helpers.rs +++ b/backend/windmill-common/src/s3_helpers.rs @@ -16,10 +16,35 @@ use object_store::{aws::AmazonS3Builder, ClientOptions}; use reqwest::header::HeaderMap; use serde::{Deserialize, Serialize}; #[cfg(feature = "parquet")] -use std::sync::Arc; +use std::sync::{Arc, Mutex}; #[cfg(feature = "parquet")] use tokio::sync::RwLock; +#[cfg(feature = "parquet")] +use crate::error::to_anyhow; +#[cfg(feature = "parquet")] +use crate::utils::rd_string; +#[cfg(feature = "parquet")] +use bytes::Bytes; +#[cfg(feature = "parquet")] +use datafusion::arrow::array::{RecordBatch, RecordBatchWriter}; +#[cfg(feature = "parquet")] +use datafusion::arrow::error::ArrowError; +#[cfg(feature = "parquet")] +use datafusion::arrow::json::writer::JsonArray; +#[cfg(feature = "parquet")] +use datafusion::arrow::{csv, json}; +#[cfg(feature = "parquet")] +use datafusion::parquet::arrow::ArrowWriter; +#[cfg(feature = "parquet")] +use futures::TryStreamExt; +#[cfg(feature = "parquet")] +use std::io::Write; +#[cfg(feature = "parquet")] +use tokio::task; +#[cfg(feature = "parquet")] +use windmill_parser_sql::S3ModeFormat; + #[cfg(feature = "parquet")] lazy_static::lazy_static! { @@ -480,3 +505,184 @@ pub fn bundle(w_id: &str, hash: &str) -> String { pub fn raw_app(w_id: &str, version: &i64) -> String { format!("/home/rfiszel/raw_app/{}/{}", w_id, version) } + +// Originally used a Arc> +// But cannot call .close() on it because it moves the value and the object is not Sized +#[cfg(feature = "parquet")] +enum RecordBatchWriterEnum { + Parquet(ArrowWriter), + Csv(csv::Writer), + Json(json::Writer), +} + +#[cfg(feature = "parquet")] +impl RecordBatchWriter for RecordBatchWriterEnum { + fn write(&mut self, batch: &RecordBatch) -> Result<(), ArrowError> { + match self { + RecordBatchWriterEnum::Parquet(w) => w.write(batch).map_err(|e| e.into()), + RecordBatchWriterEnum::Csv(w) => w.write(batch), + RecordBatchWriterEnum::Json(w) => w.write(batch), + } + } + + fn close(self) -> Result<(), ArrowError> { + match self { + RecordBatchWriterEnum::Parquet(w) => w.close().map_err(|e| e.into()).map(drop), + RecordBatchWriterEnum::Csv(w) => w.close(), + RecordBatchWriterEnum::Json(w) => w.close(), + } + } +} + +#[cfg(feature = "parquet")] +struct ChannelWriter { + sender: tokio::sync::mpsc::Sender>, +} + +#[cfg(feature = "parquet")] +impl Write for ChannelWriter { + fn write(&mut self, buf: &[u8]) -> std::io::Result { + let data: Bytes = buf.to_vec().into(); + self.sender.blocking_send(Ok(data)).map_err(|e| { + std::io::Error::new( + std::io::ErrorKind::BrokenPipe, + format!("Channel send error: {}", e), + ) + })?; + Ok(buf.len()) + } + + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } +} + +#[cfg(not(feature = "parquet"))] +pub async fn convert_json_line_stream>( + mut _stream: impl futures::TryStreamExt> + Unpin, + _output_format: windmill_parser_sql::S3ModeFormat, +) -> anyhow::Result>> { + Ok(async_stream::stream! { + yield Err(anyhow::anyhow!("Parquet feature is not enabled. Cannot convert JSON line stream.")); + }) +} + +#[cfg(feature = "parquet")] +pub async fn convert_json_line_stream>( + mut stream: impl TryStreamExt> + Unpin, + output_format: S3ModeFormat, +) -> anyhow::Result>> { + const MAX_MPSC_SIZE: usize = 1000; + + use datafusion::{execution::context::SessionContext, prelude::NdJsonReadOptions}; + use futures::StreamExt; + use std::path::PathBuf; + use tokio::io::AsyncWriteExt; + + let mut path = PathBuf::from(std::env::temp_dir()); + path.push(format!("{}.json", rd_string(8))); + let path_str = path + .to_str() + .ok_or_else(|| anyhow::anyhow!("Invalid path"))?; + + // Write the stream to a temporary file + let mut file: tokio::fs::File = tokio::fs::File::create(&path).await.map_err(to_anyhow)?; + + while let Some(chunk) = stream.next().await { + match chunk { + Ok(chunk) => { + // Convert the chunk to bytes and write it to the file + let b: bytes::Bytes = serde_json::to_string(&chunk)?.into(); + file.write_all(&b).await?; + file.write_all(b"\n").await?; + } + Err(e) => { + tokio::fs::remove_file(&path).await?; + return Err(e.into()); + } + } + } + + file.flush().await?; + file.sync_all().await?; + drop(file); + + let ctx = SessionContext::new(); + ctx.register_json( + "my_table", + path_str, + NdJsonReadOptions { ..Default::default() }, + ) + .await + .map_err(to_anyhow)?; + + let df = ctx.sql("SELECT * FROM my_table").await.map_err(to_anyhow)?; + let schema = df.schema().clone().into(); + let mut datafusion_stream = df.execute_stream().await.map_err(to_anyhow)?; + + let (tx, rx) = tokio::sync::mpsc::channel(MAX_MPSC_SIZE); + let writer: Arc>> = + Arc::new(Mutex::new(Some(match output_format { + S3ModeFormat::Parquet => RecordBatchWriterEnum::Parquet( + ArrowWriter::try_new(ChannelWriter { sender: tx.clone() }, Arc::new(schema), None) + .map_err(to_anyhow)?, + ), + + S3ModeFormat::Csv => { + RecordBatchWriterEnum::Csv(csv::Writer::new(ChannelWriter { sender: tx.clone() })) + } + S3ModeFormat::Json => { + RecordBatchWriterEnum::Json(json::Writer::<_, JsonArray>::new(ChannelWriter { + sender: tx.clone(), + })) + } + }))); + + // This spawn is so that the data is sent in the background. Else the function would deadlock + // when hitting the mpsc channel limit + task::spawn(async move { + while let Some(batch_result) = datafusion_stream.next().await { + let batch: RecordBatch = match batch_result { + Ok(batch) => batch, + Err(e) => { + tracing::error!("Error in datafusion stream: {:?}", &e); + match tx.send(Err(e.into())).await { + Ok(_) => {} + Err(e) => tracing::error!("Failed to write error to channel: {:?}", &e), + } + break; + } + }; + let writer = writer.clone(); + // Writer calls blocking_send which would crash if called from the async context + let write_result = task::spawn_blocking(move || { + // SAFETY: We await so the code is actually sequential, lock unwrap cannot panic + // Second unwrap is ok because we initialized the option with Some + writer.lock().unwrap().as_mut().unwrap().write(&batch) + }) + .await; + match write_result { + Ok(Ok(_)) => {} + Ok(Err(e)) => { + tracing::error!("Error writing batch: {:?}", &e); + match tx.send(Err(e.into())).await { + Ok(_) => {} + Err(e) => tracing::error!("Failed to write error to channel: {:?}", &e), + } + } + Err(e) => tracing::error!("Error in blocking task: {:?}", &e), + }; + } + task::spawn_blocking(move || { + writer.lock().unwrap().take().unwrap().close()?; + drop(writer); + Ok::<_, anyhow::Error>(()) + }) + .await??; + drop(ctx); + tokio::fs::remove_file(&path).await?; + Ok::<_, anyhow::Error>(()) + }); + + Ok(tokio_stream::wrappers::ReceiverStream::new(rx)) +} diff --git a/backend/windmill-worker/Cargo.toml b/backend/windmill-worker/Cargo.toml index e4de400d84..63295208a2 100644 --- a/backend/windmill-worker/Cargo.toml +++ b/backend/windmill-worker/Cargo.toml @@ -90,6 +90,7 @@ deno_tls = { workspace = true, optional = true } deno_permissions = { workspace = true, optional = true } deno_io = { workspace = true, optional = true } deno_error = { workspace = true, optional = true } +async-stream.workspace = true postgres-native-tls.workspace = true native-tls.workspace = true diff --git a/backend/windmill-worker/src/bigquery_executor.rs b/backend/windmill-worker/src/bigquery_executor.rs index 79308b2ddb..7a6c0ff92d 100644 --- a/backend/windmill-worker/src/bigquery_executor.rs +++ b/backend/windmill-worker/src/bigquery_executor.rs @@ -1,20 +1,24 @@ use std::collections::HashMap; use futures::future::BoxFuture; -use futures::FutureExt; +use futures::{FutureExt, StreamExt}; use reqwest::Client; use serde_json::{json, value::RawValue, Value}; use windmill_common::error::to_anyhow; +use windmill_common::s3_helpers::convert_json_line_stream; use windmill_common::worker::Connection; use windmill_common::{error::Error, worker::to_raw_value}; use windmill_parser_sql::{ - parse_bigquery_sig, parse_db_resource, parse_sql_blocks, parse_sql_statement_named_params, + parse_bigquery_sig, parse_db_resource, parse_s3_mode, parse_sql_blocks, + parse_sql_statement_named_params, }; use windmill_queue::CanceledBy; use serde::Deserialize; -use crate::common::{build_http_client, OccupancyMetrics}; +use crate::common::{ + build_http_client, s3_mode_args_to_worker_data, OccupancyMetrics, S3ModeWorkerData, +}; use crate::handle_child::run_future_with_polling_update_job_poller; use crate::sanitized_sql_params::sanitize_and_interpolate_unsafe_sql_args; use crate::{ @@ -31,6 +35,16 @@ struct BigqueryResponse { totalRows: Option, schema: Option, jobComplete: bool, + pageToken: Option, + jobReference: Option, +} + +#[allow(non_snake_case)] +#[derive(Deserialize, Clone)] +struct BigQueryResponseJobReference { + jobId: String, + projectId: String, + location: Option, } #[derive(Deserialize)] @@ -74,6 +88,7 @@ fn do_bigquery_inner<'a>( column_order: Option<&'a mut Option>>, skip_collect: bool, http_client: &'a Client, + s3: Option, ) -> windmill_common::error::Result>>> { let param_names = parse_sql_statement_named_params(query, '@'); @@ -120,69 +135,80 @@ fn do_bigquery_inner<'a>( e.to_string() )) })?; + let rows = handle_bigquery_response(&result, &s3, column_order).await?; - if !result.jobComplete { - return Err(Error::ExecutionErr( - "BigQuery API did not answer query in time".to_string(), - )); + if let Some(s3) = s3 { + let cloned_s3 = s3.clone(); + let cloned_http_client = http_client.clone(); + let cloned_token = token.to_string(); + let rows_stream = async_stream::stream! { + for row in rows.iter() { + yield Ok::<_, windmill_common::error::Error>(row.clone()); + } + let mut next_page_token = result.pageToken; + let Some(job_reference) = result.jobReference.clone() else { + return; + }; + while let Some(ref next_page_token_value) = next_page_token { + let response2 = cloned_http_client + .get( + format!("https://bigquery.googleapis.com/bigquery/v2/projects/{}/queries/{}", job_reference.projectId, job_reference.jobId), + ) + .bearer_auth(cloned_token.as_str()) + .query(&[ + ("pageToken", next_page_token_value.as_str()), + ("maxResults", "10000"), + ("timeoutMs", timeout_ms.to_string().as_str()), + ("location", job_reference.location.as_ref().unwrap_or(&"US".to_string()).as_str()), + ]) + .send() + .await + .map_err(|e| { + Error::ExecutionErr(format!("Could not send query to BigQuery API: {}", e)) + })?; + + if let Err(e) = response2.error_for_status_ref() { + match response2.json::().await { + Ok(bq_err) => { + yield Err(Error::ExecutionErr(format!( + "Error from BigQuery API: {}", + bq_err.error.message + ))) + .map_err(to_anyhow)?; + return; + }, + Err(_) => { + yield Err(Error::ExecutionErr(format!( + "Error from BigQuery API could not be parsed: {}", + e.to_string() + ))) + .map_err(to_anyhow)?; + return; + }, + } + } + + let result2 = response2.json::().await.map_err(|e| { + Error::ExecutionErr(format!( + "BigQuery API response could not be parsed: {}", + e.to_string() + )) + })?; + let rows = handle_bigquery_response(&result2, &Some(cloned_s3.clone()), None).await?; + for row in rows.into_iter() { + yield Ok::<_, windmill_common::error::Error>(row); + } + next_page_token = result2.pageToken; + } + }; + + let stream = + convert_json_line_stream(rows_stream.boxed(), s3.format).await?; + s3.upload(stream.boxed()).await?; + + return Ok(to_raw_value(&s3.object_key)); } - if result.rows.is_none() || result.rows.as_ref().unwrap().len() == 0 { - return Ok(serde_json::from_str("[]").unwrap()); - } - - if result.schema.is_none() { - return Err(Error::ExecutionErr( - "Incomplete response from BigQuery API".to_string(), - )); - } - - if result - .totalRows - .unwrap_or(json!("")) - .as_str() - .unwrap_or("") - .parse::() - .unwrap_or(0) - > 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 - .schema - .as_ref() - .unwrap() - .fields - .iter() - .map(|x| x.name.clone()) - .collect::>(), - ); - } - - let rows = result - .rows - .unwrap() - .iter() - .map(|row| { - let mut row_map = serde_json::Map::new(); - row.f - .iter() - .zip(result.schema.as_ref().unwrap().fields.iter()) - .for_each(|(field, schema)| { - row_map.insert( - schema.name.clone(), - parse_val(&field.v, &schema.r#type, &schema), - ); - }); - Value::from(row_map) - }) - .collect::>(); - Ok(to_raw_value(&rows)) } } @@ -204,6 +230,79 @@ fn do_bigquery_inner<'a>( Ok(result_f.boxed()) } +async fn handle_bigquery_response<'a>( + result: &BigqueryResponse, + s3: &Option, + column_order: Option<&'a mut Option>>, +) -> windmill_common::error::Result> { + if !result.jobComplete { + return Err(Error::ExecutionErr( + "BigQuery API did not answer query in time".to_string(), + )); + } + + if result.rows.is_none() || result.rows.as_ref().unwrap().len() == 0 { + return Ok(serde_json::from_str("[]").unwrap()); + } + + if result.schema.is_none() { + return Err(Error::ExecutionErr( + "Incomplete response from BigQuery API".to_string(), + )); + } + + if s3.is_none() + && result + .totalRows + .as_ref() + .unwrap_or(&json!("")) + .as_str() + .unwrap_or("") + .parse::() + .unwrap_or(0) + > 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 + .schema + .as_ref() + .unwrap() + .fields + .iter() + .map(|x| x.name.clone()) + .collect::>(), + ); + } + + let rows = result + .rows + .as_ref() + .unwrap() + .iter() + .map(|row| { + let mut row_map = serde_json::Map::new(); + row.f + .iter() + .zip(result.schema.as_ref().unwrap().fields.iter()) + .for_each(|(field, schema)| { + row_map.insert( + schema.name.clone(), + parse_val(&field.v, &schema.r#type, &schema), + ); + }); + Value::from(row_map) + }) + .collect::>(); + Ok(rows) +} + use windmill_queue::MiniPulledJob; pub async fn do_bigquery( @@ -220,6 +319,7 @@ pub async fn do_bigquery( let bigquery_args = build_args_values(job, client, conn).await?; let inline_db_res_path = parse_db_resource(&query); + let s3 = parse_s3_mode(&query)?.map(|s3| s3_mode_args_to_worker_data(s3, client.clone(), job)); let db_arg = if let Some(inline_db_res_path) = inline_db_res_path { Some( @@ -332,6 +432,7 @@ pub async fn do_bigquery( None, annotations.return_last_result && i < queries.len() - 1, &http_client, + s3.clone(), ) }) .collect::>>()?; @@ -361,6 +462,7 @@ pub async fn do_bigquery( Some(column_order), false, &http_client, + s3, )? }; diff --git a/backend/windmill-worker/src/common.rs b/backend/windmill-worker/src/common.rs index 2699b89d02..086cee7e39 100644 --- a/backend/windmill-worker/src/common.rs +++ b/backend/windmill-worker/src/common.rs @@ -31,6 +31,7 @@ use windmill_common::{ }; use anyhow::{anyhow, bail, Result}; +use windmill_parser_sql::{s3_mode_extension, S3ModeArgs, S3ModeFormat}; use windmill_queue::MiniPulledJob; use std::ops::AsyncFn; @@ -1579,3 +1580,55 @@ pub async fn par_install_language_dependencies<'a>( } Ok(()) } + +#[derive(Clone)] +pub struct S3ModeWorkerData { + pub client: AuthedClient, + pub object_key: String, + pub format: S3ModeFormat, + pub storage: Option, + pub workspace_id: String, +} + +impl S3ModeWorkerData { + pub async fn upload(&self, stream: S) -> error::Result + where + S: futures::stream::TryStream + Send + 'static, + S::Error: Into>, + bytes::Bytes: From, + { + self.client + .upload_s3_file( + self.workspace_id.as_str(), + self.object_key.clone(), + self.storage.clone(), + stream, + ) + .await + } +} + +pub fn s3_mode_args_to_worker_data( + s3: S3ModeArgs, + client: AuthedClient, + job: &MiniPulledJob, +) -> S3ModeWorkerData { + S3ModeWorkerData { + client, + storage: s3.storage, + format: s3.format, + object_key: format!( + "{}/{}.{}", + s3.prefix.unwrap_or_else(|| format!( + "wmill_datalake/{}", + job.runnable_path + .as_ref() + .map(|s| s.as_str()) + .unwrap_or("unknown_script") + )), + job.id, + s3_mode_extension(s3.format) + ), + workspace_id: job.workspace_id.clone(), + } +} diff --git a/backend/windmill-worker/src/mssql_executor.rs b/backend/windmill-worker/src/mssql_executor.rs index 303685595b..4c727b786e 100644 --- a/backend/windmill-worker/src/mssql_executor.rs +++ b/backend/windmill-worker/src/mssql_executor.rs @@ -1,5 +1,6 @@ use base64::{engine::general_purpose, Engine as _}; use chrono::{DateTime, NaiveDate, NaiveDateTime, NaiveTime, Utc}; +use futures::StreamExt; use regex::Regex; use serde::Deserialize; use serde_json::value::RawValue; @@ -8,16 +9,17 @@ use tiberius::{AuthMethod, Client, ColumnData, Config, FromSqlOwned, Query, Row, use tokio::net::TcpStream; use tokio_util::compat::TokioAsyncWriteCompatExt; use uuid::Uuid; +use windmill_common::s3_helpers::convert_json_line_stream; use windmill_common::{ error::{self, to_anyhow, Error}, utils::empty_as_none, worker::{to_raw_value, Connection}, }; -use windmill_parser_sql::{parse_db_resource, parse_mssql_sig}; +use windmill_parser_sql::{parse_db_resource, parse_mssql_sig, parse_s3_mode}; use windmill_queue::MiniPulledJob; use windmill_queue::{append_logs, CanceledBy}; -use crate::common::{build_args_values, OccupancyMetrics}; +use crate::common::{build_args_values, s3_mode_args_to_worker_data, OccupancyMetrics}; use crate::handle_child::run_future_with_polling_update_job_poller; use crate::sanitized_sql_params::sanitize_and_interpolate_unsafe_sql_args; use crate::AuthedClient; @@ -63,6 +65,7 @@ pub async fn do_mssql( let mssql_args = build_args_values(job, client, conn).await?; let inline_db_res_path = parse_db_resource(&query); + let s3 = parse_s3_mode(&query)?.map(|s3| s3_mode_args_to_worker_data(s3, client.clone(), job)); let db_arg = if let Some(inline_db_res_path) = inline_db_res_path { Some( @@ -197,27 +200,42 @@ pub async fn do_mssql( // A response to a query is a stream of data, that must be // polled to the end before querying again. Using streams allows // fetching data in an asynchronous manner, if needed. - let stream = prepared_query.query(&mut client).await.map_err(to_anyhow)?; - let results = stream.into_results().await.map_err(to_anyhow)?; - let len = results.len(); - let mut json_results = vec![]; - for (i, statement_result) in results.into_iter().enumerate() { - if annotations.return_last_result && i < len - 1 { - continue; - } - let mut json_rows = vec![]; - for row in statement_result { - let row = row_to_json(row)?; - json_rows.push(row); - } - json_results.push(json_rows); - } + if let Some(s3) = s3 { + let rows_stream = async_stream::stream! { + let mut stream = prepared_query.query(&mut client).await.map_err(to_anyhow)?.into_row_stream().map(|row| { + row_to_json(row.map_err(to_anyhow)?).map_err(to_anyhow) + }); + while let Some(row) = stream.next().await { + yield row; + } + }; - if annotations.return_last_result && json_results.len() > 0 { - Ok(to_raw_value(&json_results.pop().unwrap())) + let stream = convert_json_line_stream(rows_stream.boxed(), s3.format).await?; + s3.upload(stream.boxed()).await?; + + Ok(serde_json::value::to_raw_value(&s3.object_key)?) } else { - Ok(to_raw_value(&json_results)) + let stream = prepared_query.query(&mut client).await.map_err(to_anyhow)?; + let results = stream.into_results().await.map_err(to_anyhow)?; + let len = results.len(); + let mut json_results = vec![]; + for (i, statement_result) in results.into_iter().enumerate() { + if annotations.return_last_result && i < len - 1 { + continue; + } + let mut json_rows = vec![]; + for row in statement_result { + let row = row_to_json(row)?; + json_rows.push(row); + } + json_results.push(json_rows); + } + if annotations.return_last_result && json_results.len() > 0 { + Ok(to_raw_value(&json_results.pop().unwrap())) + } else { + Ok(to_raw_value(&json_results)) + } } }; diff --git a/backend/windmill-worker/src/mysql_executor.rs b/backend/windmill-worker/src/mysql_executor.rs index 9c95977e23..3bfe3afb49 100644 --- a/backend/windmill-worker/src/mysql_executor.rs +++ b/backend/windmill-worker/src/mysql_executor.rs @@ -1,7 +1,8 @@ use std::{collections::HashMap, sync::Arc}; +use anyhow::anyhow; use base64::Engine; -use futures::{future::BoxFuture, FutureExt}; +use futures::{future::BoxFuture, FutureExt, StreamExt}; use itertools::Itertools; use mysql_async::{ consts::ColumnType, prelude::*, FromValueError, OptsBuilder, Params, Row, SslOpts, @@ -13,17 +14,18 @@ use std::str::FromStr; use tokio::sync::Mutex; use windmill_common::{ error::{to_anyhow, Error}, + s3_helpers::convert_json_line_stream, worker::{to_raw_value, Connection}, }; use windmill_parser_sql::{ - parse_db_resource, parse_mysql_sig, parse_sql_blocks, parse_sql_statement_named_params, - RE_ARG_MYSQL_NAMED, + parse_db_resource, parse_mysql_sig, parse_s3_mode, parse_sql_blocks, + parse_sql_statement_named_params, RE_ARG_MYSQL_NAMED, }; use windmill_queue::CanceledBy; use windmill_queue::MiniPulledJob; use crate::{ - common::{build_args_values, OccupancyMetrics}, + common::{build_args_values, s3_mode_args_to_worker_data, OccupancyMetrics, S3ModeWorkerData}, handle_child::run_future_with_polling_update_job_poller, sanitized_sql_params::sanitize_and_interpolate_unsafe_sql_args, AuthedClient, @@ -39,12 +41,13 @@ struct MysqlDatabase { ssl: Option, } -pub fn do_mysql_inner<'a>( +fn do_mysql_inner<'a>( query: &'a str, all_statement_values: &Params, conn: Arc>, column_order: Option<&'a mut Option>>, skip_collect: bool, + s3: Option, ) -> windmill_common::error::Result>>> { let param_names = parse_sql_statement_named_params(query, ':') .into_iter() @@ -71,6 +74,38 @@ pub fn do_mysql_inner<'a>( .map_err(to_anyhow)?; Ok(to_raw_value(&Value::Array(vec![]))) + } else if let Some(ref s3) = s3 { + let query = query.to_string(); + let rows_stream = async_stream::stream! { + let mut conn = conn.lock().await; + let mut result = match conn.exec_iter(query, statement_values).await.map_err(to_anyhow) { + Ok(result) => result, + Err(e) => { + yield Err(anyhow!("Error executing query: {:?}", e)); + return; + } + }; + loop { + let row = result.next().await; + match row { + Ok(Some(row)) => { + yield Ok(convert_row_to_value(row)); + } + Ok(None) => { + break; + } + Err(e) => { + yield Err(anyhow!("Error fetching row: {:?}", e)); + return; + } + } + } + }; + + let stream = convert_json_line_stream(rows_stream.boxed(), s3.format).await?; + s3.upload(stream.boxed()).await?; + + Ok(serde_json::value::to_raw_value(&s3.object_key)?) } else { let rows: Vec = conn .lock() @@ -118,6 +153,7 @@ pub async fn do_mysql( let job_args = build_args_values(job, client, conn).await?; let inline_db_res_path = parse_db_resource(&query); + let s3 = parse_s3_mode(&query)?.map(|s3| s3_mode_args_to_worker_data(s3, client.clone(), job)); let db_arg = if let Some(inline_db_res_path) = inline_db_res_path { Some( @@ -252,6 +288,7 @@ pub async fn do_mysql( conn_a.clone(), None, annotations.return_last_result && i < queries.len() - 1, + s3.clone(), ) }) .collect::>>()?; @@ -277,6 +314,7 @@ pub async fn do_mysql( conn_a.clone(), Some(column_order), false, + s3, )? }; diff --git a/backend/windmill-worker/src/pg_executor.rs b/backend/windmill-worker/src/pg_executor.rs index 8cc97455c0..e12ca71626 100644 --- a/backend/windmill-worker/src/pg_executor.rs +++ b/backend/windmill-worker/src/pg_executor.rs @@ -8,7 +8,7 @@ use anyhow::Context; use base64::{engine, Engine as _}; use chrono::Utc; use futures::future::BoxFuture; -use futures::{FutureExt, TryStreamExt}; +use futures::{FutureExt, StreamExt, TryFutureExt, TryStreamExt}; use itertools::Itertools; use native_tls::{Certificate, TlsConnector}; use postgres_native_tls::MakeTlsConnector; @@ -27,14 +27,18 @@ use tokio_postgres::{ use uuid::Uuid; use windmill_common::error::to_anyhow; use windmill_common::error::{self, Error}; +use windmill_common::s3_helpers::convert_json_line_stream; use windmill_common::worker::{to_raw_value, Connection, CLOUD_HOSTED}; use windmill_parser::{Arg, Typ}; use windmill_parser_sql::{ - parse_db_resource, parse_pg_statement_arg_indices, parse_pgsql_sig, parse_sql_blocks, + parse_db_resource, parse_pg_statement_arg_indices, parse_pgsql_sig, parse_s3_mode, + parse_sql_blocks, }; use windmill_queue::{CanceledBy, MiniPulledJob}; -use crate::common::{build_args_values, sizeof_val, OccupancyMetrics}; +use crate::common::{ + build_args_values, s3_mode_args_to_worker_data, sizeof_val, OccupancyMetrics, S3ModeWorkerData, +}; use crate::handle_child::run_future_with_polling_update_job_poller; use crate::sanitized_sql_params::sanitize_and_interpolate_unsafe_sql_args; use crate::{AuthedClient, MAX_RESULT_SIZE}; @@ -68,6 +72,7 @@ fn do_postgresql_inner<'a>( column_order: Option<&'a mut Option>>, siz: &'a AtomicUsize, skip_collect: bool, + s3: Option, ) -> error::Result>>> { let mut query_params = vec![]; @@ -106,6 +111,20 @@ fn do_postgresql_inner<'a>( .execute_raw(&query, query_params) .await .map_err(to_anyhow)?; + } else if let Some(ref s3) = s3 { + let rows_stream = client + .query_raw(&query, query_params) + .map_err(to_anyhow) + .await? + .map_err(to_anyhow) + .map(|row_result| { + row_result.and_then(|row| postgres_row_to_json_value(row).map_err(to_anyhow)) + }); + + let stream = convert_json_line_stream(rows_stream.boxed(), s3.format).await?; + s3.upload(stream.boxed()).await?; + + return Ok(serde_json::value::to_raw_value(&s3.object_key)?); } else { let rows = client .query_raw(&query, query_params) @@ -172,6 +191,8 @@ pub async fn do_postgresql( let inline_db_res_path = parse_db_resource(&query); + let s3 = parse_s3_mode(&query)?.map(|s3| s3_mode_args_to_worker_data(s3, client.clone(), job)); + let db_arg = if let Some(inline_db_res_path) = inline_db_res_path { Some( client @@ -321,6 +342,7 @@ pub async fn do_postgresql( None, &size, annotations.return_last_result && i < queries.len() - 1, + s3.clone(), ) }) .collect::>>()?; @@ -347,6 +369,7 @@ pub async fn do_postgresql( Some(column_order), &size, false, + s3, )? }; diff --git a/backend/windmill-worker/src/snowflake_executor.rs b/backend/windmill-worker/src/snowflake_executor.rs index 59ceccbc60..84d2d0776d 100644 --- a/backend/windmill-worker/src/snowflake_executor.rs +++ b/backend/windmill-worker/src/snowflake_executor.rs @@ -2,21 +2,28 @@ use base64::{engine, Engine as _}; use chrono::Datelike; use core::fmt::Write; use futures::future::BoxFuture; -use futures::FutureExt; +use futures::{FutureExt, StreamExt, TryStreamExt}; use jsonwebtoken::{encode, Algorithm, EncodingKey, Header}; use reqwest::{Client, Response}; use serde_json::{json, value::RawValue, Value}; use sha2::{Digest, Sha256}; use std::collections::HashMap; +use windmill_common::error::to_anyhow; +use windmill_common::s3_helpers::convert_json_line_stream; use windmill_common::worker::Connection; use windmill_common::{error::Error, worker::to_raw_value}; -use windmill_parser_sql::{parse_db_resource, parse_snowflake_sig, parse_sql_blocks}; +use windmill_parser_sql::{ + parse_db_resource, parse_s3_mode, parse_snowflake_sig, parse_sql_blocks, +}; use windmill_queue::{CanceledBy, MiniPulledJob, HTTP_CLIENT}; use serde::{Deserialize, Serialize}; -use crate::common::{build_http_client, resolve_job_timeout, OccupancyMetrics}; +use crate::common::{ + build_http_client, resolve_job_timeout, s3_mode_args_to_worker_data, OccupancyMetrics, + S3ModeWorkerData, +}; use crate::handle_child::run_future_with_polling_update_job_poller; use crate::sanitized_sql_params::sanitize_and_interpolate_unsafe_sql_args; use crate::{common::build_args_values, AuthedClient}; @@ -123,6 +130,7 @@ fn do_snowflake_inner<'a>( column_order: Option<&'a mut Option>>, skip_collect: bool, http_client: &'a Client, + s3: Option, ) -> windmill_common::error::Result>>> { let sig = parse_snowflake_sig(&query) .map_err(|x| Error::ExecutionErr(x.to_string()))? @@ -174,7 +182,7 @@ fn do_snowflake_inner<'a>( .parse_snowflake_response::() .await?; - if response.resultSetMetaData.numRows > 10000 { + if s3.is_none() && 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(), @@ -191,54 +199,72 @@ fn do_snowflake_inner<'a>( ); } - let mut rows = response.data; + // Clones are because, in s3 mode, reqwest::Body::wrap_stream requires the stream to be + // 'static even though it doesn't make sense to be in our case since the request is + // awaited and the stream is fully read before the function returns. + // Turns out it is a real pain to trick the compiler, even using unsafe + let cloned_account_identifier: String = account_identifier.to_string(); + let cloned_token = token.to_string(); - 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 mut request = HTTP_CLIENT - .get(url) - .bearer_auth(token) - .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::() - .await?; - - rows.extend(response.data); + let rows_stream = async_stream::stream! { + for row in response.data { + yield Ok::, windmill_common::error::Error>(row); } + + if response.resultSetMetaData.partitionInfo.len() > 1 { + for idx in 1..response.resultSetMetaData.partitionInfo.len() { + let url = format!( + "https://{}.snowflakecomputing.com/api/v2/statements/{}", + cloned_account_identifier.to_uppercase(), + response.statementHandle + ); + let mut request = HTTP_CLIENT + .get(url) + .bearer_auth(cloned_token.as_str()) + .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::() + .await?; + + for row in response.data { + yield Ok(row); + } + } + } + }; + + let rows_stream = rows_stream.map_ok(move |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 + }); + + if let Some(s3) = s3 { + let rows_stream = + rows_stream.map(|r| serde_json::value::to_value(&r?).map_err(to_anyhow)); + let stream = convert_json_line_stream(rows_stream.boxed(), s3.format).await?; + s3.upload(stream.boxed()).await?; + Ok(to_raw_value(&s3.object_key)) + } else { + let rows = rows_stream + .collect::>() + .await + .into_iter() + .collect::, _>>()?; + Ok(to_raw_value(&rows)) } - - 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) } }; @@ -259,6 +285,7 @@ pub async fn do_snowflake( let snowflake_args = build_args_values(job, client, conn).await?; let inline_db_res_path = parse_db_resource(&query); + let s3 = parse_s3_mode(&query)?.map(|s3| s3_mode_args_to_worker_data(s3, client.clone(), job)); let db_arg = if let Some(inline_db_res_path) = inline_db_res_path { Some( @@ -390,6 +417,7 @@ pub async fn do_snowflake( None, annotations.return_last_result && i < queries.len() - 1, &http_client, + s3.clone(), ) }) .collect::>>()?; @@ -419,6 +447,7 @@ pub async fn do_snowflake( Some(column_order), false, &http_client, + s3.clone(), )? }; let r = run_future_with_polling_update_job_poller( diff --git a/backend/windmill-worker/src/worker.rs b/backend/windmill-worker/src/worker.rs index 54820d52eb..8590603c2f 100644 --- a/backend/windmill-worker/src/worker.rs +++ b/backend/windmill-worker/src/worker.rs @@ -39,7 +39,7 @@ use windmill_common::METRICS_DEBUG_ENABLED; #[cfg(feature = "prometheus")] use windmill_common::METRICS_ENABLED; -use reqwest::Response; +use reqwest::{Body, Response}; use serde::{de::DeserializeOwned, Deserialize, Serialize}; use sqlx::types::Json; use std::{ @@ -520,6 +520,46 @@ impl AuthedClient { _ => Err(anyhow::anyhow!(response.text().await.unwrap_or_default())), } } + + pub async fn upload_s3_file( + &self, + workspace_id: &str, + object_key: String, + storage: Option, + body: S, + ) -> error::Result + where + S: futures::stream::TryStream + Send + 'static, + S::Error: Into>, + bytes::Bytes: From, + { + let mut query = vec![("file_key", object_key)]; + if let Some(storage) = storage { + query.push(("storage", storage)); + } + self.force_client + .as_ref() + .unwrap_or(&HTTP_CLIENT) + .post(format!( + "{}/api/w/{}/job_helpers/upload_s3_file", + self.base_internal_url, workspace_id + )) + .query(&query) + .header( + reqwest::header::ACCEPT, + reqwest::header::HeaderValue::from_static("application/json"), + ) + .header( + reqwest::header::AUTHORIZATION, + reqwest::header::HeaderValue::from_str(&format!("Bearer {}", self.token)) + .map_err(|e| error::Error::BadConfig(e.to_string()))?, + ) + .body(Body::wrap_stream(body)) + .send() + .await + .context(format!("Sent upload_s3_file request",)) + .map_err(error::Error::from) + } } #[derive(Clone)]