mod expressions;
mod helpers;
mod naming;
mod params;
mod query_fingerprint;
mod scope;
mod statements;
mod type_conversion;
mod types;
pub use type_conversion::sql_type_to_neutral;
pub use types::{
AnalyzedColumn, AnalyzedParam, AnalyzedQuery, CompositeFieldInfo, CompositeInfo, EnumInfo, GroupByConfig,
NestedFieldInfo, NestedStructInfo,
};
use ahash::{AHashMap, AHashSet};
use crate::catalog::Catalog;
use crate::dialect::SqlDialect;
use crate::errors::ScytheError;
use crate::parser::{Query, QueryCommand};
use helpers::{detect_select_star_source, find_nested_placeholder_id};
use types::Analyzer;
pub fn analyze(catalog: &Catalog, query: &Query) -> Result<AnalyzedQuery, ScytheError> {
let mut analyzer = Analyzer {
catalog,
params: Vec::new(),
ctes: AHashMap::new(),
type_errors: Vec::new(),
positional_param_counter: 0,
pending_nested: Vec::new(),
next_nested_id: 0,
};
let (columns, _) = analyzer.analyze_statement(&query.stmt)?;
if let Some(err_msg) = analyzer.type_errors.first() {
return Err(ScytheError::type_mismatch(err_msg.clone()));
}
let mut columns = columns;
for col in &mut columns {
if query.annotations.nullable_overrides.iter().any(|o| o == &col.name) {
col.nullable = true;
}
if query.annotations.nonnull_overrides.iter().any(|o| o == &col.name) {
col.nullable = false;
}
if let Some(mapping) = query.annotations.json_mappings.iter().find(|m| m.column == col.name) {
col.neutral_type = format!("json_typed<{}>", mapping.rust_type);
}
}
let nested_structs = resolve_nested_struct_names(
catalog,
&query.name,
std::mem::take(&mut analyzer.pending_nested),
&mut columns,
);
analyzer.params.sort_by_key(|p| p.position);
analyzer.params.dedup_by_key(|p| p.position);
let mut params: Vec<AnalyzedParam> = analyzer
.params
.iter()
.map(|p| {
let name = query
.annotations
.positional_param_docs
.iter()
.find(|doc| doc.position == p.position)
.map(|doc| doc.name.clone())
.unwrap_or_else(|| p.name.clone().unwrap_or_else(|| format!("p{}", p.position)));
let neutral_type = p.neutral_type.clone().unwrap_or_else(|| "unknown".to_string());
AnalyzedParam {
name,
neutral_type,
nullable: p.nullable,
position: p.position,
}
})
.collect();
for opt_name in &query.annotations.optional_params {
for p in &mut params {
if p.name == *opt_name {
p.nullable = true;
}
}
}
for opt_name in &query.annotations.optional_params {
if !params.iter().any(|p| p.name == *opt_name) {
return Err(ScytheError::invalid_annotation(format!(
"@optional references unknown parameter '{}'",
opt_name
)));
}
}
{
let mut name_counts: ahash::AHashMap<String, usize> = ahash::AHashMap::new();
for p in ¶ms {
*name_counts.entry(p.name.clone()).or_insert(0) += 1;
}
let mut name_seen: ahash::AHashMap<String, usize> = ahash::AHashMap::new();
for p in &mut params {
if name_counts.get(&p.name).copied().unwrap_or(0) > 1 {
let idx = name_seen.entry(p.name.clone()).or_insert(0);
*idx += 1;
p.name = format!("{}_{}", p.name, idx);
}
}
}
let source_table = detect_select_star_source(&query.stmt);
let nested_field_types: Vec<&str> = nested_structs
.iter()
.flat_map(|nested| nested.fields.iter())
.map(|field| field.neutral_type.as_str())
.collect();
let mut composites = Vec::new();
let mut seen_composites: AHashSet<String> = AHashSet::new();
for neutral_type in columns
.iter()
.map(|c| c.neutral_type.as_str())
.chain(nested_field_types.iter().copied())
{
if let Some(comp_name) = neutral_type.strip_prefix("composite::")
&& seen_composites.insert(comp_name.to_string())
&& let Some(comp) = catalog.get_composite(comp_name)
{
composites.push(CompositeInfo {
sql_name: comp_name.to_string(),
fields: comp
.fields
.iter()
.map(|f| CompositeFieldInfo {
name: f.name.clone(),
neutral_type: sql_type_to_neutral(&f.sql_type, catalog).into_owned(),
})
.collect(),
});
}
}
let mut enums = Vec::new();
let mut seen_enums: AHashSet<String> = AHashSet::new();
let all_types: Vec<&str> = columns
.iter()
.map(|c| c.neutral_type.as_str())
.chain(params.iter().map(|p| p.neutral_type.as_str()))
.chain(nested_field_types.iter().copied())
.collect();
for nt in &all_types {
if let Some(enum_name) = nt.strip_prefix("enum::")
&& seen_enums.insert(enum_name.to_string())
&& let Some(enum_type) = catalog.get_enum(enum_name)
{
enums.push(EnumInfo {
sql_name: enum_name.to_string(),
values: enum_type.values.clone(),
});
}
}
let group_by = if query.command == QueryCommand::Grouped {
if let Some(ref group_by_value) = query.annotations.group_by {
let (table, key_column) = if let Some(dot_pos) = group_by_value.find('.') {
(
group_by_value[..dot_pos].to_string(),
group_by_value[dot_pos + 1..].to_string(),
)
} else {
return Err(ScytheError::invalid_annotation(format!(
"@group_by must be in 'table.column' format, got: {}",
group_by_value
)));
};
let parent_table_columns: Vec<String> = catalog
.get_table(&table)
.map(|t| t.columns.iter().map(|c| c.name.clone()).collect())
.unwrap_or_default();
let mut parent_columns = Vec::new();
let mut child_columns = Vec::new();
for col in &columns {
if parent_table_columns.contains(&col.name) {
parent_columns.push(col.clone());
} else {
child_columns.push(col.clone());
}
}
Some(types::GroupByConfig {
table,
key_column,
parent_columns,
child_columns,
})
} else {
None
}
} else {
None
};
Ok(AnalyzedQuery {
name: query.name.clone(),
command: query.command.clone(),
sql: query.sql.clone(),
columns,
params,
deprecated: query.annotations.deprecated.clone(),
source_table,
composites,
enums,
optional_params: query.annotations.optional_params.clone(),
group_by,
custom: query.annotations.custom.clone(),
nested_structs,
})
}
fn resolve_nested_struct_names(
catalog: &Catalog,
query_name: &str,
pending: Vec<types::PendingNestedStruct>,
columns: &mut [AnalyzedColumn],
) -> Vec<NestedStructInfo> {
if pending.is_empty() || catalog.dialect() != SqlDialect::PostgreSQL {
return Vec::new();
}
let snake_query = naming::to_snake_case(query_name).into_owned();
let mut resolved: AHashMap<u32, String> = AHashMap::new();
let mut structs: Vec<NestedStructInfo> = Vec::new();
for column in columns.iter_mut() {
while let Some(id) = find_nested_placeholder_id(&column.neutral_type) {
let final_name = if let Some(name) = resolved.get(&id) {
name.clone()
} else {
let fields = pending
.iter()
.find(|p| p.id == id)
.map(|p| p.fields.clone())
.unwrap_or_default();
let base = format!("{snake_query}_row_{}", column.name);
let name = assign_nested_struct_name(&base, fields, catalog, &mut structs);
resolved.insert(id, name.clone());
name
};
let pascal = naming::to_pascal_case(&final_name);
column.neutral_type = column.neutral_type.replacen(&format!("__nested__{id}"), &pascal, 1);
}
}
structs
}
fn assign_nested_struct_name(
base: &str,
fields: Vec<NestedFieldInfo>,
catalog: &Catalog,
structs: &mut Vec<NestedStructInfo>,
) -> String {
let mut suffix: u32 = 0;
loop {
let candidate = if suffix == 0 {
base.to_string()
} else {
format!("{base}_{suffix}")
};
if let Some(existing) = structs.iter().find(|s| s.name == candidate) {
if existing.fields == fields {
return candidate;
}
suffix += 1;
continue;
}
if catalog.get_composite(&candidate).is_none() && catalog.get_enum(&candidate).is_none() {
structs.push(NestedStructInfo {
name: candidate.clone(),
fields,
});
return candidate;
}
suffix += 1;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parser::parse_query;
fn make_catalog() -> Catalog {
Catalog::from_ddl(&[
"CREATE TABLE users (
id SERIAL PRIMARY KEY,
name TEXT NOT NULL,
email VARCHAR(255) NOT NULL,
age INTEGER,
active BOOLEAN NOT NULL DEFAULT true,
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW(),
bio TEXT,
score NUMERIC
);",
"CREATE TABLE posts (
id SERIAL PRIMARY KEY,
user_id INTEGER NOT NULL REFERENCES users(id),
title TEXT NOT NULL,
body TEXT,
published BOOLEAN NOT NULL DEFAULT false,
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW()
);",
"CREATE TABLE comments (
id SERIAL PRIMARY KEY,
post_id INTEGER NOT NULL,
user_id INTEGER NOT NULL,
body TEXT NOT NULL,
created_at TIMESTAMP WITH TIME ZONE NOT NULL DEFAULT NOW()
);",
])
.unwrap()
}
#[test]
fn test_simple_select() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUser
-- @returns :one
SELECT id, name, email FROM users WHERE id = $1;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns.len(), 3);
assert_eq!(result.columns[0].name, "id");
assert_eq!(result.columns[0].neutral_type, "int32");
assert!(!result.columns[0].nullable);
assert_eq!(result.columns[1].name, "name");
assert_eq!(result.columns[1].neutral_type, "string");
assert_eq!(result.columns[2].name, "email");
assert_eq!(result.columns[2].neutral_type, "string");
assert_eq!(result.params.len(), 1);
assert_eq!(result.params[0].position, 1);
assert_eq!(result.params[0].neutral_type, "int32");
assert_eq!(result.params[0].name, "id");
}
#[test]
fn test_select_star() {
let catalog = make_catalog();
let query = parse_query(
"-- @name ListUsers
-- @returns :many
SELECT * FROM users;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns.len(), 8);
}
#[test]
fn test_left_join_nullability() {
let catalog = make_catalog();
let query = parse_query(
"-- @name UsersWithPosts
-- @returns :many
SELECT u.id, u.name, p.title, p.body FROM users u LEFT JOIN posts p ON u.id = p.user_id;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns.len(), 4);
assert!(!result.columns[0].nullable);
assert!(!result.columns[1].nullable);
assert!(result.columns[2].nullable);
assert!(result.columns[3].nullable);
}
#[test]
fn test_aggregate_functions() {
let catalog = make_catalog();
let query = parse_query(
"-- @name UserStats
-- @returns :one
SELECT COUNT(*) as total, AVG(age) as avg_age, MAX(score) as max_score FROM users;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns.len(), 3);
assert_eq!(result.columns[0].neutral_type, "int64");
assert!(!result.columns[0].nullable);
assert_eq!(result.columns[1].neutral_type, "decimal");
assert!(result.columns[1].nullable);
assert!(result.columns[2].nullable);
}
#[test]
fn test_insert_returning() {
let catalog = make_catalog();
let query = parse_query(
"-- @name CreateUser
-- @returns :one
INSERT INTO users (name, email) VALUES ($1, $2) RETURNING id, name, email;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns.len(), 3);
assert_eq!(result.columns[0].name, "id");
assert_eq!(result.columns[0].neutral_type, "int32");
assert_eq!(result.params.len(), 2);
assert_eq!(result.params[0].name, "name");
assert_eq!(result.params[0].neutral_type, "string");
assert_eq!(result.params[1].name, "email");
assert_eq!(result.params[1].neutral_type, "string");
}
#[test]
fn test_insert_without_column_list_binds_to_catalog_columns() {
let catalog = make_catalog();
let query = parse_query(
"-- @name CreateUserFull
-- @returns :exec
INSERT INTO users VALUES ($1, $2, $3, $4, $5, $6, $7, $8);",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 8, "all 8 placeholders must be registered");
let names: Vec<&str> = result.params.iter().map(|p| p.name.as_str()).collect();
assert_eq!(
names,
["id", "name", "email", "age", "active", "created_at", "bio", "score"]
);
assert_eq!(result.params[0].neutral_type, "int32");
assert_eq!(result.params[1].neutral_type, "string");
assert_eq!(result.params[3].neutral_type, "int32");
assert_eq!(result.params[4].neutral_type, "bool");
assert_eq!(result.params[5].neutral_type, "datetime_tz");
assert_eq!(result.params[7].neutral_type, "decimal");
assert!(!result.params[1].nullable, "name is NOT NULL");
assert!(result.params[3].nullable, "age is nullable");
assert!(result.params[6].nullable, "bio is nullable");
}
#[test]
fn test_insert_without_column_list_unknown_table_registers_inferred_params() {
let catalog = make_catalog();
let query = parse_query(
"-- @name InsertNoSchema
-- @returns :exec
INSERT INTO t VALUES ($1, $2::text);",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 2, "$1 and $2 must both be registered");
assert_eq!(result.params[0].position, 1);
assert_eq!(result.params[0].name, "p1");
assert_eq!(result.params[1].position, 2);
assert_eq!(result.params[1].neutral_type, "string", "$2::text must infer as string");
}
#[test]
fn test_insert_explicit_columns_unknown_table_keeps_declared_names() {
let catalog = make_catalog();
let query = parse_query(
"-- @name InsertNoSchemaCols
-- @returns :exec
INSERT INTO t (a, b) VALUES ($1, $2);",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 2);
assert_eq!(result.params[0].name, "a");
assert_eq!(result.params[1].name, "b");
}
#[test]
fn test_insert_coalesce_param_collected() {
let catalog = make_catalog();
let query = parse_query(
"-- @name InsertBio
-- @returns :exec
INSERT INTO users (bio) VALUES (COALESCE($1, 'unknown'));",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 1, "$1 inside COALESCE must be registered");
assert_eq!(result.params[0].name, "bio");
assert_eq!(result.params[0].neutral_type, "string");
assert!(result.params[0].nullable, "bio is nullable");
}
#[test]
fn test_insert_case_param_collected() {
let catalog = make_catalog();
let query = parse_query(
"-- @name InsertName
-- @returns :exec
INSERT INTO users (name) VALUES (CASE WHEN $1 THEN 'x' ELSE 'y' END);",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 1, "$1 inside CASE must be registered");
assert_eq!(result.params[0].name, "name");
assert_eq!(result.params[0].neutral_type, "string");
}
#[test]
fn test_update_function_arg_param_collected() {
let catalog = make_catalog();
let query = parse_query(
"-- @name RenameUser
-- @returns :exec
UPDATE users SET name = LOWER(CONCAT($1, '_suffix')) WHERE id = $2;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 2, "$1 in nested function args must be registered");
assert_eq!(result.params[0].name, "name");
assert_eq!(result.params[0].neutral_type, "string");
assert_eq!(result.params[1].name, "id");
assert_eq!(result.params[1].neutral_type, "int32");
}
#[test]
fn test_coalesce_nullability() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetBio
-- @returns :one
SELECT COALESCE(bio, 'No bio') as bio FROM users WHERE id = $1;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[0].neutral_type, "string");
assert!(!result.columns[0].nullable);
}
#[test]
fn test_case_expression() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetStatus
-- @returns :many
SELECT name, CASE WHEN active THEN 'active' ELSE 'inactive' END as status FROM users;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[1].name, "status");
assert_eq!(result.columns[1].neutral_type, "string");
assert!(!result.columns[1].nullable);
}
#[test]
fn test_nullif() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetScore
-- @returns :many
SELECT NULLIF(score, 0) as adjusted_score FROM users;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[0].neutral_type, "decimal");
assert!(result.columns[0].nullable);
}
#[test]
fn test_cast_expression() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetAgeText
-- @returns :many
SELECT CAST(age AS TEXT) as age_text FROM users;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[0].neutral_type, "string");
}
#[test]
fn test_param_inside_derived_table_propagates() {
let catalog = make_catalog();
let query = parse_query(
"-- @name BucketCounts
-- @returns :many
SELECT b.bucket, count(*) AS n
FROM posts p
CROSS JOIN (SELECT $1::text AS bucket) b
GROUP BY 1;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(
result.params.len(),
1,
"$1 inside derived table must appear in analyzed.params; got {:?}",
result.params
);
assert_eq!(result.params[0].position, 1);
assert_eq!(result.params[0].neutral_type, "string");
}
#[test]
fn test_positional_param_name_override() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUser
-- @returns :one
-- @param $1 user_id: the primary key
SELECT id, name FROM users WHERE id = $1;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 1);
assert_eq!(result.params[0].name, "user_id");
assert_eq!(result.params[0].position, 1);
}
#[test]
fn test_positional_param_override_does_not_affect_unrelated_params() {
let catalog = make_catalog();
let query = parse_query(
"-- @name UpdateUser
-- @returns :exec
-- @param $2 target_id
UPDATE users SET name = $1 WHERE id = $2;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 2);
assert_eq!(result.params[0].name, "name");
assert_eq!(result.params[1].name, "target_id");
}
#[test]
fn test_update_set_arithmetic_expr_collects_all_params() {
let catalog = make_catalog();
let query = parse_query(
"-- @name IncrementUserAge
-- @returns :exec
UPDATE users SET age = age + $2 WHERE id = $1;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.params.len(), 2, "both $1 and $2 must be present");
let positions: Vec<i64> = result.params.iter().map(|p| p.position).collect();
assert!(positions.contains(&1), "missing $1; got {positions:?}");
assert!(positions.contains(&2), "missing $2; got {positions:?}");
}
#[test]
fn test_annotation_overrides() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUser
-- @returns :one
-- @nullable name
-- @nonnull age
SELECT name, age FROM users WHERE id = $1;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert!(result.columns[0].nullable);
assert!(!result.columns[1].nullable);
}
fn nested_column(name: &str, neutral_type: &str) -> AnalyzedColumn {
AnalyzedColumn {
name: name.to_string(),
neutral_type: neutral_type.to_string(),
..Default::default()
}
}
#[test]
fn test_resolve_nested_struct_names_derives_name_from_query_and_column() {
let catalog = make_catalog();
let pending = vec![types::PendingNestedStruct {
id: 0,
fields: vec![NestedFieldInfo {
name: "title".to_string(),
neutral_type: "string".to_string(),
nullable: false,
}],
}];
let mut columns = vec![nested_column("orders", "json_nested<array<__nested__0>>")];
let structs = resolve_nested_struct_names(&catalog, "GetUserOrders", pending, &mut columns);
assert_eq!(structs.len(), 1);
assert_eq!(structs[0].name, "get_user_orders_row_orders");
assert_eq!(structs[0].fields.len(), 1);
assert_eq!(structs[0].fields[0].name, "title");
}
#[test]
fn test_resolve_nested_struct_names_pascal_case_matches_neutral_type() {
let catalog = make_catalog();
let pending = vec![types::PendingNestedStruct {
id: 3,
fields: Vec::new(),
}];
let mut columns = vec![nested_column("orders", "json_nested<array<__nested__3>>")];
let structs = resolve_nested_struct_names(&catalog, "GetUserOrders", pending, &mut columns);
let expected_pascal = naming::to_pascal_case(&structs[0].name);
assert_eq!(
columns[0].neutral_type,
format!("json_nested<array<{expected_pascal}>>")
);
}
#[test]
fn test_resolve_nested_struct_names_collision_with_catalog_composite_suffixes() {
let catalog = Catalog::from_ddl(&["CREATE TYPE get_user_row_profile AS (x TEXT);"]).unwrap();
let pending = vec![types::PendingNestedStruct {
id: 0,
fields: vec![NestedFieldInfo {
name: "bio".to_string(),
neutral_type: "string".to_string(),
nullable: true,
}],
}];
let mut columns = vec![nested_column("profile", "json_nested<__nested__0>")];
let structs = resolve_nested_struct_names(&catalog, "GetUser", pending, &mut columns);
assert_eq!(structs.len(), 1);
assert_eq!(
structs[0].name, "get_user_row_profile_1",
"must suffix rather than collide with the catalog composite \"get_user_row_profile\""
);
assert_eq!(columns[0].neutral_type, "json_nested<GetUserRowProfile1>");
}
#[test]
fn test_resolve_nested_struct_names_duplicate_column_names_dedupe_identical_shape() {
let catalog = make_catalog();
let shared_fields = vec![NestedFieldInfo {
name: "id".to_string(),
neutral_type: "int32".to_string(),
nullable: false,
}];
let pending = vec![
types::PendingNestedStruct {
id: 0,
fields: shared_fields.clone(),
},
types::PendingNestedStruct {
id: 1,
fields: shared_fields,
},
];
let mut columns = vec![
nested_column("items", "json_nested<array<__nested__0>>"),
nested_column("items", "json_nested<array<__nested__1>>"),
];
let structs = resolve_nested_struct_names(&catalog, "GetOrder", pending, &mut columns);
assert_eq!(
structs.len(),
1,
"two columns with the same name and identical field shape must dedupe to one struct"
);
assert_eq!(columns[0].neutral_type, columns[1].neutral_type);
}
#[test]
fn test_resolve_nested_struct_names_duplicate_column_names_differing_shape_suffixes() {
let catalog = make_catalog();
let pending = vec![
types::PendingNestedStruct {
id: 0,
fields: vec![NestedFieldInfo {
name: "id".to_string(),
neutral_type: "int32".to_string(),
nullable: false,
}],
},
types::PendingNestedStruct {
id: 1,
fields: vec![NestedFieldInfo {
name: "name".to_string(),
neutral_type: "string".to_string(),
nullable: false,
}],
},
];
let mut columns = vec![
nested_column("items", "json_nested<array<__nested__0>>"),
nested_column("items", "json_nested<array<__nested__1>>"),
];
let structs = resolve_nested_struct_names(&catalog, "GetOrder", pending, &mut columns);
assert_eq!(
structs.len(),
2,
"same-named columns with different field shapes must not collapse into one struct"
);
assert_eq!(structs[0].name, "get_order_row_items");
assert_eq!(structs[1].name, "get_order_row_items_1");
assert_ne!(columns[0].neutral_type, columns[1].neutral_type);
}
#[test]
fn test_resolve_nested_struct_names_mysql_dialect_produces_none() {
let catalog =
Catalog::from_ddl_with_dialect(&["CREATE TABLE orders (id INTEGER NOT NULL);"], &SqlDialect::MySQL)
.unwrap();
let pending = vec![types::PendingNestedStruct {
id: 0,
fields: vec![NestedFieldInfo {
name: "id".to_string(),
neutral_type: "int32".to_string(),
nullable: false,
}],
}];
let original_neutral_type = "json_nested<array<__nested__0>>".to_string();
let mut columns = vec![nested_column("orders", &original_neutral_type)];
let structs = resolve_nested_struct_names(&catalog, "GetUserOrders", pending, &mut columns);
assert!(
structs.is_empty(),
"non-PostgreSQL dialects must never produce nested_structs"
);
assert_eq!(
columns[0].neutral_type, original_neutral_type,
"the placeholder must be left untouched, not partially substituted"
);
}
#[test]
fn test_json_agg_wildcard_produces_nested_struct() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserPosts
-- @returns :many
SELECT u.id, json_agg(p.*) AS posts FROM users u JOIN posts p ON u.id = p.user_id GROUP BY u.id;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns.len(), 2);
assert_eq!(result.columns[1].name, "posts");
assert_eq!(
result.nested_structs.len(),
1,
"exactly one nested struct must be produced"
);
let nested = &result.nested_structs[0];
assert_eq!(nested.name, "get_user_posts_row_posts");
assert_eq!(
result.columns[1].neutral_type, "json_nested<array<GetUserPostsRowPosts>>",
"json_agg wraps the resolved name in array<> and must match the PascalCase of nested.name exactly"
);
assert!(result.columns[1].nullable);
let field_names: Vec<&str> = nested.fields.iter().map(|f| f.name.as_str()).collect();
assert_eq!(
field_names,
["id", "user_id", "title", "body", "published", "created_at"]
);
let title_field = nested.fields.iter().find(|f| f.name == "title").unwrap();
assert_eq!(title_field.neutral_type, "string");
assert!(!title_field.nullable, "posts.title is NOT NULL and the join is INNER");
let body_field = nested.fields.iter().find(|f| f.name == "body").unwrap();
assert!(body_field.nullable, "posts.body has no NOT NULL constraint");
}
#[test]
fn test_json_agg_left_join_makes_array_elements_nullable() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserPostsOuter
-- @returns :many
SELECT u.id, json_agg(p.*) AS posts FROM users u LEFT JOIN posts p ON u.id = p.user_id GROUP BY u.id;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(
result.columns[1].neutral_type, "json_nested<array<nullable<GetUserPostsOuterRowPosts>>>",
"json_agg over a LEFT JOIN with no match yields [null], so the element type must be nullable"
);
let nested = &result.nested_structs[0];
let id_field = nested.fields.iter().find(|f| f.name == "id").unwrap();
assert!(
!id_field.nullable,
"posts.id is NOT NULL; inside an object json_agg actually emitted it can never be null, so \
widening the field instead of the element would model a value PostgreSQL never produces"
);
let body_field = nested.fields.iter().find(|f| f.name == "body").unwrap();
assert!(body_field.nullable, "posts.body has no NOT NULL constraint");
}
#[test]
fn test_json_agg_inner_join_elements_are_not_nullable() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserPostsInner
-- @returns :many
SELECT u.id, json_agg(p.*) AS posts FROM users u JOIN posts p ON u.id = p.user_id GROUP BY u.id;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(
result.columns[1].neutral_type,
"json_nested<array<GetUserPostsInnerRowPosts>>"
);
}
#[test]
fn test_row_to_json_left_join_does_not_wrap_element() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserPostJson
-- @returns :many
SELECT u.id, row_to_json(p.*) AS post FROM users u LEFT JOIN posts p ON u.id = p.user_id;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[1].neutral_type, "json_nested<GetUserPostJsonRowPost>");
assert!(result.columns[1].nullable, "the column itself carries the NULL");
}
#[test]
fn test_row_to_json_wildcard_produces_nested_struct_without_array() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetPostAsJson
-- @returns :many
SELECT row_to_json(p.*) AS post FROM posts p;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.nested_structs.len(), 1);
assert_eq!(
result.columns[0].neutral_type, "json_nested<GetPostAsJsonRowPost>",
"row_to_json must not wrap in array<> -- it emits one object per output row, not an aggregate"
);
}
#[test]
fn test_json_agg_scalar_argument_falls_back_to_plain_json() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserNames
-- @returns :one
SELECT json_agg(u.name) AS names FROM users u;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(
result.columns[0].neutral_type, "json",
"a scalar/column argument is not a relation shape -- must match pre-existing json_agg behaviour exactly"
);
assert!(result.nested_structs.is_empty());
}
#[test]
fn test_json_agg_bare_wildcard_falls_back_to_plain_json() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserCount
-- @returns :one
SELECT json_agg(*) AS everything FROM users;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[0].neutral_type, "json");
assert!(result.nested_structs.is_empty());
}
#[test]
fn test_json_agg_non_postgres_dialect_falls_back_to_plain_json() {
let catalog = Catalog::from_ddl_with_dialect(
&["CREATE TABLE posts (id INTEGER NOT NULL, title TEXT NOT NULL);"],
&crate::dialect::SqlDialect::MySQL,
)
.unwrap();
let query = parse_query(
"-- @name GetPosts
-- @returns :many
SELECT json_agg(p.*) AS posts FROM posts p;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(
result.columns[0].neutral_type, "json",
"the dialect gate must produce byte-identical output to today's json_agg behaviour on non-PostgreSQL catalogs"
);
assert!(result.nested_structs.is_empty());
}
#[test]
fn test_string_agg_never_produces_nested_struct() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserNamesJoined
-- @returns :one
SELECT string_agg(u.name, ',') AS names FROM users u;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[0].neutral_type, "string");
assert!(result.nested_structs.is_empty());
}
#[test]
fn test_jsonb_agg_wildcard_stays_plain_json() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserPostsB
-- @returns :many
SELECT jsonb_agg(p.*) AS posts FROM posts p;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[0].neutral_type, "json");
assert!(result.nested_structs.is_empty());
}
#[test]
fn test_nested_of_nested_aggregate_is_rejected_with_clear_diagnostic() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetAllUserPosts
-- @returns :many
WITH user_posts AS (
SELECT u.id AS user_id, json_agg(p.*) AS posts
FROM users u JOIN posts p ON p.user_id = u.id
GROUP BY u.id
)
SELECT up.user_id, json_agg(up.*) AS all_posts FROM user_posts up;",
)
.unwrap();
let err = analyze(&catalog, &query).unwrap_err();
assert!(
err.message.contains("nested aggregate over nested aggregate"),
"expected a clear nested-of-nested diagnostic, got: {}",
err.message
);
assert!(
err.message.contains("posts"),
"diagnostic should name the offending field, got: {}",
err.message
);
}
#[test]
fn test_union_arms_with_identical_nested_shape_widen_to_one_struct() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserPostsEitherWay
-- @returns :many
SELECT u.id, json_agg(p.*) AS posts FROM users u JOIN posts p ON p.user_id = u.id GROUP BY u.id
UNION
SELECT u2.id, json_agg(p2.*) AS posts FROM users u2 JOIN posts p2 ON p2.user_id = u2.id GROUP BY u2.id;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(
result.nested_structs.len(),
1,
"both arms describe the same posts shape and must widen to a single struct"
);
assert!(result.columns[1].neutral_type.starts_with("json_nested<array<"));
}
#[test]
fn test_union_arms_with_differing_nested_shape_is_rejected_with_clear_diagnostic() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserPostsOrComments
-- @returns :many
SELECT u.id, json_agg(p.*) AS posts FROM users u JOIN posts p ON p.user_id = u.id GROUP BY u.id
UNION
SELECT u2.id, json_agg(c.*) AS posts FROM users u2 JOIN comments c ON c.user_id = u2.id GROUP BY u2.id;",
)
.unwrap();
let err = analyze(&catalog, &query).unwrap_err();
assert!(
err.message.contains("different row shapes"),
"expected a clear shape-mismatch diagnostic, got: {}",
err.message
);
}
#[test]
fn test_union_nested_arm_against_plain_json_arm_is_rejected() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetUserPostsOrEmpty
-- @returns :many
SELECT json_agg(p.*) AS posts FROM posts p
UNION
SELECT '[]'::json AS posts;",
)
.unwrap();
let err = analyze(&catalog, &query).unwrap_err();
assert!(
err.message.contains("nested aggregate") && err.message.contains("left arm"),
"expected a nested-vs-non-nested diagnostic naming the offending side, got: {}",
err.message
);
}
#[test]
fn test_union_plain_json_arm_against_nested_arm_is_rejected() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetEmptyOrUserPosts
-- @returns :many
SELECT '[]'::json AS posts
UNION
SELECT json_agg(p.*) AS posts FROM posts p;",
)
.unwrap();
let err = analyze(&catalog, &query).unwrap_err();
assert!(
err.message.contains("nested aggregate") && err.message.contains("right arm"),
"expected a nested-vs-non-nested diagnostic naming the offending side, got: {}",
err.message
);
}
#[test]
fn test_array_agg_wildcard_never_produces_nested_struct() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetPostIdsAgg
-- @returns :one
SELECT array_agg(p.id) AS ids FROM posts p;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.columns[0].neutral_type, "array<int32>");
assert!(result.nested_structs.is_empty());
}
#[test]
fn test_json_agg_redshift_engine_falls_back_to_plain_json() {
let catalog = make_catalog().with_engine("redshift");
let query = parse_query(
"-- @name GetPostsRedshift
-- @returns :many
SELECT json_agg(p.*) AS posts FROM posts p;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(
result.columns[0].neutral_type, "json",
"a PostgreSQL-dialect catalog on the Redshift engine must not infer a nested struct"
);
assert!(result.nested_structs.is_empty());
}
#[test]
fn test_json_agg_postgresql_engine_still_infers() {
let catalog = make_catalog().with_engine("postgresql");
let query = parse_query(
"-- @name GetPostsPg
-- @returns :many
SELECT json_agg(p.*) AS posts FROM posts p;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
assert_eq!(result.nested_structs.len(), 1);
}
#[test]
fn test_cte_column_alias_list_names_literal_columns() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetPair
-- @returns :many
WITH t(a, b) AS (SELECT 1, 2) SELECT a, b FROM t;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
let names: Vec<&str> = result.columns.iter().map(|c| c.name.as_str()).collect();
assert_eq!(names, ["a", "b"]);
assert!(result.columns.iter().all(|c| c.neutral_type == "int64"));
}
#[test]
fn test_cte_column_alias_list_consumed_by_select_star() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetPairStar
-- @returns :many
WITH t(a, b) AS (SELECT 1, 2) SELECT * FROM t;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
let names: Vec<&str> = result.columns.iter().map(|c| c.name.as_str()).collect();
assert_eq!(names, ["a", "b"]);
}
#[test]
fn test_cte_column_alias_count_mismatch_is_rejected() {
let catalog = make_catalog();
let query = parse_query(
"-- @name GetMismatch
-- @returns :many
WITH t(a) AS (SELECT 1, 2) SELECT * FROM t;",
)
.unwrap();
let err = analyze(&catalog, &query).unwrap_err();
assert!(
err.message
.contains("CTE column alias list has 1 entries but the CTE body produces 2 columns"),
"expected a column-alias-count diagnostic, got: {}",
err.message
);
}
#[test]
fn test_recursive_cte_with_column_alias_list_names_columns() {
let catalog = make_catalog();
let query = parse_query(
"-- @name CountDown
-- @returns :many
WITH RECURSIVE t(n) AS (SELECT 1 UNION ALL SELECT n + 1 FROM t WHERE n < 10) SELECT n FROM t;",
)
.unwrap();
let result = analyze(&catalog, &query).unwrap();
let names: Vec<&str> = result.columns.iter().map(|c| c.name.as_str()).collect();
assert_eq!(names, ["n"]);
assert_eq!(result.columns[0].neutral_type, "int64");
}
}