use std::sync::Arc;
use async_trait::async_trait;
use dataflow_rs::engine::functions::AsyncFunctionHandler;
use dataflow_rs::engine::task_context::TaskContext;
use dataflow_rs::engine::task_outcome::TaskOutcome;
use serde_json::Value;
use super::connector_helpers::{
ConnectorCall, apply_output, bind_json_params, reject_mongo_connector, require_db_connector,
resolve_bind_params, timed_query, to_connect_error,
};
use super::schema::{FieldKind, FieldSchema};
use crate::connector::ConnectorRegistry;
use crate::connector::pool_cache::SqlPoolCache;
const NAME: &str = "db_write";
pub struct DbWriteHandler {
pub pool_cache: Arc<SqlPoolCache>,
pub registry: Arc<ConnectorRegistry>,
}
#[async_trait]
impl AsyncFunctionHandler for DbWriteHandler {
type Input = Value;
async fn execute(
&self,
ctx: &mut TaskContext<'_>,
input: &Value,
) -> dataflow_rs::Result<TaskOutcome> {
let call = ConnectorCall::begin(NAME, input, ctx)?;
let query = call.require_str(input, "query")?;
let params = resolve_bind_params(input, call.name, ctx)?;
call.run(&self.registry, async {
let connector_config = call.resolve(&self.registry, Some("raw_write")).await?;
let db_config = require_db_connector(&connector_config, call.connector)?;
reject_mongo_connector(call.name, call.connector, db_config)?;
let pool = self
.pool_cache
.get_pool(call.connector, db_config)
.await
.map_err(to_connect_error)?;
let sqlx_query = bind_json_params(sqlx::query(query), ¶ms);
let result = timed_query(
db_config.query_timeout_ms,
call.name,
sqlx_query.execute(&pool),
)
.await?;
let output = serde_json::json!({
"rows_affected": result.rows_affected(),
});
apply_output(ctx, call.output, output);
Ok(TaskOutcome::Success)
})
.await
}
}
pub(super) const DB_WRITE_FIELDS: &[FieldSchema] = &[
FieldSchema {
name: "connector",
description: "Name of the SQL connector to execute against.",
kind: FieldKind::String,
required: true,
resolvable: false,
alias: None,
},
FieldSchema {
name: "query",
description: "INSERT/UPDATE/DELETE statement.",
kind: FieldKind::String,
required: true,
resolvable: false,
alias: None,
},
FieldSchema {
name: "params",
description: "Array of values to bind to query placeholders, in order. Accepts {\"var\": \"path\"} to read the value from the message.",
kind: FieldKind::Array,
required: false,
resolvable: true,
alias: None,
},
FieldSchema {
name: "output",
description: "Dotted path where the rows-affected count is written.",
kind: FieldKind::String,
required: false,
resolvable: false,
alias: None,
},
];