feat: sql jobs outputting to s3 + streaming for high-number of rows (#5704)

* stream to s3 boilerplate

* S3 works with new syntax

* snowflake s3 streaming support

* postgres s3 support

* fix postgres stream format

* mysql s3 streaming

* mssql s3 streaming

* new s3 mode syntax

* optional folder param

* rename folder to prefix

* json_stream_arr_values

* cargo toml rollback

* convert_ndjson with datafusion

* format conversion kinda works

* Fixed not finishing the datafusion writer

* support for pg and mssql

* fix file ext

* bigquery conversion and works with s3 streaming

* fix s3 flag parser

* snowflake s3 streaming support

* factor out duplicate code

* remove anyhow

* Err case for parse s3 mode

* Send error to mpsc

* bigquery s3 streaming fix for huge queries

* remove extra stuff

* snowflake s3 streaming support

* small regex mistake

* cfg(not(feature = "parquet"))

* fix CI (unused import)

* error handling fix (graphite)
This commit is contained in:
Diego Imbert
2025-05-13 10:19:44 +02:00
committed by GitHub
parent 76258b7b1a
commit c7886ea07a
12 changed files with 717 additions and 143 deletions
+5
View File
@@ -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",
@@ -120,6 +120,60 @@ pub fn parse_db_resource(code: &str) -> Option<String> {
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<String>,
pub storage: Option<String>,
pub format: S3ModeFormat,
}
pub fn parse_s3_mode(code: &str) -> anyhow::Result<Option<S3ModeArgs>> {
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();
+5 -1
View File
@@ -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
+207 -1
View File
@@ -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<Mutex<dyn RecordBatchWriter + Send>>
// But cannot call .close() on it because it moves the value and the object is not Sized
#[cfg(feature = "parquet")]
enum RecordBatchWriterEnum {
Parquet(ArrowWriter<ChannelWriter>),
Csv(csv::Writer<ChannelWriter>),
Json(json::Writer<ChannelWriter, JsonArray>),
}
#[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<anyhow::Result<Bytes>>,
}
#[cfg(feature = "parquet")]
impl Write for ChannelWriter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
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<E: Into<anyhow::Error>>(
mut _stream: impl futures::TryStreamExt<Item = Result<serde_json::Value, E>> + Unpin,
_output_format: windmill_parser_sql::S3ModeFormat,
) -> anyhow::Result<impl futures::TryStreamExt<Item = anyhow::Result<bytes::Bytes>>> {
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<E: Into<anyhow::Error>>(
mut stream: impl TryStreamExt<Item = Result<serde_json::Value, E>> + Unpin,
output_format: S3ModeFormat,
) -> anyhow::Result<impl TryStreamExt<Item = anyhow::Result<bytes::Bytes>>> {
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<Mutex<Option<RecordBatchWriterEnum>>> =
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))
}
+1
View File
@@ -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
+165 -63
View File
@@ -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<Value>,
schema: Option<BigqueryResponseSchema>,
jobComplete: bool,
pageToken: Option<String>,
jobReference: Option<BigQueryResponseJobReference>,
}
#[allow(non_snake_case)]
#[derive(Deserialize, Clone)]
struct BigQueryResponseJobReference {
jobId: String,
projectId: String,
location: Option<String>,
}
#[derive(Deserialize)]
@@ -74,6 +88,7 @@ fn do_bigquery_inner<'a>(
column_order: Option<&'a mut Option<Vec<String>>>,
skip_collect: bool,
http_client: &'a Client,
s3: Option<S3ModeWorkerData>,
) -> windmill_common::error::Result<BoxFuture<'a, windmill_common::error::Result<Box<RawValue>>>> {
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::<BigqueryErrorResponse>().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::<BigqueryResponse>().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::<i64>()
.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::<Vec<String>>(),
);
}
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::<Vec<_>>();
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<S3ModeWorkerData>,
column_order: Option<&'a mut Option<Vec<String>>>,
) -> windmill_common::error::Result<Vec<Value>> {
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::<i64>()
.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::<Vec<String>>(),
);
}
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::<Vec<_>>();
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::<windmill_common::error::Result<Vec<_>>>()?;
@@ -361,6 +462,7 @@ pub async fn do_bigquery(
Some(column_order),
false,
&http_client,
s3,
)?
};
+53
View File
@@ -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<String>,
pub workspace_id: String,
}
impl S3ModeWorkerData {
pub async fn upload<S>(&self, stream: S) -> error::Result<reqwest::Response>
where
S: futures::stream::TryStream + Send + 'static,
S::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
bytes::Bytes: From<S::Ok>,
{
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(),
}
}
+38 -20
View File
@@ -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))
}
}
};
+43 -5
View File
@@ -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<bool>,
}
pub fn do_mysql_inner<'a>(
fn do_mysql_inner<'a>(
query: &'a str,
all_statement_values: &Params,
conn: Arc<Mutex<mysql_async::Conn>>,
column_order: Option<&'a mut Option<Vec<String>>>,
skip_collect: bool,
s3: Option<S3ModeWorkerData>,
) -> windmill_common::error::Result<BoxFuture<'a, windmill_common::error::Result<Box<RawValue>>>> {
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<Row> = 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::<windmill_common::error::Result<Vec<_>>>()?;
@@ -277,6 +314,7 @@ pub async fn do_mysql(
conn_a.clone(),
Some(column_order),
false,
s3,
)?
};
+26 -3
View File
@@ -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<Vec<String>>>,
siz: &'a AtomicUsize,
skip_collect: bool,
s3: Option<S3ModeWorkerData>,
) -> error::Result<BoxFuture<'a, error::Result<Box<RawValue>>>> {
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::<error::Result<Vec<_>>>()?;
@@ -347,6 +369,7 @@ pub async fn do_postgresql(
Some(column_order),
&size,
false,
s3,
)?
};
@@ -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<Vec<String>>>,
skip_collect: bool,
http_client: &'a Client,
s3: Option<S3ModeWorkerData>,
) -> windmill_common::error::Result<BoxFuture<'a, windmill_common::error::Result<Box<RawValue>>>> {
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::<SnowflakeResponse>()
.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::<SnowflakeDataOnlyResponse>()
.await?;
rows.extend(response.data);
let rows_stream = async_stream::stream! {
for row in response.data {
yield Ok::<Vec<Value>, 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::<SnowflakeDataOnlyResponse>()
.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::<Vec<_>>()
.await
.into_iter()
.collect::<Result<Vec<_>, _>>()?;
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::<Vec<_>>(),
);
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::<windmill_common::error::Result<Vec<_>>>()?;
@@ -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(
+41 -1
View File
@@ -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<S>(
&self,
workspace_id: &str,
object_key: String,
storage: Option<String>,
body: S,
) -> error::Result<Response>
where
S: futures::stream::TryStream + Send + 'static,
S::Error: Into<Box<dyn std::error::Error + Send + Sync>>,
bytes::Bytes: From<S::Ok>,
{
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)]