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
35 changes: 33 additions & 2 deletions crates/sqlcx-core/src/parser/joins.rs
Original file line number Diff line number Diff line change
Expand Up @@ -106,8 +106,9 @@ static ON_SEP_RE: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"(?i)\s+ON\s+")
static AS_SEP_RE: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"(?i)\s+AS\s+").unwrap());

// Cheap predicate: matches the JOIN keyword anywhere in a string.
// Dialect parsers should NOT run this against full SQL — JOINs inside
// subqueries would false-positive. Use [`has_outer_join`] instead.
// Kept private — callers should use [`has_outer_join`] instead, which
// scopes the match to the outer FROM body so subquery JOINs don't
// false-positive.
static JOIN_DETECT_RE: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"(?i)\bJOIN\b").unwrap());

/// Returns true if the query's *outer* FROM clause contains a JOIN.
Expand All @@ -121,6 +122,36 @@ pub fn has_outer_join(sql: &str) -> bool {
JOIN_DETECT_RE.is_match(from_body)
}

/// Resolve a SELECT column list against a multi-table JOIN context.
/// Shared across dialect parsers (postgres, mysql, sqlite): they detect
/// the JOIN via [`has_outer_join`], pull the columns-part out of the
/// SELECT, and call this function to build the typed `ColumnDef` list.
///
/// Rejects `SELECT *` across joins with a v1.2 pointer — listing
/// qualified columns explicitly is required in v1.1.
pub fn resolve_multi_table_columns(
cols_part: &str,
sql: &str,
schema_tables: &[TableDef],
source_file: &str,
) -> Result<Vec<ColumnDef>> {
if cols_part.trim() == "*" {
return Err(SqlcxError::ParseError {
file: source_file.to_string(),
message:
"SELECT * across multi-table JOINs is not supported in v1.1 — list qualified columns explicitly (users.id, orgs.slug). `SELECT *` across joins ships in v1.2."
.to_string(),
});
}

let alias_map = parse_join_clauses(sql, schema_tables, source_file)?;

cols_part
.split(',')
.map(|s| resolve_multi_table_select_column(s.trim(), &alias_map, source_file))
.collect()
}

/// Walk a query's FROM clause and return the alias → table mapping.
/// Returns an empty map (no join detected) when the query has no FROM clause.
/// Returns an error for OUTER / USING / NATURAL / CROSS joins with a message
Expand Down
54 changes: 53 additions & 1 deletion crates/sqlcx-core/src/parser/mysql.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ use regex::Regex;
use crate::annotations::extract_annotations;
use crate::error::Result;
use crate::ir::{ColumnDef, EnumDef, QueryDef, SqlType, SqlTypeCategory, TableDef};
use crate::parser::joins::{has_outer_join, resolve_multi_table_columns};
use crate::parser::{
build_params, ensure_supported_select_expr, make_unknown_column, split_column_defs,
split_query_blocks, DatabaseParser,
Expand Down Expand Up @@ -442,6 +443,7 @@ fn find_from_table<'a>(sql: &str, tables: &'a [TableDef]) -> Option<&'a TableDef
fn resolve_return_columns(
sql: &str,
table: Option<&TableDef>,
schema_tables: &[TableDef],
source_file: &str,
) -> Result<Vec<ColumnDef>> {
if !SELECT_RE.is_match(sql) {
Expand All @@ -453,6 +455,13 @@ fn resolve_return_columns(
};
let cols_part = cap[1].trim();

// Multi-table JOIN path: route qualified columns through the shared
// resolver when the outer FROM contains a JOIN. `has_outer_join` scopes
// the check to the outer FROM body so subquery JOINs don't false-trigger.
if has_outer_join(sql) {
return resolve_multi_table_columns(cols_part, sql, schema_tables, source_file);
}

if cols_part == "*" {
return Ok(table.map(|t| t.columns.clone()).unwrap_or_default());
}
Expand Down Expand Up @@ -533,7 +542,7 @@ impl DatabaseParser for MySqlParser {
let param_indices = extract_param_indices(&block.sql);
let inferred_cols = infer_param_columns(&block.sql);
let params = build_params(&block.comments, table, param_indices, inferred_cols);
let returns = resolve_return_columns(&block.sql, table, source_file)?;
let returns = resolve_return_columns(&block.sql, table, tables, source_file)?;

let clean_sql = block
.sql
Expand Down Expand Up @@ -689,4 +698,47 @@ mod tests {
assert_eq!(dr.params[0].name, "start_date");
assert_eq!(dr.params[1].name, "end_date");
}

// ── INNER JOIN path tests ────────────────────────────────────────────────

fn join_schema() -> &'static str {
"CREATE TABLE users (id INT PRIMARY KEY, name VARCHAR(255) NOT NULL, org_id INT NOT NULL);\n\
CREATE TABLE orgs (id INT PRIMARY KEY, slug VARCHAR(255) NOT NULL);"
}

#[test]
fn inner_join_resolves_qualified_columns() {
let parser = MySqlParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: GetUserWithOrg :one\nSELECT users.name, orgs.slug FROM users INNER JOIN orgs ON users.org_id = orgs.id WHERE users.id = ?;";
let queries = parser.parse_queries(sql, &tables, &enums, "q.sql").unwrap();
assert_eq!(queries[0].returns.len(), 2);
assert_eq!(queries[0].returns[0].source_table.as_deref(), Some("users"));
assert_eq!(queries[0].returns[1].source_table.as_deref(), Some("orgs"));
}

#[test]
fn inner_join_rejects_select_star() {
let parser = MySqlParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql =
"-- name: All :many\nSELECT * FROM users INNER JOIN orgs ON users.org_id = orgs.id;";
let err = parser
.parse_queries(sql, &tables, &enums, "q.sql")
.unwrap_err();
assert!(err
.to_string()
.contains("SELECT * across multi-table JOINs"));
}

#[test]
fn left_join_rejected_with_v12_pointer() {
let parser = MySqlParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: WithLeft :many\nSELECT users.id FROM users LEFT JOIN orgs ON users.org_id = orgs.id;";
let err = parser
.parse_queries(sql, &tables, &enums, "q.sql")
.unwrap_err();
assert!(err.to_string().contains("v1.1 supports INNER JOIN only"));
}
}
113 changes: 112 additions & 1 deletion crates/sqlcx-core/src/parser/postgres.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ use regex::Regex;
use crate::annotations::extract_annotations;
use crate::error::Result;
use crate::ir::{ColumnDef, EnumDef, QueryDef, SqlType, SqlTypeCategory, TableDef};
use crate::parser::joins::{has_outer_join, resolve_multi_table_columns};
use crate::parser::{
build_params, ensure_supported_select_expr, make_unknown_column, split_column_defs,
split_query_blocks, DatabaseParser,
Expand Down Expand Up @@ -445,6 +446,7 @@ fn resolve_returning_columns(sql: &str, table: Option<&TableDef>) -> Option<Vec<
fn resolve_return_columns(
sql: &str,
table: Option<&TableDef>,
schema_tables: &[TableDef],
source_file: &str,
) -> Result<Vec<ColumnDef>> {
// Check RETURNING clause first
Expand All @@ -461,6 +463,15 @@ fn resolve_return_columns(
};
let cols_part = cap[1].trim();

// Multi-table JOIN path: when the outer FROM contains a JOIN, route
// each select expression through the shared multi-table resolver.
// `has_outer_join` scopes the check to the outer FROM body so that
// subqueries with JOINs (e.g. `WHERE id IN (SELECT ... JOIN ...)`)
// don't false-trigger.
if has_outer_join(sql) {
return resolve_multi_table_columns(cols_part, sql, schema_tables, source_file);
}

if cols_part == "*" {
return Ok(table.map(|t| t.columns.clone()).unwrap_or_default());
}
Expand Down Expand Up @@ -541,7 +552,7 @@ impl DatabaseParser for PostgresParser {
let param_indices = extract_param_indices(&block.sql);
let inferred_cols = infer_param_columns(&block.sql);
let params = build_params(&block.comments, table, param_indices, inferred_cols);
let returns = resolve_return_columns(&block.sql, table, source_file)?;
let returns = resolve_return_columns(&block.sql, table, tables, source_file)?;

let clean_sql = block
.sql
Expand Down Expand Up @@ -771,4 +782,104 @@ mod tests {
let parser = crate::parser::resolve_parser("oracle");
assert!(parser.is_err());
}

// ── INNER JOIN path tests ────────────────────────────────────────────────

fn join_schema() -> &'static str {
r#"
CREATE TABLE users (
id INTEGER PRIMARY KEY,
name TEXT NOT NULL,
org_id INTEGER NOT NULL
);
CREATE TABLE orgs (
id INTEGER PRIMARY KEY,
slug TEXT NOT NULL
);
"#
}

#[test]
fn inner_join_resolves_qualified_columns() {
let parser = PostgresParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: GetUserWithOrg :one\nSELECT users.name, orgs.slug FROM users INNER JOIN orgs ON users.org_id = orgs.id WHERE users.id = $1;";
let queries = parser.parse_queries(sql, &tables, &enums, "q.sql").unwrap();
assert_eq!(queries.len(), 1);
let q = &queries[0];
assert_eq!(q.returns.len(), 2);
assert_eq!(q.returns[0].name, "name");
assert_eq!(q.returns[0].source_table.as_deref(), Some("users"));
assert_eq!(q.returns[1].name, "slug");
assert_eq!(q.returns[1].source_table.as_deref(), Some("orgs"));
}

#[test]
fn inner_join_accepts_aliases_and_as() {
let parser = PostgresParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: Listing :many\nSELECT u.id AS user_id, o.slug AS org_slug FROM users u INNER JOIN orgs o ON u.org_id = o.id;";
let queries = parser.parse_queries(sql, &tables, &enums, "q.sql").unwrap();
let q = &queries[0];
assert_eq!(q.returns[0].name, "id");
assert_eq!(q.returns[0].alias.as_deref(), Some("user_id"));
assert_eq!(q.returns[0].source_table.as_deref(), Some("users"));
assert_eq!(q.returns[1].alias.as_deref(), Some("org_slug"));
assert_eq!(q.returns[1].source_table.as_deref(), Some("orgs"));
}

#[test]
fn inner_join_rejects_select_star() {
let parser = PostgresParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: Everything :many\nSELECT * FROM users INNER JOIN orgs ON users.org_id = orgs.id;";
let err = parser
.parse_queries(sql, &tables, &enums, "q.sql")
.unwrap_err();
assert!(err
.to_string()
.contains("SELECT * across multi-table JOINs"));
}

#[test]
fn left_join_rejected_with_v12_pointer() {
let parser = PostgresParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: WithLeft :many\nSELECT users.id FROM users LEFT JOIN orgs ON users.org_id = orgs.id;";
let err = parser
.parse_queries(sql, &tables, &enums, "q.sql")
.unwrap_err();
assert!(err.to_string().contains("v1.1 supports INNER JOIN only"));
}

#[test]
fn single_table_path_still_rejects_qualified_selects() {
// Queries without JOIN go through the existing single-table path,
// which still rejects qualified selects via ensure_supported_select_expr.
// (PR #32 is the separate effort that relaxes this for single-table queries.)
let parser = PostgresParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: Bad :one\nSELECT users.id FROM users WHERE users.id = $1;";
let err = parser
.parse_queries(sql, &tables, &enums, "q.sql")
.unwrap_err();
assert!(err
.to_string()
.contains("qualified select expressions are not supported"));
}

#[test]
fn join_in_subquery_does_not_route_outer_to_multi_table() {
// The outer FROM is single-table (`users`). The JOIN lives inside
// a subquery. The outer query must use the single-table path — if
// we routed to the multi-table resolver, the unqualified outer
// `id` select would fail with "requires qualified columns".
let parser = PostgresParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: SubquerySafe :many\nSELECT id FROM users WHERE id IN (SELECT users.id FROM users INNER JOIN orgs ON users.org_id = orgs.id);";
let queries = parser.parse_queries(sql, &tables, &enums, "q.sql").unwrap();
assert_eq!(queries[0].returns.len(), 1);
assert_eq!(queries[0].returns[0].name, "id");
assert_eq!(queries[0].returns[0].source_table, None);
}
}
54 changes: 53 additions & 1 deletion crates/sqlcx-core/src/parser/sqlite.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ use regex::Regex;
use crate::annotations::extract_annotations;
use crate::error::Result;
use crate::ir::{ColumnDef, EnumDef, QueryDef, SqlType, SqlTypeCategory, TableDef};
use crate::parser::joins::{has_outer_join, resolve_multi_table_columns};
use crate::parser::{
build_params, ensure_supported_select_expr, make_unknown_column, split_column_defs,
split_query_blocks, DatabaseParser,
Expand Down Expand Up @@ -353,6 +354,7 @@ fn find_from_table<'a>(sql: &str, tables: &'a [TableDef]) -> Option<&'a TableDef
fn resolve_return_columns(
sql: &str,
table: Option<&TableDef>,
schema_tables: &[TableDef],
source_file: &str,
) -> Result<Vec<ColumnDef>> {
if !SELECT_RE.is_match(sql) {
Expand All @@ -364,6 +366,13 @@ fn resolve_return_columns(
};
let cols_part = cap[1].trim();

// Multi-table JOIN path: route qualified columns through the shared
// resolver when the outer FROM contains a JOIN. `has_outer_join` scopes
// the check to the outer FROM body so subquery JOINs don't false-trigger.
if has_outer_join(sql) {
return resolve_multi_table_columns(cols_part, sql, schema_tables, source_file);
}

if cols_part == "*" {
return Ok(table.map(|t| t.columns.clone()).unwrap_or_default());
}
Expand Down Expand Up @@ -444,7 +453,7 @@ impl DatabaseParser for SqliteParser {
let param_indices = extract_param_indices(&block.sql);
let inferred_cols = infer_param_columns(&block.sql);
let params = build_params(&block.comments, table, param_indices, inferred_cols);
let returns = resolve_return_columns(&block.sql, table, source_file)?;
let returns = resolve_return_columns(&block.sql, table, tables, source_file)?;

let clean_sql = block
.sql
Expand Down Expand Up @@ -600,4 +609,47 @@ mod tests {
assert_eq!(dr.params[0].name, "start_date");
assert_eq!(dr.params[1].name, "end_date");
}

// ── INNER JOIN path tests ────────────────────────────────────────────────

fn join_schema() -> &'static str {
"CREATE TABLE users (id INTEGER PRIMARY KEY, name TEXT NOT NULL, org_id INTEGER NOT NULL);\n\
CREATE TABLE orgs (id INTEGER PRIMARY KEY, slug TEXT NOT NULL);"
}

#[test]
fn inner_join_resolves_qualified_columns() {
let parser = SqliteParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: GetUserWithOrg :one\nSELECT users.name, orgs.slug FROM users INNER JOIN orgs ON users.org_id = orgs.id WHERE users.id = ?;";
let queries = parser.parse_queries(sql, &tables, &enums, "q.sql").unwrap();
assert_eq!(queries[0].returns.len(), 2);
assert_eq!(queries[0].returns[0].source_table.as_deref(), Some("users"));
assert_eq!(queries[0].returns[1].source_table.as_deref(), Some("orgs"));
}

#[test]
fn inner_join_rejects_select_star() {
let parser = SqliteParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql =
"-- name: All :many\nSELECT * FROM users INNER JOIN orgs ON users.org_id = orgs.id;";
let err = parser
.parse_queries(sql, &tables, &enums, "q.sql")
.unwrap_err();
assert!(err
.to_string()
.contains("SELECT * across multi-table JOINs"));
}

#[test]
fn left_join_rejected_with_v12_pointer() {
let parser = SqliteParser::new();
let (tables, enums) = parser.parse_schema(join_schema()).unwrap();
let sql = "-- name: WithLeft :many\nSELECT users.id FROM users LEFT JOIN orgs ON users.org_id = orgs.id;";
let err = parser
.parse_queries(sql, &tables, &enums, "q.sql")
.unwrap_err();
assert!(err.to_string().contains("v1.1 supports INNER JOIN only"));
}
}
13 changes: 9 additions & 4 deletions crates/sqlcx/tests/cli.rs
Original file line number Diff line number Diff line change
Expand Up @@ -320,7 +320,10 @@ fn cli_generate_prunes_stale_query_files() {
}

#[test]
fn cli_generate_rejects_qualified_selects() {
fn cli_generate_accepts_multi_table_inner_join() {
// JOIN queries with qualified columns now succeed via the multi-table
// resolver path. Single-table qualified selects are still rejected —
// that's a separate effort (PR #32).
let dir = tempfile::tempdir().unwrap();
let sql_dir = dir.path().join("sql");
let queries_dir = sql_dir.join("queries");
Expand Down Expand Up @@ -350,9 +353,11 @@ fn cli_generate_rejects_qualified_selects() {
.output()
.unwrap();

assert!(!output.status.success());
assert!(String::from_utf8_lossy(&output.stderr)
.contains("qualified select expressions are not supported yet"));
assert!(
output.status.success(),
"expected success, stderr: {}",
String::from_utf8_lossy(&output.stderr)
);
}

#[test]
Expand Down
Loading