Skip to main content

mysql_common/
named_params.rs

1// Copyright (c) 2017 Anatoly Ikorsky
2//
3// Licensed under the Apache License, Version 2.0
4// <LICENSE-APACHE or http://www.apache.org/licenses/LICENSE-2.0> or the MIT
5// license <LICENSE-MIT or http://opensource.org/licenses/MIT>, at your
6// option. All files in the project carrying such notice may not be copied,
7// modified, or distributed except according to those terms.
8
9use std::borrow::Cow;
10
11/// Appears if a statement have both named and positional parameters.
12#[derive(Debug, Clone, Copy, Eq, PartialEq)]
13pub struct MixedParamsError;
14
15enum ParserState {
16    TopLevel,
17    // (string_delimiter, in_escape)
18    InStringLiteral(u8, bool),
19    MaybeInNamedParam,
20    InNamedParam,
21    InSharpComment,
22    MaybeInDoubleDashComment1,
23    MaybeInDoubleDashComment2,
24    InDoubleDashComment,
25    MaybeInCComment1,
26    MaybeInCComment2,
27    InCComment,
28    MaybeExitCComment,
29    InQuotedIdentifier,
30}
31
32use self::ParserState::*;
33
34/// Parsed named params (see [`ParsedNamedParams::parse`]).
35#[derive(Debug, Clone, Eq, PartialEq)]
36pub struct ParsedNamedParams<'a> {
37    query: Cow<'a, [u8]>,
38    params: Vec<Cow<'a, [u8]>>,
39}
40
41impl<'a> ParsedNamedParams<'a> {
42    /// Parse named params in the given query.
43    ///
44    /// Parameters must be named according to the following convention:
45    ///
46    /// * parameter name must start with either `_` or `a..z`
47    /// * parameter name may continue with `_`, `a..z` and `0..9`
48    pub fn parse(query: &'a [u8]) -> Result<Self, MixedParamsError> {
49        let mut state = TopLevel;
50        let mut have_positional = false;
51        let mut cur_param = 0;
52        // Vec<(colon_offset, start_offset, end_offset)>
53        let mut params = Vec::new();
54        for (i, c) in query.iter().enumerate() {
55            let mut rematch = false;
56            match state {
57                TopLevel => match c {
58                    b':' => state = MaybeInNamedParam,
59                    b'/' => state = MaybeInCComment1,
60                    b'-' => state = MaybeInDoubleDashComment1,
61                    b'#' => state = InSharpComment,
62                    b'\'' => state = InStringLiteral(b'\'', false),
63                    b'"' => state = InStringLiteral(b'"', false),
64                    b'?' => have_positional = true,
65                    b'`' => state = InQuotedIdentifier,
66                    _ => (),
67                },
68                InStringLiteral(separator, in_escape) => match c {
69                    _ if in_escape => state = InStringLiteral(separator, false),
70                    x if *x == separator => state = TopLevel,
71                    x if *x == b'\\' => state = InStringLiteral(separator, true),
72                    _ => state = InStringLiteral(separator, false),
73                },
74                MaybeInNamedParam => match c {
75                    b'a'..=b'z' | b'_' => {
76                        params.push((i - 1, i, 0));
77                        state = InNamedParam;
78                    }
79                    _ => rematch = true,
80                },
81                InNamedParam => {
82                    if !matches!(c, b'a'..=b'z' | b'0'..=b'9' | b'_') {
83                        params[cur_param].2 = i;
84                        cur_param += 1;
85                        rematch = true;
86                    }
87                }
88                InSharpComment => {
89                    if *c == b'\n' {
90                        state = TopLevel
91                    }
92                }
93                MaybeInDoubleDashComment1 => match c {
94                    b'-' => state = MaybeInDoubleDashComment2,
95                    _ => state = TopLevel,
96                },
97                MaybeInDoubleDashComment2 => {
98                    if c.is_ascii_whitespace() && *c != b'\n' {
99                        state = InDoubleDashComment
100                    } else {
101                        state = TopLevel
102                    }
103                }
104                InDoubleDashComment => {
105                    if *c == b'\n' {
106                        state = TopLevel
107                    }
108                }
109                MaybeInCComment1 => match c {
110                    b'*' => state = MaybeInCComment2,
111                    _ => state = TopLevel,
112                },
113                MaybeInCComment2 => match c {
114                    b'!' | b'+' => state = TopLevel, // extensions and optimizer hints
115                    _ => state = InCComment,
116                },
117                InCComment => {
118                    if *c == b'*' {
119                        state = MaybeExitCComment
120                    }
121                }
122                MaybeExitCComment => match c {
123                    b'/' => state = TopLevel,
124                    _ => state = InCComment,
125                },
126                InQuotedIdentifier => {
127                    if *c == b'`' {
128                        state = TopLevel
129                    }
130                }
131            }
132            if rematch {
133                match c {
134                    b':' => state = MaybeInNamedParam,
135                    b'\'' => state = InStringLiteral(b'\'', false),
136                    b'"' => state = InStringLiteral(b'"', false),
137                    _ => state = TopLevel,
138                }
139            }
140        }
141
142        if let InNamedParam = state {
143            params[cur_param].2 = query.len();
144        }
145
146        if !params.is_empty() {
147            if have_positional {
148                return Err(MixedParamsError);
149            }
150            let mut real_query = Vec::with_capacity(query.len());
151            let mut last = 0;
152            let mut out_params = Vec::with_capacity(params.len());
153            for (colon_offset, start, end) in params {
154                real_query.extend(&query[last..colon_offset]);
155                real_query.push(b'?');
156                last = end;
157                out_params.push(Cow::Borrowed(&query[start..end]));
158            }
159            real_query.extend(&query[last..]);
160            Ok(Self {
161                query: Cow::Owned(real_query),
162                params: out_params,
163            })
164        } else {
165            Ok(Self {
166                query: Cow::Borrowed(query),
167                params: vec![],
168            })
169        }
170    }
171
172    /// Returns a query string to pass to MySql (named parameters have been replaced with `?`).
173    pub fn query(&self) -> &[u8] {
174        &self.query
175    }
176
177    /// Names of named parameters in order of appearance.
178    ///
179    /// # Note
180    ///
181    /// * the returned slice might be empty if original query contained
182    ///   no named parameters.
183    /// * same name may appear multiple times.
184    pub fn params(&self) -> &[Cow<'a, [u8]>] {
185        &self.params
186    }
187}
188
189#[cfg(test)]
190mod test {
191    use super::*;
192
193    macro_rules! cows {
194        ($($l:expr_2021),+ $(,)?) => { &[$(Cow::Borrowed(&$l[..]),)*] };
195    }
196
197    #[test]
198    fn should_parse_named_params() {
199        let result = ParsedNamedParams::parse(b":a :b").unwrap();
200        assert_eq!(result.query(), b"? ?");
201        assert_eq!(result.params(), cows!(b"a", b"b"));
202
203        let result = ParsedNamedParams::parse(b"SELECT (:a-10)").unwrap();
204        assert_eq!(result.query(), b"SELECT (?-10)");
205        assert_eq!(result.params(), cows!(b"a"));
206
207        let result = ParsedNamedParams::parse(br#"SELECT '"\':a' "'\"':c" :b"#).unwrap();
208        assert_eq!(result.query(), br#"SELECT '"\':a' "'\"':c" ?"#);
209        assert_eq!(result.params(), cows!(b"b"));
210
211        let result = ParsedNamedParams::parse(br":a_Aa:b").unwrap();
212        assert_eq!(result.query(), b"?Aa?");
213        assert_eq!(result.params(), cows!(b"a_", b"b"));
214
215        let result = ParsedNamedParams::parse(br"::b").unwrap();
216        assert_eq!(result.query(), b":?");
217        assert_eq!(result.params(), cows!(b"b"));
218
219        ParsedNamedParams::parse(b":a ?").unwrap_err();
220    }
221
222    #[test]
223    fn should_allow_numbers_in_param_name() {
224        let result = ParsedNamedParams::parse(b":a1 :a2").unwrap();
225        assert_eq!(result.query(), b"? ?");
226        assert_eq!(result.params(), cows!(b"a1", b"a2"));
227
228        let result = ParsedNamedParams::parse(b":1a :2a").unwrap();
229        assert_eq!(result.query(), b":1a :2a");
230        assert!(result.params().is_empty());
231    }
232
233    #[test]
234    fn special_characters_in_query() {
235        let result =
236            ParsedNamedParams::parse("SELECT 1 FROM été WHERE thing = :param;".as_bytes()).unwrap();
237        assert_eq!(
238            result.query(),
239            "SELECT 1 FROM été WHERE thing = ?;".as_bytes()
240        );
241        assert_eq!(result.params(), cows!(b"param"));
242    }
243
244    #[test]
245    fn comments_with_question_marks() {
246        let result = ParsedNamedParams::parse(
247            "SELECT 1 FROM my_table WHERE thing = :param;/* question\n  mark '?' in multiline\n\
248            comment? */\n# ??- sharp comment -??\n-- dash-dash?\n/*! extention param :param2 */\n\
249            /*+ optimizer hint :param3 */; select :foo; # another comment?"
250                .as_bytes(),
251        )
252        .unwrap();
253        assert_eq!(
254            result.query(),
255            b"SELECT 1 FROM my_table WHERE thing = ?;/* question\n  mark '?' in multiline\n\
256        comment? */\n# ??- sharp comment -??\n-- dash-dash?\n/*! extention param ? */\n\
257        /*+ optimizer hint ? */; select ?; # another comment?"
258        );
259        assert_eq!(
260            result.params(),
261            cows!(b"param", b"param2", b"param3", b"foo"),
262        );
263    }
264
265    #[test]
266    fn quoted_identifier() {
267        let result = ParsedNamedParams::parse(b"INSERT INTO `my:table` VALUES (?)").unwrap();
268        assert_eq!(result.query(), b"INSERT INTO `my:table` VALUES (?)");
269        assert!(result.params().is_empty());
270
271        let result = ParsedNamedParams::parse(b"INSERT INTO `my:table` VALUES (:foo)").unwrap();
272        assert_eq!(result.query(), b"INSERT INTO `my:table` VALUES (?)");
273        assert_eq!(result.params(), cows!(b"foo"));
274    }
275
276    #[test]
277    fn issue_179_string_literal_extended_due_to_improper_escape_sequence_handling() {
278        let result =
279            ParsedNamedParams::parse(br"INSERT INTO `my:table` VALUES (:one, ':two\\', :three)")
280                .unwrap();
281        assert_eq!(
282            result.query(),
283            br"INSERT INTO `my:table` VALUES (?, ':two\\', ?)"
284        );
285        assert_eq!(
286            result.params(),
287            &[Cow::Borrowed(&b"one"[..]), Cow::Borrowed(b"three")]
288        );
289    }
290
291    #[cfg(feature = "nightly")]
292    mod bench {
293        use super::*;
294
295        #[bench]
296        fn parse_ten_named_params(bencher: &mut test::Bencher) {
297            bencher.iter(|| {
298                let result = ParsedNamedParams::parse(
299                    r#"
300                SELECT :one, :two, :three, :four, :five, :six, :seven, :eight, :nine, :ten
301                "#,
302                )
303                .unwrap();
304                test::black_box(result);
305            });
306        }
307
308        #[bench]
309        fn parse_zero_named_params(bencher: &mut test::Bencher) {
310            bencher.iter(|| {
311                let result = ParsedNamedParams::parse(
312                    r"
313                SELECT one, two, three, four, five, six, seven, eight, nine, ten
314                ",
315                )
316                .unwrap();
317                test::black_box(result);
318            });
319        }
320    }
321}