use crate::driver::DriverType;
use crate::errors::AkitaError;
use crate::mapper::PaginationOptions;
use crate::sql::{BatchInsertData, DatabaseDialect, SqlBuilder};
use crate::{database_err, empty_data_err};
use akita_core::{
AkitaValue, FieldName, FieldType, GetFields, GetTableName, IdentifierType, IntoAkitaValue,
Params, QueryData, TableName, Wrapper,
};
use std::collections::HashSet;
pub(crate) struct SqlServerBuilder {
pub version: String, pub quoted_identifier: bool,
pub use_named_params: bool, }
impl Default for SqlServerBuilder {
fn default() -> Self {
Self {
version: "2016".to_string(),
quoted_identifier: true,
use_named_params: true,
}
}
}
impl SqlBuilder for SqlServerBuilder {
fn dialect(&self) -> DatabaseDialect {
DatabaseDialect::SQLServer
}
fn quote_identifier(&self, identifier: &str) -> String {
if self.quoted_identifier {
format!("[{}]", identifier.replace(']', "]]"))
} else {
format!("\"{}\"", identifier.replace('"', "\"\""))
}
}
fn quote_table(&self, table: &str) -> String {
if table.starts_with('[') && table.ends_with(']') {
return table.to_string();
}
if table.starts_with('"') && table.ends_with('"') {
return table.to_string();
}
table
.split('.')
.map(|part| {
if (part.starts_with('[') && part.ends_with(']'))
|| (part.starts_with('"') && part.ends_with('"'))
{
part.to_string()
} else {
self.quote_identifier(part)
}
})
.collect::<Vec<String>>()
.join(".")
}
fn process_placeholders(&self, sql: &str) -> String {
let mut result = String::new();
let mut param_index = 1;
for ch in sql.chars() {
if ch == '?' {
let param_name = format!("@p{}", param_index);
result.push_str(¶m_name);
param_index += 1;
} else {
result.push(ch);
}
}
result
}
fn build_insert_sql(
&self,
table: &TableName,
columns: Vec<FieldName>,
datas: Vec<AkitaValue>,
) -> crate::errors::Result<(String, Vec<AkitaValue>)> {
if datas.is_empty() {
return Err(empty_data_err!());
}
let column_names: Vec<(String, FieldName)> = columns
.into_iter()
.filter(|c| c.exist)
.filter(|c| !c.is_auto_increment()) .map(|c| {
let col_name = c.alias.as_ref().unwrap_or(&c.name);
(self.quote_identifier(col_name), c)
})
.collect();
if column_names.is_empty() {
return Err(database_err!(
"No columns to insert after filtering auto-increment fields"
));
}
let mut placeholders = Vec::new();
let mut params = Vec::new();
for data in datas.into_iter() {
for (i, (_col_name, field)) in column_names.iter().enumerate() {
let col_name = field.alias.as_ref().unwrap_or(&field.name);
let param_name = format!("@p{}", i + 1);
placeholders.push(param_name);
let mut value = data
.get_obj_value(col_name)
.cloned()
.unwrap_or(AkitaValue::Null);
if let Some(fill) = &field.fill {
match fill.mode.as_str() {
"insert" | "default" => {
value = fill.value.clone().unwrap_or_default();
}
_ => {}
}
}
value = self.identifier_generator_value(field, value);
params.push(value);
}
}
let column_names_str = column_names
.iter()
.map(|(c, _)| c.to_string())
.collect::<Vec<_>>()
.join(", ");
let sql = format!(
"INSERT INTO {} ({}) VALUES ({})",
self.quote_table(&table.complete_name()),
column_names_str,
placeholders.join(", ")
);
Ok((sql, params))
}
fn build_batch_insert_sql(
&self,
data: &BatchInsertData,
) -> crate::errors::Result<(String, Vec<AkitaValue>)> {
if data.columns.is_empty() || data.rows.is_empty() {
return Err(empty_data_err!());
}
let mut valid_column_indices = Vec::new();
let mut quoted_names = Vec::new();
for (idx, col) in data.columns.iter().enumerate() {
if col.exist && !col.is_auto_increment() {
valid_column_indices.push(idx);
let col_name = col.alias.as_ref().unwrap_or(&col.name);
quoted_names.push(self.quote_identifier(col_name));
}
}
if valid_column_indices.is_empty() {
return Err(empty_data_err!());
}
let mut all_params = Vec::with_capacity(valid_column_indices.len() * data.rows.len());
let mut all_placeholders = Vec::with_capacity(data.rows.len());
let mut param_num = 1;
for row in &data.rows {
let mut row_placeholders = Vec::with_capacity(valid_column_indices.len());
let mut row_params = Vec::with_capacity(valid_column_indices.len());
for &col_idx in &valid_column_indices {
let value = row.get(col_idx).cloned().unwrap_or(AkitaValue::Null);
row_placeholders.push(format!("@p{}", param_num));
row_params.push(value);
param_num += 1;
}
all_placeholders.push(format!("({})", row_placeholders.join(", ")));
all_params.extend(row_params);
}
let sql = format!(
"INSERT INTO {} ({}) VALUES {}",
self.quote_table(&data.table.complete_name()),
quoted_names.join(", "),
all_placeholders.join(", ")
);
Ok((sql, all_params))
}
fn build_query_sql(&self, wrapper: &Wrapper) -> (String, Vec<AkitaValue>) {
let data = wrapper.get_query_data();
if data.from.is_none() {
return ("".to_string(), vec![]);
}
let mut sql_parts = Vec::new();
let select_clause =
if self.version < "2012".to_string() && data.limit.is_some() && data.offset.is_none() {
let limit = data.limit.unwrap();
let columns = if data.select == "*" {
"*".to_string()
} else {
self.build_column_list(&data.select)
};
if data.distinct {
format!("SELECT DISTINCT TOP {} {}", limit, columns)
} else {
format!("SELECT TOP {} {}", limit, columns)
}
} else {
self.build_select_clause(&data)
};
sql_parts.push(select_clause);
sql_parts.push(format!(
"FROM {}",
self.build_from_clause(data.from.as_ref().unwrap())
));
let joins = self.build_join_clauses(&data.joins);
if !joins.is_empty() {
sql_parts.push(joins);
}
if !data.where_clause.is_empty() {
sql_parts.push(format!(
"WHERE {}",
self.build_where_clause(&data.where_clause)
));
}
if !data.group_by.is_empty() {
sql_parts.push(format!(
"GROUP BY {}",
self.build_group_by_clause(&data.group_by)
));
}
if !data.having.is_empty() {
sql_parts.push(format!("HAVING {}", self.build_having_clause(&data.having)));
}
self.build_sqlserver_pagination(&mut sql_parts, &data);
let sql = sql_parts.join(" ");
let final_sql = self.process_placeholders(&sql);
let params = wrapper.get_parameters();
(final_sql, params)
}
fn build_delete_sql(&self, table: &TableName, wrapper: &Wrapper) -> String {
let where_clause = wrapper.build_where_clause();
let mut sql = if let Some(limit_val) = wrapper.get_limit() {
let sql = format!(
"DELETE TOP({}) FROM {}",
limit_val,
self.quote_table(&table.complete_name())
);
sql
} else {
let sql = format!("DELETE FROM {}", self.quote_table(&table.complete_name()));
sql
};
if !where_clause.trim().is_empty() {
sql.push_str(&format!(
" WHERE {}",
self.build_where_clause(&where_clause)
));
}
self.process_placeholders(&sql)
}
fn build_update_sql(&self, table: &TableName, wrapper: &Wrapper) -> Option<String> {
let set_clause = wrapper
.get_set_operations()
.iter()
.map(|op| match &op.value {
AkitaValue::RawSql(sql_expr) => {
format!("{} = {}", self.quote_identifier(&op.column), sql_expr)
}
AkitaValue::Column(col_name) => {
format!("{} = {}", self.quote_identifier(&op.column), col_name)
}
_ => format!("{} = ?", self.quote_identifier(&op.column)),
})
.collect::<Vec<_>>()
.join(", ");
let where_clause = wrapper.build_where_clause();
let mut sql = format!("UPDATE {} SET {}", &table.complete_name(), set_clause);
if !where_clause.is_empty() {
sql.push_str(&format!(" WHERE {}", where_clause));
}
if let Some(limit_val) = wrapper.get_limit() {
sql.push_str(&format!(" LIMIT {}", limit_val));
}
Some(self.process_placeholders(&sql))
}
fn build_column_list(&self, columns: &str) -> String {
if columns == "*" {
return "*".to_string();
}
columns
.split(',')
.map(|col| col.trim())
.filter(|col| !col.is_empty())
.map(|col| {
if col.contains(" AS ") {
let parts: Vec<&str> = col.split(" AS ").collect();
if parts.len() == 2 {
return format!(
"{} AS {}",
self.quote_identifier(parts[0].trim()),
self.quote_identifier(parts[1].trim())
);
}
}
if col.contains('.') {
let parts: Vec<&str> = col.split('.').collect();
if parts.len() == 2 {
return format!(
"{}.{}",
self.quote_identifier(parts[0]),
self.quote_identifier(parts[1])
);
}
}
self.quote_identifier(col)
})
.collect::<Vec<_>>()
.join(", ")
}
fn is_reserved_keyword(&self, identifier: &str) -> bool {
let keywords = [
"ADD",
"ALL",
"ALTER",
"AND",
"ANY",
"AS",
"ASC",
"AUTHORIZATION",
"BACKUP",
"BEGIN",
"BETWEEN",
"BREAK",
"BROWSE",
"BULK",
"BY",
"CASCADE",
"CASE",
"CHECK",
"CHECKPOINT",
"CLOSE",
"CLUSTERED",
"COALESCE",
"COLLATE",
"COLUMN",
"COMMIT",
"COMPUTE",
"CONSTRAINT",
"CONTAINS",
"CONTAINSTABLE",
"CONTINUE",
"CONVERT",
"CREATE",
"CROSS",
"CURRENT",
"CURRENT_DATE",
"CURRENT_TIME",
"CURRENT_TIMESTAMP",
"CURRENT_USER",
"CURSOR",
"DATABASE",
"DBCC",
"DEALLOCATE",
"DECLARE",
"DEFAULT",
"DELETE",
"DENY",
"DESC",
"DISK",
"DISTINCT",
"DISTRIBUTED",
"DOUBLE",
"DROP",
"DUMP",
"ELSE",
"END",
"ERRLVL",
"ESCAPE",
"EXCEPT",
"EXEC",
"EXECUTE",
"EXISTS",
"EXIT",
"EXTERNAL",
"FETCH",
"FILE",
"FILLFACTOR",
"FOR",
"FOREIGN",
"FREETEXT",
"FREETEXTTABLE",
"FROM",
"FULL",
"FUNCTION",
"GOTO",
"GRANT",
"GROUP",
"HAVING",
"HOLDLOCK",
"IDENTITY",
"IDENTITY_INSERT",
"IDENTITYCOL",
"IF",
"IN",
"INDEX",
"INNER",
"INSERT",
"INTERSECT",
"INTO",
"IS",
"JOIN",
"KEY",
"KILL",
"LEFT",
"LIKE",
"LINENO",
"LOAD",
"MERGE",
"NATIONAL",
"NOCHECK",
"NONCLUSTERED",
"NOT",
"NULL",
"NULLIF",
"OF",
"OFF",
"OFFSETS",
"ON",
"OPEN",
"OPENDATASOURCE",
"OPENQUERY",
"OPENROWSET",
"OPENXML",
"OPTION",
"OR",
"ORDER",
"OUTER",
"OVER",
"PERCENT",
"PIVOT",
"PLAN",
"PRECISION",
"PRIMARY",
"PRINT",
"PROC",
"PROCEDURE",
"PUBLIC",
"RAISERROR",
"READ",
"READTEXT",
"RECONFIGURE",
"REFERENCES",
"REPLICATION",
"RESTORE",
"RESTRICT",
"RETURN",
"REVERT",
"REVOKE",
"RIGHT",
"ROLLBACK",
"ROWCOUNT",
"ROWGUIDCOL",
"RULE",
"SAVE",
"SCHEMA",
"SECURITYAUDIT",
"SELECT",
"SEMANTICKEYPHRASETABLE",
"SEMANTICSIMILARITYDETAILSTABLE",
"SEMANTICSIMILARITYTABLE",
"SESSION_USER",
"SET",
"SETUSER",
"SHUTDOWN",
"SOME",
"STATISTICS",
"SYSTEM_USER",
"TABLE",
"TABLESAMPLE",
"TEXTSIZE",
"THEN",
"TO",
"TOP",
"TRAN",
"TRANSACTION",
"TRIGGER",
"TRUNCATE",
"TRY_CONVERT",
"TSEQUAL",
"UNION",
"UNIQUE",
"UNPIVOT",
"UPDATE",
"UPDATETEXT",
"USE",
"USER",
"VALUES",
"VARYING",
"VIEW",
"WAITFOR",
"WHEN",
"WHERE",
"WHILE",
"WITH",
"WITHIN GROUP",
"WRITETEXT",
];
keywords.contains(&identifier.to_uppercase().as_str())
}
}
impl SqlServerBuilder {
fn make_param_name(&self, _column_name: &str, index: usize) -> String {
format!("@p{}", index)
}
fn normalize_param_name(&self, name: &str) -> String {
let mut result = String::new();
for ch in name.chars() {
if ch.is_alphanumeric() || ch == '_' {
result.push(ch);
} else {
result.push('_');
}
}
if result.chars().next().map_or(false, |c| c.is_numeric()) {
format!("p{}", result)
} else {
result
}
}
fn process_query_placeholders(&self, sql: &str, _column_names: &[&str]) -> String {
self.process_placeholders(sql)
}
fn build_sqlserver_pagination(&self, sql_parts: &mut Vec<String>, data: &QueryData) {
if self.version >= "2012".to_string() && (data.limit.is_some() || data.offset.is_some()) {
if data.order_by.is_empty() {
sql_parts.push("ORDER BY (SELECT NULL)".to_string());
} else {
sql_parts.push(format!(
"ORDER BY {}",
self.build_order_by_clause(&data.order_by)
));
}
if let Some(offset) = data.offset {
sql_parts.push(format!("OFFSET {} ROWS", offset));
} else if data.limit.is_some() {
sql_parts.push("OFFSET 0 ROWS".to_string());
}
if let Some(limit) = data.limit {
sql_parts.push(format!("FETCH NEXT {} ROWS ONLY", limit));
}
} else if !data.order_by.is_empty() {
sql_parts.push(format!(
"ORDER BY {}",
self.build_order_by_clause(&data.order_by)
));
}
}
fn build_insert_returning(&self, _table: &str, id_column: &str) -> Option<String> {
Some(format!(
" OUTPUT INSERTED.{}",
self.quote_identifier(id_column)
))
}
}
#[test]
fn test_mssql_sqlbuilder() {
let builder = SqlServerBuilder::default();
let field_id = FieldName {
name: "user_id".to_string(),
table: "user".to_string().into(),
alias: None,
exist: true,
select: false,
fill: None,
field_type: FieldType::TableId(IdentifierType::Auto),
};
let columns = vec![
field_id.clone(),
FieldName::from("user_name"),
FieldName::from("email_address"),
];
let mut imap = indexmap::IndexMap::new();
imap.insert("id".to_string(), AkitaValue::Int(1));
imap.insert(
"user_name".to_string(),
AkitaValue::Text("John".to_string()),
);
imap.insert(
"email_address".to_string(),
AkitaValue::Text("john@example.com".to_string()),
);
let data = AkitaValue::Object(imap);
let (sql, params) = builder
.build_insert_sql(&TableName::from("users"), columns, vec![data])
.unwrap();
println!(
"build_insert_sql mssql :{} \nparams:{}",
sql,
Params::Positional(params)
);
let wrapper = Wrapper::new()
.table("users")
.eq("user_id", 1)
.like("user_name", "%john%");
let (query_sql, query_params) = builder.build_query_sql(&wrapper);
println!(
"build_query_sql mssql :{} \n params:{}",
query_sql,
Params::Positional(query_params)
);
let columns = vec![field_id, FieldName::from("user_name")];
let rows = vec![
vec![AkitaValue::Int(1), AkitaValue::Text("John".to_string())],
vec![AkitaValue::Int(2), AkitaValue::Text("Jane".to_string())],
];
let batch_data = BatchInsertData {
table: TableName::from("users"),
columns,
rows,
id_field: None,
};
let (batch_sql, batch_params) = builder.build_batch_insert_sql(&batch_data).unwrap();
println!(
"batch_sql mssql :{} \n params:{}",
batch_sql,
Params::Positional(batch_params)
);
}