diff --git a/crates/squawk_ide/src/code_actions/mod.rs b/crates/squawk_ide/src/code_actions/mod.rs index d0acbc29..0c9dfb4b 100644 --- a/crates/squawk_ide/src/code_actions/mod.rs +++ b/crates/squawk_ide/src/code_actions/mod.rs @@ -22,6 +22,7 @@ mod rewrite_create_table_as_as_select_into; mod rewrite_double_colon_to_cast; mod rewrite_extract_as_function_call; mod rewrite_from; +mod rewrite_function_param_default_as_equals; mod rewrite_in_as_expression; mod rewrite_integer_radix; mod rewrite_is_normalized_as_function_call; @@ -70,6 +71,7 @@ use rewrite_create_table_as_as_select_into::rewrite_create_table_as_as_select_in use rewrite_double_colon_to_cast::rewrite_double_colon_to_cast; use rewrite_extract_as_function_call::rewrite_extract_as_function_call; use rewrite_from::rewrite_from; +use rewrite_function_param_default_as_equals::rewrite_function_param_default_as_equals; use rewrite_in_as_expression::rewrite_in_as_expression; use rewrite_integer_radix::rewrite_integer_radix; use rewrite_is_normalized_as_function_call::rewrite_is_normalized_as_function_call; @@ -118,6 +120,7 @@ pub fn code_actions(db: &dyn Db, position: InFile) -> Option, + actions: &mut Vec, +) -> Option<()> { + let token = token_from_offset(db, position)?; + let param_default = token.parent_ancestors().find_map(ast::ParamDefault::cast)?; + param_default + .syntax() + .ancestors() + .find_map(ast::CreateFunction::cast)?; + let default_token = param_default.default_token()?; + + actions.push(CodeAction { + title: "Rewrite `DEFAULT` as `=`".to_owned(), + edits: vec![Edit::replace(default_token.text_range(), "=".to_owned())], + kind: ActionKind::RefactorRewrite, + }); + + Some(()) +} + +#[cfg(test)] +mod test { + use insta::assert_snapshot; + + use crate::code_actions::test_utils::{apply_code_action, code_action_not_applicable}; + + use super::rewrite_function_param_default_as_equals; + + #[test] + fn rewrites_default_as_equals() { + assert_snapshot!( + apply_code_action( + rewrite_function_param_default_as_equals, + "create function f(a int def$0ault 1) returns int language sql as $$ select a $$;", + ), + @"create function f(a int = 1) returns int language sql as $$ select a $$;" + ); + } + + #[test] + fn applies_when_cursor_is_on_default_expression() { + assert_snapshot!( + apply_code_action( + rewrite_function_param_default_as_equals, + "create function f(a text DEFAULT lower($0'x')) returns text language sql as $$ select a $$;", + ), + @"create function f(a text = lower('x')) returns text language sql as $$ select a $$;" + ); + } + + #[test] + fn preserves_comments() { + assert_snapshot!( + apply_code_action( + rewrite_function_param_default_as_equals, + "create function f(a int DEFAULT$0 /* value */ 1) returns int language sql as $$ select a $$;", + ), + @"create function f(a int = /* value */ 1) returns int language sql as $$ select a $$;" + ); + } + + #[test] + fn not_applicable_to_equals() { + assert!(code_action_not_applicable( + rewrite_function_param_default_as_equals, + "create function f(a int =$0 1) returns int language sql as $$ select a $$;" + )); + } + + #[test] + fn not_applicable_to_column_default() { + assert!(code_action_not_applicable( + rewrite_function_param_default_as_equals, + "create table t(a int DEF$0AULT 1);" + )); + } +}