1use alloc::string::String;
25use core::fmt::{Debug, Write};
26use core::hash::Hash;
27
28use crate::builders::operation::Operation;
29use crate::builders::{ChangesetFormat, DiffSetBuilder, PatchsetFormat};
30use crate::encoding::{MaybeValue, Value};
31use crate::schema::NamedColumns;
32
33type ChangesetUpdatePairs<S, B> = [(MaybeValue<S, B>, MaybeValue<S, B>)];
35
36fn quote_identifier(name: &str) -> String {
40 let mut out = String::with_capacity(name.len() + 2);
41 out.push('"');
42 for c in name.chars() {
43 if c == '"' {
44 out.push_str("\"\"");
45 } else {
46 out.push(c);
47 }
48 }
49 out.push('"');
50 out
51}
52
53pub trait ColumnNames: NamedColumns {
57 fn column_name(&self, index: usize) -> Option<&str>;
61}
62
63impl ColumnNames for crate::SimpleTable {
64 fn column_name(&self, index: usize) -> Option<&str> {
65 self.column_name(index)
66 }
67}
68
69fn format_insert<T: ColumnNames, S: AsRef<str>, B: AsRef<[u8]>>(
71 table: &T,
72 values: &[Value<S, B>],
73) -> String {
74 let mut sql = String::new();
75 write!(sql, "INSERT INTO {}", quote_identifier(table.name())).unwrap();
76
77 sql.push_str(" (");
79 for i in 0..table.number_of_columns() {
80 if i > 0 {
81 sql.push_str(", ");
82 }
83 if let Some(name) = table.column_name(i) {
84 sql.push_str("e_identifier(name));
85 } else {
86 write!(sql, "\"col{i}\"").unwrap();
87 }
88 }
89 sql.push_str(") VALUES (");
90
91 for (i, val) in values.iter().enumerate() {
93 if i > 0 {
94 sql.push_str(", ");
95 }
96 write!(sql, "{val}").unwrap();
97 }
98 sql.push(')');
99 sql
100}
101
102fn format_delete_changeset<T: ColumnNames, S: AsRef<str>, B: AsRef<[u8]>>(
104 table: &T,
105 values: &[Value<S, B>],
106) -> String {
107 let mut sql = String::new();
108 write!(sql, "DELETE FROM {}", quote_identifier(table.name())).unwrap();
109 sql.push_str(" WHERE ");
110
111 let mut first = true;
113 for (col_idx, value) in values.iter().enumerate() {
114 if table.primary_key_index(col_idx).is_some() {
115 if !first {
116 sql.push_str(" AND ");
117 }
118 first = false;
119 if let Some(name) = table.column_name(col_idx) {
120 sql.push_str("e_identifier(name));
121 } else {
122 write!(sql, "\"col{col_idx}\"").unwrap();
123 }
124 sql.push_str(" = ");
125 write!(sql, "{value}").unwrap();
126 }
127 }
128 sql
129}
130
131fn format_delete_patchset<T: ColumnNames, S: AsRef<str>, B: AsRef<[u8]>>(
133 table: &T,
134 pk: &[Value<S, B>],
135) -> String {
136 let mut sql = String::new();
137 write!(sql, "DELETE FROM {}", quote_identifier(table.name())).unwrap();
138 sql.push_str(" WHERE ");
139
140 let pk_indices = table.primary_key_columns();
142
143 let mut first = true;
144 for (pk_ordinal, &col_idx) in pk_indices.iter().enumerate() {
145 if !first {
146 sql.push_str(" AND ");
147 }
148 first = false;
149 if let Some(name) = table.column_name(col_idx) {
150 sql.push_str("e_identifier(name));
151 } else {
152 write!(sql, "\"col{col_idx}\"").unwrap();
153 }
154 sql.push_str(" = ");
155 write!(sql, "{}", pk[pk_ordinal]).unwrap();
156 }
157 sql
158}
159
160fn format_update_changeset<T: ColumnNames, S: AsRef<str>, B: AsRef<[u8]>>(
162 table: &T,
163 pairs: &ChangesetUpdatePairs<S, B>,
164) -> String {
165 let mut sql = String::new();
166 write!(sql, "UPDATE {}", quote_identifier(table.name())).unwrap();
167 sql.push_str(" SET ");
168
169 let mut first_set = true;
171 for (col_idx, (_old, new)) in pairs.iter().enumerate() {
172 if let Some(new_val) = new {
173 if table.primary_key_index(col_idx).is_some() {
175 continue;
176 }
177 if !first_set {
178 sql.push_str(", ");
179 }
180 first_set = false;
181 if let Some(name) = table.column_name(col_idx) {
182 sql.push_str("e_identifier(name));
183 } else {
184 write!(sql, "\"col{col_idx}\"").unwrap();
185 }
186 sql.push_str(" = ");
187 write!(sql, "{new_val}").unwrap();
188 }
189 }
190
191 sql.push_str(" WHERE ");
193 let mut first_where = true;
194 for (col_idx, (old, _new)) in pairs.iter().enumerate() {
195 if table.primary_key_index(col_idx).is_some() {
196 if let Some(old_val) = old {
197 if !first_where {
198 sql.push_str(" AND ");
199 }
200 first_where = false;
201 if let Some(name) = table.column_name(col_idx) {
202 sql.push_str("e_identifier(name));
203 } else {
204 write!(sql, "\"col{col_idx}\"").unwrap();
205 }
206 sql.push_str(" = ");
207 write!(sql, "{old_val}").unwrap();
208 }
209 }
210 }
211 sql
212}
213
214fn format_update_patchset<T: ColumnNames, S: AsRef<str>, B: AsRef<[u8]>>(
216 table: &T,
217 pk: &[Value<S, B>],
218 pairs: &[((), MaybeValue<S, B>)],
219) -> String {
220 let mut sql = String::new();
221 write!(sql, "UPDATE {}", quote_identifier(table.name())).unwrap();
222 sql.push_str(" SET ");
223
224 let mut first_set = true;
226 for (col_idx, ((), new)) in pairs.iter().enumerate() {
227 if let Some(new_val) = new {
228 if table.primary_key_index(col_idx).is_some() {
230 continue;
231 }
232 if !first_set {
233 sql.push_str(", ");
234 }
235 first_set = false;
236 if let Some(name) = table.column_name(col_idx) {
237 sql.push_str("e_identifier(name));
238 } else {
239 write!(sql, "\"col{col_idx}\"").unwrap();
240 }
241 sql.push_str(" = ");
242 write!(sql, "{new_val}").unwrap();
243 }
244 }
245
246 sql.push_str(" WHERE ");
248 let pk_indices = table.primary_key_columns();
249
250 let mut first_where = true;
251 for (pk_ordinal, &col_idx) in pk_indices.iter().enumerate() {
252 if !first_where {
253 sql.push_str(" AND ");
254 }
255 first_where = false;
256 if let Some(name) = table.column_name(col_idx) {
257 sql.push_str("e_identifier(name));
258 } else {
259 write!(sql, "\"col{col_idx}\"").unwrap();
260 }
261 sql.push_str(" = ");
262 write!(sql, "{}", pk[pk_ordinal]).unwrap();
263 }
264 sql
265}
266
267impl<
272 T: ColumnNames,
273 S: AsRef<str> + Clone + Debug + Hash + Eq,
274 B: AsRef<[u8]> + Clone + Debug + Hash + Eq,
275> DiffSetBuilder<ChangesetFormat, T, S, B>
276{
277 pub fn sql_statements(&self) -> impl Iterator<Item = String> + '_ {
300 self.tables.iter().flat_map(|(table, rows)| {
301 rows.values().map(move |op| match op {
302 Operation::Insert { values, .. } => format_insert(table, values),
303 Operation::Delete { data: values, .. } => format_delete_changeset(table, values),
304 Operation::Update { values, .. } => format_update_changeset(table, values),
305 })
306 })
307 }
308}
309
310impl<T: ColumnNames, S: AsRef<str> + Clone + Hash + Eq, B: AsRef<[u8]> + Clone + Hash + Eq>
311 DiffSetBuilder<PatchsetFormat, T, S, B>
312{
313 pub fn sql_statements(&self) -> impl Iterator<Item = String> + '_ {
336 self.tables.iter().flat_map(|(table, rows)| {
337 rows.iter().map(move |(pk, op)| match op {
338 Operation::Insert { values, .. } => format_insert(table, values),
339 Operation::Delete { data: (), .. } => format_delete_patchset(table, pk),
340 Operation::Update { values, .. } => format_update_patchset(table, pk, values),
341 })
342 })
343 }
344}
345
346#[cfg(test)]
347mod tests {
348 use super::*;
349 use crate::{
350 ChangeDelete, ChangeSet, DiffOps, Insert, PatchDelete, PatchSet, SimpleTable, Update,
351 };
352 use alloc::vec::Vec;
353
354 #[test]
355 fn test_quote_identifier_simple() {
356 assert_eq!(quote_identifier("users"), r#""users""#);
357 }
358
359 #[test]
360 fn test_quote_identifier_with_quotes() {
361 assert_eq!(quote_identifier(r#"user"name"#), r#""user""name""#);
362 }
363
364 #[test]
365 fn test_changeset_insert_sql() {
366 let table = SimpleTable::new("users", &["id", "name"], &[0]);
367 let insert = Insert::from(table.clone())
368 .set(0, 1i64)
369 .unwrap()
370 .set(1, "Alice")
371 .unwrap();
372
373 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().insert(insert);
374 let stmts: Vec<_> = cs.sql_statements().collect();
375
376 assert_eq!(stmts.len(), 1);
377 assert_eq!(
378 stmts[0],
379 r#"INSERT INTO "users" ("id", "name") VALUES (1, 'Alice')"#
380 );
381 }
382
383 #[test]
384 fn test_patchset_insert_sql() {
385 let table = SimpleTable::new("users", &["id", "name"], &[0]);
386 let insert = Insert::from(table.clone())
387 .set(0, 1i64)
388 .unwrap()
389 .set(1, "Alice")
390 .unwrap();
391
392 let ps = PatchSet::<SimpleTable, String, Vec<u8>>::new().insert(insert);
393 let stmts: Vec<_> = ps.sql_statements().collect();
394
395 assert_eq!(stmts.len(), 1);
396 assert_eq!(
397 stmts[0],
398 r#"INSERT INTO "users" ("id", "name") VALUES (1, 'Alice')"#
399 );
400 }
401
402 #[test]
403 fn test_changeset_delete_sql() {
404 let table = SimpleTable::new("users", &["id", "name"], &[0]);
405 let delete = ChangeDelete::from(table.clone())
406 .set(0, 42i64)
407 .unwrap()
408 .set(1, "Bob")
409 .unwrap();
410
411 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().delete(delete);
412 let stmts: Vec<_> = cs.sql_statements().collect();
413
414 assert_eq!(stmts.len(), 1);
415 assert_eq!(stmts[0], r#"DELETE FROM "users" WHERE "id" = 42"#);
416 }
417
418 #[test]
419 fn test_patchset_delete_sql() {
420 let table = SimpleTable::new("users", &["id", "name"], &[0]);
421 let delete = PatchDelete::new(table.clone(), alloc::vec![Value::Integer(42)]);
422
423 let ps = PatchSet::<SimpleTable, String, Vec<u8>>::new().delete(delete);
424 let stmts: Vec<_> = ps.sql_statements().collect();
425
426 assert_eq!(stmts.len(), 1);
427 assert_eq!(stmts[0], r#"DELETE FROM "users" WHERE "id" = 42"#);
428 }
429
430 #[test]
431 fn test_changeset_update_sql() {
432 let table = SimpleTable::new("users", &["id", "name"], &[0]);
433 let update = Update::<SimpleTable, ChangesetFormat, String, Vec<u8>>::from(table.clone())
434 .set(0, 1i64, 1i64)
435 .unwrap()
436 .set(1, "Alice", "Alicia")
437 .unwrap();
438
439 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().update(update);
440 let stmts: Vec<_> = cs.sql_statements().collect();
441
442 assert_eq!(stmts.len(), 1);
443 assert_eq!(
444 stmts[0],
445 r#"UPDATE "users" SET "name" = 'Alicia' WHERE "id" = 1"#
446 );
447 }
448
449 #[test]
450 fn test_patchset_update_sql() {
451 let table = SimpleTable::new("users", &["id", "name"], &[0]);
452 let update = Update::<SimpleTable, PatchsetFormat, String, Vec<u8>>::from(table.clone())
453 .set(0, 1i64)
454 .unwrap()
455 .set(1, "Alicia")
456 .unwrap();
457
458 let ps = PatchSet::<SimpleTable, String, Vec<u8>>::new().update(update);
459 let stmts: Vec<_> = ps.sql_statements().collect();
460
461 assert_eq!(stmts.len(), 1);
462 assert_eq!(
463 stmts[0],
464 r#"UPDATE "users" SET "name" = 'Alicia' WHERE "id" = 1"#
465 );
466 }
467
468 #[test]
469 fn test_sql_escapes_quotes_in_strings() {
470 let table = SimpleTable::new("users", &["id", "name"], &[0]);
471 let insert = Insert::from(table.clone())
472 .set(0, 1i64)
473 .unwrap()
474 .set(1, "O'Brien")
475 .unwrap();
476
477 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().insert(insert);
478 let stmts: Vec<_> = cs.sql_statements().collect();
479
480 assert_eq!(
481 stmts[0],
482 r#"INSERT INTO "users" ("id", "name") VALUES (1, 'O''Brien')"#
483 );
484 }
485
486 #[test]
487 fn test_sql_escapes_quotes_in_identifiers() {
488 let table = SimpleTable::new(r#"user"table"#, &["id", r#"user"name"#], &[0]);
489 let insert = Insert::from(table.clone())
490 .set(0, 1i64)
491 .unwrap()
492 .set(1, "Alice")
493 .unwrap();
494
495 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().insert(insert);
496 let stmts: Vec<_> = cs.sql_statements().collect();
497
498 assert_eq!(
499 stmts[0],
500 r#"INSERT INTO "user""table" ("id", "user""name") VALUES (1, 'Alice')"#
501 );
502 }
503
504 #[test]
505 fn test_multiple_operations() {
506 let table = SimpleTable::new("users", &["id", "name"], &[0]);
507
508 let insert1 = Insert::from(table.clone())
509 .set(0, 1i64)
510 .unwrap()
511 .set(1, "Alice")
512 .unwrap();
513
514 let insert2 = Insert::from(table.clone())
515 .set(0, 2i64)
516 .unwrap()
517 .set(1, "Bob")
518 .unwrap();
519
520 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new()
521 .insert(insert1)
522 .insert(insert2);
523
524 let stmts: Vec<_> = cs.sql_statements().collect();
525 assert_eq!(stmts.len(), 2);
526 }
527
528 #[test]
529 fn test_composite_pk_delete() {
530 let table = SimpleTable::new("order_items", &["order_id", "item_id", "qty"], &[0, 1]);
531 let delete = ChangeDelete::from(table.clone())
532 .set(0, 100i64)
533 .unwrap()
534 .set(1, 5i64)
535 .unwrap()
536 .set(2, 10i64)
537 .unwrap();
538
539 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().delete(delete);
540 let stmts: Vec<_> = cs.sql_statements().collect();
541
542 assert_eq!(
543 stmts[0],
544 r#"DELETE FROM "order_items" WHERE "order_id" = 100 AND "item_id" = 5"#
545 );
546 }
547
548 #[derive(Debug, Clone, PartialEq, Eq, Hash)]
552 struct AnonColsTable {
553 name: alloc::string::String,
554 num_columns: usize,
555 pk_column: usize,
556 }
557
558 impl crate::DynTable for AnonColsTable {
559 fn name(&self) -> &str {
560 &self.name
561 }
562 fn number_of_columns(&self) -> usize {
563 self.num_columns
564 }
565 fn write_pk_flags(&self, buf: &mut [u8]) {
566 for (i, b) in buf.iter_mut().enumerate() {
567 *b = u8::from(i == self.pk_column);
568 }
569 }
570 }
571
572 impl crate::SchemaWithPK for AnonColsTable {
573 fn number_of_primary_keys(&self) -> usize {
574 1
575 }
576 fn primary_key_index(&self, col_idx: usize) -> Option<usize> {
577 (col_idx == self.pk_column).then_some(0)
578 }
579 fn extract_pk<S: Clone, B: Clone>(
580 &self,
581 values: &impl crate::IndexableValues<Text = S, Binary = B>,
582 ) -> Vec<Value<S, B>> {
583 alloc::vec![values.get(self.pk_column).unwrap()]
584 }
585 }
586
587 impl crate::schema::NamedColumns for AnonColsTable {
588 fn column_index(&self, _column_name: &str) -> Option<usize> {
589 None
590 }
591 }
592
593 impl crate::ColumnNames for AnonColsTable {
594 fn column_name(&self, _index: usize) -> Option<&str> {
597 None
598 }
599 }
600
601 #[test]
602 fn test_sql_output_uses_col_index_fallback() {
603 let table = AnonColsTable {
604 name: "t".into(),
605 num_columns: 2,
606 pk_column: 0,
607 };
608 let insert = Insert::from(table.clone())
609 .set(0, 42i64)
610 .unwrap()
611 .set(1, "hello")
612 .unwrap();
613 let cs = ChangeSet::<AnonColsTable, String, Vec<u8>>::new().insert(insert);
614 let stmts: Vec<_> = cs.sql_statements().collect();
615 assert_eq!(
616 stmts[0],
617 r#"INSERT INTO "t" ("col0", "col1") VALUES (42, 'hello')"#
618 );
619 }
620
621 #[test]
622 fn test_sql_output_col_index_fallback_in_patchset_delete() {
623 let table = AnonColsTable {
624 name: "t".into(),
625 num_columns: 2,
626 pk_column: 0,
627 };
628 let delete: PatchDelete<_, String, Vec<u8>> =
629 PatchDelete::new(table.clone(), alloc::vec![Value::Integer(7)]);
630 let ps = PatchSet::<AnonColsTable, String, Vec<u8>>::new().delete(delete);
631 let stmts: Vec<_> = ps.sql_statements().collect();
632 assert_eq!(stmts[0], r#"DELETE FROM "t" WHERE "col0" = 7"#);
633 }
634}