diff --git a/Cargo.lock b/Cargo.lock index d247b553..35c2838d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1569,7 +1569,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "33d852cb9b869c2a9b3df2f71a3074817f01e1844f839a144f5fcef059a4eb5d" dependencies = [ "libc", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -3325,7 +3325,7 @@ dependencies = [ "once_cell", "socket2 0.5.8", "tracing", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -3775,7 +3775,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -3832,7 +3832,7 @@ dependencies = [ "security-framework", "security-framework-sys", "webpki-root-certs 0.26.8", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -4278,9 +4278,9 @@ dependencies = [ [[package]] name = "sqltk" -version = "0.10.0" +version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4cc758d129a43b8feae9878eac0edf7d494730d5cb7ea8775b54704c83d0148d" +checksum = "94e94ce76e309c4b9ba2b13911ff771bcb568230e590350899f31742e8f84c84" dependencies = [ "bigdecimal", "sqltk-parser", @@ -4288,9 +4288,9 @@ dependencies = [ [[package]] name = "sqltk-parser" -version = "0.56.0-cipherstash.2" +version = "0.56.0-cipherstash.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "624fb59f490cedfefad2feadd5cd0ff275878501a4b200796def81da30f7c1d4" +checksum = "9aa59a63c3ef0a03491f6ef943b20150ae34fcc067250b8bf5e77302184ac8eb" dependencies = [ "bigdecimal", "log", @@ -4355,7 +4355,7 @@ dependencies = [ "cfg-if", "libc", "psm", - "windows-sys 0.52.0", + "windows-sys 0.59.0", ] [[package]] @@ -5496,7 +5496,7 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cf221c93e13a30d793f7645a0e7762c55d169dbb0a49671918a2319d289b10bb" dependencies = [ - "windows-sys 0.48.0", + "windows-sys 0.59.0", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 2ab66143..df773f32 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -42,7 +42,7 @@ strip = "none" debug = true [workspace.dependencies] -sqltk = { version = "0.10.0" } +sqltk = { version = "0.11.0" } cipherstash-client = { version = "=0.42.0" } cipherstash-config = { version = "=0.42.0" } cts-common = { version = "=0.42.0" } diff --git a/mise.toml b/mise.toml index 95770d8d..44c6902c 100644 --- a/mise.toml +++ b/mise.toml @@ -1,10 +1,10 @@ [settings] -# Config for test environments -# Can be invoked with: mise --env tcp run +# Config for test environments. Anchor these paths to this config so they are +# trusted regardless of the task's working directory. trusted_config_paths = [ - "./tests/mise.toml", - "./tests/mise.tcp.toml", - "./tests/mise.tls.toml", + "{{config_root}}/tests/mise.toml", + "{{config_root}}/tests/mise.tcp.toml", + "{{config_root}}/tests/mise.tls.toml", ] [task_config] diff --git a/packages/cipherstash-proxy-integration/src/insert/insert_on_conflict.rs b/packages/cipherstash-proxy-integration/src/insert/insert_on_conflict.rs new file mode 100644 index 00000000..8bd5a6ae --- /dev/null +++ b/packages/cipherstash-proxy-integration/src/insert/insert_on_conflict.rs @@ -0,0 +1,25 @@ +#[cfg(test)] +mod tests { + use crate::common::{assert_encrypted_text, clear, execute_query, query_by, random_id, trace}; + + #[tokio::test] + async fn conflict_update_encrypts_excluded_value() { + trace(); + clear().await; + + let id = random_id(); + let initial = "initial value".to_string(); + let updated = "updated value".to_string(); + let sql = "INSERT INTO encrypted (id, encrypted_text) VALUES ($1, $2) \ + ON CONFLICT (id) DO UPDATE SET encrypted_text = excluded.encrypted_text"; + + execute_query(sql, &[&id, &initial]).await; + execute_query(sql, &[&id, &updated]).await; + + assert_eq!( + query_by::("SELECT encrypted_text FROM encrypted WHERE id = $1", &id).await, + vec![updated.clone()] + ); + assert_encrypted_text(id, "encrypted_text", &updated).await; + } +} diff --git a/packages/cipherstash-proxy-integration/src/insert/insert_with_params.rs b/packages/cipherstash-proxy-integration/src/insert/insert_with_params.rs index a8445261..50e37553 100644 --- a/packages/cipherstash-proxy-integration/src/insert/insert_with_params.rs +++ b/packages/cipherstash-proxy-integration/src/insert/insert_with_params.rs @@ -1,9 +1,8 @@ #[cfg(test)] mod tests { - use crate::common::{clear, insert, query, random_id, random_limited, trace}; + use crate::common::{clear, insert, random_id, random_limited, trace}; use chrono::NaiveDate; use rand::{seq::IndexedRandom, Rng}; - use serde_json::Value; use tokio_postgres::types::ToSql; use tracing::info; @@ -38,8 +37,8 @@ mod tests { /// Return as a tuple of two vecs: /// - first vec contains column names /// - second vec contains values of the corresponding column type - pub fn generate_columns_with_values() -> (Vec, Vec>) { - let columns = vec![ + pub fn generate_columns_with_values() -> (Vec, Vec>) { + let columns = [ ("i16", "int2"), ("i32", "int4"), ("i64", "int8"), @@ -68,14 +67,6 @@ mod tests { (columns, values) } - pub async fn query tokio_postgres::types::FromSql<'a> + Send + Sync>( - sql: &str, - ) -> Vec { - let client = connect_with_tls(*PROXY).await; - let rows = client.query(sql, &[]).await.unwrap(); - rows.iter().map(|row| row.get(0)).collect::>() - } - #[tokio::test] pub async fn test_everything_all_at_once() { trace(); @@ -99,12 +90,6 @@ mod tests { info!(sql); insert(&sql, ¶ms).await; - - let sql = format!("SELECT {columns} FROM encrypted WHERE id = $1"); - - // let actual = query_by::<$type>(&sql, &id).await; - - // assert_eq!(expected, actual); } // test_insert_with_params!(insert_with_params_int2, i16, int2); diff --git a/packages/cipherstash-proxy-integration/src/insert/mod.rs b/packages/cipherstash-proxy-integration/src/insert/mod.rs index d2ff2bc3..632a751c 100644 --- a/packages/cipherstash-proxy-integration/src/insert/mod.rs +++ b/packages/cipherstash-proxy-integration/src/insert/mod.rs @@ -1,5 +1,7 @@ mod insert_domain_type; +mod insert_on_conflict; mod insert_with_literal; mod insert_with_null_literal; mod insert_with_null_param; mod insert_with_param; +mod insert_with_params; diff --git a/packages/eql-mapper/src/function_arg.rs b/packages/eql-mapper/src/function_arg.rs new file mode 100644 index 00000000..a6ceb02c --- /dev/null +++ b/packages/eql-mapper/src/function_arg.rs @@ -0,0 +1,28 @@ +use sqltk::parser::ast::{Expr, FunctionArg, FunctionArgExpr}; + +pub(crate) fn function_arg_expr(arg: &FunctionArg) -> &FunctionArgExpr { + match arg { + FunctionArg::Named { arg, .. } => arg, + FunctionArg::ExprNamed { arg, .. } => arg, + FunctionArg::Unnamed(arg) => arg, + } +} + +pub(crate) fn function_arg_value(arg: &FunctionArg) -> Option<&Expr> { + match function_arg_expr(arg) { + FunctionArgExpr::Expr(expr) => Some(expr), + FunctionArgExpr::QualifiedWildcard(_) | FunctionArgExpr::Wildcard => None, + } +} + +pub(crate) fn function_arg_value_mut(arg: &mut FunctionArg) -> Option<&mut Expr> { + let arg = match arg { + FunctionArg::Named { arg, .. } => arg, + FunctionArg::ExprNamed { arg, .. } => arg, + FunctionArg::Unnamed(arg) => arg, + }; + match arg { + FunctionArgExpr::Expr(expr) => Some(expr), + FunctionArgExpr::QualifiedWildcard(_) | FunctionArgExpr::Wildcard => None, + } +} diff --git a/packages/eql-mapper/src/importer.rs b/packages/eql-mapper/src/importer.rs index 0fcdf8af..c407b129 100644 --- a/packages/eql-mapper/src/importer.rs +++ b/packages/eql-mapper/src/importer.rs @@ -5,8 +5,8 @@ use crate::{ Relation, ScopeError, ScopeTracker, }; use sqltk::parser::ast::{ - Cte, Ident, Insert, ObjectNamePart, OnConflict, OnConflictAction, OnInsert, TableAlias, - TableFactor, TableObject, + Cte, Ident, Insert, ObjectNamePart, OnConflict, OnConflictAction, TableAlias, TableFactor, + TableObject, }; use sqltk::{Break, Visitable, Visitor}; use std::{cell::RefCell, fmt::Debug, marker::PhantomData, ops::ControlFlow, rc::Rc, sync::Arc}; @@ -18,6 +18,8 @@ pub struct Importer<'ast> { table_resolver: Arc, registry: Rc>>, scope_tracker: Rc>>, + insert_projections: Vec>, + shadowed_excluded_relations: Vec>>, _ast: PhantomData<&'ast ()>, } @@ -31,21 +33,27 @@ impl<'ast> Importer<'ast> { registry: registry.into(), table_resolver: table_resolver.into(), scope_tracker: scope.into(), + insert_projections: Vec::new(), + shadowed_excluded_relations: Vec::new(), _ast: PhantomData, } } - fn update_scope_for_insert_statement(&mut self, insert: &Insert) -> Result<(), ImportError> { + fn update_scope_for_insert_statement( + &mut self, + insert: &Insert, + ) -> Result, ImportError> { if let Insert { table: TableObject::TableName(table_name), table_alias, - on, .. } = insert { let table = self.table_resolver.resolve_table(table_name)?; - let projection = Projection::new_from_schema_table(table.clone())?; + let projection = Arc::new(Type::Value(Value::Projection( + Projection::new_from_schema_table(table.clone())?, + ))); // The relation is named — by its alias when one is written, by the // table name otherwise — so that qualified references (`t.col` in @@ -57,30 +65,10 @@ impl<'ast> Importer<'ast> { self.scope_tracker.borrow_mut().add_relation(Relation { name, - projection_type: Type::Value(Value::Projection(projection.clone())).into(), + projection_type: Arc::clone(&projection), })?; - // `ON CONFLICT DO UPDATE` can read the row proposed for insertion - // through the `excluded` pseudo-table, which projects exactly the - // target table's columns. Bringing it into scope is what gives - // `excluded.` a type — including the column's EQL type, so an - // upsert like `SET enc = excluded.enc` is fully constrained. - // - // An unqualified column reference in the `DO UPDATE` expressions is - // now ambiguous (both relations project it), which mirrors - // PostgreSQL's own `column reference is ambiguous` error there. - if let Some(OnInsert::OnConflict(OnConflict { - action: OnConflictAction::DoUpdate(_), - .. - })) = on - { - self.scope_tracker.borrow_mut().add_relation(Relation { - name: Some(Ident::new("excluded")), - projection_type: Type::Value(Value::Projection(projection)).into(), - })?; - } - - Ok(()) + Ok(projection) } else { Err(ImportError::Unsupported( "unsupported TableObject variant in Insert".to_string(), @@ -322,8 +310,8 @@ pub enum ImportError { #[error(transparent)] ScopeError(#[from] ScopeError), - #[error("Expected projection")] - ExpectedProjection, + #[error("Importer traversal invariant failed: {0}")] + TraversalInvariant(&'static str), #[error(transparent)] TypeError(#[from] TypeError), @@ -342,8 +330,33 @@ impl<'ast> Visitor<'ast> for Importer<'ast> { // 2. Child nodes of the `Insert` need to resolve identifiers in the context of the scope, so exit would be too // late. if let Some(insert) = node.downcast_ref::() { - if let Err(err) = self.update_scope_for_insert_statement(insert) { - return ControlFlow::Break(Break::Err(err)); + match self.update_scope_for_insert_statement(insert) { + Ok(projection) => self.insert_projections.push(projection), + Err(err) => return ControlFlow::Break(Break::Err(err)), + } + } + + // `excluded` exists only inside `ON CONFLICT DO UPDATE`. Adding it at + // the clause boundary keeps it visible to assignments and the WHERE + // predicate, but not to the INSERT source or RETURNING clause. + if let Some(on_conflict) = node.downcast_ref::() { + if on_conflict_is_update(on_conflict) { + let Some(projection_type) = self.insert_projections.last().cloned() else { + return ControlFlow::Break(Break::Err(ImportError::TraversalInvariant( + "ON CONFLICT DO UPDATE has no enclosing INSERT projection", + ))); + }; + + match self + .scope_tracker + .borrow_mut() + .add_shadowing_relation(Relation { + name: Some(Ident::new("excluded")), + projection_type, + }) { + Ok(shadowed) => self.shadowed_excluded_relations.push(shadowed), + Err(err) => return ControlFlow::Break(Break::Err(err.into())), + } } } @@ -351,6 +364,34 @@ impl<'ast> Visitor<'ast> for Importer<'ast> { } fn exit(&mut self, node: &'ast N) -> ControlFlow> { + if let Some(on_conflict) = node.downcast_ref::() { + if on_conflict_is_update(on_conflict) { + // Remove the pseudo-relation added on entry before traversal + // continues into the INSERT's RETURNING clause, restoring a + // target table binding that it temporarily shadowed. + let Some(shadowed) = self.shadowed_excluded_relations.pop() else { + return ControlFlow::Break(Break::Err(ImportError::TraversalInvariant( + "ON CONFLICT DO UPDATE exited without a shadow record", + ))); + }; + if let Err(err) = self + .scope_tracker + .borrow_mut() + .remove_shadowing_relation(&Ident::new("excluded"), shadowed) + { + return ControlFlow::Break(Break::Err(err.into())); + } + } + } + + if let Some(_insert) = node.downcast_ref::() { + if self.insert_projections.pop().is_none() { + return ControlFlow::Break(Break::Err(ImportError::TraversalInvariant( + "INSERT exited without a matching projection", + ))); + } + } + if let Some(cte) = node.downcast_ref::() { if let Err(err) = self.update_scope_for_cte(cte) { return ControlFlow::Break(Break::Err(err)); @@ -366,3 +407,7 @@ impl<'ast> Visitor<'ast> for Importer<'ast> { ControlFlow::Continue(()) } } + +fn on_conflict_is_update(on_conflict: &OnConflict) -> bool { + matches!(on_conflict.action, OnConflictAction::DoUpdate(_)) +} diff --git a/packages/eql-mapper/src/inference/infer_type_impls/expr.rs b/packages/eql-mapper/src/inference/infer_type_impls/expr.rs index 5efc6526..525fc6ec 100644 --- a/packages/eql-mapper/src/inference/infer_type_impls/expr.rs +++ b/packages/eql-mapper/src/inference/infer_type_impls/expr.rs @@ -87,6 +87,14 @@ impl<'ast> InferType<'ast, Expr> for TypeInferencer<'ast> { // Resolve an identifier using the scope, except if it happens to to be the DEFAULT keyword // in which case we resolve it to a fresh type variable. Expr::Identifier(ident) => { + if self + .named_function_arg_labels + .borrow() + .contains(&sqltk::NodeKey::new(expr_val)) + { + self.unify_node_with_type(expr_val, Type::native())?; + return Ok(()); + } // sqltk_parser treats the `DEFAULT` keyword in expression position as an identifier. if IdentCase(ident) == IdentCase(&Ident::new("default")) { self.unify_node_with_type(expr_val, self.fresh_tvar())?; @@ -767,22 +775,24 @@ impl<'ast> InferType<'ast, Expr> for TypeInferencer<'ast> { for access_expr in access_chain.iter() { match access_expr { - AccessExpr::Subscript(Subscript::Index { index }) => { - access_ty = self.fresh_tvar(); - root_ty = Type::array(access_ty.clone()); - self.unify_node_with_type(index, Type::native())?; - } - AccessExpr::Subscript(Subscript::Slice { - lower_bound, - upper_bound, - stride, - }) => { - self.unify_node_with_type(lower_bound, Type::native())?; - self.unify_node_with_type(upper_bound, Type::native())?; - self.unify_node_with_type(stride, Type::native())?; - access_ty = self.fresh_tvar(); - root_ty = Type::array(access_ty.clone()); - } + AccessExpr::Subscript(subscript) => match subscript.as_ref() { + Subscript::Index { index } => { + access_ty = self.fresh_tvar(); + root_ty = Type::array(access_ty.clone()); + self.unify_node_with_type(index, Type::native())?; + } + Subscript::Slice { + lower_bound, + upper_bound, + stride, + } => { + self.unify_node_with_type(lower_bound, Type::native())?; + self.unify_node_with_type(upper_bound, Type::native())?; + self.unify_node_with_type(stride, Type::native())?; + access_ty = self.fresh_tvar(); + root_ty = Type::array(access_ty.clone()); + } + }, AccessExpr::Dot(_) => { return Err(TypeError::UnsupportedSqlFeature( "field access of compound value".into(), diff --git a/packages/eql-mapper/src/inference/infer_type_impls/function.rs b/packages/eql-mapper/src/inference/infer_type_impls/function.rs index 3bebb8e3..5bd4e317 100644 --- a/packages/eql-mapper/src/inference/infer_type_impls/function.rs +++ b/packages/eql-mapper/src/inference/infer_type_impls/function.rs @@ -1,10 +1,11 @@ use eql_mapper_macros::trace_infer; use sqltk::parser::ast::{ - DuplicateTreatment, Function, FunctionArg, FunctionArgExpr, FunctionArgumentClause, - FunctionArguments, + DuplicateTreatment, Function, FunctionArg, FunctionArgumentClause, FunctionArguments, }; +use sqltk::NodeKey; use crate::{ + function_arg::function_arg_value, get_sql_function, inference::infer_type::InferType, unifier::{Type, Value}, @@ -24,6 +25,19 @@ use crate::{ /// [`WindowSpec`]: sqltk::parser::ast::WindowSpec #[trace_infer] impl<'ast> InferType<'ast, Function> for TypeInferencer<'ast> { + fn infer_enter(&mut self, function: &'ast Function) -> Result<(), TypeError> { + if let FunctionArguments::List(list) = &function.args { + for arg in &list.args { + if let FunctionArg::ExprNamed { name, .. } = arg { + self.named_function_arg_labels + .borrow_mut() + .insert(NodeKey::new(name.as_ref())); + } + } + } + Ok(()) + } + fn infer_exit(&mut self, function: &'ast Function) -> Result<(), TypeError> { if !matches!(function.parameters, FunctionArguments::None) { return Err(TypeError::UnsupportedSqlFeature( @@ -39,16 +53,7 @@ impl<'ast> InferType<'ast, Function> for TypeInferencer<'ast> { // silently returns the row count. if list.duplicate_treatment == Some(DuplicateTreatment::Distinct) { for arg in &list.args { - if let FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) - | FunctionArg::Named { - arg: FunctionArgExpr::Expr(expr), - .. - } - | FunctionArg::ExprNamed { - arg: FunctionArgExpr::Expr(expr), - .. - } = arg - { + if let Some(expr) = function_arg_value(arg) { self.unify_node_with_bound(expr, EqlTrait::Eq)?; } } diff --git a/packages/eql-mapper/src/inference/infer_type_impls/select.rs b/packages/eql-mapper/src/inference/infer_type_impls/select.rs index 6fd13003..6a457432 100644 --- a/packages/eql-mapper/src/inference/infer_type_impls/select.rs +++ b/packages/eql-mapper/src/inference/infer_type_impls/select.rs @@ -71,7 +71,7 @@ impl<'ast> InferType<'ast, Select> for TypeInferencer<'ast> { constraint, } => { self.unify_node_with_type(match_condition, Type::native())?; - Some(constraint) + Some(constraint.as_ref()) } JoinOperator::CrossJoin diff --git a/packages/eql-mapper/src/inference/infer_type_impls/select_items.rs b/packages/eql-mapper/src/inference/infer_type_impls/select_items.rs index 7d3108de..f36fd496 100644 --- a/packages/eql-mapper/src/inference/infer_type_impls/select_items.rs +++ b/packages/eql-mapper/src/inference/infer_type_impls/select_items.rs @@ -50,7 +50,7 @@ impl<'ast> InferType<'ast, Vec> for TypeInferencer<'ast> { opt_except: None, opt_replace: None, opt_rename: None, - } = options + } = options.as_ref() else { return Err(TypeError::UnsupportedSqlFeature( "options on wildcard".into(), diff --git a/packages/eql-mapper/src/inference/mod.rs b/packages/eql-mapper/src/inference/mod.rs index cfd31304..063581b8 100644 --- a/packages/eql-mapper/src/inference/mod.rs +++ b/packages/eql-mapper/src/inference/mod.rs @@ -94,6 +94,10 @@ pub struct TypeInferencer<'ast> { /// back up. fusable_json_chains: RefCell>>, + /// Expressions used as PostgreSQL named-argument labels (`name => value`). + /// They are syntax, not value expressions, and must not resolve as columns. + named_function_arg_labels: RefCell>>, + _ast: PhantomData<&'ast ()>, } @@ -112,6 +116,7 @@ impl<'ast> TypeInferencer<'ast> { json_accessor_paths: RefCell::new(JsonAccessorPaths::default()), query_operands: RefCell::new(QueryOperands::default()), fusable_json_chains: RefCell::new(HashSet::new()), + named_function_arg_labels: RefCell::new(HashSet::new()), _ast: PhantomData, } } diff --git a/packages/eql-mapper/src/inference/sql_types/sql_function_types.rs b/packages/eql-mapper/src/inference/sql_types/sql_function_types.rs index d8a7b972..9391690f 100644 --- a/packages/eql-mapper/src/inference/sql_types/sql_function_types.rs +++ b/packages/eql-mapper/src/inference/sql_types/sql_function_types.rs @@ -1,8 +1,9 @@ use std::sync::{Arc, LazyLock}; -use sqltk::parser::ast::{Function, FunctionArg, FunctionArgExpr, FunctionArguments, Ident}; +use sqltk::parser::ast::{Function, FunctionArguments, Ident}; use crate::{ + function_arg::function_arg_expr, unifier::{FunctionDecl, Type, Unifier}, IdentCase, TypeError, TypeInferencer, }; @@ -27,14 +28,6 @@ impl SqlFunction { } } -fn get_function_arg_expr(fn_arg: &FunctionArg) -> &FunctionArgExpr { - match fn_arg { - FunctionArg::Named { arg, .. } => arg, - FunctionArg::ExprNamed { arg, .. } => arg, - FunctionArg::Unnamed(arg) => arg, - } -} - impl SqlFunction { pub(crate) fn apply_constraints<'ast>( &self, @@ -61,7 +54,7 @@ impl SqlFunction { let args: Vec> = list .args .iter() - .map(|arg| inferencer.get_node_type(get_function_arg_expr(arg))) + .map(|arg| inferencer.get_node_type(function_arg_expr(arg))) .collect(); rule.inner .apply(&mut inferencer.unifier.borrow_mut(), &args, ret_type)? @@ -89,7 +82,7 @@ impl SqlFunction { let args: Vec> = list .args .iter() - .map(|arg| inferencer.get_node_type(get_function_arg_expr(arg))) + .map(|arg| inferencer.get_node_type(function_arg_expr(arg))) .collect(); NativeFunction::new(args.len() as u8).apply_constraints( &mut inferencer.unifier.borrow_mut(), diff --git a/packages/eql-mapper/src/lib.rs b/packages/eql-mapper/src/lib.rs index d97b7b8c..28dbfbbe 100644 --- a/packages/eql-mapper/src/lib.rs +++ b/packages/eql-mapper/src/lib.rs @@ -3,6 +3,7 @@ mod dep; mod display_helpers; mod eql_mapper; +mod function_arg; mod importer; mod inference; mod iterator_ext; @@ -40,7 +41,7 @@ pub(crate) use transformation_rules::*; #[cfg(test)] mod test { - use super::{test_helpers::*, type_check}; + use super::{test_helpers::*, type_check, EqlMapperError, ScopeError, TypeError}; use crate::{ projection, schema, test_helpers, unifier::{ @@ -2450,6 +2451,31 @@ mod test { } } + #[test] + fn rewrite_standard_sql_fn_with_expr_named_args() { + let schema = resolver(schema! { + tables: { + employees: { + eql_col (EQL: JsonLike), + } + } + }); + let statement = + parse("SELECT jsonb_path_exists(value => eql_col, path => '$.secret') FROM employees"); + let typed = type_check(schema, &statement).unwrap(); + let transformed = typed + .transform(test_helpers::dummy_encrypted_json_selector( + &statement, + vec![ast::Value::SingleQuotedString("$.secret".into())], + )) + .unwrap(); + + assert_eq!( + transformed.to_string(), + "SELECT eql_v3.jsonb_path_exists(value => eql_col, path => '') FROM employees" + ); + } + #[test] fn supports_named_arrays() { let schema = resolver(schema! { @@ -2722,6 +2748,27 @@ mod test { } } + #[test] + fn eql_v3_function_with_expr_named_arg_casts_full_payload() { + let schema = resolver(schema! { + tables: { + patients: { + id, + notes (EQL: JsonLike + Contain), + } + } + }); + let statement = parse( + "SELECT id FROM patients WHERE eql_v3.jsonb_contains(value => notes, query => $1)", + ); + let typed = type_check(schema, &statement).unwrap(); + + assert_eq!( + typed.transform(HashMap::new()).unwrap().to_string(), + "SELECT id FROM patients WHERE eql_v3.jsonb_contains(value => notes, query => $1::JSONB::public.eql_v3_text_search)" + ); + } + #[test] fn containment_operator_transforms_to_function() { let schema = resolver(schema! { @@ -4962,6 +5009,193 @@ mod test { } } + #[test] + fn insert_on_conflict_returning_cannot_reference_excluded() { + let schema = resolver(schema! { + tables: { + employees: { + id, + salary (EQL), + } + } + }); + + let statement = parse( + "INSERT INTO employees (id, salary) VALUES (1, 20000) \ + ON CONFLICT (id) DO UPDATE SET salary = excluded.salary \ + RETURNING excluded.salary", + ); + + assert_eq!( + type_check(schema, &statement).unwrap_err(), + EqlMapperError::Type(TypeError::ScopeError(ScopeError::NoMatch( + "excluded.salary".into() + ))) + ); + } + + #[test] + fn insert_source_cannot_reference_excluded() { + let schema = resolver(schema! { + tables: { + employees: { + id, + salary (EQL), + } + } + }); + let statement = parse( + "INSERT INTO employees (id, salary) \ + SELECT excluded.id, excluded.salary FROM employees \ + ON CONFLICT (id) DO UPDATE SET salary = excluded.salary", + ); + + assert_eq!( + type_check(schema, &statement).unwrap_err(), + EqlMapperError::Type(TypeError::ScopeError(ScopeError::NoMatch( + "excluded.id".into() + ))) + ); + } + + #[test] + fn insert_on_conflict_returning_cannot_reference_excluded_wildcard() { + let schema = resolver(schema! { + tables: { + employees: { + id, + salary (EQL), + } + } + }); + let statement = parse( + "INSERT INTO employees (id, salary) VALUES (1, 20000) \ + ON CONFLICT (id) DO UPDATE SET salary = excluded.salary \ + RETURNING excluded.*", + ); + + assert_eq!( + type_check(schema, &statement).unwrap_err(), + EqlMapperError::Type(TypeError::ScopeError(ScopeError::NoMatch( + "excluded".into() + ))) + ); + } + + #[test] + fn insert_into_table_named_excluded_is_valid() { + let schema = resolver(schema! { + tables: { + excluded: { + id, + salary, + } + } + }); + let statement = parse( + "INSERT INTO excluded (id, salary) VALUES (1, 20000) \ + ON CONFLICT (id) DO UPDATE SET salary = excluded.salary", + ); + + type_check(schema, &statement).unwrap(); + } + + #[test] + fn insert_on_conflict_update_keeps_unqualified_columns_ambiguous() { + let schema = resolver(schema! { + tables: { + employees: { + id, + salary, + } + } + }); + let statement = parse( + "INSERT INTO employees (id, salary) VALUES (1, 20000) \ + ON CONFLICT (id) DO UPDATE SET salary = salary", + ); + + assert_eq!( + type_check(schema, &statement).unwrap_err(), + EqlMapperError::Type(TypeError::ScopeError(ScopeError::AmbiguousMatch( + "salary".into() + ))) + ); + } + + #[test] + fn insert_on_conflict_do_nothing_does_not_add_excluded() { + let schema = resolver(schema! { + tables: { + employees: { + id, + salary, + } + } + }); + + for sql in [ + "INSERT INTO employees (id, salary) VALUES (1, 20000) ON CONFLICT (id) DO NOTHING", + "INSERT INTO employees (id, salary) VALUES (1, 20000) ON CONFLICT (id) DO NOTHING RETURNING *", + ] { + type_check(Arc::clone(&schema), &parse(sql)).unwrap(); + } + } + + #[test] + fn insert_on_conflict_returning_unqualified_column_is_not_ambiguous() { + let schema = resolver(schema! { + tables: { + employees: { + id, + salary (EQL), + } + } + }); + + let statement = parse( + "INSERT INTO employees (id, salary) VALUES (1, 20000) \ + ON CONFLICT (id) DO UPDATE SET salary = excluded.salary \ + RETURNING salary", + ); + + let typed = type_check(schema, &statement) + .expect("the target table must be the only relation visible to RETURNING"); + + assert_eq!( + typed.projection, + projection![(EQL(employees.salary) as salary)] + ); + } + + #[test] + fn insert_on_conflict_returning_wildcard_only_projects_target_table() { + let schema = resolver(schema! { + tables: { + employees: { + id, + salary (EQL), + } + } + }); + + let statement = parse( + "INSERT INTO employees (id, salary) VALUES (1, 20000) \ + ON CONFLICT (id) DO UPDATE SET salary = excluded.salary \ + RETURNING *", + ); + + let typed = type_check(schema, &statement).expect("RETURNING * must type check"); + + assert_eq!( + typed.projection, + projection![ + (NATIVE(employees.id) as id), + (EQL(employees.salary) as salary) + ] + ); + } + /// A conflict only fires off a unique index, and uniqueness of an /// encrypted column would be judged on the randomised ciphertext — the /// conflict would never fire. Rejected explicitly. @@ -5180,6 +5414,24 @@ mod test { } } + #[test] + fn count_distinct_expr_named_arg_uses_eq_term() { + let schema = resolver(schema! { + tables: { + employees: { + salary (EQL: Eq), + } + } + }); + let statement = parse("SELECT count(DISTINCT value => salary) FROM employees"); + let typed = type_check(schema, &statement).unwrap(); + + assert_eq!( + typed.transform(HashMap::new()).unwrap().to_string(), + "SELECT count(DISTINCT value => eql_v3.eq_term(salary)) FROM employees" + ); + } + /// The `Eq` bound on `DISTINCT` aggregate arguments must reject a column /// whose domain carries no equality term at all. (`Ord` implies `Eq` in /// this model — equality falls back to the ordering term — so the diff --git a/packages/eql-mapper/src/scope_tracker.rs b/packages/eql-mapper/src/scope_tracker.rs index d6f867c9..b15d77e7 100644 --- a/packages/eql-mapper/src/scope_tracker.rs +++ b/packages/eql-mapper/src/scope_tracker.rs @@ -67,6 +67,28 @@ impl<'ast> ScopeTracker<'ast> { self.current_scope()?.borrow_mut().add_relation(relation) } + /// Add a relation that temporarily shadows the last relation with the same name. + pub(crate) fn add_shadowing_relation( + &mut self, + relation: Relation, + ) -> Result>, ScopeError> { + Ok(self + .current_scope()? + .borrow_mut() + .add_shadowing_relation(relation)) + } + + /// Remove a temporary relation and restore the relation it shadowed, if any. + pub(crate) fn remove_shadowing_relation( + &mut self, + name: &Ident, + shadowed: Option>, + ) -> Result<(), ScopeError> { + self.current_scope()? + .borrow_mut() + .remove_shadowing_relation(name, shadowed) + } + pub(crate) fn resolve_relation(&self, name: &ObjectName) -> Result, ScopeError> { self.current_scope()?.borrow().resolve_relation(name) } @@ -234,6 +256,40 @@ impl<'ast> Scope<'ast> { Ok(()) } + fn add_shadowing_relation(&mut self, relation: Relation) -> Option> { + let shadowed = relation.name.as_ref().and_then(|name| { + let name = IdentCase(name); + self.relations + .iter() + .rposition(|relation| { + relation.name.as_ref().map(IdentCase::from).as_ref() == Some(&name) + }) + .map(|index| self.relations.remove(index)) + }); + self.relations.push(Rc::new(relation)); + shadowed + } + + fn remove_shadowing_relation( + &mut self, + name: &Ident, + shadowed: Option>, + ) -> Result<(), ScopeError> { + let name = IdentCase(name); + match self.relations.iter().rposition(|relation| { + relation.name.as_ref().map(IdentCase::from).as_ref() == Some(&name) + }) { + Some(index) => { + self.relations.remove(index); + if let Some(shadowed) = shadowed { + self.relations.push(shadowed); + } + Ok(()) + } + None => Err(ScopeError::NoMatch(name.to_string())), + } + } + pub(crate) fn resolve_relation(&self, name: &ObjectName) -> Result, ScopeError> { if name.0.len() > 1 { return Err(ScopeError::UnsupportedSqlFeature( diff --git a/packages/eql-mapper/src/transformation_rules/cast_full_payload_operands.rs b/packages/eql-mapper/src/transformation_rules/cast_full_payload_operands.rs index 0531d077..8c63dbf4 100644 --- a/packages/eql-mapper/src/transformation_rules/cast_full_payload_operands.rs +++ b/packages/eql-mapper/src/transformation_rules/cast_full_payload_operands.rs @@ -1,11 +1,10 @@ use std::collections::HashMap; use std::sync::Arc; -use sqltk::parser::ast::{ - Assignment, Expr, Function, FunctionArg, FunctionArgExpr, FunctionArguments, Values, -}; +use sqltk::parser::ast::{Assignment, Expr, Function, FunctionArguments, Values}; use sqltk::{NodeKey, NodePath, Visitable}; +use crate::function_arg::{function_arg_value, function_arg_value_mut}; use crate::unifier::{Type, Value}; use crate::EqlMapperError; @@ -63,18 +62,7 @@ impl<'ast> CastFullPayloadOperands<'ast> { _ => None, }; - args.into_iter().flatten().filter_map(|arg| match arg { - FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) - | FunctionArg::Named { - arg: FunctionArgExpr::Expr(expr), - .. - } - | FunctionArg::ExprNamed { - arg: FunctionArgExpr::Expr(expr), - .. - } => Some(expr), - _ => None, - }) + args.into_iter().flatten().filter_map(function_arg_value) } fn args_mut(function: &mut Function) -> impl Iterator { @@ -83,18 +71,9 @@ impl<'ast> CastFullPayloadOperands<'ast> { _ => None, }; - args.into_iter().flatten().filter_map(|arg| match arg { - FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) - | FunctionArg::Named { - arg: FunctionArgExpr::Expr(expr), - .. - } - | FunctionArg::ExprNamed { - arg: FunctionArgExpr::Expr(expr), - .. - } => Some(expr), - _ => None, - }) + args.into_iter() + .flatten() + .filter_map(function_arg_value_mut) } /// Whether `function` is an `eql_v3.*` call — the only functions whose diff --git a/packages/eql-mapper/src/transformation_rules/collapse_json_accessor_chain.rs b/packages/eql-mapper/src/transformation_rules/collapse_json_accessor_chain.rs index 033b9bfa..5c1becaf 100644 --- a/packages/eql-mapper/src/transformation_rules/collapse_json_accessor_chain.rs +++ b/packages/eql-mapper/src/transformation_rules/collapse_json_accessor_chain.rs @@ -107,8 +107,8 @@ impl<'ast> CollapseJsonAccessorChain<'ast> { uses_odbc_syntax: false, args: FunctionArguments::List(FunctionArgumentList { args: vec![ - FunctionArg::Unnamed(FunctionArgExpr::Expr(container)), - FunctionArg::Unnamed(FunctionArgExpr::Expr(selector)), + FunctionArg::Unnamed(FunctionArgExpr::Expr(Box::new(container))), + FunctionArg::Unnamed(FunctionArgExpr::Expr(Box::new(selector))), ], duplicate_treatment: None, clauses: vec![], diff --git a/packages/eql-mapper/src/transformation_rules/helpers.rs b/packages/eql-mapper/src/transformation_rules/helpers.rs index 3eea0c01..81461ac1 100644 --- a/packages/eql-mapper/src/transformation_rules/helpers.rs +++ b/packages/eql-mapper/src/transformation_rules/helpers.rs @@ -164,7 +164,7 @@ pub(crate) fn eql_v3_term_call(fn_name: &str, arg: Expr) -> Expr { ]), uses_odbc_syntax: false, args: FunctionArguments::List(FunctionArgumentList { - args: vec![FunctionArg::Unnamed(FunctionArgExpr::Expr(arg))], + args: vec![FunctionArg::Unnamed(FunctionArgExpr::Expr(Box::new(arg)))], duplicate_treatment: None, clauses: vec![], }), diff --git a/packages/eql-mapper/src/transformation_rules/rewrite_containment_ops.rs b/packages/eql-mapper/src/transformation_rules/rewrite_containment_ops.rs index aa83a8b2..0acfe477 100644 --- a/packages/eql-mapper/src/transformation_rules/rewrite_containment_ops.rs +++ b/packages/eql-mapper/src/transformation_rules/rewrite_containment_ops.rs @@ -72,8 +72,8 @@ impl<'ast> RewriteContainmentOps<'ast> { uses_odbc_syntax: false, args: FunctionArguments::List(FunctionArgumentList { args: vec![ - FunctionArg::Unnamed(FunctionArgExpr::Expr(left)), - FunctionArg::Unnamed(FunctionArgExpr::Expr(right)), + FunctionArg::Unnamed(FunctionArgExpr::Expr(Box::new(left))), + FunctionArg::Unnamed(FunctionArgExpr::Expr(Box::new(right))), ], duplicate_treatment: None, clauses: vec![], diff --git a/packages/eql-mapper/src/transformation_rules/rewrite_eql_aggregate_distinct.rs b/packages/eql-mapper/src/transformation_rules/rewrite_eql_aggregate_distinct.rs index 77000856..d90b152c 100644 --- a/packages/eql-mapper/src/transformation_rules/rewrite_eql_aggregate_distinct.rs +++ b/packages/eql-mapper/src/transformation_rules/rewrite_eql_aggregate_distinct.rs @@ -3,11 +3,12 @@ use std::mem; use std::sync::Arc; use sqltk::parser::ast::{ - DuplicateTreatment, Expr, Function, FunctionArg, FunctionArgExpr, FunctionArguments, - ObjectName, ObjectNamePart, Value as SqltkValue, + DuplicateTreatment, Expr, Function, FunctionArguments, ObjectName, ObjectNamePart, + Value as SqltkValue, }; use sqltk::{NodeKey, NodePath, Visitable}; +use crate::function_arg::{function_arg_value, function_arg_value_mut}; use crate::unifier::{DomainIdentity, Type, Value}; use crate::EqlMapperError; @@ -67,18 +68,7 @@ impl<'ast> RewriteEqlAggregateDistinct<'ast> { list.args .iter() - .map(|arg| match arg { - FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) - | FunctionArg::Named { - arg: FunctionArgExpr::Expr(expr), - .. - } - | FunctionArg::ExprNamed { - arg: FunctionArgExpr::Expr(expr), - .. - } => self.eql_identity_of(expr), - _ => None, - }) + .map(|arg| function_arg_value(arg).and_then(|expr| self.eql_identity_of(expr))) .collect() } @@ -143,16 +133,7 @@ impl<'ast> TransformationRule<'ast> for RewriteEqlAggregateDistinct<'ast> { ))); }; - if let FunctionArg::Unnamed(FunctionArgExpr::Expr(expr)) - | FunctionArg::Named { - arg: FunctionArgExpr::Expr(expr), - .. - } - | FunctionArg::ExprNamed { - arg: FunctionArgExpr::Expr(expr), - .. - } = arg - { + if let Some(expr) = function_arg_value_mut(arg) { let counted = mem::replace(expr, Expr::Value(SqltkValue::Null.into())); *expr = eql_v3_term_call(term_fn, counted); } diff --git a/packages/eql-mapper/src/transformation_rules/rewrite_json_value_selector_eq.rs b/packages/eql-mapper/src/transformation_rules/rewrite_json_value_selector_eq.rs index 6e6ff97b..bc8f71cf 100644 --- a/packages/eql-mapper/src/transformation_rules/rewrite_json_value_selector_eq.rs +++ b/packages/eql-mapper/src/transformation_rules/rewrite_json_value_selector_eq.rs @@ -108,8 +108,8 @@ impl<'ast> RewriteJsonValueSelectorEq<'ast> { uses_odbc_syntax: false, args: FunctionArguments::List(FunctionArgumentList { args: vec![ - FunctionArg::Unnamed(FunctionArgExpr::Expr(container)), - FunctionArg::Unnamed(FunctionArgExpr::Expr(needle)), + FunctionArg::Unnamed(FunctionArgExpr::Expr(Box::new(container))), + FunctionArg::Unnamed(FunctionArgExpr::Expr(Box::new(needle))), ], duplicate_treatment: None, clauses: vec![], diff --git a/packages/eql-mapper/src/transformation_rules/rewrite_standard_sql_fns_on_eql_types.rs b/packages/eql-mapper/src/transformation_rules/rewrite_standard_sql_fns_on_eql_types.rs index 25ad4a1f..55657551 100644 --- a/packages/eql-mapper/src/transformation_rules/rewrite_standard_sql_fns_on_eql_types.rs +++ b/packages/eql-mapper/src/transformation_rules/rewrite_standard_sql_fns_on_eql_types.rs @@ -1,8 +1,9 @@ use std::{collections::HashMap, sync::Arc}; -use sqltk::parser::ast::{Expr, Function, FunctionArg, FunctionArguments}; +use sqltk::parser::ast::{Expr, Function, FunctionArguments}; use sqltk::{AsNodeKey, NodeKey, NodePath, Visitable}; +use crate::function_arg::function_arg_expr; use crate::unifier::{Type, Value}; use crate::{get_eql_v3_function_name, get_sql_function, EqlMapperError}; @@ -33,13 +34,11 @@ impl<'ast> RewriteStandardSqlFnsOnEqlTypes<'ast> { self.node_types.get(&query.as_node_key()), Some(Type::Value(Value::Eql(_))) ), - FunctionArguments::List(list) => list.args.iter().any(|arg| match arg { - FunctionArg::Named { arg, .. } - | FunctionArg::ExprNamed { arg, .. } - | FunctionArg::Unnamed(arg) => matches!( - self.node_types.get(&arg.as_node_key()), + FunctionArguments::List(list) => list.args.iter().any(|arg| { + matches!( + self.node_types.get(&function_arg_expr(arg).as_node_key()), Some(Type::Value(Value::Eql(_))) - ), + ) }), } }