use serde_json::{Map, Value as Json};
use crate::config::WriteConfig;
use crate::query::error::QueryError;
use crate::query::ir::{self, Cond};
use crate::query::lower::{Params, lower_with};
use crate::query::schema::EntityRegistry;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum WriteOp {
Insert,
Update,
Delete,
Upsert,
}
impl WriteOp {
pub fn as_str(self) -> &'static str {
match self {
WriteOp::Insert => "insert",
WriteOp::Update => "update",
WriteOp::Delete => "delete",
WriteOp::Upsert => "upsert",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConflictAction {
Update,
Nothing,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ResolvedConflict {
pub targets: Vec<String>,
pub action: ConflictAction,
}
#[derive(Debug, Clone, PartialEq)]
pub enum ResolvedWrite {
Insert {
table: String,
columns: Vec<String>,
rows: Vec<Vec<ir::Value>>,
returning: Vec<String>,
},
Update {
table: String,
set: Vec<(String, ir::Value)>,
cond: Option<Cond>,
returning: Vec<String>,
},
Delete {
table: String,
cond: Option<Cond>,
returning: Vec<String>,
},
Upsert {
table: String,
columns: Vec<String>,
rows: Vec<Vec<ir::Value>>,
set: Vec<(String, ir::Value)>,
conflict: ResolvedConflict,
returning: Vec<String>,
},
}
impl ResolvedWrite {
pub fn op(&self) -> WriteOp {
match self {
ResolvedWrite::Insert { .. } => WriteOp::Insert,
ResolvedWrite::Update { .. } => WriteOp::Update,
ResolvedWrite::Delete { .. } => WriteOp::Delete,
ResolvedWrite::Upsert { .. } => WriteOp::Upsert,
}
}
pub fn returning(&self) -> &[String] {
match self {
ResolvedWrite::Insert { returning, .. }
| ResolvedWrite::Update { returning, .. }
| ResolvedWrite::Delete { returning, .. }
| ResolvedWrite::Upsert { returning, .. } => returning,
}
}
pub fn is_multi_row(&self) -> bool {
match self {
ResolvedWrite::Insert { rows, .. } | ResolvedWrite::Upsert { rows, .. } => {
rows.len() > 1
}
ResolvedWrite::Update { .. } | ResolvedWrite::Delete { .. } => false,
}
}
pub fn effective_filter(&self) -> bool {
match self {
ResolvedWrite::Update { cond, .. } | ResolvedWrite::Delete { cond, .. } => {
cond.as_ref().is_some_and(|c| !c.is_always_true())
}
ResolvedWrite::Insert { .. } | ResolvedWrite::Upsert { .. } => false,
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum WriteError {
MissingField { field: String, op: String },
UnfilteredMutation { op: String },
UnfilteredNotAllowed { op: String },
TooManyRows { requested: usize, max: u64 },
Query(QueryError),
}
impl WriteError {
fn invalid_envelope(msg: impl Into<String>) -> Self {
WriteError::Query(QueryError::InvalidEnvelope(msg.into()))
}
}
impl std::fmt::Display for WriteError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
WriteError::MissingField { field, op } => {
write!(f, "'{op}' requires a '{field}' field")
}
WriteError::UnfilteredMutation { op } => write!(
f,
"'{op}' has no filter; set \"all\": true to intentionally affect every row"
),
WriteError::UnfilteredNotAllowed { op } => write!(
f,
"unfiltered '{op}' is disabled (enable write.allow_unfiltered to permit it)"
),
WriteError::TooManyRows { requested, max } => write!(
f,
"insert of {requested} rows exceeds the configured maximum {max}"
),
WriteError::Query(e) => write!(f, "{e}"),
}
}
}
impl std::error::Error for WriteError {}
impl From<QueryError> for WriteError {
fn from(e: QueryError) -> Self {
WriteError::Query(e)
}
}
impl From<WriteError> for dataflow_rs::engine::error::DataflowError {
fn from(e: WriteError) -> Self {
if matches!(&e, WriteError::Query(q) if q.is_connector_detail()) {
return crate::errors::connector_detail_error(e);
}
dataflow_rs::engine::error::DataflowError::Validation(e.to_string())
}
}
const ENVELOPE_KEYS: [&str; 8] = [
"op",
"target",
"values",
"set",
"filter",
"on_conflict",
"returning",
"all",
];
pub fn resolve_write(
input: &Json,
params: &Params,
reg: &EntityRegistry,
cfg: &WriteConfig,
) -> Result<ResolvedWrite, WriteError> {
let obj = input.as_object().ok_or_else(|| {
WriteError::invalid_envelope("write envelope must be a JSON object".to_string())
})?;
if let Some(unknown) = obj.keys().find(|k| !ENVELOPE_KEYS.contains(&k.as_str())) {
return Err(WriteError::invalid_envelope(format!(
"unknown key '{unknown}' in write envelope (expected \
op/target/values/set/filter/on_conflict/returning/all)"
)));
}
let op = match input.get("op").and_then(|v| v.as_str()) {
Some("insert") => WriteOp::Insert,
Some("update") => WriteOp::Update,
Some("delete") => WriteOp::Delete,
Some("upsert") => WriteOp::Upsert,
Some(other) => {
return Err(WriteError::invalid_envelope(format!(
"unknown op '{other}' (expected insert/update/delete/upsert)"
)));
}
None => {
return Err(WriteError::invalid_envelope(
"missing required string field 'op'".to_string(),
));
}
};
let target = input
.get("target")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.ok_or_else(|| {
WriteError::invalid_envelope("missing required string field 'target'".to_string())
})?
.to_string();
let table = reg.physical_table(&target)?;
let (columns, rows) = parse_rows(input.get("values"), params, reg, &target)?;
let set = parse_set(input.get("set"), params, reg, &target)?;
let filter_node = input.get("filter").filter(|v| !v.is_null());
let cond = match filter_node {
Some(f) => Some(lower_with(f, params, reg, &target)?),
None => None,
};
let conflict = parse_conflict(input.get("on_conflict"), reg, &target)?;
let returning = parse_returning(input.get("returning"), reg, &target)?;
let all = input.get("all").and_then(|v| v.as_bool()).unwrap_or(false);
let missing = |field: &str| WriteError::MissingField {
field: field.to_string(),
op: op.as_str().to_string(),
};
let w = match op {
WriteOp::Insert => {
if rows.is_empty() {
return Err(missing("values"));
}
ResolvedWrite::Insert {
table,
columns,
rows,
returning,
}
}
WriteOp::Update => {
if set.is_empty() {
return Err(missing("set"));
}
ResolvedWrite::Update {
table,
set,
cond,
returning,
}
}
WriteOp::Delete => ResolvedWrite::Delete {
table,
cond,
returning,
},
WriteOp::Upsert => {
if rows.is_empty() {
return Err(missing("values"));
}
let conflict = conflict.ok_or_else(|| missing("on_conflict"))?;
ResolvedWrite::Upsert {
table,
columns,
rows,
set,
conflict,
returning,
}
}
};
if let ResolvedWrite::Insert { rows, .. } | ResolvedWrite::Upsert { rows, .. } = &w
&& rows.len() as u64 > cfg.max_rows
{
return Err(WriteError::TooManyRows {
requested: rows.len(),
max: cfg.max_rows,
});
}
if matches!(op, WriteOp::Update | WriteOp::Delete) && !w.effective_filter() {
if !all {
return Err(WriteError::UnfilteredMutation {
op: op.as_str().to_string(),
});
}
if !cfg.allow_unfiltered {
return Err(WriteError::UnfilteredNotAllowed {
op: op.as_str().to_string(),
});
}
}
Ok(w)
}
fn parse_rows(
node: Option<&Json>,
params: &Params,
reg: &EntityRegistry,
entity: &str,
) -> Result<(Vec<String>, Vec<Vec<ir::Value>>), WriteError> {
let raw_rows: Vec<&Map<String, Json>> = match node {
None | Some(Json::Null) => return Ok((Vec::new(), Vec::new())),
Some(Json::Object(m)) => vec![m],
Some(Json::Array(a)) => {
let mut out = Vec::with_capacity(a.len());
for (i, r) in a.iter().enumerate() {
out.push(r.as_object().ok_or_else(|| {
WriteError::invalid_envelope(format!("values[{i}] must be an object"))
})?);
}
out
}
Some(_) => {
return Err(WriteError::invalid_envelope(
"'values' must be an object or an array of objects".to_string(),
));
}
};
if raw_rows.is_empty() {
return Ok((Vec::new(), Vec::new()));
}
let logical: Vec<String> = raw_rows[0].keys().cloned().collect();
if logical.is_empty() {
return Err(WriteError::invalid_envelope(
"'values' rows must have at least one column".to_string(),
));
}
let columns: Vec<String> = logical
.iter()
.map(|c| reg.resolve_write_column(entity, c, "values"))
.collect::<Result<_, _>>()?;
let mut rows = Vec::with_capacity(raw_rows.len());
for (i, r) in raw_rows.iter().enumerate() {
if r.len() != logical.len() || logical.iter().any(|k| !r.contains_key(k)) {
return Err(WriteError::invalid_envelope(format!(
"values[{i}] must have the same columns as the first row"
)));
}
let mut vals = Vec::with_capacity(logical.len());
for c in &logical {
vals.push(resolve_value_node(
&r[c],
params,
&format!("values[{i}].{c}"),
)?);
}
rows.push(vals);
}
Ok((columns, rows))
}
fn parse_set(
node: Option<&Json>,
params: &Params,
reg: &EntityRegistry,
entity: &str,
) -> Result<Vec<(String, ir::Value)>, WriteError> {
let map = match node {
None | Some(Json::Null) => return Ok(Vec::new()),
Some(Json::Object(m)) => m,
Some(_) => {
return Err(WriteError::invalid_envelope(
"'set' must be an object of column → value".to_string(),
));
}
};
let mut out = Vec::with_capacity(map.len());
for (col, v) in map {
let phys = reg.resolve_write_column(entity, col, "set")?;
out.push((phys, resolve_value_node(v, params, &format!("set.{col}"))?));
}
Ok(out)
}
fn parse_conflict(
node: Option<&Json>,
reg: &EntityRegistry,
entity: &str,
) -> Result<Option<ResolvedConflict>, WriteError> {
let map = match node {
None | Some(Json::Null) => return Ok(None),
Some(Json::Object(m)) => m,
Some(_) => {
return Err(WriteError::invalid_envelope(
"'on_conflict' must be an object".to_string(),
));
}
};
if let Some(unknown) = map
.keys()
.find(|k| !matches!(k.as_str(), "target" | "action"))
{
return Err(WriteError::invalid_envelope(format!(
"unknown key '{unknown}' in on_conflict (expected target/action)"
)));
}
let targets_raw = map
.get("target")
.and_then(|v| v.as_array())
.ok_or_else(|| {
WriteError::invalid_envelope(
"on_conflict.target must be an array of columns".to_string(),
)
})?;
if targets_raw.is_empty() {
return Err(WriteError::invalid_envelope(
"on_conflict.target must name at least one column".to_string(),
));
}
let mut targets = Vec::with_capacity(targets_raw.len());
for t in targets_raw {
let name = t.as_str().ok_or_else(|| {
WriteError::invalid_envelope("on_conflict.target entries must be strings".to_string())
})?;
targets.push(reg.resolve_write_column(entity, name, "on_conflict.target")?);
}
let action = match map.get("action").and_then(|v| v.as_str()) {
None | Some("update") => ConflictAction::Update,
Some("nothing") => ConflictAction::Nothing,
Some(other) => {
return Err(WriteError::invalid_envelope(format!(
"on_conflict.action '{other}' must be \"update\" or \"nothing\""
)));
}
};
Ok(Some(ResolvedConflict { targets, action }))
}
fn parse_returning(
node: Option<&Json>,
reg: &EntityRegistry,
entity: &str,
) -> Result<Vec<String>, WriteError> {
let arr = match node {
None | Some(Json::Null) => return Ok(Vec::new()),
Some(Json::Array(a)) => a,
Some(_) => {
return Err(WriteError::invalid_envelope(
"'returning' must be an array of column names".to_string(),
));
}
};
let mut out = Vec::with_capacity(arr.len());
for (i, c) in arr.iter().enumerate() {
let name = c.as_str().ok_or_else(|| {
WriteError::invalid_envelope(format!("returning[{i}] must be a string"))
})?;
let field = reg.resolve_field(entity, name, &format!("returning[{i}]"))?;
out.push(field.physical);
}
Ok(out)
}
fn resolve_value_node(node: &Json, params: &Params, at: &str) -> Result<ir::Value, WriteError> {
if let Json::Object(m) = node
&& m.len() == 1
&& let Some(p) = m.get("param")
{
let name = p.as_str().ok_or_else(|| {
WriteError::invalid_envelope(format!("{at}: param name must be a string"))
})?;
let resolved = params.get(name).ok_or_else(|| {
WriteError::Query(QueryError::MissingParam {
name: name.to_string(),
at: at.to_string(),
})
})?;
return json_to_value(resolved, at);
}
json_to_value(node, at)
}
fn json_to_value(j: &Json, at: &str) -> Result<ir::Value, WriteError> {
Ok(match j {
Json::Null => ir::Value::Null,
Json::Bool(b) => ir::Value::Bool(*b),
Json::Number(n) => {
if let Some(i) = n.as_i64() {
ir::Value::Int(i)
} else if let Some(f) = n.as_f64() {
ir::Value::Float(f)
} else {
ir::Value::Str(n.to_string())
}
}
Json::String(s) => ir::Value::Str(s.clone()),
Json::Array(_) | Json::Object(_) => {
return Err(WriteError::Query(QueryError::NotRepresentable {
what: "an array/object column value".to_string(),
at: at.to_string(),
}));
}
})
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn permissive() -> WriteConfig {
WriteConfig {
max_rows: 1000,
allow_unfiltered: true,
}
}
fn resolve(input: Json) -> Result<ResolvedWrite, WriteError> {
resolve_write(
&input,
&Params::new(),
&EntityRegistry::identity(),
&permissive(),
)
}
#[test]
fn missing_op_is_invalid_envelope() {
let err = resolve(json!({ "target": "orders" })).expect_err("no op");
assert!(matches!(
err,
WriteError::Query(QueryError::InvalidEnvelope(_))
));
assert!(err.to_string().contains("op"), "{err}");
}
#[test]
fn unknown_op_is_invalid_envelope_naming_the_op() {
let err = resolve(json!({ "op": "truncate", "target": "orders" })).expect_err("bad op");
let msg = err.to_string();
assert!(msg.contains("truncate"), "{msg}");
assert!(msg.contains("insert/update/delete/upsert"), "{msg}");
}
#[test]
fn missing_target_is_invalid_envelope() {
let err = resolve(json!({ "op": "insert", "values": {"a": 1} })).expect_err("no target");
assert!(err.to_string().contains("target"), "{err}");
}
#[test]
fn empty_target_is_invalid_envelope() {
let err = resolve(json!({ "op": "insert", "target": "", "values": {"a": 1} }))
.expect_err("empty target");
assert!(err.to_string().contains("target"), "{err}");
}
#[test]
fn unknown_envelope_keys_are_rejected_naming_the_key() {
for bad in ["retuning", "vaules", "flter"] {
let mut input = json!({ "op": "insert", "target": "orders", "values": {"a": 1} });
input[bad] = json!(["id"]);
let err = resolve(input).expect_err("unknown key must be rejected");
assert!(
matches!(err, WriteError::Query(QueryError::InvalidEnvelope(_))),
"{err}"
);
assert!(err.to_string().contains(bad), "{err}");
}
}
#[test]
fn unknown_on_conflict_keys_are_rejected() {
let err = resolve(json!({
"op": "upsert", "target": "users",
"values": { "email": "a@x.io" },
"on_conflict": { "target": ["email"], "action": "update", "do": "nothing" }
}))
.expect_err("unknown on_conflict key must be rejected");
assert!(err.to_string().contains("'do'"), "{err}");
assert!(err.to_string().contains("on_conflict"), "{err}");
}
#[test]
fn insert_without_values_is_missing_field() {
let err = resolve(json!({ "op": "insert", "target": "orders" })).expect_err("no values");
assert!(
matches!(&err, WriteError::MissingField { field, op } if field == "values" && op == "insert"),
"{err}"
);
}
#[test]
fn update_without_set_is_missing_field() {
let err = resolve(json!({ "op": "update", "target": "orders" })).expect_err("no set");
assert!(
matches!(&err, WriteError::MissingField { field, op } if field == "set" && op == "update"),
"{err}"
);
}
#[test]
fn upsert_without_on_conflict_is_missing_field() {
let err = resolve(json!({ "op": "upsert", "target": "orders", "values": {"id": 1} }))
.expect_err("no on_conflict");
assert!(
matches!(&err, WriteError::MissingField { field, op } if field == "on_conflict" && op == "upsert"),
"{err}"
);
}
#[test]
fn nested_object_column_value_is_not_representable() {
let err = resolve(json!({
"op": "insert",
"target": "orders",
"values": { "meta": { "nested": true } }
}))
.expect_err("object value");
assert!(
matches!(err, WriteError::Query(QueryError::NotRepresentable { .. })),
"{err}"
);
assert!(err.to_string().contains("meta"), "{err}");
}
#[test]
fn array_column_value_is_not_representable() {
let err = resolve(json!({
"op": "insert",
"target": "orders",
"values": { "tags": ["a", "b"] }
}))
.expect_err("array value");
assert!(
matches!(err, WriteError::Query(QueryError::NotRepresentable { .. })),
"{err}"
);
}
#[test]
fn unknown_param_reference_errors_with_the_name() {
let err = resolve(json!({
"op": "insert",
"target": "orders",
"values": { "total": { "param": "missing_param" } }
}))
.expect_err("unknown param");
assert!(err.to_string().contains("missing_param"), "{err}");
}
#[test]
fn param_reference_resolves_to_the_named_value() {
let mut params = Params::new();
params.insert("amount".to_string(), json!(42));
let resolved = resolve_write(
&json!({
"op": "insert",
"target": "orders",
"values": { "total": { "param": "amount" } }
}),
¶ms,
&EntityRegistry::identity(),
&permissive(),
)
.expect("resolves");
let ResolvedWrite::Insert { rows, .. } = resolved else {
unreachable!("insert resolves to the Insert variant");
};
assert_eq!(rows.len(), 1);
assert!(matches!(rows[0][0], ir::Value::Int(42)));
}
#[test]
fn unfiltered_delete_without_all_is_rejected_by_resolve_write() {
let err = resolve(json!({ "op": "delete", "target": "orders" }))
.expect_err("no filter, no acknowledgement");
assert!(
matches!(&err, WriteError::UnfilteredMutation { op } if op == "delete"),
"{err}"
);
}
#[test]
fn unfiltered_delete_with_all_still_needs_the_config_opt_in() {
let err = resolve_write(
&json!({ "op": "delete", "target": "orders", "all": true }),
&Params::new(),
&EntityRegistry::identity(),
&WriteConfig {
max_rows: 1000,
allow_unfiltered: false,
},
)
.expect_err("config forbids unfiltered mutations");
assert!(
matches!(&err, WriteError::UnfilteredNotAllowed { op } if op == "delete"),
"{err}"
);
}
#[test]
fn unfiltered_delete_with_both_opt_ins_resolves() {
let resolved =
resolve(json!({ "op": "delete", "target": "orders", "all": true })).expect("resolves");
assert!(!resolved.effective_filter());
}
#[test]
fn a_bulk_insert_over_max_rows_is_rejected_by_resolve_write() {
let err = resolve_write(
&json!({
"op": "insert",
"target": "orders",
"values": [ {"a": 1}, {"a": 2}, {"a": 3} ]
}),
&Params::new(),
&EntityRegistry::identity(),
&WriteConfig {
max_rows: 2,
allow_unfiltered: false,
},
)
.expect_err("3 rows over a cap of 2");
assert!(
matches!(
err,
WriteError::TooManyRows {
requested: 3,
max: 2
}
),
"{err}"
);
}
#[test]
fn a_filter_that_folds_to_true_does_not_count_as_filtered() {
for vacuous in [
json!({ "and": [] }),
json!({ "!": { "or": [] } }),
json!({ "and": [{ "and": [] }] }),
] {
let err = resolve(json!({
"op": "delete",
"target": "orders",
"filter": vacuous.clone(),
}))
.expect_err("a vacuous filter restricts nothing");
assert!(
matches!(err, WriteError::UnfilteredMutation { .. }),
"filter {vacuous} must not satisfy the guard"
);
}
}
#[test]
fn a_real_filter_counts_as_filtered() {
let resolved = resolve(json!({
"op": "delete",
"target": "orders",
"filter": { "==": [{ "field": "status" }, "cancelled"] },
}))
.expect("resolves");
assert!(resolved.effective_filter());
}
#[test]
fn an_unsatisfiable_filter_still_counts_as_filtered() {
let resolved = resolve(json!({
"op": "delete",
"target": "orders",
"filter": { "or": [] },
}))
.expect("resolves");
assert!(resolved.effective_filter());
}
}