use std::sync::Arc;
use async_trait::async_trait;
use dataflow_rs::engine::error::DataflowError;
use dataflow_rs::engine::task_context::TaskContext;
use dataflow_rs::engine::task_outcome::TaskOutcome;
use mongodb::bson::Document;
use serde_json::{Map, Value, json};
use super::connector_handler::{ConnectorHandler, Produced};
use super::connector_helpers::{
ConnectorCall, QueryBudget, QueryFailure, acquire_conn, build_entity_registry, decode_failure,
encode_failure, es_request, es_write_error, is_mongo, require_op_allowed, resolve_params,
resolve_row_format, send_es, timed_query, to_connect_error, to_exec_error,
};
use super::mongo_common::{delete_envelope, empty_insert_envelope, update_envelope};
use super::schema::{FieldKind, FieldSchema};
use super::templated_input::TemplatedInput;
use crate::connector::mongo_pool::MongoPoolCache;
use crate::connector::pool_cache::SqlPoolCache;
use crate::connector::{ConnectorConfig, ConnectorRegistry, EsConnectorConfig};
use crate::engine::HandlerError;
use crate::query::backend::es::{EsWrite, bulk_ndjson};
use crate::query::backend::mongo::MongoWrite;
use crate::query::write::{ResolvedWrite, WriteOp};
use crate::query::{self, SqlDialect};
use crate::storage::detect_backend;
const NAME: &str = <DataWriteHandler as ConnectorHandler>::NAME;
pub(super) const PARTIAL_STATUS: u16 = 207;
pub struct DataWriteHandler {
pub pool_cache: Arc<SqlPoolCache>,
pub mongo_pool_cache: Arc<MongoPoolCache>,
pub http_client: reqwest::Client,
pub registry: Arc<ConnectorRegistry>,
pub write_config: crate::config::WriteConfig,
pub max_returning: usize,
}
pub struct DataWrite {
params: Map<String, Value>,
format: crate::connector::sql_decode::RowFormat,
schema: Option<Value>,
envelope: Value,
database: Option<String>,
}
#[async_trait]
impl ConnectorHandler for DataWriteHandler {
const NAME: &'static str = "data_write";
type Kind = crate::connector::DataBackend;
type Input = TemplatedInput;
type Parsed = DataWrite;
fn registry(&self) -> &Arc<ConnectorRegistry> {
&self.registry
}
fn parse(
&self,
_call: &ConnectorCall<'_>,
input: &TemplatedInput,
ctx: &TaskContext<'_>,
) -> Result<Self::Parsed, HandlerError> {
let envelope = input
.get("write")
.ok_or_else(|| {
DataflowError::from(query::write::WriteError::Query(
query::QueryError::InvalidEnvelope(
"missing `write`: the mutation envelope (op/target/values/set/filter/\
on_conflict/returning/all) is nested under `write`, alongside the \
handler's `connector`/`schema`/`params`/`database`/`output`"
.to_string(),
),
))
})?
.clone();
Ok(DataWrite {
params: resolve_params(input, <Self as ConnectorHandler>::NAME, ctx),
format: resolve_row_format(input, <Self as ConnectorHandler>::NAME, ctx)?,
schema: input.get("schema").cloned(),
envelope,
database: input
.get("database")
.and_then(Value::as_str)
.map(str::to_string),
})
}
async fn run(
&self,
parsed: Self::Parsed,
conn: &ConnectorConfig,
call: &ConnectorCall<'_>,
_input: &TemplatedInput,
_ctx: &mut TaskContext<'_>,
) -> Result<Produced, HandlerError> {
let registry = build_entity_registry(parsed.schema.as_ref(), conn, call.connector)?;
let resolved = query::write::resolve_write(
&parsed.envelope,
&parsed.params,
®istry,
&self.write_config,
)
.map_err(DataflowError::from)?;
if let Some(gates) = conn.operation_gates() {
require_op_allowed(gates, resolved.op().as_str(), call.connector)?;
}
let (result, outcome) = match conn {
ConnectorConfig::Db(db) if is_mongo(&db.connection_string) => {
let database = parsed.database.as_deref().ok_or_else(|| {
DataflowError::Validation(format!("{} requires 'database' field", call.name))
})?;
let mw =
query::backend::mongo::render_write(&resolved).map_err(DataflowError::from)?;
let client = self
.mongo_pool_cache
.get_client(call.connector, db)
.await
.map_err(to_connect_error)?;
timed_query(
db.query_timeout_ms,
call.name,
execute_mongo(&client, database, mw),
)
.await?
}
ConnectorConfig::Db(db) => {
let dialect: SqlDialect = detect_backend(&db.connection_string)
.map_err(to_exec_error)?
.into();
let (sql, values) = query::backend::sql::render_write(&resolved, dialect)
.map_err(DataflowError::from)?;
let pool = self
.pool_cache
.get_pool(call.connector, db)
.await
.map_err(to_connect_error)?;
execute_sql(
&pool,
&sql,
values,
&resolved,
db.query_timeout_ms,
parsed.format,
self.max_returning,
)
.await?
}
ConnectorConfig::Es(es) => {
let ew =
query::backend::es::render_write(&resolved).map_err(DataflowError::from)?;
run_es_write(&self.http_client, es, ew).await?
}
other => unreachable!(
"DataBackend admitted a '{}' connector",
other.connector_type()
),
};
Ok(Produced::with_outcome(result, outcome))
}
}
#[allow(clippy::too_many_arguments)]
async fn execute_sql(
pool: &crate::connector::pool_cache::SqlPool,
sql: &str,
values: sea_query_sqlx::SqlxValues,
w: &ResolvedWrite,
timeout_ms: Option<u64>,
format: crate::connector::sql_decode::RowFormat,
max_returning: usize,
) -> Result<(Value, TaskOutcome), DataflowError> {
let budget = QueryBudget::start(timeout_ms);
let mut out = crate::connector::pool_cache::dispatch_sql_pool!(
pool, p, rows_to_json, _bind, typed_args, write_result => {
let scalars = crate::connector::sql_encode::scalars_from_sea(&values.0);
let fallback = crate::connector::sql_encode::sea_args_for(p, values);
if w.is_multi_row() || w.returns_unbounded_rows() {
let mut tx = budget.run(NAME, p.begin()).await?;
let bound = budget
.run(NAME, async {
typed_args(&mut tx, sql, scalars.as_deref())
.await
.map_err(|e| QueryFailure::Classified(encode_failure(NAME, e)))
})
.await?;
let (args, persistent) = decide(bound, fallback);
let out = run_write_statement(
&mut *tx, sql, args, persistent, w, &budget, rows_to_json, write_result,
format, max_returning,
)
.await?;
budget.run(NAME, tx.commit()).await?;
out
} else {
let mut conn = acquire_conn(&budget, NAME, p).await?;
let bound = budget
.run(NAME, async {
typed_args(&mut conn, sql, scalars.as_deref())
.await
.map_err(|e| QueryFailure::Classified(encode_failure(NAME, e)))
})
.await?;
let (args, persistent) = decide(bound, fallback);
run_write_statement(
&mut *conn, sql, args, persistent, w, &budget, rows_to_json, write_result,
format, max_returning,
)
.await?
}
}
);
out["status"] = json!("ok");
Ok((out, TaskOutcome::Success))
}
fn decide<A>(bound: crate::connector::sql_encode::Bound<A>, fallback: A) -> (A, bool) {
match bound {
crate::connector::sql_encode::Bound::Typed(args) => (args, true),
crate::connector::sql_encode::Bound::Fallback { cache } => (fallback, cache),
}
}
#[allow(clippy::too_many_arguments)]
async fn run_write_statement<'e, E, R, F, G>(
executor: E,
sql: &'e str,
args: R::Arguments,
persistent: bool,
w: &ResolvedWrite,
budget: &QueryBudget,
rows_to_json: F,
write_result: G,
format: crate::connector::sql_decode::RowFormat,
max_returning: usize,
) -> Result<Value, DataflowError>
where
E: sqlx::Executor<'e, Database = R>,
R: sqlx::Database + sqlx::database::HasStatementCache,
R::Arguments: sqlx::IntoArguments<R>,
F: Fn(
&[R::Row],
crate::connector::sql_decode::RowFormat,
) -> Result<Vec<Value>, crate::connector::sql_decode::DecodeError>,
G: Fn(&R::QueryResult) -> (u64, Option<i64>),
{
if w.returning().is_empty() {
let res = budget
.run(
NAME,
sqlx::query_with(sqlx::AssertSqlSafe(sql), args)
.persistent(persistent)
.execute(executor),
)
.await?;
let (rows_affected, last_insert_id) = write_result(&res);
let mut out = json!({ "rows_affected": rows_affected });
if matches!(w.op(), WriteOp::Insert | WriteOp::Upsert)
&& let Some(id) = last_insert_id
{
out["last_insert_id"] = json!(id);
}
Ok(out)
} else {
let rows = budget
.run(NAME, async {
use futures::TryStreamExt;
let mut stream = sqlx::query_with(sqlx::AssertSqlSafe(sql), args)
.persistent(persistent)
.fetch(executor);
let mut rows = Vec::new();
while let Some(row) = stream.try_next().await? {
if rows.len() >= max_returning {
return Err(super::connector_helpers::QueryFailure::Limit(format!(
"{NAME} 'returning' set exceeds query.max_limit ({max_returning} rows) — narrow the filter, drop 'returning', or raise the cap"
)));
}
rows.push(row);
}
Ok(rows)
})
.await?;
let returning = rows_to_json(&rows, format).map_err(|e| decode_failure(NAME, e))?;
let count = returning.len();
Ok(json!({ "rows_affected": count, "returning": returning }))
}
}
async fn execute_mongo(
client: &mongodb::Client,
database: &str,
mw: MongoWrite,
) -> Result<(Value, TaskOutcome), DataflowError> {
let db = client.database(database);
match mw {
MongoWrite::Insert { collection, docs } => {
if docs.is_empty() {
return Ok((empty_insert_envelope(), TaskOutcome::Success));
}
let sent = docs.len();
let coll = db.collection::<Document>(&collection);
match coll.insert_many(docs).await {
Ok(res) => {
let mut pairs: Vec<(usize, mongodb::bson::Bson)> =
res.inserted_ids.into_iter().collect();
pairs.sort_by_key(|(i, _)| *i);
let ids: Vec<Option<Value>> = pairs
.into_iter()
.map(|(_, b)| serde_json::to_value(b).ok())
.collect();
bulk_result(query::bulk::BulkOutcome::all_ok(ids), "MongoDB insert")
}
Err(e) => match mongo_write_errors(&e) {
Some(failed) => {
let outcome = query::backend::mongo::insert_outcome(sent, &failed);
if outcome.nothing_applied()
&& let Some(err) =
super::mongo_common::all_duplicate_key(&failed, "MongoDB insert")
{
return Err(err);
}
bulk_result(outcome, "MongoDB insert")
}
None => Err(super::mongo_common::mongo_error(e)),
},
}
}
MongoWrite::Update {
collection,
filter,
update,
upsert,
multi,
} => {
let coll = db.collection::<Document>(&collection);
let res = if multi {
coll.update_many(filter, update)
.await
.map_err(super::mongo_common::mongo_error)?
} else {
coll.update_one(filter, update)
.upsert(upsert)
.await
.map_err(super::mongo_common::mongo_error)?
};
Ok((update_envelope(&res), TaskOutcome::Success))
}
MongoWrite::Delete { collection, filter } => {
let coll = db.collection::<Document>(&collection);
let res = coll
.delete_many(filter)
.await
.map_err(super::mongo_common::mongo_error)?;
Ok((delete_envelope(res.deleted_count), TaskOutcome::Success))
}
}
}
pub(super) fn mongo_write_errors(e: &mongodb::error::Error) -> Option<Vec<(usize, Value)>> {
let mongodb::error::ErrorKind::InsertMany(bulk) = e.kind.as_ref() else {
return None;
};
let errors = bulk.write_errors.as_ref().filter(|w| !w.is_empty())?;
Some(
errors
.iter()
.map(|we| {
(
we.index,
json!({ "code": we.code, "message": we.message.clone() }),
)
})
.collect(),
)
}
pub(super) fn bulk_result(
outcome: query::bulk::BulkOutcome,
what: &str,
) -> Result<(Value, TaskOutcome), DataflowError> {
if outcome.nothing_applied() {
let first = outcome.first_error().cloned().unwrap_or(Value::Null);
return Err(DataflowError::function_execution(
format!("{what} failed, no documents were written: {first}"),
None,
));
}
let task = if outcome.is_partial() {
TaskOutcome::Status(PARTIAL_STATUS)
} else {
TaskOutcome::Success
};
Ok((outcome.to_json(), task))
}
fn es_url(base: &str, segments: &[&str], query: &[(&str, &str)]) -> Result<String, DataflowError> {
let mut url = url::Url::parse(base).map_err(to_exec_error)?;
url.path_segments_mut()
.map_err(|_| DataflowError::Validation("es connector url cannot be a base".to_string()))?
.pop_if_empty()
.extend(segments);
for (k, v) in query {
url.query_pairs_mut().append_pair(k, v);
}
Ok(url.into())
}
async fn run_es_write(
client: &reqwest::Client,
es: &EsConnectorConfig,
ew: EsWrite,
) -> Result<(Value, TaskOutcome), DataflowError> {
match ew {
EsWrite::BulkInsert { index, docs } => {
if docs.is_empty() {
return Ok((
json!({ "status": "ok", "inserted": 0, "ids": [] }),
TaskOutcome::Success,
));
}
let sent = docs.len();
let url = es_url(&es.url, &[&index, "_bulk"], &[("refresh", "wait_for")])?;
let req = es_request(client, es, reqwest::Method::POST, &url)
.await?
.header("Content-Type", "application/x-ndjson")
.body(bulk_ndjson(&docs));
let (status, body) = send_es(req, es.max_response_size).await?;
if !status.is_success() {
return Err(es_write_error(status, &body));
}
bulk_result(
query::backend::es::bulk_outcome(&body, sent),
"Elasticsearch bulk insert",
)
}
EsWrite::UpdateByQuery { index, body } => {
let resp = run_by_query(client, es, &index, "_update_by_query", &body).await?;
Ok((
json!({ "status": "ok", "matched": resp["total"], "modified": resp["updated"] }),
TaskOutcome::Success,
))
}
EsWrite::DeleteByQuery { index, body } => {
let resp = run_by_query(client, es, &index, "_delete_by_query", &body).await?;
Ok((
json!({ "status": "ok", "deleted": resp["deleted"] }),
TaskOutcome::Success,
))
}
EsWrite::UpdateDoc { index, id, body } => {
let url = es_url(
&es.url,
&[&index, "_update", &id],
&[("refresh", "wait_for")],
)?;
let req = es_request(client, es, reqwest::Method::POST, &url)
.await?
.json(&body);
let (status, resp) = send_es(req, es.max_response_size).await?;
if !status.is_success() {
return Err(es_write_error(status, &resp));
}
let out = match resp.get("result").and_then(|r| r.as_str()) {
Some("created") => {
json!({ "status": "ok", "matched": 0, "modified": 0, "upserted_id": id })
}
Some("noop") => json!({ "status": "ok", "matched": 1, "modified": 0 }),
_ => json!({ "status": "ok", "matched": 1, "modified": 1 }),
};
Ok((out, TaskOutcome::Success))
}
EsWrite::CreateDoc { index, id, doc } => {
let url = es_url(
&es.url,
&[&index, "_doc", &id],
&[("op_type", "create"), ("refresh", "wait_for")],
)?;
let req = es_request(client, es, reqwest::Method::PUT, &url)
.await?
.json(&doc);
let (status, resp) = send_es(req, es.max_response_size).await?;
if status == reqwest::StatusCode::CONFLICT {
return Ok((
json!({ "status": "ok", "matched": 1, "modified": 0 }),
TaskOutcome::Success,
));
}
if !status.is_success() {
return Err(es_write_error(status, &resp));
}
Ok((
json!({ "status": "ok", "matched": 0, "modified": 0, "upserted_id": id }),
TaskOutcome::Success,
))
}
}
}
async fn run_by_query(
client: &reqwest::Client,
es: &EsConnectorConfig,
index: &str,
endpoint: &str,
body: &Value,
) -> Result<Value, DataflowError> {
let url = es_url(&es.url, &[index, endpoint], &[("refresh", "true")])?;
let req = es_request(client, es, reqwest::Method::POST, &url)
.await?
.json(body);
let (status, resp) = send_es(req, es.max_response_size).await?;
if !status.is_success() {
return Err(es_write_error(status, &resp));
}
let failed = resp
.get("failures")
.and_then(|f| f.as_array())
.is_some_and(|f| !f.is_empty())
|| resp
.get("version_conflicts")
.and_then(|v| v.as_u64())
.unwrap_or(0)
> 0;
if failed {
return Err(DataflowError::function_execution(
format!("Elasticsearch {endpoint} had failures: {resp}"),
None,
));
}
Ok(resp)
}
pub(super) const DATA_WRITE_FIELDS: &[FieldSchema] = &[
FieldSchema {
name: "connector",
description: "Name of the db (SQL/MongoDB) or es (Elasticsearch) connector to write to.",
kind: FieldKind::String,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "write",
description: "Backend-neutral mutation envelope: op/target/values/set/filter/on_conflict/returning/all.",
kind: FieldKind::Object,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "params",
description: "Object of named values folded into {\"param\": ..} nodes in values/set/filter. \
A value of {\"var\": \"path\"} is read from the message context.",
kind: FieldKind::Object,
resolvable: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "schema",
description: "Inline entity schema (renames, allowlist, writable flag). Undeclared \
entities and columns are rejected, so a write without one reaches \
nothing; pass {\"unmapped\": \"identity\"} for pre-1.0 pass-through.",
kind: FieldKind::Object,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "database",
description: "MongoDB database name. Optional here because the same task shape is \
valid against SQL and Elasticsearch, which need no database key; \
required — and checked at workflow activation — once the referenced \
connector is a MongoDB one (F52).",
kind: FieldKind::String,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "numeric_as",
description: "How an arbitrary-precision decimal column is rendered: \"number\" (default) or \"string\". A number is computable in JSONLogic and rounds beyond 2^53 or on most decimal fractions; a string keeps every digit, which is what a money column needs. SQL backends only.",
kind: FieldKind::String,
resolvable: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "binary_as",
description: "How a binary column is rendered: \"auto\" (default), \"hex\", \"base64\" or \"text\". Auto reads the bytes as text when they are valid UTF-8 and as hex when they are not, so its result shape depends on the data; name an encoding for a column that is genuinely binary. SQL backends only.",
kind: FieldKind::String,
resolvable: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "output",
description: "Dotted path in the message where the write result is written. Defaults to \"data\".",
kind: FieldKind::String,
template_at: &[""],
..FieldSchema::DEFAULT
},
];
pub(super) const DATA_WRITE_ENVELOPE_FIELDS: &[FieldSchema] = &[
FieldSchema {
name: "op",
description: "Mutation kind: \"insert\", \"update\", \"delete\", or \"upsert\".",
kind: FieldKind::String,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "target",
description: "Logical entity to write to (schema-resolved to a table/collection).",
kind: FieldKind::String,
required: true,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "values",
description: "Row object, or array of row objects (bulk), for insert/upsert.",
kind: FieldKind::Any,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "set",
description: "Object of column → value assignments for update/upsert.",
kind: FieldKind::Object,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "filter",
description: "Query-dialect filter selecting rows for update/delete. An \
update/delete without it is rejected unless \"all\": true is set.",
kind: FieldKind::Object,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "on_conflict",
description: "Upsert conflict clause: { \"target\": [cols], \"action\": \"update\"|\"nothing\" }.",
kind: FieldKind::Object,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "returning",
description: "Column names to return from mutated rows (Postgres/SQLite only).",
kind: FieldKind::Array,
..FieldSchema::DEFAULT
},
FieldSchema {
name: "all",
description: "Acknowledge an intentionally unfiltered update/delete (affects every row).",
kind: FieldKind::Bool,
..FieldSchema::DEFAULT
},
];
#[cfg(test)]
mod tests {
use super::*;
use crate::query::bulk::{BulkOutcome, ItemOutcome};
#[test]
fn a_clean_bulk_is_a_plain_success() {
let (out, task) = bulk_result(
BulkOutcome::all_ok(vec![Some(json!("a")), Some(json!("b"))]),
"MongoDB insert",
)
.expect("a clean bulk is not an error");
assert_eq!(task, TaskOutcome::Success);
assert_eq!(out["status"], "ok", "{out}");
assert_eq!(out["inserted"], 2, "{out}");
}
#[test]
fn a_partial_bulk_is_multi_status_not_an_error() {
let (out, task) = bulk_result(
BulkOutcome {
items: vec![
ItemOutcome::ok(0, Some(json!("a"))),
ItemOutcome::error(1, json!({ "code": 11000 })),
ItemOutcome::skipped(2),
],
},
"MongoDB insert",
)
.expect("a partial write must not fail the task");
assert_eq!(
task,
TaskOutcome::Status(PARTIAL_STATUS),
"a partial write must be visible in the audit trail as 207"
);
assert_eq!(out["status"], "partial", "{out}");
assert_eq!(out["items"][1]["index"], 1, "{out}");
}
#[test]
fn a_bulk_that_applied_nothing_is_a_hard_error() {
let err = bulk_result(
BulkOutcome {
items: vec![
ItemOutcome::error(0, json!({ "message": "duplicate key" })),
ItemOutcome::skipped(1),
],
},
"MongoDB insert",
)
.expect_err("a bulk that wrote nothing is a failure");
let msg = err.to_string();
assert!(msg.contains("MongoDB insert"), "{msg}");
assert!(msg.contains("duplicate key"), "must name why: {msg}");
}
}