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}