Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions crates/squawk_ide/src/code_actions/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -118,6 +120,7 @@ pub fn code_actions(db: &dyn Db, position: InFile<TextSize>) -> Option<Vec<CodeA
rewrite_table_as_select(db, position, &mut actions);
rewrite_select_as_table(db, position, &mut actions);
rewrite_from(db, position, &mut actions);
rewrite_function_param_default_as_equals(db, position, &mut actions);
rewrite_integer_radix(db, position, &mut actions);
rewrite_leading_from(db, position, &mut actions);
rewrite_values_as_select(db, position, &mut actions);
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
use rowan::TextSize;
use salsa::Database as Db;
use squawk_linter::Edit;
use squawk_syntax::ast::{self, AstNode};

use crate::{file::InFile, offsets::token_from_offset};

use super::{ActionKind, CodeAction};

pub(super) fn rewrite_function_param_default_as_equals(
db: &dyn Db,
position: InFile<TextSize>,
actions: &mut Vec<CodeAction>,
) -> 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);"
));
}
}
Loading