1use std::fmt::Write as _;
18
19use crate::bind::BoundExpr;
20use crate::keyword::Keyword;
21use crate::lexer::{Lexer, Punctuator, Span, Token, TokenKind};
22
23const MOST_VALUES: usize = 64;
29
30const LONGEST_TEXT: usize = 1 << 16;
36
37#[derive(Clone, Debug, PartialEq)]
39pub enum LiftedValue {
40 Integer(i64),
42 Real(f64),
44 Text(Vec<u8>),
46}
47
48#[derive(Clone, Debug, PartialEq)]
50pub struct Lifted {
51 pub text: String,
53 pub values: Vec<LiftedValue>,
55 pub row_width: Option<usize>,
65}
66
67pub fn lift_insert_literals(sql: &str) -> Option<Lifted> {
76 let source = sql.as_bytes();
77 let mut lexer = Lexer::new(source);
78 let first = lexer.next_token().ok()?;
79 if !matches!(
80 first.keyword(),
81 Some(Keyword::INSERT) | Some(Keyword::REPLACE)
82 ) {
83 return None;
84 }
85 let values_end = skip_to_values(&mut lexer)?;
86 let mut rows = Vec::with_capacity(source.len().saturating_sub(values_end) / 3 + 8);
90 loop {
91 let token = lexer.next_token().ok()?;
92 match token.kind {
93 TokenKind::EndOfInput => break,
94 TokenKind::Parameter => return None,
95 _ => rows.push(token),
96 }
97 }
98 let (lifts, widths) = lone_literals(source, &rows)?;
99 if let Some(width) = one_width_of_lone_literals(&widths) {
100 return Some(template_of(sql, values_end, width, lifts));
101 }
102 if lifts.is_empty() || lifts.len() > MOST_VALUES {
103 return None;
104 }
105 Some(rewrite(sql, lifts))
106}
107
108fn one_width_of_lone_literals(widths: &[(usize, usize)]) -> Option<usize> {
113 let (width, _) = *widths.first()?;
114 let shared = widths.len() >= 2
115 && (1..=MOST_VALUES).contains(&width)
116 && widths
117 .iter()
118 .all(|(elements, lifted)| *elements == width && *lifted == width);
119 shared.then_some(width)
120}
121
122fn template_of(
130 sql: &str,
131 values_end: usize,
132 width: usize,
133 lifts: Vec<(Span, LiftedValue)>,
134) -> Lifted {
135 let mut text = String::with_capacity(values_end.saturating_add(width.saturating_mul(5)));
136 text.push_str(sql.get(..values_end).unwrap_or(""));
137 text.push_str(" (");
138 for number in 1..=width {
139 if number > 1 {
140 text.push_str(", ");
141 }
142 let _ = write!(text, "?{number}");
143 }
144 text.push(')');
145 Lifted {
146 text,
147 values: lifts.into_iter().map(|(_, value)| value).collect(),
148 row_width: Some(width),
149 }
150}
151
152fn skip_to_values(lexer: &mut Lexer<'_>) -> Option<usize> {
156 let mut depth = 0usize;
157 loop {
158 let token = lexer.next_token().ok()?;
159 match token.kind {
160 TokenKind::EndOfInput | TokenKind::Parameter => return None,
161 TokenKind::Punctuator(Punctuator::LeftParen) => depth = depth.saturating_add(1),
162 TokenKind::Punctuator(Punctuator::RightParen) => depth = depth.checked_sub(1)?,
163 _ if depth == 0 && token.keyword() == Some(Keyword::VALUES) => {
164 return Some(token.span.end as usize)
165 }
166 _ => {}
167 }
168 }
169}
170
171fn lone_literals(source: &[u8], tokens: &[Token]) -> Option<LoneLiterals> {
179 let mut lifts = Vec::with_capacity(tokens.len() / 2 + 1);
182 let mut widths = Vec::with_capacity(tokens.len() / 3 + 1);
183 let mut at = 0usize;
184 loop {
185 if !tokens.get(at)?.is(Punctuator::LeftParen) {
186 return None;
187 }
188 let before = lifts.len();
189 let (next, elements) = row_elements(source, tokens, at.saturating_add(1), &mut lifts)?;
190 widths.push((elements, lifts.len().saturating_sub(before)));
191 at = next;
192 match tokens.get(at) {
193 None => return Some((lifts, widths)),
194 Some(token) if token.is(Punctuator::Comma) => at = at.saturating_add(1),
195 Some(token)
196 if token.is(Punctuator::Semicolon) && at.saturating_add(1) == tokens.len() =>
197 {
198 return Some((lifts, widths));
199 }
200 Some(_) => return None,
201 }
202 }
203}
204
205type LoneLiterals = (Vec<(Span, LiftedValue)>, Vec<(usize, usize)>);
208
209fn row_elements(
218 source: &[u8],
219 tokens: &[Token],
220 start: usize,
221 lifts: &mut Vec<(Span, LiftedValue)>,
222) -> Option<(usize, usize)> {
223 let mut depth = 1usize;
224 let mut element = start;
225 let mut at = start;
226 let mut elements = 0usize;
227 while depth > 0 {
228 let token = *tokens.get(at)?;
229 if token.is(Punctuator::LeftParen) {
230 depth = depth.saturating_add(1);
231 } else if token.is(Punctuator::RightParen) {
232 depth = depth.saturating_sub(1);
233 }
234 if depth == 0 || (depth == 1 && token.is(Punctuator::Comma)) {
235 if let Some(lifted) = lone_literal(source, tokens.get(element..at)?)? {
236 lifts.push(lifted);
237 }
238 elements = elements.saturating_add(1);
239 element = at.saturating_add(1);
240 }
241 at = at.saturating_add(1);
242 }
243 Some((at, elements))
244}
245
246fn lone_literal(source: &[u8], element: &[Token]) -> Option<Option<(Span, LiftedValue)>> {
255 let (negated, token) = match element {
256 [token] => (false, *token),
257 [minus, token] if minus.is(Punctuator::Minus) => (true, *token),
258 _ => return Some(None),
259 };
260 let span = match negated {
261 true => element.first()?.span.to(token.span),
262 false => token.span,
263 };
264 let value = match token.kind {
265 TokenKind::Integer => integer_value(token.text(source), negated, token.span)?,
266 TokenKind::Float => {
267 let value = crate::bind::literal::real_literal(token.text(source));
268 LiftedValue::Real(if negated { -value } else { value })
269 }
270 TokenKind::String if !negated => {
271 let text = crate::lexer::string_text(source, token);
272 if text.len() > LONGEST_TEXT {
273 return Some(None);
274 }
275 LiftedValue::Text(text.into_owned())
276 }
277 _ => return Some(None),
278 };
279 Some(Some((span, value)))
280}
281
282fn integer_value(digits: &[u8], negated: bool, span: Span) -> Option<LiftedValue> {
288 let mut text = Vec::with_capacity(digits.len().saturating_add(1));
289 if negated {
290 text.push(b'-');
291 }
292 text.extend_from_slice(digits);
293 match crate::bind::literal::checked_integer_literal(&text, span).ok()? {
294 BoundExpr::Integer(value) => Some(LiftedValue::Integer(value)),
295 BoundExpr::Real(value) => Some(LiftedValue::Real(value)),
296 _ => None,
297 }
298}
299
300fn rewrite(sql: &str, lifts: Vec<(Span, LiftedValue)>) -> Lifted {
305 let mut text = String::with_capacity(sql.len());
306 let mut values = Vec::with_capacity(lifts.len());
307 let mut copied = 0usize;
308 for (number, (span, value)) in lifts.into_iter().enumerate() {
309 text.push_str(sql.get(copied..span.start as usize).unwrap_or(""));
310 let _ = write!(text, "?{}", number.saturating_add(1));
313 copied = span.end as usize;
314 values.push(value);
315 }
316 text.push_str(sql.get(copied..).unwrap_or(""));
317 Lifted {
318 text,
319 values,
320 row_width: None,
321 }
322}
323
324#[cfg(test)]
325mod tests {
326 use super::*;
327
328 #[test]
329 fn a_row_of_lone_literals_becomes_parameters() {
330 let lifted = lift_insert_literals("INSERT INTO t (a, b, c) VALUES ('it''s', -5, 2.5);")
331 .expect("lifted");
332 assert_eq!(lifted.text, "INSERT INTO t (a, b, c) VALUES (?1, ?2, ?3);");
333 assert_eq!(
334 lifted.values,
335 vec![
336 LiftedValue::Text(b"it's".to_vec()),
337 LiftedValue::Integer(-5),
338 LiftedValue::Real(2.5)
339 ]
340 );
341 }
342
343 #[test]
344 fn literals_inside_expressions_stay_in_the_text() {
345 let lifted =
346 lift_insert_literals("INSERT INTO t VALUES (upper('a'), 1 + 2, (SELECT 3), 4)")
347 .expect("lifted");
348 assert_eq!(
349 lifted.text,
350 "INSERT INTO t VALUES (upper('a'), 1 + 2, (SELECT 3), ?1)"
351 );
352 assert_eq!(lifted.values, vec![LiftedValue::Integer(4)]);
353 }
354
355 #[test]
356 fn the_extremes_convert_as_the_binder_converts_them() {
357 let lifted = lift_insert_literals(
358 "INSERT INTO t VALUES (-9223372036854775808, 9223372036854775808, 0x10, -0.0)",
359 )
360 .expect("lifted");
361 assert_eq!(lifted.values.first(), Some(&LiftedValue::Integer(i64::MIN)));
362 assert_eq!(
363 lifted.values.get(1),
364 Some(&LiftedValue::Real(9_223_372_036_854_775_808.0))
365 );
366 assert_eq!(lifted.values.get(2), Some(&LiftedValue::Integer(16)));
367 match lifted.values.get(3) {
368 Some(LiftedValue::Real(value)) => assert!(value.is_sign_negative()),
369 other => panic!("expected a negative zero, got {other:?}"),
370 }
371 }
372
373 #[test]
374 fn statements_it_cannot_rewrite_safely_are_left_alone() {
375 for sql in [
376 "SELECT 1",
377 "INSERT INTO t VALUES (?1, 2)",
378 "INSERT INTO t VALUES (1) ON CONFLICT DO NOTHING",
379 "INSERT INTO t VALUES (1) RETURNING a",
380 "INSERT INTO t DEFAULT VALUES",
381 "INSERT INTO t SELECT 1",
382 "INSERT INTO t VALUES (upper('a'))",
383 "INSERT INTO t VALUES (0x10000000000000000)",
384 "WITH c AS (SELECT 1) INSERT INTO t VALUES (1)",
385 ] {
386 assert_eq!(lift_insert_literals(sql), None, "{sql}");
387 }
388 }
389
390 #[test]
391 fn several_rows_are_numbered_in_order() {
392 let lifted =
393 lift_insert_literals("REPLACE INTO t VALUES (1, 'a'), (2, NULL)").expect("lifted");
394 assert_eq!(lifted.text, "REPLACE INTO t VALUES (?1, ?2), (?3, NULL)");
395 assert_eq!(lifted.values.len(), 3);
396 assert_eq!(lifted.row_width, None);
397 }
398
399 #[test]
400 fn rows_of_lone_literals_become_one_row_template() {
401 let lifted =
402 lift_insert_literals("INSERT INTO t (a, b) VALUES (1, 'a'), (-2, 'b'), (3.5, 'c');")
403 .expect("lifted");
404 assert_eq!(lifted.text, "INSERT INTO t (a, b) VALUES (?1, ?2)");
405 assert_eq!(lifted.row_width, Some(2));
406 assert_eq!(
407 lifted.values,
408 vec![
409 LiftedValue::Integer(1),
410 LiftedValue::Text(b"a".to_vec()),
411 LiftedValue::Integer(-2),
412 LiftedValue::Text(b"b".to_vec()),
413 LiftedValue::Real(3.5),
414 LiftedValue::Text(b"c".to_vec()),
415 ]
416 );
417 let uneven = lift_insert_literals("INSERT INTO t VALUES (1, 2), (3)").expect("lifted");
420 assert_eq!(uneven.row_width, None);
421 }
422}