use std::collections::HashMap;
use crate::{
ast::{DataType, Expr, InsertSource, InsertStmt, Value},
binder::{BindError, Binder, bound::insert::BoundInsertStmt, expr::eval_expr},
catalog::objects::ColumnEntry,
common::symbol::Symbol,
};
impl<'c> Binder<'c> {
pub fn bind_insert_table(
&self,
db: Symbol,
default_schema: Symbol,
stmt: InsertStmt,
) -> Result<BoundInsertStmt, BindError> {
if stmt.on_conflict.is_some() {
return Err(BindError::UnsupportedOnConflict);
}
if !stmt.returning.is_empty() {
return Err(BindError::UnsupportedReturning);
}
let (schema, table_name) = stmt.table.resolve_schema_table(default_schema);
let table = self
.catalog
.get_table(db, schema, table_name)
.map_err(|_| BindError::TableNotFound(table_name))?;
let table_columns = table.columns.clone();
let rows = match stmt.source {
InsertSource::Values(rows) => rows,
InsertSource::Select(_) | InsertSource::DefaultValues => {
return Err(BindError::UnsupportedInsertSource);
}
};
let target_columns: Vec<&ColumnEntry> = if stmt.columns.is_empty() {
table_columns.iter().collect()
} else {
let mut resolved = Vec::with_capacity(stmt.columns.len());
for col_name in &stmt.columns {
let entry = table_columns
.iter()
.find(|c| c.name == *col_name)
.ok_or(BindError::ColumnNotFound(*col_name))?;
resolved.push(entry);
}
resolved
};
let mut bound_rows = Vec::with_capacity(rows.len());
let col_map: HashMap<Symbol, usize> = table_columns
.iter()
.enumerate()
.map(|(i, col)| (col.name, i))
.collect();
for col in &table_columns {
let was_targeted = target_columns.iter().any(|c| c.name == col.name);
if !was_targeted && (!col.nullable || col.is_primary_key) {
return Err(BindError::MissingNotNullColumn(col.name));
}
}
let defaults: Vec<Option<Value>> = table_columns
.iter()
.map(|col| match &col.default {
Some(Expr::Literal(v)) => Some(v.clone()),
_ => None,
})
.collect();
for (row_idx, row) in rows.into_iter().enumerate() {
if row.len() != target_columns.len() {
return Err(BindError::ColumnCountMismatch {
expected: target_columns.len(),
found: row.len(),
});
}
let mut user_values = Vec::with_capacity(row.len());
for expr in row {
user_values.push(eval_expr(&expr, &self.catalog.interner)?);
}
for (target_col, value) in target_columns.iter().zip(user_values.iter()) {
check_type_compat(&target_col.data_type, value, target_col.name, row_idx)?;
}
let mut full_row: Vec<Value> = defaults
.iter()
.map(|d| d.clone().unwrap_or(Value::Null))
.collect();
for (target_col, value) in target_columns.iter().zip(user_values.into_iter()) {
let pos = col_map[&target_col.name];
full_row[pos] = value;
}
bound_rows.push(full_row);
}
Ok(BoundInsertStmt {
db,
schema,
table: table_name,
rows: bound_rows,
})
}
}
fn check_type_compat(
data_type: &DataType,
value: &Value,
col: Symbol,
row_idx: usize,
) -> Result<(), BindError> {
match (data_type, value) {
(_, Value::Null) => Ok(()),
(DataType::Float | DataType::Double, Value::Int(_)) => Ok(()),
(DataType::SmallInt, Value::Int(n)) => {
if *n >= i16::MIN as i64 && *n <= i16::MAX as i64 {
Ok(())
} else {
Err(BindError::TypeMismatch {
col,
row: row_idx,
expected: "SMALLINT (-32768..32767)",
got: "integer out of range",
})
}
}
(DataType::Int, Value::Int(_)) => Ok(()),
(DataType::BigInt, Value::Int(_)) => Ok(()),
(DataType::Boolean, Value::Boolean(_)) => Ok(()),
(DataType::Float, Value::Float(_)) => Ok(()),
(DataType::Double, Value::Float(_)) => Ok(()),
(DataType::VarChar(_) | DataType::Char(_) | DataType::Text, Value::String(_)) => Ok(()),
(expected, actual) => Err(BindError::TypeMismatch {
col,
row: row_idx,
expected: type_name(expected),
got: value_kind(actual),
}),
}
}
fn type_name(dt: &DataType) -> &'static str {
match dt {
DataType::SmallInt => "SMALLINT",
DataType::Int => "INT",
DataType::BigInt => "BIGINT",
DataType::Float => "FLOAT",
DataType::Double => "DOUBLE",
DataType::Boolean => "BOOLEAN",
DataType::VarChar(_) => "VARCHAR",
DataType::Char(_) => "CHAR",
DataType::Text => "TEXT",
_ => "unsupported type",
}
}
fn value_kind(v: &Value) -> &'static str {
match v {
Value::Int(_) => "integer",
Value::Float(_) => "float",
Value::String(_) => "string",
Value::Boolean(_) => "boolean",
Value::BitString(_) => "bit string",
Value::HexString(_) => "hex string",
Value::Null => "NULL",
}
}