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
46 changes: 46 additions & 0 deletions crates/squawk_ide/src/goto_definition.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1406,6 +1406,16 @@ select * from (
");
}

#[test]
fn goto_recursive_cte_star_column_count_cycle_not_found() {
goto_not_found(
"
with recursive t as (select * from t)
select x$0 from t as t(a);
",
);
}

#[test]
fn goto_cte_forward_ref_ignored_for_qualified_star() {
assert_snapshot!(goto("
Expand Down Expand Up @@ -9567,6 +9577,42 @@ select a$0 from (select * from t) u(x);
);
}

#[test]
fn goto_cte_star_column_count_through_nested_ctes() {
assert_snapshot!(goto("
with s as (select 1 a, 2 b),
u as (select * from s),
t as (select 3 c),
v as (select * from u, t)
select c$0 from v as x(p, q);
"), @"
╭▸
4 │ t as (select 3 c),
│ ─ 2. destination
5 │ v as (select * from u, t)
6 │ select c from v as x(p, q);
╰╴ ─ 1. source
");
}

#[test]
fn goto_view_star_column_count_through_nested_views() {
assert_snapshot!(goto("
create table s (a int, b int);
create view u as select * from s;
create table t (c int);
create view v as select * from u, t;
select c$0 from v as x(p, q);
"), @"
╭▸
4 │ create table t (c int);
│ ─ 2. destination
5 │ create view v as select * from u, t;
6 │ select c from v as x(p, q);
╰╴ ─ 1. source
");
}

#[test]
fn goto_subquery_column_with_as() {
assert_snapshot!(goto("
Expand Down
158 changes: 128 additions & 30 deletions crates/squawk_ide/src/resolve.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@ use crate::{
db::{ancestors_with_embedded, binders},
};
use rowan::TextSize;
use rustc_hash::FxHashSet;
use smallvec::{SmallVec, smallvec};
use squawk_syntax::{
SyntaxKind, SyntaxNode, SyntaxNodePtr,
Expand Down Expand Up @@ -3988,31 +3989,57 @@ fn count_columns_for_path(db: &dyn Db, path: InFile<&ast::PathRef>) -> Option<us
count_columns_for_table_name(db, &table_name, &schemas, file)
}

type ColumnCountVisited = FxHashSet<InFile<SyntaxNodePtr>>;

fn count_columns_for_cte_or_table_name(
db: &dyn Db,
table_name: &Name,
schemas: &ResolvedSchemas,
file: FileId,
name_ref: &impl ast::NameLike,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
if schemas.unqualified()
&& let Some(with_table) = ast_nav::find_cte_with_table(name_ref, table_name)
{
return count_columns_for_with_table(db, file, with_table);
return count_columns_for_with_table(db, file, with_table, visited);
}

count_columns_for_table_name(db, table_name, schemas, file)
count_columns_for_table_name_impl(db, table_name, schemas, file, visited)
}

fn count_columns_for_table_name(
db: &dyn Db,
table_name: &Name,
schemas: &ResolvedSchemas,
file: FileId,
) -> Option<usize> {
count_columns_for_table_name_impl(db, table_name, schemas, file, &mut FxHashSet::default())
}

fn count_columns_for_table_name_impl(
db: &dyn Db,
table_name: &Name,
schemas: &ResolvedSchemas,
file: FileId,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
let (table_like_ptr, kind) = resolve_view_or_table(db, table_name, schemas, file)?;
let file = table_like_ptr.file_id;
if !visited.insert(table_like_ptr) {
return None;
}
let result = count_columns_for_table_like(db, table_like_ptr, kind, visited);
visited.remove(&table_like_ptr);
result
}

fn count_columns_for_table_like(
db: &dyn Db,
table_like_ptr: InFile<SyntaxNodePtr>,
kind: LocationKind,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
let file = table_like_ptr.file_id;
match kind {
LocationKind::Table => {
let table_like_node = table_like_ptr.to_node(db);
Expand Down Expand Up @@ -4049,14 +4076,19 @@ fn count_columns_for_table_name(
db,
file,
&create_table_as.query()?.select_variant()?,
visited,
);
}

if let Some(select_into) = table_like_node.ancestors().find_map(ast::SelectInto::cast) {
if let Some(target_list) = select_into.select_clause().and_then(|c| c.target_list())
{
return Some(target_list.targets().count());
}
let target_list = select_into.select_clause()?.target_list()?;
return count_columns_for_target_list(
db,
file,
&target_list,
select_into.from_clause().as_ref(),
visited,
);
}

None
Expand All @@ -4070,7 +4102,7 @@ fn count_columns_for_table_name(
return Some(column_list.column_names().count());
}

count_columns_for_select_variant(db, file, &create_view.query()?)
count_columns_for_select_variant(db, file, &create_view.query()?, visited)
}
_ => None,
}
Expand All @@ -4080,26 +4112,43 @@ fn count_columns_for_with_table(
db: &dyn Db,
file: FileId,
with_table: ast::WithTable,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
if let Some(column_list) = with_table.column_list() {
return Some(column_list.column_names().count());
}

count_columns_for_with_query(db, file, with_table.query()?)
let query = with_table.query()?;
let with_table_ptr = InFile::new(file, SyntaxNodePtr::new(with_table.syntax()));
if !visited.insert(with_table_ptr) {
return None;
}
let result = count_columns_for_with_query(db, file, query, visited);
visited.remove(&with_table_ptr);
result
}

fn count_columns_for_with_query(db: &dyn Db, file: FileId, query: ast::WithQuery) -> Option<usize> {
fn count_columns_for_with_query(
db: &dyn Db,
file: FileId,
query: ast::WithQuery,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
match query {
ast::WithQuery::CompoundSelect(compound_select) => {
count_columns_for_select_variant(db, file, &compound_select.lhs()?)
count_columns_for_select_variant(db, file, &compound_select.lhs()?, visited)
}
ast::WithQuery::ParenSelect(paren_select) => {
count_columns_for_select_variant(db, file, &paren_select.select()?)
count_columns_for_select_variant(db, file, &paren_select.select()?, visited)
}
ast::WithQuery::Select(select) => Some(count_columns_for_target_list(
ast::WithQuery::Select(select) => count_columns_for_target_list(
db,
file,
&select.select_clause()?.target_list()?,
)),
ast::WithQuery::Table(table) => count_columns_for_table_query(db, file, &table),
select.from_clause().as_ref(),
visited,
),
ast::WithQuery::Table(table) => count_columns_for_table_query(db, file, &table, visited),
ast::WithQuery::Values(values) => count_columns_for_values(&values),
ast::WithQuery::Delete(_)
| ast::WithQuery::Insert(_)
Expand All @@ -4112,36 +4161,65 @@ fn count_columns_for_select_variant(
db: &dyn Db,
file: FileId,
select_variant: &ast::SelectVariant,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
match select_variant {
ast::SelectVariant::CompoundSelect(compound_select) => {
count_columns_for_select_variant(db, file, &compound_select.lhs()?)
count_columns_for_select_variant(db, file, &compound_select.lhs()?, visited)
}
ast::SelectVariant::ParenSelect(paren_select) => {
count_columns_for_select_variant(db, file, &paren_select.select()?)
count_columns_for_select_variant(db, file, &paren_select.select()?, visited)
}
ast::SelectVariant::Select(select) => Some(count_columns_for_target_list(
ast::SelectVariant::Select(select) => count_columns_for_target_list(
db,
file,
&select.select_clause()?.target_list()?,
)),
ast::SelectVariant::SelectInto(select_into) => Some(count_columns_for_target_list(
select.from_clause().as_ref(),
visited,
),
ast::SelectVariant::SelectInto(select_into) => count_columns_for_target_list(
db,
file,
&select_into.select_clause()?.target_list()?,
)),
ast::SelectVariant::Table(table) => count_columns_for_table_query(db, file, table),
select_into.from_clause().as_ref(),
visited,
),
ast::SelectVariant::Table(table) => count_columns_for_table_query(db, file, table, visited),
ast::SelectVariant::Values(values) => count_columns_for_values(values),
}
}

fn count_columns_for_target_list(target_list: &ast::TargetList) -> usize {
target_list.targets().count()
fn count_columns_for_target_list(
db: &dyn Db,
file: FileId,
target_list: &ast::TargetList,
from_clause: Option<&ast::FromClause>,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
let mut total = 0usize;
for target in target_list.targets() {
let count = if let Some(from_clause) = from_clause {
count_columns_for_target_impl(db, InFile::new(file, &target), from_clause, visited)?
} else {
1
};
total = total.saturating_add(count);
}
Some(total)
}

fn count_columns_for_table_query(db: &dyn Db, file: FileId, table: &ast::Table) -> Option<usize> {
fn count_columns_for_table_query(
db: &dyn Db,
file: FileId,
table: &ast::Table,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
let path = table.relation_name()?.relation_name_ref()?.path_ref()?;
let (schema, table_name) = name::schema_and_name_path(&path)?;
let table_name_ref = relation_name_ref_from_table(table)?;
let position = table_name_ref.syntax().text_range().start();
let schemas = bind(db, file).resolved_schemas(position, schema.as_ref());
count_columns_for_cte_or_table_name(db, &table_name, &schemas, file, &table_name_ref)
count_columns_for_cte_or_table_name(db, &table_name, &schemas, file, &table_name_ref, visited)
}

fn count_columns_for_values(values: &ast::Values) -> Option<usize> {
Expand Down Expand Up @@ -4751,18 +4829,27 @@ fn count_columns_for_target(
db: &dyn Db,
target: InFile<&ast::Target>,
from_clause: &ast::FromClause,
) -> Option<usize> {
count_columns_for_target_impl(db, target, from_clause, &mut FxHashSet::default())
}

fn count_columns_for_target_impl(
db: &dyn Db,
target: InFile<&ast::Target>,
from_clause: &ast::FromClause,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
let file = target.file_id;
let target = target.value;
if target.star_token().is_some() {
return count_columns_for_from_clause(db, InFile::new(file, from_clause));
return count_columns_for_from_clause(db, InFile::new(file, from_clause), visited);
}

if let Some(ast::Expr::FieldExpr(field_expr)) = target.expr()
&& let Some(table_name) = qualified_star_table_name(&field_expr)
&& let Some(from_item) = find_from_item_in_from_clause(from_clause, &table_name)
{
return count_columns_for_from_item(db, InFile::new(file, &from_item));
return count_columns_for_from_item_impl(db, InFile::new(file, &from_item), visited);
}

Some(1)
Expand All @@ -4771,14 +4858,17 @@ fn count_columns_for_target(
fn count_columns_for_from_clause(
db: &dyn Db,
from_clause: InFile<&ast::FromClause>,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
let file = from_clause.file_id;
let from_clause = from_clause.value;
let mut total: usize = 0;
let mut found = false;

for from_item in ast_nav::iter_from_clause(from_clause) {
if let Some(count) = count_columns_for_from_item(db, InFile::new(file, &from_item)) {
if let Some(count) =
count_columns_for_from_item_impl(db, InFile::new(file, &from_item), visited)
{
total = total.saturating_add(count);
found = true;
}
Expand All @@ -4788,13 +4878,21 @@ fn count_columns_for_from_clause(
}

fn count_columns_for_from_item(db: &dyn Db, from_item: InFile<&ast::FromItem>) -> Option<usize> {
count_columns_for_from_item_impl(db, from_item, &mut FxHashSet::default())
}

fn count_columns_for_from_item_impl(
db: &dyn Db,
from_item: InFile<&ast::FromItem>,
visited: &mut ColumnCountVisited,
) -> Option<usize> {
let file = from_item.file_id;
let from_item = from_item.value;
let (schema, table_name) = name::schema_and_table_from_from_item(from_item)?;
let scope_name_ref = relation_name_ref_from_from_item(from_item)?;
let position = scope_name_ref.syntax().text_range().start();
let schemas = bind(db, file).resolved_schemas(position, schema.as_ref());
count_columns_for_cte_or_table_name(db, &table_name, &schemas, file, &scope_name_ref)
count_columns_for_cte_or_table_name(db, &table_name, &schemas, file, &scope_name_ref, visited)
}

fn resolve_from_clause_column_after_index(
Expand Down
Loading