use std::collections::HashMap;
use crate::{
condition::Condition,
mssql::util::{generate_where_condition_str, remove_quotes_and_backslashes},
};
use log::{debug, info};
use crate::table::Table;
use super::{select::SelectQueryBuilder, Connection};
pub fn update<T: Table + Default>(conn: &mut Connection, table: T) -> UpdateQueryBuilder<T> {
UpdateQueryBuilder::new(conn, table)
}
pub struct UpdateQueryBuilder<'a, T: Table + Default> {
conn: &'a mut Connection,
table: Option<T>,
columns: Vec<String>,
sub_queries: HashMap<String, SelectQueryBuilder<'a, T>>,
where_condition: Option<Condition<'a>>,
}
impl<'a, T: Table + Default> UpdateQueryBuilder<'a, T> {
pub fn new(conn: &'a mut Connection, table: T) -> Self {
UpdateQueryBuilder {
conn,
table: Some(table),
columns: Vec::new(),
sub_queries: HashMap::new(),
where_condition: None,
}
}
pub fn set(mut self, columns: Vec<String>) -> Self {
self.columns = columns;
self
}
pub fn set_subqueries(mut self, columns: HashMap<String, SelectQueryBuilder<'a, T>>) -> Self {
self.sub_queries = columns;
self
}
pub fn where_clause(mut self, condition: Condition<'a>) -> Self {
self.where_condition = Some(condition);
self
}
pub async fn build(self) -> Result<String, String> {
let table_name = self
.table
.as_ref()
.map(|t| t.get_name().to_string())
.unwrap_or("".to_string());
let table_name_str = remove_quotes_and_backslashes(&table_name);
let set = if let Some(table) = &self.table {
let mut set_fields = Vec::new();
let fields = table.get_column_fields();
let values = table.get_column_values();
for column in &self.columns {
if let Some(index) = fields.iter().position(|c| column == c) {
let value = values.get(index).cloned().unwrap_or_default();
let formatted_value = if value.is_empty() {
"NULL".to_string()
} else if value.parse::<f64>().is_ok() {
value
} else {
format!("'{}'", value)
};
set_fields.push(format!("{} = {}", column, formatted_value));
} else {
eprintln!("Column '{}' does not exist in the table", column);
}
}
for (column_name, sub_query) in &self.sub_queries {
let formatted_value = format!("({})", sub_query.build_query());
set_fields.push(format!("{} = {}", column_name, formatted_value));
}
set_fields.join(", ")
} else {
String::new()
};
let where_condition_str = generate_where_condition_str(self.where_condition);
let query = format!(
"UPDATE {} SET {} {}",
table_name_str, set, where_condition_str,
);
debug!("{}", query);
match self.conn.client.execute(query.as_str(), &[]).await {
Ok(_) => Ok("Success!".to_string()),
Err(_) => Err("Could not execute...".to_string()),
}
}
}