#![allow(dead_code)]
use std::fmt::Debug;
use std::marker::PhantomData;
use crate::error::QueryResult;
use crate::filter::{Filter, FilterValue};
use crate::sql::quote_identifier;
use crate::traits::{Model, QueryEngine};
#[derive(Debug, Clone)]
pub enum NestedWrite<T: Model> {
Create(Vec<NestedCreateData<T>>),
CreateOrConnect(Vec<NestedCreateOrConnectData<T>>),
Connect(Vec<Filter>),
Disconnect(Vec<Filter>),
Set(Vec<Filter>),
Delete(Vec<Filter>),
Update(Vec<NestedUpdateData<T>>),
Upsert(Vec<NestedUpsertData<T>>),
UpdateMany(NestedUpdateManyData<T>),
DeleteMany(Filter),
}
impl<T: Model> NestedWrite<T> {
pub fn create(data: NestedCreateData<T>) -> Self {
Self::Create(vec![data])
}
pub fn create_many(data: Vec<NestedCreateData<T>>) -> Self {
Self::Create(data)
}
pub fn connect_one(filter: impl Into<Filter>) -> Self {
Self::Connect(vec![filter.into()])
}
pub fn connect(filters: Vec<impl Into<Filter>>) -> Self {
Self::Connect(filters.into_iter().map(Into::into).collect())
}
pub fn disconnect_one(filter: impl Into<Filter>) -> Self {
Self::Disconnect(vec![filter.into()])
}
pub fn disconnect(filters: Vec<impl Into<Filter>>) -> Self {
Self::Disconnect(filters.into_iter().map(Into::into).collect())
}
pub fn set(filters: Vec<impl Into<Filter>>) -> Self {
Self::Set(filters.into_iter().map(Into::into).collect())
}
pub fn delete(filters: Vec<impl Into<Filter>>) -> Self {
Self::Delete(filters.into_iter().map(Into::into).collect())
}
pub fn delete_many(filter: impl Into<Filter>) -> Self {
Self::DeleteMany(filter.into())
}
}
#[derive(Debug, Clone)]
pub struct NestedCreateData<T: Model> {
pub data: Vec<(String, FilterValue)>,
_model: PhantomData<T>,
}
impl<T: Model> NestedCreateData<T> {
pub fn new(data: Vec<(String, FilterValue)>) -> Self {
Self {
data,
_model: PhantomData,
}
}
pub fn from_pairs(
pairs: impl IntoIterator<Item = (impl Into<String>, impl Into<FilterValue>)>,
) -> Self {
Self::new(
pairs
.into_iter()
.map(|(k, v)| (k.into(), v.into()))
.collect(),
)
}
}
impl<T: Model> Default for NestedCreateData<T> {
fn default() -> Self {
Self::new(Vec::new())
}
}
#[derive(Debug, Clone)]
pub struct NestedCreateOrConnectData<T: Model> {
pub filter: Filter,
pub create: NestedCreateData<T>,
}
impl<T: Model> NestedCreateOrConnectData<T> {
pub fn new(filter: impl Into<Filter>, create: NestedCreateData<T>) -> Self {
Self {
filter: filter.into(),
create,
}
}
}
#[derive(Debug, Clone)]
pub struct NestedUpdateData<T: Model> {
pub filter: Filter,
pub data: Vec<(String, FilterValue)>,
_model: PhantomData<T>,
}
impl<T: Model> NestedUpdateData<T> {
pub fn new(filter: impl Into<Filter>, data: Vec<(String, FilterValue)>) -> Self {
Self {
filter: filter.into(),
data,
_model: PhantomData,
}
}
pub fn from_pairs(
filter: impl Into<Filter>,
pairs: impl IntoIterator<Item = (impl Into<String>, impl Into<FilterValue>)>,
) -> Self {
Self::new(
filter,
pairs
.into_iter()
.map(|(k, v)| (k.into(), v.into()))
.collect(),
)
}
}
#[derive(Debug, Clone)]
pub struct NestedUpsertData<T: Model> {
pub filter: Filter,
pub create: NestedCreateData<T>,
pub update: Vec<(String, FilterValue)>,
_model: PhantomData<T>,
}
impl<T: Model> NestedUpsertData<T> {
pub fn new(
filter: impl Into<Filter>,
create: NestedCreateData<T>,
update: Vec<(String, FilterValue)>,
) -> Self {
Self {
filter: filter.into(),
create,
update,
_model: PhantomData,
}
}
}
#[derive(Debug, Clone)]
pub struct NestedUpdateManyData<T: Model> {
pub filter: Filter,
pub data: Vec<(String, FilterValue)>,
_model: PhantomData<T>,
}
impl<T: Model> NestedUpdateManyData<T> {
pub fn new(filter: impl Into<Filter>, data: Vec<(String, FilterValue)>) -> Self {
Self {
filter: filter.into(),
data,
_model: PhantomData,
}
}
}
#[derive(Debug)]
pub struct NestedWriteBuilder {
parent_table: String,
parent_pk: Vec<String>,
related_table: String,
foreign_key: String,
is_one_to_many: bool,
join_table: Option<JoinTableInfo>,
}
#[derive(Debug, Clone)]
pub struct JoinTableInfo {
pub table_name: String,
pub parent_column: String,
pub related_column: String,
}
impl NestedWriteBuilder {
pub fn one_to_many(
parent_table: impl Into<String>,
parent_pk: Vec<String>,
related_table: impl Into<String>,
foreign_key: impl Into<String>,
) -> Self {
Self {
parent_table: parent_table.into(),
parent_pk,
related_table: related_table.into(),
foreign_key: foreign_key.into(),
is_one_to_many: true,
join_table: None,
}
}
pub fn many_to_many(
parent_table: impl Into<String>,
parent_pk: Vec<String>,
related_table: impl Into<String>,
join_table: JoinTableInfo,
) -> Self {
Self {
parent_table: parent_table.into(),
parent_pk,
related_table: related_table.into(),
foreign_key: String::new(), is_one_to_many: false,
join_table: Some(join_table),
}
}
pub fn build_connect_sql<T: Model>(
&self,
parent_id: &FilterValue,
filters: &[Filter],
) -> Vec<(String, Vec<FilterValue>)> {
let mut statements = Vec::new();
if self.is_one_to_many {
for filter in filters {
let (where_sql, mut params) = filter.to_sql(0, &crate::dialect::Postgres);
let sql = format!(
"UPDATE {} SET {} = ${} WHERE {}",
quote_identifier(&self.related_table),
quote_identifier(&self.foreign_key),
params.len() + 1,
where_sql
);
params.push(parent_id.clone());
statements.push((sql, params));
}
} else if let Some(join) = &self.join_table {
for filter in filters {
let (where_sql, mut params) = filter.to_sql(0, &crate::dialect::Postgres);
let select_sql = format!(
"SELECT {} FROM {} WHERE {}",
quote_identifier(T::PRIMARY_KEY.first().unwrap_or(&"id")),
quote_identifier(&self.related_table),
where_sql
);
let insert_sql = format!(
"INSERT INTO {} ({}, {}) SELECT ${}, {} FROM {} WHERE {} ON CONFLICT DO NOTHING",
quote_identifier(&join.table_name),
quote_identifier(&join.parent_column),
quote_identifier(&join.related_column),
params.len() + 1,
quote_identifier(T::PRIMARY_KEY.first().unwrap_or(&"id")),
quote_identifier(&self.related_table),
where_sql
);
params.push(parent_id.clone());
statements.push((insert_sql, params));
let _ = select_sql;
}
}
statements
}
pub fn build_disconnect_sql(
&self,
parent_id: &FilterValue,
filters: &[Filter],
) -> Vec<(String, Vec<FilterValue>)> {
let mut statements = Vec::new();
if self.is_one_to_many {
for filter in filters {
let (where_sql, mut params) = filter.to_sql(0, &crate::dialect::Postgres);
let sql = format!(
"UPDATE {} SET {} = NULL WHERE {} AND {} = ${}",
quote_identifier(&self.related_table),
quote_identifier(&self.foreign_key),
where_sql,
quote_identifier(&self.foreign_key),
params.len() + 1
);
params.push(parent_id.clone());
statements.push((sql, params));
}
} else if let Some(join) = &self.join_table {
for filter in filters {
let (where_sql, mut params) = filter.to_sql(1, &crate::dialect::Postgres);
let sql = format!(
"DELETE FROM {} WHERE {} = $1 AND {} IN (SELECT id FROM {} WHERE {})",
quote_identifier(&join.table_name),
quote_identifier(&join.parent_column),
quote_identifier(&join.related_column),
quote_identifier(&self.related_table),
where_sql
);
let mut final_params = vec![parent_id.clone()];
final_params.extend(params);
params = final_params;
statements.push((sql, params));
}
}
statements
}
pub fn build_set_sql<T: Model>(
&self,
parent_id: &FilterValue,
filters: &[Filter],
) -> Vec<(String, Vec<FilterValue>)> {
let mut statements = Vec::new();
if self.is_one_to_many {
let sql = format!(
"UPDATE {} SET {} = NULL WHERE {} = $1",
quote_identifier(&self.related_table),
quote_identifier(&self.foreign_key),
quote_identifier(&self.foreign_key)
);
statements.push((sql, vec![parent_id.clone()]));
} else if let Some(join) = &self.join_table {
let sql = format!(
"DELETE FROM {} WHERE {} = $1",
quote_identifier(&join.table_name),
quote_identifier(&join.parent_column)
);
statements.push((sql, vec![parent_id.clone()]));
}
statements.extend(self.build_connect_sql::<T>(parent_id, filters));
statements
}
pub fn build_create_sql<T: Model>(
&self,
parent_id: &FilterValue,
creates: &[NestedCreateData<T>],
) -> Vec<(String, Vec<FilterValue>)> {
let mut statements = Vec::with_capacity(creates.len());
let quoted_table = quote_identifier(&self.related_table);
for create in creates {
let row_len = create.data.len() + 1;
let mut columns: Vec<String> = Vec::with_capacity(row_len);
let mut values: Vec<FilterValue> = Vec::with_capacity(row_len);
for (k, v) in &create.data {
columns.push(k.clone());
values.push(v.clone());
}
columns.push(self.foreign_key.clone());
values.push(parent_id.clone());
let mut col_list = String::new();
let mut placeholders = String::new();
for (i, c) in columns.iter().enumerate() {
if i > 0 {
col_list.push_str(", ");
placeholders.push_str(", ");
}
col_list.push_str("e_identifier(c));
use std::fmt::Write;
let _ = write!(placeholders, "${}", i + 1);
}
let sql = format!(
"INSERT INTO {} ({}) VALUES ({}) RETURNING *",
quoted_table, col_list, placeholders,
);
statements.push((sql, values));
}
statements
}
pub fn build_delete_sql(
&self,
parent_id: &FilterValue,
filters: &[Filter],
) -> Vec<(String, Vec<FilterValue>)> {
let mut statements = Vec::new();
for filter in filters {
let (where_sql, mut params) = filter.to_sql(0, &crate::dialect::Postgres);
let sql = format!(
"DELETE FROM {} WHERE {} AND {} = ${}",
quote_identifier(&self.related_table),
where_sql,
quote_identifier(&self.foreign_key),
params.len() + 1
);
params.push(parent_id.clone());
statements.push((sql, params));
}
statements
}
}
#[derive(Debug, Default)]
pub struct NestedWriteOperations {
pub pre_statements: Vec<(String, Vec<FilterValue>)>,
pub post_statements: Vec<(String, Vec<FilterValue>)>,
}
impl NestedWriteOperations {
pub fn new() -> Self {
Self::default()
}
pub fn add_pre(&mut self, sql: String, params: Vec<FilterValue>) {
self.pre_statements.push((sql, params));
}
pub fn add_post(&mut self, sql: String, params: Vec<FilterValue>) {
self.post_statements.push((sql, params));
}
pub fn extend(&mut self, other: Self) {
self.pre_statements.extend(other.pre_statements);
self.post_statements.extend(other.post_statements);
}
pub fn is_empty(&self) -> bool {
self.pre_statements.is_empty() && self.post_statements.is_empty()
}
pub fn len(&self) -> usize {
self.pre_statements.len() + self.post_statements.len()
}
}
#[derive(Debug, Clone)]
pub enum NestedWriteOp {
Create {
relation: &'static str,
target_table: &'static str,
foreign_key: &'static str,
payload: Vec<Vec<(String, FilterValue)>>,
},
Connect {
relation: &'static str,
target_table: &'static str,
foreign_key: &'static str,
target_pk: &'static str,
pk: FilterValue,
},
Disconnect {
relation: &'static str,
target_table: &'static str,
foreign_key: &'static str,
target_pk: &'static str,
pk: FilterValue,
},
Delete {
relation: &'static str,
target_table: &'static str,
target_pk: &'static str,
pk: FilterValue,
},
DeleteMany {
relation: &'static str,
target_table: &'static str,
foreign_key: &'static str,
filter: Filter,
},
Update {
relation: &'static str,
target_table: &'static str,
target_pk: &'static str,
pk: FilterValue,
payload: Vec<(String, crate::inputs::WriteOp)>,
},
UpdateMany {
relation: &'static str,
target_table: &'static str,
foreign_key: &'static str,
filter: Filter,
payload: Vec<(String, crate::inputs::WriteOp)>,
},
Upsert {
relation: &'static str,
target_table: &'static str,
foreign_key: &'static str,
target_pk: &'static str,
pk: FilterValue,
create_payload: Vec<(String, FilterValue)>,
update_payload: Vec<(String, crate::inputs::WriteOp)>,
},
ConnectOrCreate {
relation: &'static str,
target_table: &'static str,
foreign_key: &'static str,
where_filter: Filter,
create_payload: Vec<(String, FilterValue)>,
},
Set {
relation: &'static str,
target_table: &'static str,
foreign_key: &'static str,
target_pk: &'static str,
set_pks: Vec<FilterValue>,
},
}
fn build_writeop_set_clause(
payload: &[(String, crate::inputs::WriteOp)],
dialect: &dyn crate::dialect::SqlDialect,
start_ph: usize,
) -> (String, Vec<FilterValue>) {
let mut fragments: Vec<String> = Vec::with_capacity(payload.len());
let mut params: Vec<FilterValue> = Vec::with_capacity(payload.len());
let mut next_ph = start_ph;
for (col, op) in payload {
let (frag, maybe_val) =
op.to_set_fragment(&dialect.quote_ident(col), &dialect.placeholder(next_ph));
fragments.push(frag);
if let Some(val) = maybe_val {
params.push(val);
next_ph += 1;
}
}
(fragments.join(", "), params)
}
impl NestedWriteOp {
pub async fn execute<E>(self, engine: &E, parent_pk: &FilterValue) -> QueryResult<()>
where
E: QueryEngine,
{
match self {
NestedWriteOp::Connect {
relation: _,
target_table,
foreign_key,
target_pk,
pk,
} => {
let dialect = engine.dialect();
let sql = format!(
"UPDATE {} SET {} = {} WHERE {} = {}",
dialect.quote_ident(target_table),
dialect.quote_ident(foreign_key),
dialect.placeholder(1),
dialect.quote_ident(target_pk),
dialect.placeholder(2),
);
engine
.execute_raw(&sql, vec![parent_pk.clone(), pk])
.await?;
Ok(())
}
NestedWriteOp::Disconnect {
relation: _,
target_table,
foreign_key,
target_pk,
pk,
} => {
let dialect = engine.dialect();
let sql = format!(
"UPDATE {} SET {} = NULL WHERE {} = {}",
dialect.quote_ident(target_table),
dialect.quote_ident(foreign_key),
dialect.quote_ident(target_pk),
dialect.placeholder(1),
);
engine.execute_raw(&sql, vec![pk]).await?;
Ok(())
}
NestedWriteOp::Delete {
relation: _,
target_table,
target_pk,
pk,
} => {
let dialect = engine.dialect();
let sql = format!(
"DELETE FROM {} WHERE {} = {}",
dialect.quote_ident(target_table),
dialect.quote_ident(target_pk),
dialect.placeholder(1),
);
let affected = engine.execute_raw(&sql, vec![pk]).await?;
if affected != 1 {
return Err(crate::error::QueryError::not_found(target_table)
.with_context("Nested Delete by PK"));
}
Ok(())
}
NestedWriteOp::DeleteMany {
relation: _,
target_table,
foreign_key,
filter,
} => {
let dialect = engine.dialect();
let is_unconstrained = matches!(filter, Filter::None);
let sql = if is_unconstrained {
format!(
"DELETE FROM {} WHERE {} = {}",
dialect.quote_ident(target_table),
dialect.quote_ident(foreign_key),
dialect.placeholder(1),
)
} else {
let (filter_sql, params_tail) = filter.to_sql(1, dialect);
let sql = format!(
"DELETE FROM {} WHERE {} = {} AND ({})",
dialect.quote_ident(target_table),
dialect.quote_ident(foreign_key),
dialect.placeholder(1),
filter_sql,
);
let mut params = Vec::with_capacity(params_tail.len() + 1);
params.push(parent_pk.clone());
params.extend(params_tail);
return engine.execute_raw(&sql, params).await.map(|_| ());
};
engine.execute_raw(&sql, vec![parent_pk.clone()]).await?;
Ok(())
}
NestedWriteOp::Update {
relation: _,
target_table,
target_pk,
pk,
payload,
} => {
if payload.is_empty() {
return Ok(());
}
let dialect = engine.dialect();
let (set_text, mut update_params) = build_writeop_set_clause(&payload, dialect, 1);
let next_placeholder = update_params.len() + 1;
update_params.push(pk);
let sql = format!(
"UPDATE {} SET {} WHERE {} = {}",
dialect.quote_ident(target_table),
set_text,
dialect.quote_ident(target_pk),
dialect.placeholder(next_placeholder),
);
let affected = engine.execute_raw(&sql, update_params).await?;
if affected != 1 {
return Err(crate::error::QueryError::not_found(target_table)
.with_context("Nested Update by PK"));
}
Ok(())
}
NestedWriteOp::UpdateMany {
relation: _,
target_table,
foreign_key,
filter,
payload,
} => {
if payload.is_empty() {
return Ok(());
}
let dialect = engine.dialect();
let (set_text, mut params) = build_writeop_set_clause(&payload, dialect, 1);
let fk_placeholder_idx = params.len() + 1;
let fk_placeholder = dialect.placeholder(fk_placeholder_idx);
params.push(parent_pk.clone());
let is_unconstrained = matches!(filter, Filter::None);
let sql = if is_unconstrained {
format!(
"UPDATE {} SET {} WHERE {} = {}",
dialect.quote_ident(target_table),
set_text,
dialect.quote_ident(foreign_key),
fk_placeholder,
)
} else {
let (filter_sql, filter_params) = filter.to_sql(fk_placeholder_idx, dialect);
params.extend(filter_params);
format!(
"UPDATE {} SET {} WHERE {} = {} AND ({})",
dialect.quote_ident(target_table),
set_text,
dialect.quote_ident(foreign_key),
fk_placeholder,
filter_sql,
)
};
engine.execute_raw(&sql, params).await?;
Ok(())
}
NestedWriteOp::Upsert {
relation: _,
target_table,
foreign_key,
target_pk,
pk,
create_payload,
update_payload,
} => {
let dialect = engine.dialect();
if !dialect.supports_upsert() {
return Err(crate::error::QueryError::unsupported(format!(
"Nested Upsert is not supported by the `{}` engine",
std::any::type_name::<dyn crate::dialect::SqlDialect>()
)));
}
if update_payload.is_empty() && create_payload.is_empty() {
return Ok(());
}
if create_payload.is_empty() {
return Err(crate::error::QueryError::invalid_input(
"create_payload",
"Nested Upsert requires at least one create column when no row to update",
));
}
let insert_arity = create_payload.len() + 1; let (probe_set_text, probe_set_params) =
build_writeop_set_clause(&update_payload, dialect, insert_arity + 1);
let build_insert_columns_and_values = || {
let mut columns: Vec<String> =
create_payload.iter().map(|(c, _)| c.clone()).collect();
let mut values: Vec<FilterValue> =
create_payload.iter().map(|(_, v)| v.clone()).collect();
columns.push(foreign_key.to_string());
values.push(parent_pk.clone());
let placeholders: Vec<String> =
(1..=values.len()).map(|i| dialect.placeholder(i)).collect();
let quoted_cols: Vec<String> =
columns.iter().map(|c| dialect.quote_ident(c)).collect();
(columns, values, placeholders, quoted_cols)
};
let conflict_cols = [target_pk];
let upsert_clause_text = dialect.upsert_clause(&conflict_cols, &probe_set_text);
if !upsert_clause_text.is_empty() {
let (_, mut values, placeholders, quoted_cols) =
build_insert_columns_and_values();
let _ = pk;
let effective_upsert_clause = if update_payload.is_empty() {
let do_nothing = dialect.upsert_do_nothing_clause(&conflict_cols);
if do_nothing.is_empty() {
return Ok(());
}
do_nothing
} else {
values.extend(probe_set_params);
upsert_clause_text
};
let sql = format!(
"INSERT INTO {} ({}) VALUES ({}){}",
dialect.quote_ident(target_table),
quoted_cols.join(", "),
placeholders.join(", "),
effective_upsert_clause,
);
engine.execute_raw(&sql, values).await?;
return Ok(());
}
if update_payload.is_empty() {
let (_, values, placeholders, quoted_cols) = build_insert_columns_and_values();
let insert_sql = format!(
"INSERT INTO {} ({}) VALUES ({})",
dialect.quote_ident(target_table),
quoted_cols.join(", "),
placeholders.join(", "),
);
engine.execute_raw(&insert_sql, values).await?;
return Ok(());
}
let (set_text, mut update_params) =
build_writeop_set_clause(&update_payload, dialect, 1);
let next_placeholder = update_params.len() + 1;
update_params.push(pk.clone());
let update_sql = format!(
"UPDATE {} SET {} WHERE {} = {}",
dialect.quote_ident(target_table),
set_text,
dialect.quote_ident(target_pk),
dialect.placeholder(next_placeholder),
);
let affected = engine.execute_raw(&update_sql, update_params).await?;
if affected > 0 {
return Ok(());
}
let (_, values, placeholders, quoted_cols) = build_insert_columns_and_values();
let insert_sql = format!(
"INSERT INTO {} ({}) VALUES ({})",
dialect.quote_ident(target_table),
quoted_cols.join(", "),
placeholders.join(", "),
);
engine.execute_raw(&insert_sql, values).await?;
Ok(())
}
NestedWriteOp::ConnectOrCreate {
relation: _,
target_table,
foreign_key,
where_filter,
create_payload,
} => {
if matches!(where_filter, Filter::None) {
return Err(crate::error::QueryError::not_found(target_table).with_context(
"Nested ConnectOrCreate: empty `where` block would match every row; supply a unique filter",
));
}
let dialect = engine.dialect();
let (filter_sql, filter_params) = where_filter.to_sql(1, dialect);
let mut update_params: Vec<FilterValue> =
Vec::with_capacity(filter_params.len() + 1);
update_params.push(parent_pk.clone());
update_params.extend(filter_params);
let update_sql = format!(
"UPDATE {} SET {} = {} WHERE {}",
dialect.quote_ident(target_table),
dialect.quote_ident(foreign_key),
dialect.placeholder(1),
filter_sql,
);
let affected = engine.execute_raw(&update_sql, update_params).await?;
if affected > 0 {
return Ok(());
}
if create_payload.is_empty() {
return Err(
crate::error::QueryError::not_found(target_table).with_context(
"Nested ConnectOrCreate: no match and create payload empty",
),
);
}
let mut columns: Vec<String> =
create_payload.iter().map(|(c, _)| c.clone()).collect();
let mut values: Vec<FilterValue> =
create_payload.into_iter().map(|(_, v)| v).collect();
columns.push(foreign_key.to_string());
values.push(parent_pk.clone());
let placeholders: Vec<String> =
(1..=values.len()).map(|i| dialect.placeholder(i)).collect();
let quoted_cols: Vec<String> =
columns.iter().map(|c| dialect.quote_ident(c)).collect();
let insert_sql = format!(
"INSERT INTO {} ({}) VALUES ({})",
dialect.quote_ident(target_table),
quoted_cols.join(", "),
placeholders.join(", "),
);
engine.execute_raw(&insert_sql, values).await?;
Ok(())
}
NestedWriteOp::Create {
relation: _,
target_table,
foreign_key,
payload,
} => {
if payload.is_empty() {
return Ok(());
}
let dialect = engine.dialect();
let first = &payload[0];
let mut columns: Vec<String> = first.iter().map(|(c, _)| c.clone()).collect();
columns.push(foreign_key.to_string());
let cols_per_row = columns.len();
let quoted_cols: Vec<String> =
columns.iter().map(|c| dialect.quote_ident(c)).collect();
let mut values: Vec<FilterValue> = Vec::with_capacity(payload.len() * cols_per_row);
let mut row_placeholders: Vec<String> = Vec::with_capacity(payload.len());
let mut next_placeholder = 1usize;
for child in payload {
let mut row_phs: Vec<String> = Vec::with_capacity(cols_per_row);
for (_, v) in child {
values.push(v);
row_phs.push(dialect.placeholder(next_placeholder));
next_placeholder += 1;
}
values.push(parent_pk.clone());
row_phs.push(dialect.placeholder(next_placeholder));
next_placeholder += 1;
row_placeholders.push(format!("({})", row_phs.join(", ")));
}
let sql = format!(
"INSERT INTO {} ({}) VALUES {}",
dialect.quote_ident(target_table),
quoted_cols.join(", "),
row_placeholders.join(", "),
);
engine.execute_raw(&sql, values).await?;
Ok(())
}
NestedWriteOp::Set {
relation: _,
target_table,
foreign_key,
target_pk,
set_pks,
} => {
let dialect = engine.dialect();
if set_pks.is_empty() {
let sql = format!(
"UPDATE {} SET {} = NULL WHERE {} = {}",
dialect.quote_ident(target_table),
dialect.quote_ident(foreign_key),
dialect.quote_ident(foreign_key),
dialect.placeholder(1),
);
engine.execute_raw(&sql, vec![parent_pk.clone()]).await?;
return Ok(());
}
let mut disconnect_params: Vec<FilterValue> = Vec::with_capacity(set_pks.len() + 1);
disconnect_params.push(parent_pk.clone());
let mut not_in_placeholders: Vec<String> = Vec::with_capacity(set_pks.len());
for (i, pk) in set_pks.iter().enumerate() {
disconnect_params.push(pk.clone());
not_in_placeholders.push(dialect.placeholder(i + 2));
}
let disconnect_sql = format!(
"UPDATE {} SET {} = NULL WHERE {} = {} AND {} NOT IN ({})",
dialect.quote_ident(target_table),
dialect.quote_ident(foreign_key),
dialect.quote_ident(foreign_key),
dialect.placeholder(1),
dialect.quote_ident(target_pk),
not_in_placeholders.join(", "),
);
engine
.execute_raw(&disconnect_sql, disconnect_params)
.await?;
let mut connect_params: Vec<FilterValue> = Vec::with_capacity(set_pks.len() + 1);
connect_params.push(parent_pk.clone());
let mut in_placeholders: Vec<String> = Vec::with_capacity(set_pks.len());
for (i, pk) in set_pks.iter().enumerate() {
connect_params.push(pk.clone());
in_placeholders.push(dialect.placeholder(i + 2));
}
let connect_sql = format!(
"UPDATE {} SET {} = {} WHERE {} IN ({})",
dialect.quote_ident(target_table),
dialect.quote_ident(foreign_key),
dialect.placeholder(1),
dialect.quote_ident(target_pk),
in_placeholders.join(", "),
);
engine.execute_raw(&connect_sql, connect_params).await?;
Ok(())
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::{Arc, Mutex};
use crate::error::QueryError;
use crate::traits::BoxFuture;
type StatementLog = Arc<Mutex<Vec<(String, Vec<FilterValue>)>>>;
#[derive(Clone, Copy)]
enum DialectKind {
Postgres,
Mysql,
Mssql,
NotSql,
}
#[derive(Clone)]
struct RecordingEngine {
recorded: StatementLog,
affected: Arc<Mutex<Vec<u64>>>,
dialect_kind: DialectKind,
}
impl RecordingEngine {
fn new() -> Self {
Self {
recorded: Arc::new(Mutex::new(Vec::new())),
affected: Arc::new(Mutex::new(Vec::new())),
dialect_kind: DialectKind::Postgres,
}
}
fn with_affected(seq: Vec<u64>) -> Self {
let mut rev = seq;
rev.reverse();
Self {
recorded: Arc::new(Mutex::new(Vec::new())),
affected: Arc::new(Mutex::new(rev)),
dialect_kind: DialectKind::Postgres,
}
}
fn mysql() -> Self {
Self {
dialect_kind: DialectKind::Mysql,
..Self::new()
}
}
fn mssql() -> Self {
Self {
dialect_kind: DialectKind::Mssql,
..Self::new()
}
}
fn mssql_with_affected(seq: Vec<u64>) -> Self {
Self {
dialect_kind: DialectKind::Mssql,
..Self::with_affected(seq)
}
}
fn notsql() -> Self {
Self {
dialect_kind: DialectKind::NotSql,
..Self::new()
}
}
fn statements(&self) -> Vec<(String, Vec<FilterValue>)> {
self.recorded.lock().unwrap().clone()
}
}
impl crate::traits::QueryEngine for RecordingEngine {
fn dialect(&self) -> &dyn crate::dialect::SqlDialect {
match self.dialect_kind {
DialectKind::Postgres => &crate::dialect::Postgres,
DialectKind::Mysql => &crate::dialect::Mysql,
DialectKind::Mssql => &crate::dialect::Mssql,
DialectKind::NotSql => &crate::dialect::NotSql,
}
}
fn query_many<T: Model + crate::row::FromRow + Send + 'static>(
&self,
_sql: &str,
_params: Vec<FilterValue>,
) -> BoxFuture<'_, QueryResult<Vec<T>>> {
Box::pin(async { Ok(Vec::new()) })
}
fn query_one<T: Model + crate::row::FromRow + Send + 'static>(
&self,
_sql: &str,
_params: Vec<FilterValue>,
) -> BoxFuture<'_, QueryResult<T>> {
Box::pin(async { Err(QueryError::not_found("test")) })
}
fn query_optional<T: Model + crate::row::FromRow + Send + 'static>(
&self,
_sql: &str,
_params: Vec<FilterValue>,
) -> BoxFuture<'_, QueryResult<Option<T>>> {
Box::pin(async { Ok(None) })
}
fn execute_insert<T: Model + crate::row::FromRow + Send + 'static>(
&self,
_sql: &str,
_params: Vec<FilterValue>,
) -> BoxFuture<'_, QueryResult<T>> {
Box::pin(async { Err(QueryError::not_found("test")) })
}
fn execute_update<T: Model + crate::row::FromRow + Send + 'static>(
&self,
_sql: &str,
_params: Vec<FilterValue>,
) -> BoxFuture<'_, QueryResult<Vec<T>>> {
Box::pin(async { Ok(Vec::new()) })
}
fn execute_delete(
&self,
_sql: &str,
_params: Vec<FilterValue>,
) -> BoxFuture<'_, QueryResult<u64>> {
Box::pin(async { Ok(0) })
}
fn execute_raw(
&self,
sql: &str,
params: Vec<FilterValue>,
) -> BoxFuture<'_, QueryResult<u64>> {
let recorded = self.recorded.clone();
let affected = self.affected.clone();
let sql = sql.to_string();
Box::pin(async move {
recorded.lock().unwrap().push((sql, params));
let next = affected.lock().unwrap().pop().unwrap_or(1);
Ok(next)
})
}
fn count(&self, _sql: &str, _params: Vec<FilterValue>) -> BoxFuture<'_, QueryResult<u64>> {
Box::pin(async { Ok(0) })
}
}
struct TestModel;
impl Model for TestModel {
const MODEL_NAME: &'static str = "Post";
const TABLE_NAME: &'static str = "posts";
const PRIMARY_KEY: &'static [&'static str] = &["id"];
const COLUMNS: &'static [&'static str] = &["id", "title", "user_id"];
}
struct TagModel;
impl Model for TagModel {
const MODEL_NAME: &'static str = "Tag";
const TABLE_NAME: &'static str = "tags";
const PRIMARY_KEY: &'static [&'static str] = &["id"];
const COLUMNS: &'static [&'static str] = &["id", "name"];
}
#[test]
fn test_nested_create_data() {
let data: NestedCreateData<TestModel> =
NestedCreateData::from_pairs([("title", FilterValue::String("Test Post".to_string()))]);
assert_eq!(data.data.len(), 1);
assert_eq!(data.data[0].0, "title");
}
#[test]
fn test_nested_write_create() {
let data: NestedCreateData<TestModel> =
NestedCreateData::from_pairs([("title", FilterValue::String("Test Post".to_string()))]);
let write: NestedWrite<TestModel> = NestedWrite::create(data);
match write {
NestedWrite::Create(creates) => assert_eq!(creates.len(), 1),
_ => panic!("Expected Create variant"),
}
}
#[test]
fn test_nested_write_connect() {
let write: NestedWrite<TestModel> = NestedWrite::connect(vec![
Filter::Equals("id".into(), FilterValue::Int(1)),
Filter::Equals("id".into(), FilterValue::Int(2)),
]);
match write {
NestedWrite::Connect(filters) => assert_eq!(filters.len(), 2),
_ => panic!("Expected Connect variant"),
}
}
#[test]
fn test_nested_write_disconnect() {
let write: NestedWrite<TestModel> =
NestedWrite::disconnect_one(Filter::Equals("id".into(), FilterValue::Int(1)));
match write {
NestedWrite::Disconnect(filters) => assert_eq!(filters.len(), 1),
_ => panic!("Expected Disconnect variant"),
}
}
#[test]
fn test_nested_write_set() {
let write: NestedWrite<TestModel> =
NestedWrite::set(vec![Filter::Equals("id".into(), FilterValue::Int(1))]);
match write {
NestedWrite::Set(filters) => assert_eq!(filters.len(), 1),
_ => panic!("Expected Set variant"),
}
}
#[test]
fn test_builder_one_to_many_connect() {
let builder =
NestedWriteBuilder::one_to_many("users", vec!["id".to_string()], "posts", "user_id");
let parent_id = FilterValue::Int(1);
let filters = vec![Filter::Equals("id".into(), FilterValue::Int(10))];
let statements = builder.build_connect_sql::<TestModel>(&parent_id, &filters);
assert_eq!(statements.len(), 1);
let (sql, params) = &statements[0];
assert_eq!(sql, "UPDATE posts SET user_id = $2 WHERE \"id\" = $1");
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::Int(10));
assert_eq!(params[1], FilterValue::Int(1));
}
#[test]
fn test_builder_one_to_many_disconnect() {
let builder =
NestedWriteBuilder::one_to_many("users", vec!["id".to_string()], "posts", "user_id");
let parent_id = FilterValue::Int(1);
let filters = vec![Filter::Equals("id".into(), FilterValue::Int(10))];
let statements = builder.build_disconnect_sql(&parent_id, &filters);
assert_eq!(statements.len(), 1);
let (sql, params) = &statements[0];
assert_eq!(
sql,
"UPDATE posts SET user_id = NULL WHERE \"id\" = $1 AND user_id = $2"
);
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::Int(10));
assert_eq!(params[1], FilterValue::Int(1));
}
#[test]
fn test_builder_many_to_many_connect() {
let builder = NestedWriteBuilder::many_to_many(
"posts",
vec!["id".to_string()],
"tags",
JoinTableInfo {
table_name: "post_tags".to_string(),
parent_column: "post_id".to_string(),
related_column: "tag_id".to_string(),
},
);
let parent_id = FilterValue::Int(1);
let filters = vec![Filter::Equals("id".into(), FilterValue::Int(10))];
let statements = builder.build_connect_sql::<TagModel>(&parent_id, &filters);
assert_eq!(statements.len(), 1);
let (sql, params) = &statements[0];
assert_eq!(
sql,
"INSERT INTO post_tags (post_id, tag_id) SELECT $2, id FROM tags WHERE \"id\" = $1 ON CONFLICT DO NOTHING"
);
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::Int(10));
assert_eq!(params[1], FilterValue::Int(1));
}
#[test]
fn test_builder_many_to_many_disconnect() {
let builder = NestedWriteBuilder::many_to_many(
"posts",
vec!["id".to_string()],
"tags",
JoinTableInfo {
table_name: "post_tags".to_string(),
parent_column: "post_id".to_string(),
related_column: "tag_id".to_string(),
},
);
let parent_id = FilterValue::Int(1);
let filters = vec![Filter::Equals("id".into(), FilterValue::Int(10))];
let statements = builder.build_disconnect_sql(&parent_id, &filters);
assert_eq!(statements.len(), 1);
let (sql, params) = &statements[0];
assert_eq!(
sql,
"DELETE FROM post_tags WHERE post_id = $1 AND tag_id IN (SELECT id FROM tags WHERE \"id\" = $2)"
);
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::Int(1));
assert_eq!(params[1], FilterValue::Int(10));
}
#[test]
fn test_builder_one_to_many_delete() {
let builder =
NestedWriteBuilder::one_to_many("users", vec!["id".to_string()], "posts", "user_id");
let parent_id = FilterValue::Int(1);
let filters = vec![Filter::Equals("id".into(), FilterValue::Int(10))];
let statements = builder.build_delete_sql(&parent_id, &filters);
assert_eq!(statements.len(), 1);
let (sql, params) = &statements[0];
assert_eq!(sql, "DELETE FROM posts WHERE \"id\" = $1 AND user_id = $2");
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::Int(10));
assert_eq!(params[1], FilterValue::Int(1));
}
#[test]
fn test_builder_create() {
let builder =
NestedWriteBuilder::one_to_many("users", vec!["id".to_string()], "posts", "user_id");
let parent_id = FilterValue::Int(1);
let creates = vec![NestedCreateData::<TestModel>::from_pairs([(
"title",
FilterValue::String("New Post".to_string()),
)])];
let statements = builder.build_create_sql::<TestModel>(&parent_id, &creates);
assert_eq!(statements.len(), 1);
let (sql, params) = &statements[0];
assert!(sql.contains("INSERT INTO"));
assert!(sql.contains("posts"));
assert!(sql.contains("RETURNING"));
assert_eq!(params.len(), 2); }
#[test]
fn test_builder_set() {
let builder =
NestedWriteBuilder::one_to_many("users", vec!["id".to_string()], "posts", "user_id");
let parent_id = FilterValue::Int(1);
let filters = vec![Filter::Equals("id".into(), FilterValue::Int(10))];
let statements = builder.build_set_sql::<TestModel>(&parent_id, &filters);
assert!(statements.len() >= 2);
let (first_sql, _) = &statements[0];
assert!(first_sql.contains("UPDATE"));
assert!(first_sql.contains("NULL"));
}
#[test]
fn test_nested_write_operations() {
let mut ops = NestedWriteOperations::new();
assert!(ops.is_empty());
assert_eq!(ops.len(), 0);
ops.add_pre("SELECT 1".to_string(), vec![]);
ops.add_post("SELECT 2".to_string(), vec![]);
assert!(!ops.is_empty());
assert_eq!(ops.len(), 2);
}
#[test]
fn test_nested_create_or_connect() {
let create_data: NestedCreateData<TestModel> =
NestedCreateData::from_pairs([("title", FilterValue::String("New Post".to_string()))]);
let create_or_connect = NestedCreateOrConnectData::new(
Filter::Equals("title".into(), FilterValue::String("Existing".to_string())),
create_data,
);
assert!(matches!(create_or_connect.filter, Filter::Equals(..)));
assert_eq!(create_or_connect.create.data.len(), 1);
}
#[test]
fn test_nested_update_data() {
let update: NestedUpdateData<TestModel> = NestedUpdateData::from_pairs(
Filter::Equals("id".into(), FilterValue::Int(1)),
[("title", FilterValue::String("Updated".to_string()))],
);
assert!(matches!(update.filter, Filter::Equals(..)));
assert_eq!(update.data.len(), 1);
assert_eq!(update.data[0].0, "title");
}
#[tokio::test]
async fn nested_op_connect_emits_update_set_where() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Connect {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(42),
};
let parent_pk = FilterValue::Int(7);
op.execute(&engine, &parent_pk).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1, "expected one UPDATE statement");
let (sql, params) = &stmts[0];
assert!(sql.contains("UPDATE"), "got: {sql}");
assert!(sql.contains("posts"), "got: {sql}");
assert!(sql.contains("author_id"), "got: {sql}");
assert!(sql.contains("SET"), "got: {sql}");
assert!(sql.contains("WHERE"), "got: {sql}");
assert!(sql.contains("$1"), "got: {sql}");
assert!(sql.contains("$2"), "got: {sql}");
assert_eq!(params, &vec![FilterValue::Int(7), FilterValue::Int(42)]);
}
#[tokio::test]
async fn nested_op_delete_many_with_filter_emits_fk_and_filter_clause() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::DeleteMany {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
filter: Filter::Equals("published".into(), FilterValue::Bool(false)),
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1);
let (sql, params) = &stmts[0];
assert_eq!(
sql,
"DELETE FROM \"posts\" WHERE \"author_id\" = $1 AND (\"published\" = $2)"
);
assert!(!sql.contains("$3"), "off-by-one placeholder: {sql}");
assert_eq!(params.len(), 2);
assert!(matches!(params[0], FilterValue::Int(7)));
assert!(matches!(params[1], FilterValue::Bool(false)));
}
#[tokio::test]
async fn nested_op_delete_many_with_filter_mysql_uses_positional_placeholders() {
let engine = RecordingEngine::mysql();
let op = NestedWriteOp::DeleteMany {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
filter: Filter::Equals("published".into(), FilterValue::Bool(false)),
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1);
let (sql, params) = &stmts[0];
assert_eq!(
sql,
"DELETE FROM `posts` WHERE `author_id` = ? AND (`published` = ?)"
);
assert!(!sql.contains('$'), "MySQL must not emit $N: {sql}");
assert_eq!(params.len(), 2);
assert!(matches!(params[0], FilterValue::Int(7)));
assert!(matches!(params[1], FilterValue::Bool(false)));
}
#[tokio::test]
async fn nested_op_delete_many_with_empty_filter_omits_and_clause() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::DeleteMany {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
filter: Filter::None,
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
let (sql, params) = &stmts[0];
assert!(sql.contains("DELETE FROM"), "got: {sql}");
assert!(
!sql.contains("AND"),
"should omit AND when filter empty: {sql}"
);
assert_eq!(params.len(), 1);
}
#[tokio::test]
async fn nested_op_delete_emits_delete_where_pk() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Delete {
relation: "posts",
target_table: "posts",
target_pk: "id",
pk: FilterValue::Int(42),
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1);
let (sql, params) = &stmts[0];
assert!(sql.contains("DELETE FROM"), "got: {sql}");
assert!(sql.contains("posts"), "got: {sql}");
assert!(sql.contains("WHERE"), "got: {sql}");
assert_eq!(params, &vec![FilterValue::Int(42)]);
}
#[tokio::test]
async fn nested_op_disconnect_emits_update_set_null() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Disconnect {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(42),
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1);
let (sql, params) = &stmts[0];
assert!(sql.contains("UPDATE"), "got: {sql}");
assert!(sql.contains("posts"), "got: {sql}");
assert!(sql.contains("author_id"), "got: {sql}");
assert!(sql.contains("NULL"), "got: {sql}");
assert!(sql.contains("WHERE"), "got: {sql}");
assert_eq!(params, &vec![FilterValue::Int(42)]);
}
#[tokio::test]
async fn nested_op_update_plain_set() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::new();
let op = NestedWriteOp::Update {
relation: "posts",
target_table: "posts",
target_pk: "id",
pk: FilterValue::Int(42),
payload: vec![(
"title".to_string(),
WriteOp::Set(FilterValue::String("renamed".to_string())),
)],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1);
let (sql, params) = &stmts[0];
assert!(sql.contains("UPDATE"), "got: {sql}");
assert!(sql.contains("posts"), "got: {sql}");
assert!(sql.contains("title"), "got: {sql}");
assert!(sql.contains("SET"), "got: {sql}");
assert!(sql.contains("WHERE"), "got: {sql}");
assert!(sql.contains("$1"), "got: {sql}");
assert!(sql.contains("$2"), "got: {sql}");
assert_eq!(params.len(), 2);
assert!(matches!(params[0], FilterValue::String(_)));
assert_eq!(params[1], FilterValue::Int(42));
}
#[tokio::test]
async fn nested_op_update_increment() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::new();
let op = NestedWriteOp::Update {
relation: "posts",
target_table: "posts",
target_pk: "id",
pk: FilterValue::Int(42),
payload: vec![("views".to_string(), WriteOp::Increment(FilterValue::Int(1)))],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
let (sql, _) = &stmts[0];
assert!(sql.contains("+"), "got: {sql}");
assert!(sql.contains("views"), "got: {sql}");
}
#[tokio::test]
async fn nested_op_update_mixed_set_and_increment() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::new();
let op = NestedWriteOp::Update {
relation: "posts",
target_table: "posts",
target_pk: "id",
pk: FilterValue::Int(42),
payload: vec![
(
"title".to_string(),
WriteOp::Set(FilterValue::String("renamed".to_string())),
),
("views".to_string(), WriteOp::Increment(FilterValue::Int(1))),
],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
let (sql, params) = &stmts[0];
assert!(sql.contains("title"), "got: {sql}");
assert!(sql.contains("views"), "got: {sql}");
assert!(sql.contains("+"), "got: {sql}");
assert!(sql.contains("$3"), "got: {sql}");
assert_eq!(params.len(), 3);
}
#[tokio::test]
async fn nested_op_update_empty_payload_is_noop() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Update {
relation: "posts",
target_table: "posts",
target_pk: "id",
pk: FilterValue::Int(42),
payload: vec![],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
assert!(
engine.statements().is_empty(),
"empty payload should emit no SQL"
);
}
#[tokio::test]
async fn nested_op_update_many_with_filter() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::new();
let op = NestedWriteOp::UpdateMany {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
filter: Filter::Equals("published".into(), FilterValue::Bool(false)),
payload: vec![("views".to_string(), WriteOp::Set(FilterValue::Int(0)))],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1);
let (sql, params) = &stmts[0];
assert_eq!(
sql,
"UPDATE \"posts\" SET \"views\" = $1 WHERE \"author_id\" = $2 AND (\"published\" = $3)"
);
assert!(!sql.contains("$4"), "skipped placeholder slot: {sql}");
assert_eq!(params.len(), 3);
assert_eq!(params[0], FilterValue::Int(0));
assert_eq!(params[1], FilterValue::Int(7));
assert_eq!(params[2], FilterValue::Bool(false));
}
#[tokio::test]
async fn nested_op_update_many_with_filter_mysql_uses_positional_placeholders() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::mysql();
let op = NestedWriteOp::UpdateMany {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
filter: Filter::Equals("published".into(), FilterValue::Bool(false)),
payload: vec![("views".to_string(), WriteOp::Set(FilterValue::Int(0)))],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1);
let (sql, params) = &stmts[0];
assert_eq!(
sql,
"UPDATE `posts` SET `views` = ? WHERE `author_id` = ? AND (`published` = ?)"
);
assert!(!sql.contains('$'), "MySQL must not emit $N: {sql}");
assert_eq!(params.len(), 3);
assert_eq!(params[0], FilterValue::Int(0));
assert_eq!(params[1], FilterValue::Int(7));
assert_eq!(params[2], FilterValue::Bool(false));
}
#[tokio::test]
async fn nested_op_update_many_with_empty_filter() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::new();
let op = NestedWriteOp::UpdateMany {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
filter: Filter::None,
payload: vec![("views".to_string(), WriteOp::Set(FilterValue::Int(0)))],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
let (sql, params) = &stmts[0];
assert!(sql.contains("UPDATE"), "got: {sql}");
assert!(
!sql.contains("AND"),
"should omit AND when filter empty: {sql}"
);
assert_eq!(params.len(), 2);
}
#[test]
fn test_nested_upsert_data() {
let create: NestedCreateData<TestModel> =
NestedCreateData::from_pairs([("title", FilterValue::String("New".to_string()))]);
let upsert: NestedUpsertData<TestModel> = NestedUpsertData::new(
Filter::Equals("id".into(), FilterValue::Int(1)),
create,
vec![(
"title".to_string(),
FilterValue::String("Updated".to_string()),
)],
);
assert!(matches!(upsert.filter, Filter::Equals(..)));
assert_eq!(upsert.create.data.len(), 1);
assert_eq!(upsert.update.len(), 1);
}
#[tokio::test]
async fn nested_op_upsert_single_statement_on_postgres() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::new();
let op = NestedWriteOp::Upsert {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(99),
create_payload: vec![("title".to_string(), FilterValue::String("new".to_string()))],
update_payload: vec![("views".to_string(), WriteOp::Increment(FilterValue::Int(1)))],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(
stmts.len(),
1,
"expected a single-statement upsert; got {stmts:#?}"
);
let (sql, params) = &stmts[0];
assert!(sql.contains("INSERT INTO"), "got: {sql}");
assert!(sql.contains("posts"), "got: {sql}");
assert!(sql.contains("ON CONFLICT (\"id\")"), "got: {sql}");
assert!(sql.contains("DO UPDATE SET"), "got: {sql}");
assert!(
sql.contains("VALUES ($1, $2)"),
"INSERT VALUES placeholders: {sql}"
);
assert!(sql.contains("$3"), "got: {sql}");
assert_eq!(params.len(), 3);
assert_eq!(params[0], FilterValue::String("new".to_string()));
assert_eq!(params[1], FilterValue::Int(7));
assert_eq!(params[2], FilterValue::Int(1));
}
#[tokio::test]
async fn nested_op_upsert_two_statement_fallback_on_mssql_update_path() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::mssql();
let op = NestedWriteOp::Upsert {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(99),
create_payload: vec![("title".to_string(), FilterValue::String("new".to_string()))],
update_payload: vec![("views".to_string(), WriteOp::Increment(FilterValue::Int(1)))],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(
stmts.len(),
1,
"expected only the UPDATE — INSERT should not have run"
);
let (sql, update_params) = &stmts[0];
assert!(sql.starts_with("UPDATE"), "got: {sql}");
assert!(!sql.contains("ON CONFLICT"), "got: {sql}");
assert!(!sql.contains("ON DUPLICATE"), "got: {sql}");
assert!(sql.contains("[posts]"), "got: {sql}");
assert!(sql.contains("@P1"), "SET clause should use @P1: {sql}");
assert!(sql.contains("@P2"), "WHERE clause should use @P2: {sql}");
assert_eq!(update_params.len(), 2);
assert_eq!(update_params[0], FilterValue::Int(1)); assert_eq!(update_params[1], FilterValue::Int(99)); }
#[tokio::test]
async fn nested_op_upsert_two_statement_fallback_on_mssql_insert_path() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::mssql_with_affected(vec![0, 1]);
let op = NestedWriteOp::Upsert {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(99),
create_payload: vec![("title".to_string(), FilterValue::String("new".to_string()))],
update_payload: vec![("views".to_string(), WriteOp::Increment(FilterValue::Int(1)))],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 2, "expected UPDATE then INSERT");
let (update_sql, _) = &stmts[0];
assert!(update_sql.starts_with("UPDATE"), "got: {update_sql}");
let (insert_sql, insert_params) = &stmts[1];
assert!(insert_sql.starts_with("INSERT INTO"), "got: {insert_sql}");
assert!(!insert_sql.contains("ON CONFLICT"), "got: {insert_sql}");
assert!(insert_sql.contains("[posts]"), "got: {insert_sql}");
assert!(insert_sql.contains("[author_id]"), "got: {insert_sql}");
assert_eq!(insert_params.len(), 2);
assert_eq!(insert_params[0], FilterValue::String("new".to_string()));
assert_eq!(insert_params[1], FilterValue::Int(7));
}
#[tokio::test]
async fn nested_op_connect_or_create_connect_path_when_affected() {
let engine = RecordingEngine::with_affected(vec![1]);
let op = NestedWriteOp::ConnectOrCreate {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
where_filter: Filter::Equals("id".into(), FilterValue::Int(42)),
create_payload: vec![(
"title".to_string(),
FilterValue::String("fallback".to_string()),
)],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(
stmts.len(),
1,
"expected only the UPDATE — INSERT should not have run"
);
let (sql, params) = &stmts[0];
assert_eq!(
sql,
"UPDATE \"posts\" SET \"author_id\" = $1 WHERE \"id\" = $2"
);
assert!(!sql.contains("$3"), "off-by-one placeholder: {sql}");
assert!(!sql.contains("INSERT"), "got: {sql}");
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::Int(7));
assert_eq!(params[1], FilterValue::Int(42));
}
#[tokio::test]
async fn nested_op_connect_or_create_connect_path_mysql_uses_positional_placeholders() {
let engine = RecordingEngine::mysql();
let op = NestedWriteOp::ConnectOrCreate {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
where_filter: Filter::Equals("id".into(), FilterValue::Int(42)),
create_payload: vec![(
"title".to_string(),
FilterValue::String("fallback".to_string()),
)],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1);
let (sql, params) = &stmts[0];
assert_eq!(sql, "UPDATE `posts` SET `author_id` = ? WHERE `id` = ?");
assert!(!sql.contains('$'), "MySQL must not emit $N: {sql}");
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::Int(7));
assert_eq!(params[1], FilterValue::Int(42));
}
#[tokio::test]
async fn nested_op_connect_or_create_create_path_when_zero_affected() {
let engine = RecordingEngine::with_affected(vec![0, 1]);
let op = NestedWriteOp::ConnectOrCreate {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
where_filter: Filter::Equals("id".into(), FilterValue::Int(42)),
create_payload: vec![(
"title".to_string(),
FilterValue::String("fallback".to_string()),
)],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 2);
let (update_sql, _) = &stmts[0];
assert!(update_sql.contains("UPDATE"), "got: {update_sql}");
let (insert_sql, insert_params) = &stmts[1];
assert!(insert_sql.contains("INSERT INTO"), "got: {insert_sql}");
assert!(insert_sql.contains("posts"), "got: {insert_sql}");
assert!(insert_sql.contains("title"), "got: {insert_sql}");
assert!(insert_sql.contains("author_id"), "got: {insert_sql}");
assert_eq!(insert_params.len(), 2);
assert_eq!(
insert_params[0],
FilterValue::String("fallback".to_string())
);
assert_eq!(insert_params[1], FilterValue::Int(7));
}
#[tokio::test]
async fn nested_op_set_with_empty_list_clears_all_children() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Set {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
set_pks: vec![],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1, "expected only the disconnect-all UPDATE");
let (sql, params) = &stmts[0];
assert!(sql.contains("UPDATE"), "got: {sql}");
assert!(sql.contains("posts"), "got: {sql}");
assert!(sql.contains("author_id"), "got: {sql}");
assert!(sql.contains("= NULL"), "got: {sql}");
assert!(!sql.contains("NOT IN"), "got: {sql}");
assert!(!sql.contains(" IN ("), "got: {sql}");
assert_eq!(params.len(), 1);
assert_eq!(params[0], FilterValue::Int(7));
}
#[tokio::test]
async fn nested_op_set_with_non_empty_list_emits_disconnect_then_connect() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Set {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
set_pks: vec![
FilterValue::Int(1),
FilterValue::Int(2),
FilterValue::Int(3),
],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 2);
let (disconnect_sql, disconnect_params) = &stmts[0];
assert!(disconnect_sql.contains("UPDATE"), "got: {disconnect_sql}");
assert!(disconnect_sql.contains("posts"), "got: {disconnect_sql}");
assert!(
disconnect_sql.contains("author_id"),
"got: {disconnect_sql}"
);
assert!(disconnect_sql.contains("= NULL"), "got: {disconnect_sql}");
assert!(disconnect_sql.contains("NOT IN"), "got: {disconnect_sql}");
assert_eq!(disconnect_params.len(), 4);
assert_eq!(disconnect_params[0], FilterValue::Int(7));
assert_eq!(disconnect_params[1], FilterValue::Int(1));
assert_eq!(disconnect_params[2], FilterValue::Int(2));
assert_eq!(disconnect_params[3], FilterValue::Int(3));
let (connect_sql, connect_params) = &stmts[1];
assert!(connect_sql.contains("UPDATE"), "got: {connect_sql}");
assert!(connect_sql.contains("posts"), "got: {connect_sql}");
assert!(connect_sql.contains("author_id"), "got: {connect_sql}");
assert!(connect_sql.contains(" IN ("), "got: {connect_sql}");
assert!(!connect_sql.contains("NOT IN"), "got: {connect_sql}");
assert_eq!(connect_params.len(), 4);
assert_eq!(connect_params[0], FilterValue::Int(7));
assert_eq!(connect_params[1], FilterValue::Int(1));
assert_eq!(connect_params[2], FilterValue::Int(2));
assert_eq!(connect_params[3], FilterValue::Int(3));
}
#[tokio::test]
async fn nested_op_set_with_single_element_uses_single_placeholder_in_lists() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Set {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
set_pks: vec![FilterValue::Int(5)],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 2);
let (disconnect_sql, _) = &stmts[0];
assert!(
disconnect_sql.contains("NOT IN ($2)"),
"got: {disconnect_sql}"
);
let (connect_sql, _) = &stmts[1];
assert!(connect_sql.contains(" IN ($2)"), "got: {connect_sql}");
assert!(!connect_sql.contains("NOT IN"), "got: {connect_sql}");
}
#[tokio::test]
async fn nested_op_set_disconnect_clears_only_current_parents_children() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Set {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
set_pks: vec![FilterValue::Int(1)],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
let (disconnect_sql, _) = &stmts[0];
assert!(
disconnect_sql.contains("author_id\" = $1"),
"expected `author_id = $1` clause; got: {disconnect_sql}"
);
}
#[tokio::test]
async fn nested_op_connect_or_create_rejects_empty_where() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::ConnectOrCreate {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
where_filter: Filter::None,
create_payload: vec![(
"title".to_string(),
FilterValue::String("fallback".to_string()),
)],
};
let err = op
.execute(&engine, &FilterValue::Int(7))
.await
.expect_err("empty where must be rejected");
let op_ctx = err.context.operation.clone().unwrap_or_default();
assert!(op_ctx.contains("ConnectOrCreate"), "got: {op_ctx}");
assert!(engine.statements().is_empty());
}
#[tokio::test]
async fn nested_op_upsert_single_statement_empty_update_payload_emits_do_nothing() {
let engine = RecordingEngine::new();
let op = NestedWriteOp::Upsert {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(99),
create_payload: vec![("title".to_string(), FilterValue::String("new".into()))],
update_payload: vec![],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(
stmts.len(),
1,
"expected one INSERT ... DO NOTHING; got {stmts:#?}"
);
let (sql, params) = &stmts[0];
assert!(sql.starts_with("INSERT INTO"), "got: {sql}");
assert!(sql.contains("posts"), "got: {sql}");
assert!(sql.contains("ON CONFLICT (\"id\")"), "got: {sql}");
assert!(sql.contains("DO NOTHING"), "got: {sql}");
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::String("new".into()));
assert_eq!(params[1], FilterValue::Int(7));
}
#[tokio::test]
async fn nested_op_upsert_two_statement_fallback_empty_update_payload_emits_bare_insert() {
let engine = RecordingEngine::mssql();
let op = NestedWriteOp::Upsert {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(99),
create_payload: vec![("title".to_string(), FilterValue::String("new".into()))],
update_payload: vec![],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(stmts.len(), 1, "expected one bare INSERT; got {stmts:#?}");
let (sql, params) = &stmts[0];
assert!(sql.starts_with("INSERT INTO"), "got: {sql}");
assert!(sql.contains("[posts]"), "got: {sql}");
assert!(!sql.contains("ON CONFLICT"), "got: {sql}");
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::String("new".into()));
assert_eq!(params[1], FilterValue::Int(7));
}
#[tokio::test]
async fn nested_op_upsert_single_statement_all_unset_update_payload() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::new();
let op = NestedWriteOp::Upsert {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(99),
create_payload: vec![("title".to_string(), FilterValue::String("new".into()))],
update_payload: vec![("deleted_at".to_string(), WriteOp::Unset)],
};
op.execute(&engine, &FilterValue::Int(7)).await.unwrap();
let stmts = engine.statements();
assert_eq!(
stmts.len(),
1,
"expected one INSERT...ON CONFLICT statement; got {stmts:#?}"
);
let (sql, params) = &stmts[0];
assert!(sql.starts_with("INSERT INTO"), "got: {sql}");
assert!(
sql.contains("ON CONFLICT (\"id\") DO UPDATE SET"),
"got: {sql}"
);
assert!(sql.contains("\"deleted_at\" = NULL"), "got: {sql}");
assert_eq!(params.len(), 2);
assert_eq!(params[0], FilterValue::String("new".into()));
assert_eq!(params[1], FilterValue::Int(7));
}
#[tokio::test]
async fn nested_op_upsert_empty_create_payload_returns_invalid_input() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::new();
let op = NestedWriteOp::Upsert {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(99),
create_payload: vec![],
update_payload: vec![("views".to_string(), WriteOp::Increment(FilterValue::Int(1)))],
};
let err = op.execute(&engine, &FilterValue::Int(7)).await.unwrap_err();
assert_eq!(
err.code,
crate::error::ErrorCode::InvalidParameter,
"expected InvalidParameter (invalid_input), got: {err:?}"
);
let msg = format!("{err}");
assert!(
msg.contains("create_payload") || msg.contains("create column"),
"msg: {msg}"
);
}
#[tokio::test]
async fn nested_op_upsert_on_notsql_returns_unsupported_not_panic() {
use crate::inputs::WriteOp;
let engine = RecordingEngine::notsql();
let op = NestedWriteOp::Upsert {
relation: "posts",
target_table: "posts",
foreign_key: "author_id",
target_pk: "id",
pk: FilterValue::Int(99),
create_payload: vec![("title".to_string(), FilterValue::String("x".into()))],
update_payload: vec![("views".to_string(), WriteOp::Increment(FilterValue::Int(1)))],
};
let err = op.execute(&engine, &FilterValue::Int(7)).await.unwrap_err();
let msg = format!("{err}");
assert!(
msg.contains("Upsert is not supported")
|| msg.contains("unsupported")
|| msg.contains("Unsupported"),
"msg: {msg}"
);
assert_eq!(
engine.statements().len(),
0,
"no SQL should have been emitted"
);
}
}