use std::sync::Arc;
use async_trait::async_trait;
use dataflow_rs::engine::error::DataflowError;
use dataflow_rs::engine::functions::AsyncFunctionHandler;
use dataflow_rs::engine::task_context::TaskContext;
use dataflow_rs::engine::task_outcome::TaskOutcome;
use mongodb::bson::Document;
use serde_json::{Value, json};
use sqlx::any::AnyRow;
use super::connector_helpers::{
ConnectorCall, QueryBudget, apply_output, build_entity_registry, es_request, es_write_error,
is_mongo, require_op_allowed, resolve_params, send_es, timed_query, to_connect_error,
to_exec_error,
};
use super::db_read::rows_to_json;
use super::schema::{FieldKind, FieldSchema};
use crate::connector::mongo_pool::MongoPoolCache;
use crate::connector::pool_cache::SqlPoolCache;
use crate::connector::{ConnectorConfig, ConnectorRegistry, EsConnectorConfig};
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 = "data_write";
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,
}
#[async_trait]
impl AsyncFunctionHandler for DataWriteHandler {
type Input = Value;
async fn execute(
&self,
ctx: &mut TaskContext<'_>,
input: &Value,
) -> dataflow_rs::Result<TaskOutcome> {
let call = ConnectorCall::begin(NAME, input, ctx)?;
let params = resolve_params(input.get("params"), ctx);
call.run(&self.registry, async {
let connector_config = call.resolve(&self.registry, None).await?;
let registry =
build_entity_registry(input.get("schema"), &connector_config, call.connector)?;
let envelope = input.get("write").ok_or_else(|| {
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(),
))
})?;
let resolved =
query::write::resolve_write(envelope, ¶ms, ®istry, &self.write_config)?;
if let Some(gates) = connector_config.operation_gates() {
require_op_allowed(gates, resolved.op().as_str(), call.connector)?;
}
let (result, outcome) = match connector_config.as_ref() {
ConnectorConfig::Db(db) if is_mongo(&db.connection_string) => {
let database = call.require_str(input, "database")?;
let mw = query::backend::mongo::render_write(&resolved)?;
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, async {
execute_mongo(&client, database, mw)
.await
.map_err(|e| e.to_string())
})
.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)?;
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).await?
}
ConnectorConfig::Es(es) => {
let ew = query::backend::es::render_write(&resolved)?;
run_es_write(&self.http_client, es, ew).await?
}
_ => {
return Err(DataflowError::Validation(format!(
"Connector '{}' is not a db or es connector",
call.connector
)));
}
};
apply_output(ctx, call.output, result);
Ok(outcome)
})
.await
}
}
async fn execute_sql(
pool: &sqlx::AnyPool,
sql: &str,
values: sea_query_sqlx::SqlxValues,
w: &ResolvedWrite,
timeout_ms: Option<u64>,
) -> Result<(Value, TaskOutcome), DataflowError> {
let budget = QueryBudget::start(timeout_ms);
let mut out = if w.is_multi_row() {
let mut tx = budget.run(NAME, pool.begin()).await?;
let out = run_write_statement(&mut *tx, sql, values, w, &budget).await?;
budget.run(NAME, tx.commit()).await?;
out
} else {
run_write_statement(pool, sql, values, w, &budget).await?
};
out["status"] = json!("ok");
Ok((out, TaskOutcome::Success))
}
async fn run_write_statement<'e, E>(
executor: E,
sql: &str,
values: sea_query_sqlx::SqlxValues,
w: &ResolvedWrite,
budget: &QueryBudget,
) -> Result<Value, DataflowError>
where
E: sqlx::Executor<'e, Database = sqlx::Any>,
{
if w.returning().is_empty() {
let res = budget
.run(NAME, sqlx::query_with(sql, values).execute(executor))
.await?;
let mut out = json!({ "rows_affected": res.rows_affected() });
if matches!(w.op(), WriteOp::Insert | WriteOp::Upsert)
&& let Some(id) = res.last_insert_id()
{
out["last_insert_id"] = json!(id);
}
Ok(out)
} else {
let rows: Vec<AnyRow> = budget
.run(NAME, sqlx::query_with(sql, values).fetch_all(executor))
.await?;
let returning = rows_to_json(&rows)?;
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((
json!({ "status": "ok", "inserted": 0, "ids": [] }),
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();
Ok((
query::bulk::BulkOutcome::all_ok(ids).to_json(),
TaskOutcome::Success,
))
}
Err(e) => match mongo_write_errors(&e) {
Some(failed) => bulk_result(
query::backend::mongo::insert_outcome(sent, &failed),
"MongoDB insert",
),
None => Err(to_exec_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(to_exec_error)?
} else {
coll.update_one(filter, update)
.upsert(upsert)
.await
.map_err(to_exec_error)?
};
let mut out = json!({
"status": "ok",
"matched": res.matched_count,
"modified": res.modified_count,
});
if let Some(id) = res.upserted_id {
out["upserted_id"] = serde_json::to_value(id).unwrap_or(Value::Null);
}
Ok((out, TaskOutcome::Success))
}
MongoWrite::Delete { collection, filter } => {
let coll = db.collection::<Document>(&collection);
let res = coll.delete_many(filter).await.map_err(to_exec_error)?;
Ok((
json!({ "status": "ok", "deleted": res.deleted_count }),
TaskOutcome::Success,
))
}
}
}
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(),
)
}
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,
resolvable: false,
alias: None,
},
FieldSchema {
name: "write",
description: "Backend-neutral mutation envelope: op/target/values/set/filter/on_conflict/returning/all.",
kind: FieldKind::Object,
required: true,
resolvable: false,
alias: None,
},
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,
required: false,
resolvable: false,
alias: None,
},
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,
required: false,
resolvable: false,
alias: None,
},
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,
required: false,
resolvable: true,
alias: None,
},
FieldSchema {
name: "output",
description: "Dotted path in the message where the write result is written. Defaults to \"data\".",
kind: FieldKind::String,
required: false,
resolvable: false,
alias: None,
},
];
pub(super) const DATA_WRITE_ENVELOPE_FIELDS: &[FieldSchema] = &[
FieldSchema {
name: "op",
description: "Mutation kind: \"insert\", \"update\", \"delete\", or \"upsert\".",
kind: FieldKind::String,
required: true,
resolvable: false,
alias: None,
},
FieldSchema {
name: "target",
description: "Logical entity to write to (schema-resolved to a table/collection).",
kind: FieldKind::String,
required: true,
resolvable: false,
alias: None,
},
FieldSchema {
name: "values",
description: "Row object, or array of row objects (bulk), for insert/upsert.",
kind: FieldKind::Any,
required: false,
resolvable: false,
alias: None,
},
FieldSchema {
name: "set",
description: "Object of column → value assignments for update/upsert.",
kind: FieldKind::Object,
required: false,
resolvable: false,
alias: None,
},
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,
required: false,
resolvable: false,
alias: None,
},
FieldSchema {
name: "on_conflict",
description: "Upsert conflict clause: { \"target\": [cols], \"action\": \"update\"|\"nothing\" }.",
kind: FieldKind::Object,
required: false,
resolvable: false,
alias: None,
},
FieldSchema {
name: "returning",
description: "Column names to return from mutated rows (Postgres/SQLite only).",
kind: FieldKind::Array,
required: false,
resolvable: false,
alias: None,
},
FieldSchema {
name: "all",
description: "Acknowledge an intentionally unfiltered update/delete (affects every row).",
kind: FieldKind::Bool,
required: false,
resolvable: false,
alias: None,
},
];
#[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}");
}
}