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 mut first = true;
141 for (pk_ordinal, col_idx) in table.primary_key_columns().enumerate() {
142 if !first {
143 sql.push_str(" AND ");
144 }
145 first = false;
146 if let Some(name) = table.column_name(col_idx) {
147 sql.push_str("e_identifier(name));
148 } else {
149 write!(sql, "\"col{col_idx}\"").unwrap();
150 }
151 sql.push_str(" = ");
152 write!(sql, "{}", pk[pk_ordinal]).unwrap();
153 }
154 sql
155}
156
157fn format_update_changeset<T: ColumnNames, S: AsRef<str>, B: AsRef<[u8]>>(
159 table: &T,
160 pairs: &ChangesetUpdatePairs<S, B>,
161) -> String {
162 let mut sql = String::new();
163 write!(sql, "UPDATE {}", quote_identifier(table.name())).unwrap();
164 sql.push_str(" SET ");
165
166 let mut first_set = true;
168 for (col_idx, (_old, new)) in pairs.iter().enumerate() {
169 if let Some(new_val) = new {
170 if table.primary_key_index(col_idx).is_some() {
172 continue;
173 }
174 if !first_set {
175 sql.push_str(", ");
176 }
177 first_set = false;
178 if let Some(name) = table.column_name(col_idx) {
179 sql.push_str("e_identifier(name));
180 } else {
181 write!(sql, "\"col{col_idx}\"").unwrap();
182 }
183 sql.push_str(" = ");
184 write!(sql, "{new_val}").unwrap();
185 }
186 }
187
188 sql.push_str(" WHERE ");
190 let mut first_where = true;
191 for (col_idx, (old, _new)) in pairs.iter().enumerate() {
192 if table.primary_key_index(col_idx).is_some() {
193 if let Some(old_val) = old {
194 if !first_where {
195 sql.push_str(" AND ");
196 }
197 first_where = false;
198 if let Some(name) = table.column_name(col_idx) {
199 sql.push_str("e_identifier(name));
200 } else {
201 write!(sql, "\"col{col_idx}\"").unwrap();
202 }
203 sql.push_str(" = ");
204 write!(sql, "{old_val}").unwrap();
205 }
206 }
207 }
208 sql
209}
210
211fn format_update_patchset<T: ColumnNames, S: AsRef<str>, B: AsRef<[u8]>>(
213 table: &T,
214 pk: &[Value<S, B>],
215 pairs: &[((), MaybeValue<S, B>)],
216) -> String {
217 let mut sql = String::new();
218 write!(sql, "UPDATE {}", quote_identifier(table.name())).unwrap();
219 sql.push_str(" SET ");
220
221 let mut first_set = true;
223 for (col_idx, ((), new)) in pairs.iter().enumerate() {
224 if let Some(new_val) = new {
225 if table.primary_key_index(col_idx).is_some() {
227 continue;
228 }
229 if !first_set {
230 sql.push_str(", ");
231 }
232 first_set = false;
233 if let Some(name) = table.column_name(col_idx) {
234 sql.push_str("e_identifier(name));
235 } else {
236 write!(sql, "\"col{col_idx}\"").unwrap();
237 }
238 sql.push_str(" = ");
239 write!(sql, "{new_val}").unwrap();
240 }
241 }
242
243 sql.push_str(" WHERE ");
244 let mut first_where = true;
245 for (pk_ordinal, col_idx) in table.primary_key_columns().enumerate() {
246 if !first_where {
247 sql.push_str(" AND ");
248 }
249 first_where = false;
250 if let Some(name) = table.column_name(col_idx) {
251 sql.push_str("e_identifier(name));
252 } else {
253 write!(sql, "\"col{col_idx}\"").unwrap();
254 }
255 sql.push_str(" = ");
256 write!(sql, "{}", pk[pk_ordinal]).unwrap();
257 }
258 sql
259}
260
261impl<
266 T: ColumnNames,
267 S: AsRef<str> + Clone + Debug + Hash + Eq,
268 B: AsRef<[u8]> + Clone + Debug + Hash + Eq,
269> DiffSetBuilder<ChangesetFormat, T, S, B>
270{
271 pub fn sql_statements(&self) -> impl Iterator<Item = String> + '_ {
294 self.tables.iter().flat_map(|(table, rows)| {
295 rows.values().map(move |op| match op {
296 Operation::Insert { values, .. } => format_insert(table, values),
297 Operation::Delete { data: values, .. } => format_delete_changeset(table, values),
298 Operation::Update { values, .. } => format_update_changeset(table, values),
299 })
300 })
301 }
302}
303
304impl<T: ColumnNames, S: AsRef<str> + Clone + Hash + Eq, B: AsRef<[u8]> + Clone + Hash + Eq>
305 DiffSetBuilder<PatchsetFormat, T, S, B>
306{
307 pub fn sql_statements(&self) -> impl Iterator<Item = String> + '_ {
330 self.tables.iter().flat_map(|(table, rows)| {
331 rows.iter().map(move |(pk, op)| match op {
332 Operation::Insert { values, .. } => format_insert(table, values),
333 Operation::Delete { data: (), .. } => format_delete_patchset(table, pk),
334 Operation::Update { values, .. } => format_update_patchset(table, pk, values),
335 })
336 })
337 }
338}
339
340#[cfg(test)]
341mod tests {
342 use super::*;
343 use crate::{
344 ChangeDelete, ChangeSet, DiffOps, Insert, PatchDelete, PatchSet, SimpleTable, Update,
345 };
346 use alloc::vec::Vec;
347
348 #[test]
349 fn test_quote_identifier_simple() {
350 assert_eq!(quote_identifier("users"), r#""users""#);
351 }
352
353 #[test]
354 fn test_quote_identifier_with_quotes() {
355 assert_eq!(quote_identifier(r#"user"name"#), r#""user""name""#);
356 }
357
358 #[test]
359 fn test_changeset_insert_sql() {
360 let table = SimpleTable::new("users", &["id", "name"], &[0]);
361 let insert = Insert::from(table.clone())
362 .set(0, 1i64)
363 .unwrap()
364 .set(1, "Alice")
365 .unwrap();
366
367 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().insert(insert);
368 let stmts: Vec<_> = cs.sql_statements().collect();
369
370 assert_eq!(stmts.len(), 1);
371 assert_eq!(
372 stmts[0],
373 r#"INSERT INTO "users" ("id", "name") VALUES (1, 'Alice')"#
374 );
375 }
376
377 #[test]
378 fn test_patchset_insert_sql() {
379 let table = SimpleTable::new("users", &["id", "name"], &[0]);
380 let insert = Insert::from(table.clone())
381 .set(0, 1i64)
382 .unwrap()
383 .set(1, "Alice")
384 .unwrap();
385
386 let ps = PatchSet::<SimpleTable, String, Vec<u8>>::new().insert(insert);
387 let stmts: Vec<_> = ps.sql_statements().collect();
388
389 assert_eq!(stmts.len(), 1);
390 assert_eq!(
391 stmts[0],
392 r#"INSERT INTO "users" ("id", "name") VALUES (1, 'Alice')"#
393 );
394 }
395
396 #[test]
397 fn test_changeset_delete_sql() {
398 let table = SimpleTable::new("users", &["id", "name"], &[0]);
399 let delete = ChangeDelete::from(table.clone())
400 .set(0, 42i64)
401 .unwrap()
402 .set(1, "Bob")
403 .unwrap();
404
405 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().delete(delete);
406 let stmts: Vec<_> = cs.sql_statements().collect();
407
408 assert_eq!(stmts.len(), 1);
409 assert_eq!(stmts[0], r#"DELETE FROM "users" WHERE "id" = 42"#);
410 }
411
412 #[test]
413 fn test_patchset_delete_sql() {
414 let table = SimpleTable::new("users", &["id", "name"], &[0]);
415 let delete = PatchDelete::new(table.clone(), alloc::vec![Value::Integer(42)]);
416
417 let ps = PatchSet::<SimpleTable, String, Vec<u8>>::new().delete(delete);
418 let stmts: Vec<_> = ps.sql_statements().collect();
419
420 assert_eq!(stmts.len(), 1);
421 assert_eq!(stmts[0], r#"DELETE FROM "users" WHERE "id" = 42"#);
422 }
423
424 #[test]
425 fn test_changeset_update_sql() {
426 let table = SimpleTable::new("users", &["id", "name"], &[0]);
427 let update = Update::<SimpleTable, ChangesetFormat, String, Vec<u8>>::from(table.clone())
428 .set(0, 1i64, 1i64)
429 .unwrap()
430 .set(1, "Alice", "Alicia")
431 .unwrap();
432
433 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().update(update);
434 let stmts: Vec<_> = cs.sql_statements().collect();
435
436 assert_eq!(stmts.len(), 1);
437 assert_eq!(
438 stmts[0],
439 r#"UPDATE "users" SET "name" = 'Alicia' WHERE "id" = 1"#
440 );
441 }
442
443 #[test]
444 fn test_patchset_update_sql() {
445 let table = SimpleTable::new("users", &["id", "name"], &[0]);
446 let update = Update::<SimpleTable, PatchsetFormat, String, Vec<u8>>::from(table.clone())
447 .set(0, 1i64)
448 .unwrap()
449 .set(1, "Alicia")
450 .unwrap();
451
452 let ps = PatchSet::<SimpleTable, String, Vec<u8>>::new().update(update);
453 let stmts: Vec<_> = ps.sql_statements().collect();
454
455 assert_eq!(stmts.len(), 1);
456 assert_eq!(
457 stmts[0],
458 r#"UPDATE "users" SET "name" = 'Alicia' WHERE "id" = 1"#
459 );
460 }
461
462 #[test]
463 fn test_sql_escapes_quotes_in_strings() {
464 let table = SimpleTable::new("users", &["id", "name"], &[0]);
465 let insert = Insert::from(table.clone())
466 .set(0, 1i64)
467 .unwrap()
468 .set(1, "O'Brien")
469 .unwrap();
470
471 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().insert(insert);
472 let stmts: Vec<_> = cs.sql_statements().collect();
473
474 assert_eq!(
475 stmts[0],
476 r#"INSERT INTO "users" ("id", "name") VALUES (1, 'O''Brien')"#
477 );
478 }
479
480 #[test]
481 fn test_sql_escapes_quotes_in_identifiers() {
482 let table = SimpleTable::new(r#"user"table"#, &["id", r#"user"name"#], &[0]);
483 let insert = Insert::from(table.clone())
484 .set(0, 1i64)
485 .unwrap()
486 .set(1, "Alice")
487 .unwrap();
488
489 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().insert(insert);
490 let stmts: Vec<_> = cs.sql_statements().collect();
491
492 assert_eq!(
493 stmts[0],
494 r#"INSERT INTO "user""table" ("id", "user""name") VALUES (1, 'Alice')"#
495 );
496 }
497
498 #[test]
499 fn test_multiple_operations() {
500 let table = SimpleTable::new("users", &["id", "name"], &[0]);
501
502 let insert1 = Insert::from(table.clone())
503 .set(0, 1i64)
504 .unwrap()
505 .set(1, "Alice")
506 .unwrap();
507
508 let insert2 = Insert::from(table.clone())
509 .set(0, 2i64)
510 .unwrap()
511 .set(1, "Bob")
512 .unwrap();
513
514 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new()
515 .insert(insert1)
516 .insert(insert2);
517
518 let stmts: Vec<_> = cs.sql_statements().collect();
519 assert_eq!(stmts.len(), 2);
520 }
521
522 #[test]
523 fn test_composite_pk_delete() {
524 let table = SimpleTable::new("order_items", &["order_id", "item_id", "qty"], &[0, 1]);
525 let delete = ChangeDelete::from(table.clone())
526 .set(0, 100i64)
527 .unwrap()
528 .set(1, 5i64)
529 .unwrap()
530 .set(2, 10i64)
531 .unwrap();
532
533 let cs = ChangeSet::<SimpleTable, String, Vec<u8>>::new().delete(delete);
534 let stmts: Vec<_> = cs.sql_statements().collect();
535
536 assert_eq!(
537 stmts[0],
538 r#"DELETE FROM "order_items" WHERE "order_id" = 100 AND "item_id" = 5"#
539 );
540 }
541
542 #[derive(Debug, Clone, PartialEq, Eq, Hash)]
546 struct AnonColsTable {
547 name: alloc::string::String,
548 num_columns: usize,
549 pk_column: usize,
550 }
551
552 impl crate::DynTable for AnonColsTable {
553 fn name(&self) -> &str {
554 &self.name
555 }
556 fn number_of_columns(&self) -> usize {
557 self.num_columns
558 }
559 fn write_pk_flags(&self, buf: &mut [u8]) {
560 for (i, b) in buf.iter_mut().enumerate() {
561 *b = u8::from(i == self.pk_column);
562 }
563 }
564 }
565
566 impl crate::SchemaWithPK for AnonColsTable {
567 fn number_of_primary_keys(&self) -> usize {
568 1
569 }
570 fn primary_key_index(&self, col_idx: usize) -> Option<usize> {
571 (col_idx == self.pk_column).then_some(0)
572 }
573 fn extract_pk<S: Clone, B: Clone>(
574 &self,
575 values: &impl crate::IndexableValues<Text = S, Binary = B>,
576 ) -> Vec<Value<S, B>> {
577 alloc::vec![values.get(self.pk_column).unwrap()]
578 }
579 }
580
581 impl crate::schema::NamedColumns for AnonColsTable {
582 fn column_index(&self, _column_name: &str) -> Option<usize> {
583 None
584 }
585 }
586
587 impl crate::ColumnNames for AnonColsTable {
588 fn column_name(&self, _index: usize) -> Option<&str> {
591 None
592 }
593 }
594
595 #[test]
596 fn test_sql_output_uses_col_index_fallback() {
597 let table = AnonColsTable {
598 name: "t".into(),
599 num_columns: 2,
600 pk_column: 0,
601 };
602 let insert = Insert::from(table.clone())
603 .set(0, 42i64)
604 .unwrap()
605 .set(1, "hello")
606 .unwrap();
607 let cs = ChangeSet::<AnonColsTable, String, Vec<u8>>::new().insert(insert);
608 let stmts: Vec<_> = cs.sql_statements().collect();
609 assert_eq!(
610 stmts[0],
611 r#"INSERT INTO "t" ("col0", "col1") VALUES (42, 'hello')"#
612 );
613 }
614
615 #[test]
616 fn test_sql_output_col_index_fallback_in_patchset_delete() {
617 let table = AnonColsTable {
618 name: "t".into(),
619 num_columns: 2,
620 pk_column: 0,
621 };
622 let delete: PatchDelete<_, String, Vec<u8>> =
623 PatchDelete::new(table.clone(), alloc::vec![Value::Integer(7)]);
624 let ps = PatchSet::<AnonColsTable, String, Vec<u8>>::new().delete(delete);
625 let stmts: Vec<_> = ps.sql_statements().collect();
626 assert_eq!(stmts[0], r#"DELETE FROM "t" WHERE "col0" = 7"#);
627 }
628}