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
94 changes: 92 additions & 2 deletions prqlc/prqlc/src/semantic/lowering.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
use std::collections::hash_map::RandomState;
use std::collections::{HashMap, HashSet};
use std::collections::{BTreeSet, HashMap, HashSet};
use std::iter::zip;

use enum_as_inner::EnumAsInner;
Expand Down Expand Up @@ -302,9 +302,12 @@ impl Lowerer {

// lower the expr
let items = self.lower_interpolations(items)?;
let columns = tuple_fields_to_relation_columns(columns);
let columns = try_extract_sql_columns(columns, &items);

let relation = rq::Relation {
kind: rq::RelationKind::SString(items),
columns: tuple_fields_to_relation_columns(columns),
columns,
};

self.table_buffer.push(TableDecl {
Expand Down Expand Up @@ -992,6 +995,93 @@ impl Lowerer {
}
}

/// Attempts to extract column names from an S-String to avoid wildcards when possible
fn try_extract_sql_columns(
columns: Vec<RelationColumn>,
items: &[InterpolateItem<rq::Expr>],
) -> Vec<RelationColumn> {
use sqlparser::ast;

let mut has_wildcard = false;

let sql_columns = items
.iter()
.map(|item| match item {
InterpolateItem::String(s) => {
let sql_ast =
sqlparser::parser::Parser::parse_sql(&sqlparser::dialect::GenericDialect {}, s)
.map_err(|err| format!("could not parse {item:?}: {err:?}"))?;
if sql_ast.len() != 1 {
return Err(format!(
"expected exactly one statement, got {}",
sql_ast.len()
));
}

let statement = sql_ast.into_iter().next().unwrap();

if let sqlparser::ast::Statement::Query(query) = statement {
if let sqlparser::ast::SetExpr::Select(select_stmt) = *query.body {
select_stmt
.projection
.into_iter()
.map(|expr| match expr {
ast::SelectItem::UnnamedExpr(expr) => {
if let ast::Expr::Identifier(ast::Ident { value, .. }) = expr {
Ok(value)
} else {
Err(format!("Only Idents are supported, got {expr:?}"))
}
}
ast::SelectItem::ExprWithAlias { alias, .. } => Ok(alias.value), // Store alias
ast::SelectItem::QualifiedWildcard(_, _)
| ast::SelectItem::Wildcard(_) => {
has_wildcard = true;
Err("columns contain a wildcard".into())
}
})
.collect::<Result<Vec<String>, String>>()
} else {
Err(format!("not a SELECT statement: {query:?}"))
}
} else {
Err(format!("not a Query: {statement:?}"))
}
}
InterpolateItem::Expr { .. } => Err(format!(
"could not extract columns from item {item:?}: not a string"
)),
})
.collect::<Result<Vec<Vec<String>>, _>>();

let sql_columns = match sql_columns {
Ok(sql_columns) => sql_columns,
Err(cause) => {
log::warn!("Could not extract SQL columns: {cause}");
return columns;
}
}
.into_iter()
.flatten()
// deduplicate extracted columns, but preserve their order
.collect::<BTreeSet<String>>();

if has_wildcard {
log::debug!("s-string contains a wildcard, skipping column extraction");
return columns;
}

columns
.into_iter()
.filter(|column| matches!(column, RelationColumn::Single(_)))
.chain(
sql_columns
.into_iter()
.map(|col| RelationColumn::Single(Some(col))),
)
.collect()
}

fn str_lit(string: String) -> rq::Expr {
rq::Expr {
kind: rq::ExprKind::Literal(Literal::String(string)),
Expand Down
42 changes: 39 additions & 3 deletions prqlc/prqlc/tests/integration/sql.rs
Original file line number Diff line number Diff line change
Expand Up @@ -968,6 +968,42 @@ fn test_sort_in_nested_append() {
);
}

#[test]
fn test_column_name_extraction_in_s_strings() {
assert_snapshot!(compile(r#"
from s"SELECT album_id, artist_id `title` FROM `albums`"
join side:left (
s"SELECT id, name FROM `artists`"
) (this.artist_id == that.id)
"#).unwrap(),
@r"
WITH table_0 AS (
SELECT
album_id,
artist_id `title`
FROM
`albums`
),
table_1 AS (
SELECT
id,
name
FROM
`artists`
)
SELECT
table_0.artist_id,
table_0.album_id,
table_0.title,
table_1.id,
table_1.name
FROM
table_0
LEFT OUTER JOIN table_1 ON table_0.artist_id = table_1.id
"
)
}

#[test]
fn test_rn_ids_are_unique() {
// this is wrong, output will have duplicate y_id and x_id
Expand Down Expand Up @@ -2791,7 +2827,7 @@ fn test_bare_s_string_01() {
rude
)
SELECT
*
insensitive
FROM
table_0
"
Expand All @@ -2813,7 +2849,7 @@ fn test_bare_s_string_02() {
rude
)
SELECT
*
insensitive
FROM
table_0
"
Expand All @@ -2839,7 +2875,7 @@ fn test_bare_s_string_03() {
bar
)
SELECT
*
foo
FROM
table_0
");
Expand Down