diff --git a/crates/squawk_ide/src/binder.rs b/crates/squawk_ide/src/binder.rs index 868504bc..09317d84 100644 --- a/crates/squawk_ide/src/binder.rs +++ b/crates/squawk_ide/src/binder.rs @@ -2,6 +2,7 @@ /// see: typescript-go/internal/binder/binder.go use la_arena::Arena; use rowan::{TextRange, TextSize}; +use rustc_hash::FxHashMap; use smallvec::SmallVec; use squawk_syntax::{SyntaxNodePtr, ast, ast::AstNode}; @@ -51,6 +52,8 @@ pub(crate) struct Binder { // If we have a `create schema foo` command then commands nested inside that // get `foo` for their schema. schema_regions: Vec<(TextRange, Schema)>, + savepoint_stack: Vec<(Name, SyntaxNodePtr)>, + savepoint_refs: FxHashMap, } impl Binder { @@ -68,6 +71,8 @@ impl Binder { }], default_schema_override: None, schema_regions: vec![], + savepoint_stack: vec![], + savepoint_refs: FxHashMap::default(), } } @@ -84,6 +89,15 @@ impl Binder { Some(self.symbols[symbol_id].ptr) } + pub(crate) fn lookup_savepoint( + &self, + savepoint_ref: &ast::SavepointRef, + ) -> Option { + self.savepoint_refs + .get(&SyntaxNodePtr::new(savepoint_ref.syntax())) + .copied() + } + pub(crate) fn resolved_schemas( &self, position: TextSize, @@ -358,6 +372,9 @@ fn bind_stmt(b: &mut Binder, stmt: ast::Stmt) { ast::Stmt::Prepare(prepare) => bind_prepare(b, prepare), ast::Stmt::Listen(listen) => bind_listen(b, listen), ast::Stmt::SavepointCreate(savepoint) => bind_savepoint(b, savepoint), + ast::Stmt::ReleaseSavepoint(release) => bind_release_savepoint(b, release), + ast::Stmt::Rollback(rollback) => bind_rollback(b, rollback), + ast::Stmt::Begin(_) | ast::Stmt::Commit(_) => b.savepoint_stack.clear(), ast::Stmt::Select(select) => bind_select(b, select), ast::Stmt::Set(set) => bind_set(b, set), ast::Stmt::CreatePolicy(create_policy) => bind_create_policy(b, create_policy), @@ -1682,9 +1699,42 @@ fn bind_savepoint(b: &mut Binder, savepoint: ast::SavepointCreate) { table: None, }); + b.savepoint_stack.push((savepoint_name.clone(), name_ptr)); + b.scope.insert(savepoint_name, savepoint_id); } +fn bind_savepoint_ref(b: &mut Binder, savepoint_ref: &ast::SavepointRef) -> Option { + let name = Name::from_node(savepoint_ref); + let idx = b.savepoint_stack.iter().rposition(|(n, _)| *n == name)?; + b.savepoint_refs.insert( + SyntaxNodePtr::new(savepoint_ref.syntax()), + b.savepoint_stack[idx].1, + ); + Some(idx) +} + +fn bind_release_savepoint(b: &mut Binder, release: ast::ReleaseSavepoint) { + let Some(savepoint_ref) = release.savepoint_ref() else { + return; + }; + + if let Some(idx) = bind_savepoint_ref(b, &savepoint_ref) { + b.savepoint_stack.truncate(idx); + } +} + +fn bind_rollback(b: &mut Binder, rollback: ast::Rollback) { + let Some(savepoint_ref) = rollback.savepoint_ref() else { + b.savepoint_stack.clear(); + return; + }; + + if let Some(idx) = bind_savepoint_ref(b, &savepoint_ref) { + b.savepoint_stack.truncate(idx + 1); + } +} + fn item_name(path: &ast::Path) -> Option { Some(Name::from_node(&path.segment()?)) } diff --git a/crates/squawk_ide/src/goto_definition.rs b/crates/squawk_ide/src/goto_definition.rs index 55691147..625d4007 100644 --- a/crates/squawk_ide/src/goto_definition.rs +++ b/crates/squawk_ide/src/goto_definition.rs @@ -754,6 +754,91 @@ rollback to savepoint sp$0; "); } + #[test] + fn goto_rollback_to_savepoint_without_savepoint_keyword() { + assert_snapshot!(goto(" +begin; +savepoint sp; +rollback to sp$0; +"), @" + ╭▸ + 3 │ savepoint sp; + │ ── 2. destination + 4 │ rollback to sp; + ╰╴ ─ 1. source + "); + } + + #[test] + fn goto_rollback_to_most_recent_savepoint_with_same_name() { + assert_snapshot!(goto(" +begin; +savepoint sp; +savepoint sp; +rollback to sp$0; +"), @" + ╭▸ + 4 │ savepoint sp; + │ ── 2. destination + 5 │ rollback to sp; + ╰╴ ─ 1. source + "); + } + + #[test] + fn goto_rollback_keeps_target_savepoint() { + assert_snapshot!(goto(" +begin; +savepoint sp; +rollback to sp; +rollback to sp$0; +"), @" + ╭▸ + 3 │ savepoint sp; + │ ── 2. destination + 4 │ rollback to sp; + 5 │ rollback to sp; + ╰╴ ─ 1. source + "); + } + + #[test] + fn goto_rollback_discards_later_savepoints() { + goto_not_found( + " +begin; +savepoint keep; +savepoint discarded; +rollback to keep; +release savepoint discarded$0; +", + ); + } + + #[test] + fn goto_commit_discards_savepoints() { + goto_not_found( + " +begin; +savepoint sp; +commit; +release savepoint sp$0; +", + ); + } + + #[test] + fn goto_bare_rollback_discards_savepoints() { + goto_not_found( + " +begin; +savepoint sp; +rollback; +release savepoint sp$0; +", + ); + } + #[test] fn goto_release_savepoint() { assert_snapshot!(goto(" diff --git a/crates/squawk_ide/src/resolve.rs b/crates/squawk_ide/src/resolve.rs index 9fc7629d..8bb4ddae 100644 --- a/crates/squawk_ide/src/resolve.rs +++ b/crates/squawk_ide/src/resolve.rs @@ -1729,7 +1729,7 @@ pub(crate) fn resolve_savepoint_ref( ) -> Option> { let file = savepoint_ref.file_id; let binder = bind(db, file); - let ptr = binder.lookup(&Name::from_node(savepoint_ref.value), SymbolKind::Savepoint)?; + let ptr = binder.lookup_savepoint(savepoint_ref.value)?; Some(smallvec![Location::new( file, ptr.text_range(),