use std::fmt::Write as _;
use crate::bind::BoundExpr;
use crate::keyword::Keyword;
use crate::lexer::{Lexer, Punctuator, Span, Token, TokenKind};
const MOST_VALUES: usize = 64;
const LONGEST_TEXT: usize = 1 << 16;
#[derive(Clone, Debug, PartialEq)]
pub enum LiftedValue {
Integer(i64),
Real(f64),
Text(Vec<u8>),
}
#[derive(Clone, Debug, PartialEq)]
pub struct Lifted {
pub text: String,
pub values: Vec<LiftedValue>,
pub row_width: Option<usize>,
}
pub fn lift_insert_literals(sql: &str) -> Option<Lifted> {
let source = sql.as_bytes();
let mut lexer = Lexer::new(source);
let first = lexer.next_token().ok()?;
if !matches!(
first.keyword(),
Some(Keyword::INSERT) | Some(Keyword::REPLACE)
) {
return None;
}
let values_end = skip_to_values(&mut lexer)?;
let mut rows = Vec::with_capacity(source.len().saturating_sub(values_end) / 3 + 8);
loop {
let token = lexer.next_token().ok()?;
match token.kind {
TokenKind::EndOfInput => break,
TokenKind::Parameter => return None,
_ => rows.push(token),
}
}
let (lifts, widths) = lone_literals(source, &rows)?;
if let Some(width) = one_width_of_lone_literals(&widths) {
return Some(template_of(sql, values_end, width, lifts));
}
if lifts.is_empty() || lifts.len() > MOST_VALUES {
return None;
}
Some(rewrite(sql, lifts))
}
fn one_width_of_lone_literals(widths: &[(usize, usize)]) -> Option<usize> {
let (width, _) = *widths.first()?;
let shared = widths.len() >= 2
&& (1..=MOST_VALUES).contains(&width)
&& widths
.iter()
.all(|(elements, lifted)| *elements == width && *lifted == width);
shared.then_some(width)
}
fn template_of(
sql: &str,
values_end: usize,
width: usize,
lifts: Vec<(Span, LiftedValue)>,
) -> Lifted {
let mut text = String::with_capacity(values_end.saturating_add(width.saturating_mul(5)));
text.push_str(sql.get(..values_end).unwrap_or(""));
text.push_str(" (");
for number in 1..=width {
if number > 1 {
text.push_str(", ");
}
let _ = write!(text, "?{number}");
}
text.push(')');
Lifted {
text,
values: lifts.into_iter().map(|(_, value)| value).collect(),
row_width: Some(width),
}
}
fn skip_to_values(lexer: &mut Lexer<'_>) -> Option<usize> {
let mut depth = 0usize;
loop {
let token = lexer.next_token().ok()?;
match token.kind {
TokenKind::EndOfInput | TokenKind::Parameter => return None,
TokenKind::Punctuator(Punctuator::LeftParen) => depth = depth.saturating_add(1),
TokenKind::Punctuator(Punctuator::RightParen) => depth = depth.checked_sub(1)?,
_ if depth == 0 && token.keyword() == Some(Keyword::VALUES) => {
return Some(token.span.end as usize)
}
_ => {}
}
}
}
fn lone_literals(source: &[u8], tokens: &[Token]) -> Option<LoneLiterals> {
let mut lifts = Vec::with_capacity(tokens.len() / 2 + 1);
let mut widths = Vec::with_capacity(tokens.len() / 3 + 1);
let mut at = 0usize;
loop {
if !tokens.get(at)?.is(Punctuator::LeftParen) {
return None;
}
let before = lifts.len();
let (next, elements) = row_elements(source, tokens, at.saturating_add(1), &mut lifts)?;
widths.push((elements, lifts.len().saturating_sub(before)));
at = next;
match tokens.get(at) {
None => return Some((lifts, widths)),
Some(token) if token.is(Punctuator::Comma) => at = at.saturating_add(1),
Some(token)
if token.is(Punctuator::Semicolon) && at.saturating_add(1) == tokens.len() =>
{
return Some((lifts, widths));
}
Some(_) => return None,
}
}
}
type LoneLiterals = (Vec<(Span, LiftedValue)>, Vec<(usize, usize)>);
fn row_elements(
source: &[u8],
tokens: &[Token],
start: usize,
lifts: &mut Vec<(Span, LiftedValue)>,
) -> Option<(usize, usize)> {
let mut depth = 1usize;
let mut element = start;
let mut at = start;
let mut elements = 0usize;
while depth > 0 {
let token = *tokens.get(at)?;
if token.is(Punctuator::LeftParen) {
depth = depth.saturating_add(1);
} else if token.is(Punctuator::RightParen) {
depth = depth.saturating_sub(1);
}
if depth == 0 || (depth == 1 && token.is(Punctuator::Comma)) {
if let Some(lifted) = lone_literal(source, tokens.get(element..at)?)? {
lifts.push(lifted);
}
elements = elements.saturating_add(1);
element = at.saturating_add(1);
}
at = at.saturating_add(1);
}
Some((at, elements))
}
fn lone_literal(source: &[u8], element: &[Token]) -> Option<Option<(Span, LiftedValue)>> {
let (negated, token) = match element {
[token] => (false, *token),
[minus, token] if minus.is(Punctuator::Minus) => (true, *token),
_ => return Some(None),
};
let span = match negated {
true => element.first()?.span.to(token.span),
false => token.span,
};
let value = match token.kind {
TokenKind::Integer => integer_value(token.text(source), negated, token.span)?,
TokenKind::Float => {
let value = crate::bind::literal::real_literal(token.text(source));
LiftedValue::Real(if negated { -value } else { value })
}
TokenKind::String if !negated => {
let text = crate::lexer::string_text(source, token);
if text.len() > LONGEST_TEXT {
return Some(None);
}
LiftedValue::Text(text.into_owned())
}
_ => return Some(None),
};
Some(Some((span, value)))
}
fn integer_value(digits: &[u8], negated: bool, span: Span) -> Option<LiftedValue> {
let mut text = Vec::with_capacity(digits.len().saturating_add(1));
if negated {
text.push(b'-');
}
text.extend_from_slice(digits);
match crate::bind::literal::checked_integer_literal(&text, span).ok()? {
BoundExpr::Integer(value) => Some(LiftedValue::Integer(value)),
BoundExpr::Real(value) => Some(LiftedValue::Real(value)),
_ => None,
}
}
fn rewrite(sql: &str, lifts: Vec<(Span, LiftedValue)>) -> Lifted {
let mut text = String::with_capacity(sql.len());
let mut values = Vec::with_capacity(lifts.len());
let mut copied = 0usize;
for (number, (span, value)) in lifts.into_iter().enumerate() {
text.push_str(sql.get(copied..span.start as usize).unwrap_or(""));
let _ = write!(text, "?{}", number.saturating_add(1));
copied = span.end as usize;
values.push(value);
}
text.push_str(sql.get(copied..).unwrap_or(""));
Lifted {
text,
values,
row_width: None,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_row_of_lone_literals_becomes_parameters() {
let lifted = lift_insert_literals("INSERT INTO t (a, b, c) VALUES ('it''s', -5, 2.5);")
.expect("lifted");
assert_eq!(lifted.text, "INSERT INTO t (a, b, c) VALUES (?1, ?2, ?3);");
assert_eq!(
lifted.values,
vec![
LiftedValue::Text(b"it's".to_vec()),
LiftedValue::Integer(-5),
LiftedValue::Real(2.5)
]
);
}
#[test]
fn literals_inside_expressions_stay_in_the_text() {
let lifted =
lift_insert_literals("INSERT INTO t VALUES (upper('a'), 1 + 2, (SELECT 3), 4)")
.expect("lifted");
assert_eq!(
lifted.text,
"INSERT INTO t VALUES (upper('a'), 1 + 2, (SELECT 3), ?1)"
);
assert_eq!(lifted.values, vec![LiftedValue::Integer(4)]);
}
#[test]
fn the_extremes_convert_as_the_binder_converts_them() {
let lifted = lift_insert_literals(
"INSERT INTO t VALUES (-9223372036854775808, 9223372036854775808, 0x10, -0.0)",
)
.expect("lifted");
assert_eq!(lifted.values.first(), Some(&LiftedValue::Integer(i64::MIN)));
assert_eq!(
lifted.values.get(1),
Some(&LiftedValue::Real(9_223_372_036_854_775_808.0))
);
assert_eq!(lifted.values.get(2), Some(&LiftedValue::Integer(16)));
match lifted.values.get(3) {
Some(LiftedValue::Real(value)) => assert!(value.is_sign_negative()),
other => panic!("expected a negative zero, got {other:?}"),
}
}
#[test]
fn statements_it_cannot_rewrite_safely_are_left_alone() {
for sql in [
"SELECT 1",
"INSERT INTO t VALUES (?1, 2)",
"INSERT INTO t VALUES (1) ON CONFLICT DO NOTHING",
"INSERT INTO t VALUES (1) RETURNING a",
"INSERT INTO t DEFAULT VALUES",
"INSERT INTO t SELECT 1",
"INSERT INTO t VALUES (upper('a'))",
"INSERT INTO t VALUES (0x10000000000000000)",
"WITH c AS (SELECT 1) INSERT INTO t VALUES (1)",
] {
assert_eq!(lift_insert_literals(sql), None, "{sql}");
}
}
#[test]
fn several_rows_are_numbered_in_order() {
let lifted =
lift_insert_literals("REPLACE INTO t VALUES (1, 'a'), (2, NULL)").expect("lifted");
assert_eq!(lifted.text, "REPLACE INTO t VALUES (?1, ?2), (?3, NULL)");
assert_eq!(lifted.values.len(), 3);
assert_eq!(lifted.row_width, None);
}
#[test]
fn rows_of_lone_literals_become_one_row_template() {
let lifted =
lift_insert_literals("INSERT INTO t (a, b) VALUES (1, 'a'), (-2, 'b'), (3.5, 'c');")
.expect("lifted");
assert_eq!(lifted.text, "INSERT INTO t (a, b) VALUES (?1, ?2)");
assert_eq!(lifted.row_width, Some(2));
assert_eq!(
lifted.values,
vec![
LiftedValue::Integer(1),
LiftedValue::Text(b"a".to_vec()),
LiftedValue::Integer(-2),
LiftedValue::Text(b"b".to_vec()),
LiftedValue::Real(3.5),
LiftedValue::Text(b"c".to_vec()),
]
);
let uneven = lift_insert_literals("INSERT INTO t VALUES (1, 2), (3)").expect("lifted");
assert_eq!(uneven.row_width, None);
}
}