Skip to main content

dbkit_core/
mutation.rs

1use std::marker::PhantomData;
2
3use crate::compile::{CompiledSql, SqlBuilder, ToSql};
4use crate::expr::{ColumnValue, Expr, Value};
5use crate::func;
6use crate::query::Select;
7use crate::schema::{Column, ColumnRef, Table};
8
9#[derive(Debug, Clone)]
10pub struct Insert<Out> {
11    table: Table,
12    columns: Vec<ColumnRef>,
13    values: Vec<Value>,
14    row_count: usize,
15    mode: InsertMode,
16    conflict: Option<InsertConflict>,
17    returning: Option<Vec<ColumnRef>>,
18    returning_all: bool,
19    _marker: PhantomData<Out>,
20}
21
22#[derive(Debug, Clone, Copy, PartialEq, Eq)]
23enum InsertMode {
24    Unset,
25    Values,
26    Rows,
27}
28
29#[derive(Debug, Clone, PartialEq, Eq)]
30enum InsertConflict {
31    DoNothing { target: Vec<ColumnRef> },
32    DoUpdate { target: Vec<ColumnRef>, updates: Vec<ColumnRef> },
33}
34
35mod private {
36    pub trait Sealed {}
37}
38
39pub trait ConflictColumns<M>: private::Sealed {
40    fn into_columns(self) -> Vec<ColumnRef>;
41}
42
43impl<M, T> ConflictColumns<M> for Column<M, T> {
44    fn into_columns(self) -> Vec<ColumnRef> {
45        vec![self.as_ref()]
46    }
47}
48impl<M, T> private::Sealed for Column<M, T> {}
49
50macro_rules! impl_conflict_columns_tuple {
51    ($(($($ty:ident:$col:ident),+)),+ $(,)?) => {
52        $(
53            impl<M, $($ty),+> ConflictColumns<M> for ($(Column<M, $ty>,)+) {
54                fn into_columns(self) -> Vec<ColumnRef> {
55                    let ($($col,)+) = self;
56                    vec![$($col.as_ref()),+]
57                }
58            }
59
60            impl<M, $($ty),+> private::Sealed for ($(Column<M, $ty>,)+) {}
61        )+
62    };
63}
64
65impl_conflict_columns_tuple!(
66    (T1:c1, T2:c2),
67    (T1:c1, T2:c2, T3:c3),
68    (T1:c1, T2:c2, T3:c3, T4:c4),
69    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5),
70    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6),
71    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7),
72    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8),
73    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9),
74    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10),
75    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11),
76    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12),
77    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13),
78    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14),
79    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15),
80    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16),
81    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17),
82    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18),
83    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19),
84    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20),
85    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21),
86    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22),
87    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23),
88    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24),
89    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24, T25:c25),
90    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24, T25:c25, T26:c26),
91    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24, T25:c25, T26:c26, T27:c27),
92    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24, T25:c25, T26:c26, T27:c27, T28:c28),
93    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24, T25:c25, T26:c26, T27:c27, T28:c28, T29:c29),
94    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24, T25:c25, T26:c26, T27:c27, T28:c28, T29:c29, T30:c30),
95    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24, T25:c25, T26:c26, T27:c27, T28:c28, T29:c29, T30:c30, T31:c31),
96    (T1:c1, T2:c2, T3:c3, T4:c4, T5:c5, T6:c6, T7:c7, T8:c8, T9:c9, T10:c10, T11:c11, T12:c12, T13:c13, T14:c14, T15:c15, T16:c16, T17:c17, T18:c18, T19:c19, T20:c20, T21:c21, T22:c22, T23:c23, T24:c24, T25:c25, T26:c26, T27:c27, T28:c28, T29:c29, T30:c30, T31:c31, T32:c32)
97);
98
99impl<Out> Insert<Out> {
100    pub fn new(table: Table) -> Self {
101        Self {
102            table,
103            columns: Vec::new(),
104            values: Vec::new(),
105            row_count: 0,
106            mode: InsertMode::Unset,
107            conflict: None,
108            returning: None,
109            returning_all: false,
110            _marker: PhantomData,
111        }
112    }
113
114    pub fn value<M, T, V>(mut self, column: Column<M, T>, value: V) -> Self
115    where
116        V: ColumnValue<T>,
117    {
118        if self.mode == InsertMode::Rows {
119            panic!("dbkit: cannot use value() after row()");
120        }
121        self.mode = InsertMode::Values;
122        if self.row_count == 0 {
123            self.row_count = 1;
124        }
125        let value = match value.into_value() {
126            Some(value) => value,
127            None => Value::Null,
128        };
129        self.columns.push(column.as_ref());
130        self.values.push(value);
131        self
132    }
133
134    pub fn row<F>(mut self, build: F) -> Self
135    where
136        F: FnOnce(InsertRow) -> InsertRow,
137    {
138        if self.mode == InsertMode::Values {
139            self.mode = InsertMode::Rows;
140        }
141        if self.mode == InsertMode::Unset {
142            self.mode = InsertMode::Rows;
143        }
144
145        let expected = if self.columns.is_empty() {
146            None
147        } else {
148            Some(self.columns.clone())
149        };
150        let row = build(InsertRow::new(expected));
151        if self.columns.is_empty() {
152            self.columns = row.columns.clone();
153        } else if row.columns != self.columns {
154            panic!("dbkit: insert row columns must match");
155        }
156        if row.values.len() != self.columns.len() {
157            panic!("dbkit: insert row value count mismatch");
158        }
159
160        if self.row_count == 0 && !self.values.is_empty() {
161            self.row_count = 1;
162        }
163        self.values.extend(row.values);
164        self.row_count += 1;
165        self
166    }
167
168    pub fn returning(mut self, columns: Vec<ColumnRef>) -> Self {
169        self.returning = Some(columns);
170        self.returning_all = false;
171        self
172    }
173
174    pub fn returning_all(mut self) -> Self {
175        self.returning = None;
176        self.returning_all = true;
177        self
178    }
179
180    pub fn on_conflict_do_nothing<M, C>(mut self, target: C) -> Self
181    where
182        C: ConflictColumns<M>,
183    {
184        self.conflict = Some(InsertConflict::DoNothing {
185            target: target.into_columns(),
186        });
187        self
188    }
189
190    pub fn on_conflict_do_update<M, C, U>(mut self, target: C, updates: U) -> Self
191    where
192        C: ConflictColumns<M>,
193        U: ConflictColumns<M>,
194    {
195        self.conflict = Some(InsertConflict::DoUpdate {
196            target: target.into_columns(),
197            updates: updates.into_columns(),
198        });
199        self
200    }
201
202    pub fn compile(&self) -> CompiledSql {
203        let mut builder = SqlBuilder::new();
204        builder.push_sql("INSERT INTO ");
205        builder.push_sql(&self.table.qualified_name());
206        builder.push_sql(" (");
207        for (idx, col) in self.columns.iter().enumerate() {
208            if idx > 0 {
209                builder.push_sql(", ");
210            }
211            builder.push_sql(col.name);
212        }
213        builder.push_sql(") VALUES (");
214        let row_len = self.columns.len();
215        let row_count = if self.row_count == 0 && !self.values.is_empty() {
216            1
217        } else {
218            self.row_count
219        };
220        if row_count == 0 {
221            builder.push_sql(")");
222        } else {
223            for row_idx in 0..row_count {
224                if row_idx > 0 {
225                    builder.push_sql(", (");
226                }
227                for col_idx in 0..row_len {
228                    if col_idx > 0 {
229                        builder.push_sql(", ");
230                    }
231                    let value = self.values[row_idx * row_len + col_idx].clone();
232                    builder.push_value(value);
233                }
234                builder.push_sql(")");
235            }
236        }
237        if let Some(conflict) = &self.conflict {
238            let target = match conflict {
239                InsertConflict::DoNothing { target } => target,
240                InsertConflict::DoUpdate { target, .. } => target,
241            };
242            builder.push_sql(" ON CONFLICT (");
243            for (idx, col) in target.iter().enumerate() {
244                if idx > 0 {
245                    builder.push_sql(", ");
246                }
247                builder.push_sql(col.name);
248            }
249            builder.push_sql(")");
250
251            match conflict {
252                InsertConflict::DoNothing { .. } => {
253                    builder.push_sql(" DO NOTHING");
254                }
255                InsertConflict::DoUpdate { updates, .. } => {
256                    builder.push_sql(" DO UPDATE SET ");
257                    for (idx, col) in updates.iter().enumerate() {
258                        if idx > 0 {
259                            builder.push_sql(", ");
260                        }
261                        builder.push_sql(col.name);
262                        builder.push_sql(" = EXCLUDED.");
263                        builder.push_sql(col.name);
264                    }
265                }
266            }
267        }
268        if self.returning_all {
269            builder.push_sql(" RETURNING ");
270            builder.push_sql(self.table.qualifier());
271            builder.push_sql(".*");
272        } else if let Some(columns) = &self.returning {
273            builder.push_sql(" RETURNING ");
274            for (idx, col) in columns.iter().enumerate() {
275                if idx > 0 {
276                    builder.push_sql(", ");
277                }
278                builder.push_column(*col);
279            }
280        }
281        builder.finish()
282    }
283}
284
285#[derive(Debug, Clone)]
286pub struct InsertRow {
287    columns: Vec<ColumnRef>,
288    values: Vec<Value>,
289    expected: Option<Vec<ColumnRef>>,
290}
291
292impl InsertRow {
293    fn new(expected: Option<Vec<ColumnRef>>) -> Self {
294        let columns = expected.clone().unwrap_or_default();
295        Self {
296            columns,
297            values: Vec::new(),
298            expected,
299        }
300    }
301
302    pub fn value<M, T, V>(mut self, column: Column<M, T>, value: V) -> Self
303    where
304        V: ColumnValue<T>,
305    {
306        let column_ref = column.as_ref();
307        if let Some(expected) = &self.expected {
308            let idx = self.values.len();
309            if idx >= expected.len() {
310                panic!("dbkit: insert row has too many values");
311            }
312            if expected[idx] != column_ref {
313                panic!("dbkit: insert row column mismatch");
314            }
315        } else {
316            self.columns.push(column_ref);
317        }
318
319        let value = match value.into_value() {
320            Some(value) => value,
321            None => Value::Null,
322        };
323        self.values.push(value);
324        self
325    }
326}
327
328#[derive(Debug, Clone)]
329pub struct Update<Out> {
330    table: Table,
331    sets: Vec<(ColumnRef, Value)>,
332    filters: Vec<Expr<bool>>,
333    returning: Option<Vec<ColumnRef>>,
334    returning_all: bool,
335    _marker: PhantomData<Out>,
336}
337
338impl<Out> Update<Out> {
339    pub fn new(table: Table) -> Self {
340        Self {
341            table,
342            sets: Vec::new(),
343            filters: Vec::new(),
344            returning: None,
345            returning_all: false,
346            _marker: PhantomData,
347        }
348    }
349
350    pub fn set<M, T, V>(mut self, column: Column<M, T>, value: V) -> Self
351    where
352        V: ColumnValue<T>,
353    {
354        let value = match value.into_value() {
355            Some(value) => value,
356            None => Value::Null,
357        };
358        self.sets.push((column.as_ref(), value));
359        self
360    }
361
362    pub fn filter(mut self, expr: Expr<bool>) -> Self {
363        self.filters.push(expr);
364        self
365    }
366
367    pub fn where_exists<SubOut, SubLoads, SubLock, SubDistinctState, SubGroupState>(
368        self,
369        subquery: Select<SubOut, SubLoads, SubLock, SubDistinctState, SubGroupState>,
370    ) -> Self {
371        self.filter(func::exists(subquery))
372    }
373
374    pub fn where_not_exists<SubOut, SubLoads, SubLock, SubDistinctState, SubGroupState>(
375        self,
376        subquery: Select<SubOut, SubLoads, SubLock, SubDistinctState, SubGroupState>,
377    ) -> Self {
378        self.filter(func::exists(subquery).not())
379    }
380
381    pub fn returning(mut self, columns: Vec<ColumnRef>) -> Self {
382        self.returning = Some(columns);
383        self.returning_all = false;
384        self
385    }
386
387    pub fn returning_all(mut self) -> Self {
388        self.returning = None;
389        self.returning_all = true;
390        self
391    }
392
393    pub fn compile(&self) -> CompiledSql {
394        let mut builder = SqlBuilder::new();
395        builder.push_sql("UPDATE ");
396        builder.push_sql(&self.table.qualified_name());
397        builder.push_sql(" SET ");
398        for (idx, (col, value)) in self.sets.iter().enumerate() {
399            if idx > 0 {
400                builder.push_sql(", ");
401            }
402            builder.push_sql(col.name);
403            builder.push_sql(" = ");
404            builder.push_value(value.clone());
405        }
406        if !self.filters.is_empty() {
407            builder.push_sql(" WHERE ");
408            for (idx, expr) in self.filters.iter().enumerate() {
409                if idx > 0 {
410                    builder.push_sql(" AND ");
411                }
412                expr.node.to_sql(&mut builder);
413            }
414        }
415        if self.returning_all {
416            builder.push_sql(" RETURNING ");
417            builder.push_sql(self.table.qualifier());
418            builder.push_sql(".*");
419        } else if let Some(columns) = &self.returning {
420            builder.push_sql(" RETURNING ");
421            for (idx, col) in columns.iter().enumerate() {
422                if idx > 0 {
423                    builder.push_sql(", ");
424                }
425                builder.push_column(*col);
426            }
427        }
428        builder.finish()
429    }
430}
431
432#[derive(Debug, Clone)]
433pub struct Delete {
434    table: Table,
435    filters: Vec<Expr<bool>>,
436    returning: Option<Vec<ColumnRef>>,
437    returning_all: bool,
438}
439
440impl Delete {
441    pub fn new(table: Table) -> Self {
442        Self {
443            table,
444            filters: Vec::new(),
445            returning: None,
446            returning_all: false,
447        }
448    }
449
450    pub fn filter(mut self, expr: Expr<bool>) -> Self {
451        self.filters.push(expr);
452        self
453    }
454
455    pub fn where_exists<SubOut, SubLoads, SubLock, SubDistinctState, SubGroupState>(
456        self,
457        subquery: Select<SubOut, SubLoads, SubLock, SubDistinctState, SubGroupState>,
458    ) -> Self {
459        self.filter(func::exists(subquery))
460    }
461
462    pub fn where_not_exists<SubOut, SubLoads, SubLock, SubDistinctState, SubGroupState>(
463        self,
464        subquery: Select<SubOut, SubLoads, SubLock, SubDistinctState, SubGroupState>,
465    ) -> Self {
466        self.filter(func::exists(subquery).not())
467    }
468
469    pub fn returning(mut self, columns: Vec<ColumnRef>) -> Self {
470        self.returning = Some(columns);
471        self.returning_all = false;
472        self
473    }
474
475    pub fn returning_all(mut self) -> Self {
476        self.returning = None;
477        self.returning_all = true;
478        self
479    }
480
481    pub fn compile(&self) -> CompiledSql {
482        let mut builder = SqlBuilder::new();
483        builder.push_sql("DELETE FROM ");
484        builder.push_sql(&self.table.qualified_name());
485        if !self.filters.is_empty() {
486            builder.push_sql(" WHERE ");
487            for (idx, expr) in self.filters.iter().enumerate() {
488                if idx > 0 {
489                    builder.push_sql(" AND ");
490                }
491                expr.node.to_sql(&mut builder);
492            }
493        }
494        if self.returning_all {
495            builder.push_sql(" RETURNING ");
496            builder.push_sql(self.table.qualifier());
497            builder.push_sql(".*");
498        } else if let Some(columns) = &self.returning {
499            builder.push_sql(" RETURNING ");
500            for (idx, col) in columns.iter().enumerate() {
501                if idx > 0 {
502                    builder.push_sql(", ");
503                }
504                builder.push_column(*col);
505            }
506        }
507        builder.finish()
508    }
509}