Skip to main content

inillucent_sql/
lift.rs

1//! Turning the literal values of an `INSERT ... VALUES` into parameters, so
2//! statements that differ only in those values share one compiled plan.
3//!
4//! Invariant: **a statement is rewritten only where a literal and a bound
5//! parameter mean the same thing.** That is a lone literal that is a whole
6//! element of a `VALUES` row: a number, a negated number or a quoted string.
7//! A literal inside an expression, a function call or a subquery stays where it
8//! is, because there the binder may read the literal itself. Anything this
9//! does not recognise answers `None`, and the statement is compiled as written.
10//!
11//! **Why** (task-2191). A script of single row inserts, which is what `.dump`
12//! writes and what a person pastes into the shell, compiled every statement
13//! from scratch. The compile was a third of each insert, and SQLite pays the
14//! same. With the values lifted out, the second statement onwards is a lexer
15//! pass and a lookup in the plan cache.
16
17use std::fmt::Write as _;
18
19use crate::bind::BoundExpr;
20use crate::keyword::Keyword;
21use crate::lexer::{Lexer, Punctuator, Span, Token, TokenKind};
22
23/// The most values one statement may have lifted.
24///
25/// A multi row `INSERT` of thousands of rows is compiled once anyway, and
26/// keying the plan cache by a text of thousands of parameters would hold a plan
27/// nothing asks for twice.
28const MOST_VALUES: usize = 64;
29
30/// The longest string literal that is lifted, in bytes.
31///
32/// The parser checks a literal's length against the engine's limit while it
33/// reads it, and a bound value is not checked the same way. Leaving long
34/// strings in the text keeps that check where it was.
35const LONGEST_TEXT: usize = 1 << 16;
36
37/// One value taken out of a statement's text.
38#[derive(Clone, Debug, PartialEq)]
39pub enum LiftedValue {
40    /// An integer literal, with its sign.
41    Integer(i64),
42    /// A real literal, or an integer literal too large for 64 bits.
43    Real(f64),
44    /// A string literal, with doubled quotes undoubled.
45    Text(Vec<u8>),
46}
47
48/// A statement with its literal values replaced by `?1`, `?2` and so on.
49#[derive(Clone, Debug, PartialEq)]
50pub struct Lifted {
51    /// The rewritten text.
52    pub text: String,
53    /// The values, `?1` first.
54    pub values: Vec<LiftedValue>,
55    /// How many values make one row, when the statement had several rows of
56    /// lone literals and `text` is the one row template `VALUES (?1, ...)`.
57    ///
58    /// **Several rows as one template** (task-2191). An `INSERT` of twenty
59    /// thousand rows built a syntax tree node, a bound expression and a
60    /// compiled expression for every literal. With every element of every row
61    /// a lone literal and every row the same width, the rows are values and
62    /// the statement is the one row form run over all of them; `values` then
63    /// holds every row in order. `None` means `values` binds `text` once.
64    pub row_width: Option<usize>,
65}
66
67/// Rewrites an `INSERT ... VALUES` whose rows hold lone literals, or answers
68/// `None` when the statement is anything else.
69///
70/// The statement must hold no parameter of its own, must end after its last
71/// row (an upsert or a `RETURNING` clause answers `None`), and must lift at
72/// least one value.
73///
74/// @param sql - one statement
75pub 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    // Sized from the text, at about one token for every three bytes, so a
87    // statement of twenty thousand rows does not grow it a dozen times over
88    // (task-2191); a one row statement still asks for a few dozen.
89    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
108/// Returns the width every row shares when there are several rows and every
109/// element of every row was lifted, and `None` otherwise.
110///
111/// @param widths - each row's element count and how many of them were lifted
112fn 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
122/// Builds the one row template `... VALUES (?1, ..., ?width)` and keeps every
123/// row's values, in order.
124///
125/// @param sql - the statement
126/// @param values_end - where the `VALUES` keyword ends
127/// @param width - how many values make one row
128/// @param lifts - every lifted value, row by row
129fn 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
152/// Moves the lexer past the `VALUES` keyword of the statement's own clause.
153///
154/// @param lexer - the lexer, just past `INSERT` or `REPLACE`
155fn 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
171/// Finds every lone literal element of the rows after `VALUES`.
172///
173/// Answers `None` when the tokens are not rows followed by an optional
174/// semicolon, or when a literal cannot be converted the way the binder would.
175///
176/// @param source - the statement's bytes
177/// @param tokens - the tokens after `VALUES`
178fn lone_literals(source: &[u8], tokens: &[Token]) -> Option<LoneLiterals> {
179    // A lifted value takes at least two tokens, itself and a comma or the
180    // closing parenthesis, and a row at least three.
181    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
205/// What [`lone_literals`] found: every lifted value with where it was, and
206/// for each row how many elements it had and how many of them were lifted.
207type LoneLiterals = (Vec<(Span, LiftedValue)>, Vec<(usize, usize)>);
208
209/// Reads one row's elements and lifts the lone literals among them.
210///
211/// Returns the position just past the row's closing parenthesis.
212///
213/// @param source - the statement's bytes
214/// @param tokens - the tokens after `VALUES`
215/// @param start - the position just past the row's opening parenthesis
216/// @param lifts - where lifted values are added
217fn 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
246/// Converts one row element when it is a lone literal.
247///
248/// Answers `Some(None)` for an element that is not one, which stays in the
249/// text, and `None` for a literal the binder would refuse, so the statement is
250/// compiled as written and reports the refusal itself.
251///
252/// @param source - the statement's bytes
253/// @param element - the element's tokens
254fn 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
282/// Converts an integer literal the way the binder does, sign folded in.
283///
284/// @param digits - the literal as written, without a sign
285/// @param negated - whether a minus came before it
286/// @param span - where the literal was written
287fn 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
300/// Builds the rewritten text, each lifted span replaced by its parameter.
301///
302/// @param sql - the statement
303/// @param lifts - the spans and values, in the order they appear
304fn 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        // Written into the text directly; `to_string` allocated a string per
311        // value only to copy it in.
312        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        // Rows of different widths are the parser's to refuse, so they are
418        // not made into a template.
419        let uneven = lift_insert_literals("INSERT INTO t VALUES (1, 2), (3)").expect("lifted");
420        assert_eq!(uneven.row_width, None);
421    }
422}