Skip to main content

operator/statement/
set.rs

1// Copyright 2023 Greptime Team
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::str::FromStr;
16use std::time::Duration;
17
18use common_time::Timezone;
19use lazy_static::lazy_static;
20use regex::Regex;
21use session::ReadPreference;
22use session::context::Channel::Postgres;
23use session::context::QueryContextRef;
24use session::session_config::{PGByteaOutputValue, PGDateOrder, PGDateTimeStyle, PGIntervalStyle};
25use snafu::{OptionExt, ResultExt, ensure};
26use sql::ast::{Expr, Ident, Value};
27use sql::statements::set_variables::SetVariables;
28use sqlparser::ast::ValueWithSpan;
29
30use crate::error::{InvalidConfigValueSnafu, InvalidSqlSnafu, NotSupportedSnafu, Result};
31
32lazy_static! {
33    // Regex rules:
34    // The string must start with a number (one or more digits).
35    // The number must be followed by one of the valid time units (ms, s, min, h, d).
36    // The string must end immediately after the unit, meaning there can be no extra
37    // characters or spaces after the valid time specification.
38    static ref PG_TIME_INPUT_REGEX: Regex = Regex::new(r"^(\d+)(ms|s|min|h|d)$").unwrap();
39}
40
41/// Sets the session WAL policy for ordinary inserts without changing table options.
42pub fn set_skip_wal(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
43    let [Expr::Value(value)] = exprs.as_slice() else {
44        return NotSupportedSnafu {
45            feat: "SET skip_wal requires exactly one boolean value",
46        }
47        .fail();
48    };
49    let skip_wal = match &value.value {
50        Value::Boolean(value) => *value,
51        Value::SingleQuotedString(value) | Value::DoubleQuotedString(value) => {
52            value.parse::<bool>().map_err(|_| {
53                NotSupportedSnafu {
54                    feat: format!("Invalid skip_wal value {value:?}: expected true or false"),
55                }
56                .build()
57            })?
58        }
59        _ => {
60            return NotSupportedSnafu {
61                feat: "SET skip_wal requires true or false",
62            }
63            .fail();
64        }
65    };
66    ctx.set_skip_wal(skip_wal);
67    Ok(())
68}
69
70pub fn set_read_preference(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
71    let read_preference_expr = exprs.first().context(NotSupportedSnafu {
72        feat: "No read preference find in set variable statement",
73    })?;
74
75    match read_preference_expr {
76        Expr::Value(ValueWithSpan {
77            value: Value::SingleQuotedString(expr),
78            ..
79        })
80        | Expr::Value(ValueWithSpan {
81            value: Value::DoubleQuotedString(expr),
82            ..
83        }) => {
84            match ReadPreference::from_str(expr.as_str().to_lowercase().as_str()) {
85                Ok(read_preference) => ctx.set_read_preference(read_preference),
86                Err(_) => {
87                    return NotSupportedSnafu {
88                        feat: format!(
89                            "Invalid read preference expr {} in set variable statement",
90                            expr,
91                        ),
92                    }
93                    .fail();
94                }
95            }
96            Ok(())
97        }
98        expr => NotSupportedSnafu {
99            feat: format!(
100                "Unsupported read preference expr {} in set variable statement",
101                expr
102            ),
103        }
104        .fail(),
105    }
106}
107
108pub fn set_timezone(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
109    let tz_expr = exprs.first().context(NotSupportedSnafu {
110        feat: "No timezone find in set variable statement",
111    })?;
112    match tz_expr {
113        Expr::Value(ValueWithSpan {
114            value: Value::SingleQuotedString(tz),
115            ..
116        })
117        | Expr::Value(ValueWithSpan {
118            value: Value::DoubleQuotedString(tz),
119            ..
120        }) => {
121            match Timezone::from_tz_string(tz.as_str()) {
122                Ok(timezone) => ctx.set_timezone(timezone),
123                Err(_) => {
124                    return NotSupportedSnafu {
125                        feat: format!("Invalid timezone expr {} in set variable statement", tz),
126                    }
127                    .fail();
128                }
129            }
130            Ok(())
131        }
132        expr => NotSupportedSnafu {
133            feat: format!(
134                "Unsupported timezone expr {} in set variable statement",
135                expr
136            ),
137        }
138        .fail(),
139    }
140}
141
142pub fn set_bytea_output(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
143    let Some((var_value, [])) = exprs.split_first() else {
144        return (NotSupportedSnafu {
145            feat: "Set variable value must have one and only one value for bytea_output",
146        })
147        .fail();
148    };
149    let Expr::Value(value) = var_value else {
150        return (NotSupportedSnafu {
151            feat: "Set variable value must be a value",
152        })
153        .fail();
154    };
155    ctx.configuration_parameter().set_postgres_bytea_output(
156        PGByteaOutputValue::try_from(value.value.clone()).context(InvalidConfigValueSnafu)?,
157    );
158    Ok(())
159}
160
161pub fn validate_client_encoding(set: SetVariables) -> Result<()> {
162    let Some((encoding, [])) = set.value.split_first() else {
163        return InvalidSqlSnafu {
164            err_msg: "must provide one and only one client encoding value",
165        }
166        .fail();
167    };
168    let encoding = match encoding {
169        Expr::Value(ValueWithSpan {
170            value: Value::SingleQuotedString(x),
171            ..
172        })
173        | Expr::Identifier(Ident {
174            value: x,
175            quote_style: _,
176            span: _,
177        }) => x.to_uppercase(),
178        _ => {
179            return InvalidSqlSnafu {
180                err_msg: format!("client encoding must be a string, actual: {:?}", encoding),
181            }
182            .fail();
183        }
184    };
185    // For the sake of simplicity, we only support "UTF8" ("UNICODE" is the alias for it,
186    // see https://www.postgresql.org/docs/current/multibyte.html#MULTIBYTE-CHARSET-SUPPORTED).
187    // "UTF8" is universal and sufficient for almost all cases.
188    // GreptimeDB itself is always using "UTF8" as the internal encoding.
189    ensure!(
190        encoding == "UTF8" || encoding == "UNICODE",
191        NotSupportedSnafu {
192            feat: format!("client encoding of '{}'", encoding)
193        }
194    );
195    Ok(())
196}
197
198// if one of original value and new value is none, return the other one
199// returns new values only when it equals to original one else return error.
200// This is only used for handling datestyle
201fn merge_datestyle_value<T>(value: Option<T>, new_value: Option<T>) -> Result<Option<T>>
202where
203    T: PartialEq,
204{
205    match (&value, &new_value) {
206        (None, _) => Ok(new_value),
207        (_, None) => Ok(value),
208        (Some(v1), Some(v2)) if v1 == v2 => Ok(new_value),
209        _ => InvalidSqlSnafu {
210            err_msg: "Conflicting \"datestyle\" specifications.",
211        }
212        .fail(),
213    }
214}
215
216fn try_parse_datestyle(expr: &Expr) -> Result<(Option<PGDateTimeStyle>, Option<PGDateOrder>)> {
217    enum ParsedDateStyle {
218        Order(PGDateOrder),
219        Style(PGDateTimeStyle),
220    }
221    fn try_parse_str(s: &str) -> Result<ParsedDateStyle> {
222        PGDateTimeStyle::try_from(s)
223            .map_or_else(
224                |_| PGDateOrder::try_from(s).map(ParsedDateStyle::Order),
225                |style| Ok(ParsedDateStyle::Style(style)),
226            )
227            .context(InvalidConfigValueSnafu)
228    }
229    match expr {
230        Expr::Identifier(Ident {
231            value: s,
232            quote_style: _,
233            span: _,
234        })
235        | Expr::Value(ValueWithSpan {
236            value: Value::SingleQuotedString(s),
237            ..
238        })
239        | Expr::Value(ValueWithSpan {
240            value: Value::DoubleQuotedString(s),
241            ..
242        }) => s
243            .split(',')
244            .map(|s| s.trim())
245            .try_fold((None, None), |(style, order), s| match try_parse_str(s)? {
246                ParsedDateStyle::Order(o) => Ok((style, merge_datestyle_value(order, Some(o))?)),
247                ParsedDateStyle::Style(s) => Ok((merge_datestyle_value(style, Some(s))?, order)),
248            }),
249        _ => NotSupportedSnafu {
250            feat: "Not supported expression for datestyle",
251        }
252        .fail(),
253    }
254}
255
256/// Set the allow query fallback configuration parameter to true or false based on the provided expressions.
257///
258pub fn set_allow_query_fallback(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
259    let allow_fallback_expr = exprs.first().context(NotSupportedSnafu {
260        feat: "No allow query fallback value find in set variable statement",
261    })?;
262    match allow_fallback_expr {
263        Expr::Value(ValueWithSpan {
264            value: Value::Boolean(allow),
265            span: _,
266        }) => {
267            ctx.configuration_parameter()
268                .set_allow_query_fallback(*allow);
269            Ok(())
270        }
271        expr => NotSupportedSnafu {
272            feat: format!(
273                "Unsupported allow query fallback expr {} in set variable statement",
274                expr
275            ),
276        }
277        .fail(),
278    }
279}
280
281pub fn set_intervalstyle(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
282    let Some((var_value, [])) = exprs.split_first() else {
283        return NotSupportedSnafu {
284            feat: "Set variable value must have one and only one value for intervalstyle",
285        }
286        .fail();
287    };
288    let intervalstyle = match var_value {
289        Expr::Identifier(Ident {
290            value,
291            quote_style: _,
292            span: _,
293        }) => PGIntervalStyle::try_from(value.as_str()).context(InvalidConfigValueSnafu)?,
294        Expr::Value(value) => {
295            PGIntervalStyle::try_from(&value.value).context(InvalidConfigValueSnafu)?
296        }
297        _ => {
298            return NotSupportedSnafu {
299                feat: "Set variable value must be a value or identifier",
300            }
301            .fail();
302        }
303    };
304    ctx.configuration_parameter()
305        .set_pg_intervalstyle_format(intervalstyle);
306    Ok(())
307}
308
309pub fn set_datestyle(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
310    // ORDER,
311    // STYLE,
312    // ORDER,ORDER
313    // ORDER,STYLE
314    // STYLE,ORDER
315    let (style, order) = exprs
316        .iter()
317        .try_fold((None, None), |(style, order), expr| {
318            let (new_style, new_order) = try_parse_datestyle(expr)?;
319            Ok((
320                merge_datestyle_value(style, new_style)?,
321                merge_datestyle_value(order, new_order)?,
322            ))
323        })?;
324
325    let (old_style, older_order) = *ctx.configuration_parameter().pg_datetime_style();
326    ctx.configuration_parameter()
327        .set_pg_datetime_style(style.unwrap_or(old_style), order.unwrap_or(older_order));
328    Ok(())
329}
330
331pub fn set_query_timeout(exprs: Vec<Expr>, ctx: QueryContextRef) -> Result<()> {
332    let timeout_expr = exprs.first().context(NotSupportedSnafu {
333        feat: "No timeout value find in set query timeout statement",
334    })?;
335    match timeout_expr {
336        Expr::Value(ValueWithSpan {
337            value: Value::Number(timeout, _),
338            ..
339        }) => {
340            match timeout.parse::<u64>() {
341                Ok(timeout) => ctx.set_query_timeout(Duration::from_millis(timeout)),
342                Err(_) => {
343                    return NotSupportedSnafu {
344                        feat: format!("Invalid timeout expr {} in set variable statement", timeout),
345                    }
346                    .fail();
347                }
348            }
349            Ok(())
350        }
351        // postgres support time units i.e. SET STATEMENT_TIMEOUT = '50ms';
352        Expr::Value(ValueWithSpan {
353            value: Value::SingleQuotedString(timeout),
354            ..
355        })
356        | Expr::Value(ValueWithSpan {
357            value: Value::DoubleQuotedString(timeout),
358            ..
359        }) => {
360            if ctx.channel() != Postgres {
361                return NotSupportedSnafu {
362                    feat: format!("Invalid timeout expr {} in set variable statement", timeout),
363                }
364                .fail();
365            }
366            let timeout = parse_pg_query_timeout_input(timeout)?;
367            ctx.set_query_timeout(Duration::from_millis(timeout));
368            Ok(())
369        }
370        expr => NotSupportedSnafu {
371            feat: format!(
372                "Unsupported timeout expr {} in set variable statement",
373                expr
374            ),
375        }
376        .fail(),
377    }
378}
379
380// support time units in ms, s, min, h, d for postgres protocol.
381// https://www.postgresql.org/docs/8.4/config-setting.html#:~:text=Valid%20memory%20units%20are%20kB,%2C%20and%20d%20(days).
382fn parse_pg_query_timeout_input(input: &str) -> Result<u64> {
383    match input.parse::<u64>() {
384        Ok(timeout) => Ok(timeout),
385        Err(_) => {
386            if let Some(captures) = PG_TIME_INPUT_REGEX.captures(input) {
387                let value = captures[1].parse::<u64>().expect("regex failed");
388                let unit = &captures[2];
389
390                match unit {
391                    "ms" => Ok(value),
392                    "s" => Ok(value * 1000),
393                    "min" => Ok(value * 60 * 1000),
394                    "h" => Ok(value * 60 * 60 * 1000),
395                    "d" => Ok(value * 24 * 60 * 60 * 1000),
396                    _ => unreachable!("regex failed"),
397                }
398            } else {
399                NotSupportedSnafu {
400                    feat: format!(
401                        "Unsupported timeout expr {} in set variable statement",
402                        input
403                    ),
404                }
405                .fail()
406            }
407        }
408    }
409}
410
411#[cfg(test)]
412mod test {
413    use session::Session;
414    use session::context::{Channel, QueryContext};
415    use sql::ast::{Expr, Value};
416
417    use super::set_skip_wal;
418    use crate::statement::set::parse_pg_query_timeout_input;
419
420    #[test]
421    fn test_set_skip_wal() {
422        let ctx = QueryContext::arc();
423        for value in [true, false] {
424            set_skip_wal(vec![Expr::Value(Value::Boolean(value).into())], ctx.clone()).unwrap();
425            assert_eq!(ctx.skip_wal(), value);
426            set_skip_wal(
427                vec![Expr::Value(
428                    Value::SingleQuotedString(value.to_string()).into(),
429                )],
430                ctx.clone(),
431            )
432            .unwrap();
433            assert_eq!(ctx.skip_wal(), value);
434        }
435        ctx.set_skip_wal(true);
436        for values in [
437            vec![],
438            vec![Expr::Value(Value::Number("1".to_string(), false).into())],
439            vec![Expr::Value(
440                Value::SingleQuotedString("invalid".to_string()).into(),
441            )],
442            vec![Expr::Value(Value::Boolean(false).into()); 2],
443        ] {
444            assert!(set_skip_wal(values, ctx.clone()).is_err());
445            assert!(ctx.skip_wal());
446        }
447    }
448
449    #[test]
450    fn test_set_skip_wal_session_isolation() {
451        for channel in [Channel::Mysql, Channel::Postgres] {
452            let session = Session::new(None, channel, Default::default(), 0);
453            let other = Session::new(None, channel, Default::default(), 1);
454            assert!(!session.new_query_context().skip_wal());
455            set_skip_wal(
456                vec![Expr::Value(Value::Boolean(true).into())],
457                session.new_query_context(),
458            )
459            .unwrap();
460            assert!(session.new_query_context().skip_wal());
461            assert!(!other.new_query_context().skip_wal());
462        }
463    }
464
465    #[test]
466    fn test_parse_pg_query_timeout_input() {
467        assert!(parse_pg_query_timeout_input("").is_err());
468        assert!(parse_pg_query_timeout_input(" 50 ms").is_err());
469        assert!(parse_pg_query_timeout_input("5s 1ms").is_err());
470        assert!(parse_pg_query_timeout_input("3a").is_err());
471        assert!(parse_pg_query_timeout_input("1.5min").is_err());
472        assert!(parse_pg_query_timeout_input("ms").is_err());
473        assert!(parse_pg_query_timeout_input("a").is_err());
474        assert!(parse_pg_query_timeout_input("-1").is_err());
475
476        assert_eq!(50, parse_pg_query_timeout_input("50").unwrap());
477        assert_eq!(12, parse_pg_query_timeout_input("12ms").unwrap());
478        assert_eq!(2000, parse_pg_query_timeout_input("2s").unwrap());
479        assert_eq!(60000, parse_pg_query_timeout_input("1min").unwrap());
480    }
481}