use crate::ast::write_payload::{
check_insert_shape, insert_columns, insert_values, simple_write_column,
};
use crate::ast::*;
use crate::transpiler::SqlGenerator;
use crate::transpiler::conditions::{ConditionToSql, output_expr_sql, returning_clause_sql};
use crate::transpiler::dialect::Dialect;
pub(crate) fn shape_error_comment(error: &str) -> String {
format!(
"/* ERROR: {} */",
error.replace("*/", "* /").replace("/*", "/ *")
)
}
pub fn build_insert(cmd: &Qail, dialect: Dialect) -> String {
if cmd.validate_applied_insert_scope().is_err() {
return crate::transpiler::INVALID_INSERT_SCOPE_SQL.to_string();
}
if let Err(error) = check_insert_shape(cmd, simple_write_column, |message| message) {
return shape_error_comment(&error);
}
let generator = dialect.generator();
let mut sql = super::cte::build_write_with_prefix(cmd, dialect);
sql.push_str("INSERT INTO ");
sql.push_str(&generator.quote_identifier(&cmd.table));
let cols: Vec<String> = insert_columns(cmd)
.into_iter()
.map(|c| render_insert_column(c, generator.as_ref()))
.collect();
if !cols.is_empty() {
sql.push_str(" (");
sql.push_str(&cols.join(", "));
sql.push(')');
}
if let Some(ref overriding) = cmd.overriding {
match overriding {
OverridingKind::SystemValue => sql.push_str(" OVERRIDING SYSTEM VALUE"),
OverridingKind::UserValue => sql.push_str(" OVERRIDING USER VALUE"),
}
}
if cmd.default_values {
sql.push_str(" DEFAULT VALUES");
} else if let Some(ref source_query) = cmd.source_query {
use crate::transpiler::ToSql;
sql.push(' ');
sql.push_str(&source_query.to_sql_with_dialect(dialect));
} else {
let values: Vec<String> = insert_values(cmd)
.iter()
.map(|c| c.to_value_sql(generator.as_ref()))
.collect();
sql.push_str(" VALUES (");
sql.push_str(&values.join(", "));
sql.push(')');
}
if let Some(on_conflict) = &cmd.on_conflict {
sql.push_str(&build_on_conflict(
on_conflict,
&dialect,
generator.as_ref(),
));
}
sql.push_str(&returning_clause_sql(cmd, generator.as_ref(), |expr| {
output_expr_sql(expr, generator.as_ref())
}));
sql
}
fn build_on_conflict(
on_conflict: &OnConflict,
_dialect: &Dialect,
generator: &dyn SqlGenerator,
) -> String {
build_on_conflict_postgres(on_conflict, generator)
}
fn build_on_conflict_postgres(on_conflict: &OnConflict, generator: &dyn SqlGenerator) -> String {
let mut sql = String::from(" ON CONFLICT");
match (&on_conflict.constraint, on_conflict.columns.is_empty()) {
(Some(_), false) => {
sql.push_str(" /* ERROR: conflict target has both columns and a constraint */");
}
(Some(constraint), true) if constraint.is_empty() || constraint.contains(['.', '\0']) => {
sql.push_str(" /* ERROR: Invalid conflict constraint name */");
}
(Some(constraint), true) => {
sql.push_str(" ON CONSTRAINT ");
sql.push_str(&generator.quote_identifier(constraint));
}
(None, true) => {}
(None, false) => {
let cols: Vec<String> = on_conflict
.columns
.iter()
.map(|c| generator.quote_identifier(c))
.collect();
sql.push_str(" (");
sql.push_str(&cols.join(", "));
sql.push(')');
}
}
match &on_conflict.action {
ConflictAction::DoNothing => {
sql.push_str(" DO NOTHING");
}
ConflictAction::DoUpdate { assignments } => {
sql.push_str(" DO UPDATE SET ");
let sets: Vec<String> = assignments
.iter()
.map(|(col, expr)| {
format!(
"{} = {}",
generator.quote_identifier(col),
render_sql_expr(expr, generator)
)
})
.collect();
sql.push_str(&sets.join(", "));
if !on_conflict.where_conditions.is_empty() {
let preds: Vec<String> = on_conflict
.where_conditions
.iter()
.map(|c| c.to_sql(generator, None))
.collect();
sql.push_str(" WHERE ");
sql.push_str(&preds.join(" AND "));
}
}
}
sql
}
fn render_insert_column(expr: &Expr, generator: &dyn SqlGenerator) -> String {
match expr {
Expr::Named(name) => generator.quote_identifier(name),
_ => "/* ERROR: Invalid insert column */".to_string(),
}
}
fn render_sql_expr(expr: &Expr, generator: &dyn SqlGenerator) -> String {
match expr {
Expr::Star => "*".to_string(),
Expr::Named(name) => render_named_expr(name, generator),
Expr::Literal(value) => value.to_string(),
Expr::Binary {
left, op, right, ..
} => match op {
op if op.is_postfix() => format!("({} {})", render_sql_expr(left, generator), op),
_ => op.infix_sql(
&render_sql_expr(left, generator),
&render_sql_expr(right, generator),
),
},
Expr::FunctionCall { name, args, .. } => {
let Some(function) = render_function_name(name) else {
return "/* ERROR: Invalid function name */".to_string();
};
match crate::transpiler::render_function_args(args, generator, |arg| {
render_sql_expr(arg, generator)
}) {
Ok(args) => format!("{function}({args})"),
Err(error) => error,
}
}
Expr::FunctionArg { .. } => crate::transpiler::MISPLACED_FUNCTION_ARG_SQL.to_string(),
Expr::Cast {
expr, target_type, ..
} => {
let Some(target_type) = checked_sql_type_fragment(target_type) else {
return "/* ERROR: Invalid cast target type */".to_string();
};
format!("{}::{}", render_sql_expr(expr, generator), target_type)
}
Expr::JsonAccess {
column,
path_segments,
..
} => render_json_access(column, path_segments, generator),
Expr::Collate {
expr, collation, ..
} => format!(
"{} COLLATE {}",
render_sql_expr(expr, generator),
render_qualified_identifier(collation, generator)
),
Expr::FieldAccess { expr, field, .. } => format!(
"({}).{}",
render_sql_expr(expr, generator),
render_qualified_identifier(field, generator)
),
_ => "/* ERROR: Invalid expression */".to_string(),
}
}
fn render_named_expr(name: &str, generator: &dyn SqlGenerator) -> String {
if name == "*"
|| name.starts_with('\'')
|| name.starts_with('"')
|| name.starts_with(':')
|| name.starts_with('$')
|| name.parse::<f64>().is_ok()
|| name.eq_ignore_ascii_case("NULL")
|| name.eq_ignore_ascii_case("TRUE")
|| name.eq_ignore_ascii_case("FALSE")
{
name.to_string()
} else {
generator.quote_identifier(name)
}
}
fn render_function_name(name: &str) -> Option<String> {
if name.is_empty()
|| name.contains('\0')
|| name.split('.').any(str::is_empty)
|| !name
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'_' || b == b'.')
{
None
} else {
Some(name.to_uppercase())
}
}
fn checked_sql_type_fragment(fragment: &str) -> Option<String> {
let fragment = fragment.trim();
if fragment.is_empty()
|| fragment.contains('\0')
|| fragment.contains(';')
|| fragment.contains('\'')
|| fragment.contains('"')
|| fragment.contains("--")
|| fragment.contains("/*")
|| fragment.contains("*/")
|| !fragment.bytes().all(|b| {
b.is_ascii_alphanumeric()
|| matches!(
b,
b'_' | b'.' | b' ' | b'(' | b')' | b',' | b'[' | b']' | b'%' | b'+' | b'-'
)
})
{
None
} else {
Some(fragment.to_string())
}
}
fn render_qualified_identifier(value: &str, generator: &dyn SqlGenerator) -> String {
if value.is_empty() || value.as_bytes().contains(&0) || value.split('.').any(str::is_empty) {
"/* ERROR: Invalid identifier */".to_string()
} else {
generator.quote_identifier(value)
}
}
fn render_json_access(
column: &str,
path_segments: &[(JsonPathSegment, bool)],
generator: &dyn SqlGenerator,
) -> String {
let mut sql = generator.quote_identifier(column);
for (segment, as_text) in path_segments {
let op = if *as_text { "->>" } else { "->" };
sql.push_str(&format!("{}{}", op, segment));
}
sql
}