Skip to main content

reinhardt_db/backends/optimization/
batch_ops.rs

1//! Batch operations for improved performance
2//!
3//! Provides efficient batch insert, update, and delete operations:
4//! - Bulk insert with COPY or multi-value INSERT
5//! - Batch updates with optimized queries
6//! - Transaction batching
7
8use crate::backends::error::Result;
9use crate::backends::types::DatabaseType;
10use async_trait::async_trait;
11
12/// Batch operations trait
13#[async_trait]
14pub trait BatchOperations {
15	/// Execute batch insert
16	async fn batch_insert(
17		&self,
18		table: &str,
19		columns: &[&str],
20		rows: Vec<Vec<String>>,
21	) -> Result<u64>;
22
23	/// Execute batch update
24	async fn batch_update(
25		&self,
26		table: &str,
27		updates: Vec<(String, Vec<(String, String)>)>, // (where_clause, [(column, value)])
28	) -> Result<u64>;
29
30	/// Execute batch delete
31	async fn batch_delete(&self, table: &str, ids: Vec<i64>) -> Result<u64>;
32}
33
34/// Identifier quoting style for different database backends
35#[derive(Debug, Clone, Copy, PartialEq, Eq)]
36pub enum QuoteStyle {
37	/// ANSI SQL: double quotes (`"identifier"`)
38	Ansi,
39	/// MySQL: backticks (`` `identifier` ``)
40	Backtick,
41}
42
43impl QuoteStyle {
44	/// Quote an identifier using this style, escaping embedded quote characters
45	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	/// Convert a [`DatabaseType`] to the appropriate [`QuoteStyle`]
55	///
56	/// - MySQL uses backtick quoting
57	/// - PostgreSQL and SQLite use ANSI double-quote quoting
58	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
66/// Builder for batch insert operations
67pub 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	/// Create a new batch insert builder
77	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	/// Set columns for insert
88	pub fn columns(mut self, columns: Vec<String>) -> Self {
89		self.columns = columns;
90		self
91	}
92
93	/// Add a row of values
94	pub fn add_row(mut self, row: Vec<String>) -> Self {
95		self.rows.push(row);
96		self
97	}
98
99	/// Set batch size (number of rows per INSERT statement)
100	pub fn batch_size(mut self, size: usize) -> Self {
101		self.batch_size = size;
102		self
103	}
104
105	/// Set the identifier quoting style for the target database backend
106	///
107	/// Defaults to [`QuoteStyle::Ansi`] (double quotes). Use
108	/// [`QuoteStyle::Backtick`] for MySQL.
109	pub fn quote_style(mut self, style: QuoteStyle) -> Self {
110		self.quote_style = style;
111		self
112	}
113
114	/// Build SQL statements for batch insert
115	///
116	/// Identifiers (table name and column names) are quoted using the
117	/// configured [`QuoteStyle`] (defaults to ANSI double quotes).
118	/// Embedded quote characters are escaped to prevent SQL injection.
119	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	/// Get total number of rows
157	pub fn row_count(&self) -> usize {
158		self.rows.len()
159	}
160}
161
162/// Internal representation for an update entry with parameterized WHERE clause
163struct UpdateEntry {
164	/// Column name used in the WHERE equality condition
165	where_column: String,
166	/// Value bound to the WHERE parameter
167	where_value: String,
168	/// Column-value pairs to SET
169	columns_values: Vec<(String, String)>,
170}
171
172/// Builder for batch update operations
173///
174/// Uses parameterized queries to prevent SQL injection. Add updates with
175/// [`add_update_parameterized`](Self::add_update_parameterized) and build
176/// statements with [`build_sql_parameterized`](Self::build_sql_parameterized).
177pub struct BatchUpdateBuilder {
178	table: String,
179	updates: Vec<UpdateEntry>,
180}
181
182impl BatchUpdateBuilder {
183	/// Create a new batch update builder
184	pub fn new(table: impl Into<String>) -> Self {
185		Self {
186			table: table.into(),
187			updates: Vec::new(),
188		}
189	}
190
191	/// Add an update operation with a parameterized WHERE clause
192	///
193	/// Uses column equality condition (`where_column = $N`) with bind parameter
194	/// to prevent SQL injection.
195	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	/// Build parameterized SQL statements for batch update
210	///
211	/// Returns a list of `(sql, params)` tuples where `sql` contains `$N`
212	/// placeholders and `params` contains the corresponding bind values.
213	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	/// Get total number of updates
244	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		// Arrange
257		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		// Act
264		let sql_statements = builder.build_sql();
265
266		// Assert
267		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		// Arrange
278		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		// Act
287		let sql_statements = builder.build_sql();
288
289		// Assert - 5 rows with batch size 2 = 3 SQL statements (2 + 2 + 1)
290		assert_eq!(sql_statements.len(), 3);
291	}
292
293	#[rstest]
294	fn test_batch_update_builder_uses_parameterized_queries() {
295		// Arrange
296		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		// Act
309		let statements = builder.build_sql_parameterized();
310
311		// Assert
312		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		// Arrange - attempt SQL injection via where_value
325		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		// Act
332		let statements = builder.build_sql_parameterized();
333
334		// Assert - the malicious value is a bind parameter, not in the SQL string
335		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		// Arrange - attempt SQL injection via column value
344		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		// Act
351		let statements = builder.build_sql_parameterized();
352
353		// Assert - the malicious value is a bind parameter, not in the SQL string
354		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		// Arrange
363		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		// Act
373		let statements = builder.build_sql_parameterized();
374
375		// Assert
376		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		// Arrange
384		let builder = BatchInsertBuilder::new("users")
385			.columns(vec!["name".to_string()])
386			.add_row(vec!["Alice'; DROP TABLE users; --".to_string()]);
387
388		// Act
389		let sql_statements = builder.build_sql();
390
391		// Assert - single quotes should be escaped
392		assert!(sql_statements[0].contains("Alice''; DROP TABLE users; --"));
393	}
394
395	#[rstest]
396	fn test_batch_insert_quotes_table_name() {
397		// Arrange - table name with double quote injection attempt
398		let builder = BatchInsertBuilder::new("users\"; DROP TABLE data; --")
399			.columns(vec!["name".to_string()])
400			.add_row(vec!["Alice".to_string()]);
401
402		// Act
403		let sql_statements = builder.build_sql();
404
405		// Assert - table name must be properly quoted and escaped
406		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		// Arrange - column name with double quote injection attempt
412		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		// Act
420		let sql_statements = builder.build_sql();
421
422		// Assert - column names must be properly quoted
423		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		// Arrange - use backtick quoting for MySQL
430		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		// Act
436		let sql_statements = builder.build_sql();
437
438		// Assert - identifiers should use backticks, not double quotes
439		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		// Arrange - backtick in identifier should be escaped by doubling
448		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		// Act
454		let sql_statements = builder.build_sql();
455
456		// Assert - embedded backticks must be doubled
457		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		// Arrange & Act & Assert
464		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		// Arrange - use DatabaseType to select quoting style
472		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		// Act
478		let sql_statements = builder.build_sql();
479
480		// Assert - MySQL should use backticks
481		assert!(sql_statements[0].contains("INSERT INTO `users`"));
482		assert!(sql_statements[0].contains("`name`"));
483	}
484}