1use 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 static ref PG_TIME_INPUT_REGEX: Regex = Regex::new(r"^(\d+)(ms|s|min|h|d)$").unwrap();
39}
40
41pub 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 ensure!(
190 encoding == "UTF8" || encoding == "UNICODE",
191 NotSupportedSnafu {
192 feat: format!("client encoding of '{}'", encoding)
193 }
194 );
195 Ok(())
196}
197
198fn 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
256pub 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 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 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
380fn 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}