reinhardt_db/backends/optimization/
batch_ops.rs1use crate::backends::error::Result;
9use crate::backends::types::DatabaseType;
10use async_trait::async_trait;
11
12#[async_trait]
14pub trait BatchOperations {
15 async fn batch_insert(
17 &self,
18 table: &str,
19 columns: &[&str],
20 rows: Vec<Vec<String>>,
21 ) -> Result<u64>;
22
23 async fn batch_update(
25 &self,
26 table: &str,
27 updates: Vec<(String, Vec<(String, String)>)>, ) -> Result<u64>;
29
30 async fn batch_delete(&self, table: &str, ids: Vec<i64>) -> Result<u64>;
32}
33
34#[derive(Debug, Clone, Copy, PartialEq, Eq)]
36pub enum QuoteStyle {
37 Ansi,
39 Backtick,
41}
42
43impl QuoteStyle {
44 fn quote_identifier(&self, ident: &str) -> String {
46 match self {
47 QuoteStyle::Ansi => format!("\"{}\"", ident.replace('"', "\"\"")),
48 QuoteStyle::Backtick => format!("`{}`", ident.replace('`', "``")),
49 }
50 }
51}
52
53impl From<DatabaseType> for QuoteStyle {
54 fn from(db_type: DatabaseType) -> Self {
59 match db_type {
60 DatabaseType::Mysql => QuoteStyle::Backtick,
61 DatabaseType::Postgres | DatabaseType::Sqlite => QuoteStyle::Ansi,
62 }
63 }
64}
65
66pub struct BatchInsertBuilder {
68 table: String,
69 columns: Vec<String>,
70 rows: Vec<Vec<String>>,
71 batch_size: usize,
72 quote_style: QuoteStyle,
73}
74
75impl BatchInsertBuilder {
76 pub fn new(table: impl Into<String>) -> Self {
78 Self {
79 table: table.into(),
80 columns: Vec::new(),
81 rows: Vec::new(),
82 batch_size: 1000,
83 quote_style: QuoteStyle::Ansi,
84 }
85 }
86
87 pub fn columns(mut self, columns: Vec<String>) -> Self {
89 self.columns = columns;
90 self
91 }
92
93 pub fn add_row(mut self, row: Vec<String>) -> Self {
95 self.rows.push(row);
96 self
97 }
98
99 pub fn batch_size(mut self, size: usize) -> Self {
101 self.batch_size = size;
102 self
103 }
104
105 pub fn quote_style(mut self, style: QuoteStyle) -> Self {
110 self.quote_style = style;
111 self
112 }
113
114 pub fn build_sql(&self) -> Vec<String> {
120 let mut statements = Vec::new();
121
122 let quoted_table = self.quote_style.quote_identifier(&self.table);
123 let quoted_columns = self
124 .columns
125 .iter()
126 .map(|c| self.quote_style.quote_identifier(c))
127 .collect::<Vec<_>>()
128 .join(", ");
129
130 for chunk in self.rows.chunks(self.batch_size) {
131 let values_list: Vec<String> = chunk
132 .iter()
133 .map(|row| {
134 let values = row
135 .iter()
136 .map(|v| format!("'{}'", v.replace('\'', "''")))
137 .collect::<Vec<_>>()
138 .join(", ");
139 format!("({})", values)
140 })
141 .collect();
142
143 let sql = format!(
144 "INSERT INTO {} ({}) VALUES {}",
145 quoted_table,
146 quoted_columns,
147 values_list.join(", ")
148 );
149
150 statements.push(sql);
151 }
152
153 statements
154 }
155
156 pub fn row_count(&self) -> usize {
158 self.rows.len()
159 }
160}
161
162struct UpdateEntry {
164 where_column: String,
166 where_value: String,
168 columns_values: Vec<(String, String)>,
170}
171
172pub struct BatchUpdateBuilder {
178 table: String,
179 updates: Vec<UpdateEntry>,
180}
181
182impl BatchUpdateBuilder {
183 pub fn new(table: impl Into<String>) -> Self {
185 Self {
186 table: table.into(),
187 updates: Vec::new(),
188 }
189 }
190
191 pub fn add_update_parameterized(
196 mut self,
197 where_column: String,
198 where_value: String,
199 columns_values: Vec<(String, String)>,
200 ) -> Self {
201 self.updates.push(UpdateEntry {
202 where_column,
203 where_value,
204 columns_values,
205 });
206 self
207 }
208
209 pub fn build_sql_parameterized(&self) -> Vec<(String, Vec<String>)> {
214 self.updates
215 .iter()
216 .map(|entry| {
217 let mut params = Vec::with_capacity(entry.columns_values.len() + 1);
218 let mut param_idx = 1usize;
219
220 let set_clause = entry
221 .columns_values
222 .iter()
223 .map(|(col, val)| {
224 let placeholder = format!("{} = ${}", col, param_idx);
225 params.push(val.clone());
226 param_idx += 1;
227 placeholder
228 })
229 .collect::<Vec<_>>()
230 .join(", ");
231
232 let sql = format!(
233 "UPDATE {} SET {} WHERE {} = ${}",
234 self.table, set_clause, entry.where_column, param_idx
235 );
236 params.push(entry.where_value.clone());
237
238 (sql, params)
239 })
240 .collect()
241 }
242
243 pub fn update_count(&self) -> usize {
245 self.updates.len()
246 }
247}
248
249#[cfg(test)]
250mod tests {
251 use super::*;
252 use rstest::rstest;
253
254 #[rstest]
255 fn test_batch_insert_builder() {
256 let builder = BatchInsertBuilder::new("users")
258 .columns(vec!["name".to_string(), "email".to_string()])
259 .add_row(vec!["Alice".to_string(), "alice@example.com".to_string()])
260 .add_row(vec!["Bob".to_string(), "bob@example.com".to_string()])
261 .batch_size(2);
262
263 let sql_statements = builder.build_sql();
265
266 assert_eq!(sql_statements.len(), 1);
268 assert!(sql_statements[0].contains("INSERT INTO \"users\""));
269 assert!(sql_statements[0].contains("\"name\""));
270 assert!(sql_statements[0].contains("\"email\""));
271 assert!(sql_statements[0].contains("Alice"));
272 assert!(sql_statements[0].contains("Bob"));
273 }
274
275 #[rstest]
276 fn test_batch_insert_chunking() {
277 let mut builder = BatchInsertBuilder::new("users")
279 .columns(vec!["name".to_string()])
280 .batch_size(2);
281
282 for i in 0..5 {
283 builder = builder.add_row(vec![format!("User{}", i)]);
284 }
285
286 let sql_statements = builder.build_sql();
288
289 assert_eq!(sql_statements.len(), 3);
291 }
292
293 #[rstest]
294 fn test_batch_update_builder_uses_parameterized_queries() {
295 let builder = BatchUpdateBuilder::new("users")
297 .add_update_parameterized(
298 "id".to_string(),
299 "1".to_string(),
300 vec![("name".to_string(), "Alice Updated".to_string())],
301 )
302 .add_update_parameterized(
303 "id".to_string(),
304 "2".to_string(),
305 vec![("name".to_string(), "Bob Updated".to_string())],
306 );
307
308 let statements = builder.build_sql_parameterized();
310
311 assert_eq!(statements.len(), 2);
313 let (sql, params) = &statements[0];
314 assert_eq!(sql, "UPDATE users SET name = $1 WHERE id = $2");
315 assert_eq!(params, &["Alice Updated", "1"]);
316
317 let (sql, params) = &statements[1];
318 assert_eq!(sql, "UPDATE users SET name = $1 WHERE id = $2");
319 assert_eq!(params, &["Bob Updated", "2"]);
320 }
321
322 #[rstest]
323 fn test_batch_update_sql_injection_in_where_value_is_parameterized() {
324 let builder = BatchUpdateBuilder::new("users").add_update_parameterized(
326 "id".to_string(),
327 "1 OR 1=1; DROP TABLE users; --".to_string(),
328 vec![("name".to_string(), "hacked".to_string())],
329 );
330
331 let statements = builder.build_sql_parameterized();
333
334 let (sql, params) = &statements[0];
336 assert_eq!(sql, "UPDATE users SET name = $1 WHERE id = $2");
337 assert!(!sql.contains("DROP TABLE"));
338 assert_eq!(params[1], "1 OR 1=1; DROP TABLE users; --");
339 }
340
341 #[rstest]
342 fn test_batch_update_sql_injection_in_set_value_is_parameterized() {
343 let builder = BatchUpdateBuilder::new("users").add_update_parameterized(
345 "id".to_string(),
346 "1".to_string(),
347 vec![("name".to_string(), "'; DROP TABLE users; --".to_string())],
348 );
349
350 let statements = builder.build_sql_parameterized();
352
353 let (sql, params) = &statements[0];
355 assert_eq!(sql, "UPDATE users SET name = $1 WHERE id = $2");
356 assert!(!sql.contains("DROP TABLE"));
357 assert_eq!(params[0], "'; DROP TABLE users; --");
358 }
359
360 #[rstest]
361 fn test_batch_update_multiple_columns() {
362 let builder = BatchUpdateBuilder::new("users").add_update_parameterized(
364 "id".to_string(),
365 "42".to_string(),
366 vec![
367 ("name".to_string(), "Alice".to_string()),
368 ("email".to_string(), "alice@example.com".to_string()),
369 ],
370 );
371
372 let statements = builder.build_sql_parameterized();
374
375 let (sql, params) = &statements[0];
377 assert_eq!(sql, "UPDATE users SET name = $1, email = $2 WHERE id = $3");
378 assert_eq!(params, &["Alice", "alice@example.com", "42"]);
379 }
380
381 #[rstest]
382 fn test_sql_injection_protection_in_insert() {
383 let builder = BatchInsertBuilder::new("users")
385 .columns(vec!["name".to_string()])
386 .add_row(vec!["Alice'; DROP TABLE users; --".to_string()]);
387
388 let sql_statements = builder.build_sql();
390
391 assert!(sql_statements[0].contains("Alice''; DROP TABLE users; --"));
393 }
394
395 #[rstest]
396 fn test_batch_insert_quotes_table_name() {
397 let builder = BatchInsertBuilder::new("users\"; DROP TABLE data; --")
399 .columns(vec!["name".to_string()])
400 .add_row(vec!["Alice".to_string()]);
401
402 let sql_statements = builder.build_sql();
404
405 assert!(sql_statements[0].starts_with("INSERT INTO \"users\"\"; DROP TABLE data; --\""));
407 }
408
409 #[rstest]
410 fn test_batch_insert_quotes_column_names() {
411 let builder = BatchInsertBuilder::new("users")
413 .columns(vec![
414 "name".to_string(),
415 "col\"; DROP TABLE users; --".to_string(),
416 ])
417 .add_row(vec!["Alice".to_string(), "value".to_string()]);
418
419 let sql_statements = builder.build_sql();
421
422 assert!(sql_statements[0].contains("\"name\""));
424 assert!(sql_statements[0].contains("\"col\"\"; DROP TABLE users; --\""));
425 }
426
427 #[rstest]
428 fn test_batch_insert_mysql_backtick_quoting() {
429 let builder = BatchInsertBuilder::new("users")
431 .columns(vec!["name".to_string(), "email".to_string()])
432 .add_row(vec!["Alice".to_string(), "alice@example.com".to_string()])
433 .quote_style(QuoteStyle::Backtick);
434
435 let sql_statements = builder.build_sql();
437
438 assert_eq!(sql_statements.len(), 1);
440 assert!(sql_statements[0].contains("INSERT INTO `users`"));
441 assert!(sql_statements[0].contains("`name`"));
442 assert!(sql_statements[0].contains("`email`"));
443 }
444
445 #[rstest]
446 fn test_batch_insert_mysql_backtick_escaping() {
447 let builder = BatchInsertBuilder::new("my`table")
449 .columns(vec!["col`name".to_string()])
450 .add_row(vec!["value".to_string()])
451 .quote_style(QuoteStyle::Backtick);
452
453 let sql_statements = builder.build_sql();
455
456 assert!(sql_statements[0].contains("INSERT INTO `my``table`"));
458 assert!(sql_statements[0].contains("`col``name`"));
459 }
460
461 #[rstest]
462 fn test_quote_style_from_database_type() {
463 assert_eq!(QuoteStyle::from(DatabaseType::Mysql), QuoteStyle::Backtick);
465 assert_eq!(QuoteStyle::from(DatabaseType::Postgres), QuoteStyle::Ansi);
466 assert_eq!(QuoteStyle::from(DatabaseType::Sqlite), QuoteStyle::Ansi);
467 }
468
469 #[rstest]
470 fn test_batch_insert_with_database_type() {
471 let builder = BatchInsertBuilder::new("users")
473 .columns(vec!["name".to_string()])
474 .add_row(vec!["Alice".to_string()])
475 .quote_style(QuoteStyle::from(DatabaseType::Mysql));
476
477 let sql_statements = builder.build_sql();
479
480 assert!(sql_statements[0].contains("INSERT INTO `users`"));
482 assert!(sql_statements[0].contains("`name`"));
483 }
484}