diff --git a/src/operator/src/statement.rs b/src/operator/src/statement.rs index 21c3ed6146..150020bc9f 100644 --- a/src/operator/src/statement.rs +++ b/src/operator/src/statement.rs @@ -63,7 +63,7 @@ use query::QueryEngineRef; use query::parser::QueryStatement; use session::context::{Channel, QueryContextBuilder, QueryContextRef}; use session::table_name::table_idents_to_full_name; -use set::{set_query_timeout, set_read_preference}; +use set::{set_query_timeout, set_read_preference, set_skip_wal}; use snafu::{OptionExt, ResultExt, ensure}; use sql::ast::ObjectNamePartExt; use sql::statements::OptionMap; @@ -534,6 +534,7 @@ impl StatementExecutor { match var_name.as_str() { "READ_PREFERENCE" => set_read_preference(set_var.value, query_ctx)?, + "SKIP_WAL" => set_skip_wal(set_var.value, query_ctx)?, "@@TIME_ZONE" | "@@SESSION.TIME_ZONE" | "TIMEZONE" | "TIME_ZONE" => { set_timezone(set_var.value, query_ctx)? diff --git a/src/operator/src/statement/set.rs b/src/operator/src/statement/set.rs index b0305dbdae..83924953e7 100644 --- a/src/operator/src/statement/set.rs +++ b/src/operator/src/statement/set.rs @@ -38,6 +38,35 @@ lazy_static! { static ref PG_TIME_INPUT_REGEX: Regex = Regex::new(r"^(\d+)(ms|s|min|h|d)$").unwrap(); } +/// Sets the session WAL policy for ordinary inserts without changing table options. +pub fn set_skip_wal(exprs: Vec, ctx: QueryContextRef) -> Result<()> { + let [Expr::Value(value)] = exprs.as_slice() else { + return NotSupportedSnafu { + feat: "SET skip_wal requires exactly one boolean value", + } + .fail(); + }; + let skip_wal = match &value.value { + Value::Boolean(value) => *value, + Value::SingleQuotedString(value) | Value::DoubleQuotedString(value) => { + value.parse::().map_err(|_| { + NotSupportedSnafu { + feat: format!("Invalid skip_wal value {value:?}: expected true or false"), + } + .build() + })? + } + _ => { + return NotSupportedSnafu { + feat: "SET skip_wal requires true or false", + } + .fail(); + } + }; + ctx.set_skip_wal(skip_wal); + Ok(()) +} + pub fn set_read_preference(exprs: Vec, ctx: QueryContextRef) -> Result<()> { let read_preference_expr = exprs.first().context(NotSupportedSnafu { feat: "No read preference find in set variable statement", @@ -381,8 +410,58 @@ fn parse_pg_query_timeout_input(input: &str) -> Result { #[cfg(test)] mod test { + use session::Session; + use session::context::{Channel, QueryContext}; + use sql::ast::{Expr, Value}; + + use super::set_skip_wal; use crate::statement::set::parse_pg_query_timeout_input; + #[test] + fn test_set_skip_wal() { + let ctx = QueryContext::arc(); + for value in [true, false] { + set_skip_wal(vec![Expr::Value(Value::Boolean(value).into())], ctx.clone()).unwrap(); + assert_eq!(ctx.skip_wal(), value); + set_skip_wal( + vec![Expr::Value( + Value::SingleQuotedString(value.to_string()).into(), + )], + ctx.clone(), + ) + .unwrap(); + assert_eq!(ctx.skip_wal(), value); + } + ctx.set_skip_wal(true); + for values in [ + vec![], + vec![Expr::Value(Value::Number("1".to_string(), false).into())], + vec![Expr::Value( + Value::SingleQuotedString("invalid".to_string()).into(), + )], + vec![Expr::Value(Value::Boolean(false).into()); 2], + ] { + assert!(set_skip_wal(values, ctx.clone()).is_err()); + assert!(ctx.skip_wal()); + } + } + + #[test] + fn test_set_skip_wal_session_isolation() { + for channel in [Channel::Mysql, Channel::Postgres] { + let session = Session::new(None, channel, Default::default(), 0); + let other = Session::new(None, channel, Default::default(), 1); + assert!(!session.new_query_context().skip_wal()); + set_skip_wal( + vec![Expr::Value(Value::Boolean(true).into())], + session.new_query_context(), + ) + .unwrap(); + assert!(session.new_query_context().skip_wal()); + assert!(!other.new_query_context().skip_wal()); + } + } + #[test] fn test_parse_pg_query_timeout_input() { assert!(parse_pg_query_timeout_input("").is_err());