use super::connection::{DatabaseBackend, DatabaseConnection};
use super::{Model, QuerySet};
use reinhardt_query::prelude::{
Alias, ColumnRef, DeleteStatement, Expr, ExprTrait, Func, InsertStatement, MySqlQueryBuilder,
PostgresQueryBuilder, Query, QueryBuilder, SelectStatement, SqliteQueryBuilder,
UpdateStatement, Values,
};
use std::collections::HashMap;
use std::marker::PhantomData;
use std::sync::Arc;
use tokio::sync::RwLock;
use uuid::Uuid;
fn build_insert_sql(stmt: &InsertStatement, backend: DatabaseBackend) -> (String, Values) {
match backend {
DatabaseBackend::Postgres => PostgresQueryBuilder.build_insert(stmt),
DatabaseBackend::MySql => MySqlQueryBuilder.build_insert(stmt),
DatabaseBackend::Sqlite => SqliteQueryBuilder.build_insert(stmt),
}
}
fn build_update_sql(stmt: &UpdateStatement, backend: DatabaseBackend) -> (String, Values) {
match backend {
DatabaseBackend::Postgres => PostgresQueryBuilder.build_update(stmt),
DatabaseBackend::MySql => MySqlQueryBuilder.build_update(stmt),
DatabaseBackend::Sqlite => SqliteQueryBuilder.build_update(stmt),
}
}
fn build_select_sql(stmt: &SelectStatement, backend: DatabaseBackend) -> (String, Values) {
match backend {
DatabaseBackend::Postgres => PostgresQueryBuilder.build_select(stmt),
DatabaseBackend::MySql => MySqlQueryBuilder.build_select(stmt),
DatabaseBackend::Sqlite => SqliteQueryBuilder.build_select(stmt),
}
}
fn select_to_string(stmt: &SelectStatement, backend: DatabaseBackend) -> String {
build_select_sql(stmt, backend).0
}
fn insert_to_string(stmt: &InsertStatement, backend: DatabaseBackend) -> String {
build_insert_sql(stmt, backend).0
}
fn build_delete_sql(stmt: &DeleteStatement, backend: DatabaseBackend) -> (String, Values) {
match backend {
DatabaseBackend::Postgres => PostgresQueryBuilder.build_delete(stmt),
DatabaseBackend::MySql => MySqlQueryBuilder.build_delete(stmt),
DatabaseBackend::Sqlite => SqliteQueryBuilder.build_delete(stmt),
}
}
static DB: once_cell::sync::OnceCell<Arc<RwLock<Option<DatabaseConnection>>>> =
once_cell::sync::OnceCell::new();
pub async fn init_database(url: &str) -> reinhardt_core::exception::Result<()> {
init_database_with_pool_size(url, None).await
}
pub async fn init_database_with_pool_size(
url: &str,
pool_size: Option<u32>,
) -> reinhardt_core::exception::Result<()> {
let conn = DatabaseConnection::connect_with_pool_size(url, pool_size).await?;
DB.get_or_init(|| Arc::new(RwLock::new(Some(conn))));
Ok(())
}
pub async fn reinitialize_database(url: &str) -> reinhardt_core::exception::Result<()> {
reinitialize_database_with_pool_size(url, None).await
}
pub async fn reinitialize_database_with_pool_size(
url: &str,
pool_size: Option<u32>,
) -> reinhardt_core::exception::Result<()> {
let conn = DatabaseConnection::connect_with_pool_size(url, pool_size).await?;
if let Some(db_cell) = DB.get() {
let mut guard = db_cell.write().await;
*guard = Some(conn);
} else {
DB.get_or_init(|| Arc::new(RwLock::new(Some(conn))));
}
Ok(())
}
pub async fn get_connection() -> reinhardt_core::exception::Result<DatabaseConnection> {
let db = DB.get().ok_or_else(|| {
reinhardt_core::exception::Error::Database("Database not initialized".to_string())
})?;
let guard = db.read().await;
guard.clone().ok_or_else(|| {
reinhardt_core::exception::Error::Database("Database connection not available".to_string())
})
}
pub struct Manager<M: Model> {
_marker: PhantomData<M>,
}
impl<M: Model> Manager<M> {
pub fn new() -> Self {
Self {
_marker: PhantomData,
}
}
pub fn all(&self) -> QuerySet<M> {
QuerySet::new()
}
pub fn filter(&self, filter: impl Into<super::query::FilterCondition>) -> QuerySet<M> {
QuerySet::new().filter(filter)
}
pub fn get(&self, pk: M::PrimaryKey) -> QuerySet<M> {
let pk_field = M::primary_key_field();
let pk_value = M::primary_key_filter_value(pk);
let filter = super::query::Filter::new(
pk_field.to_string(),
super::query::FilterOperator::Eq,
pk_value,
);
QuerySet::new().filter(filter)
}
pub fn limit(&self, limit: usize) -> QuerySet<M> {
QuerySet::new().limit(limit)
}
pub fn order_by(&self, fields: &[&str]) -> QuerySet<M> {
QuerySet::new().order_by(fields)
}
pub fn annotate(&self, annotation: super::annotation::Annotation) -> QuerySet<M> {
QuerySet::new().annotate(annotation)
}
pub fn defer(&self, fields: &[&str]) -> QuerySet<M> {
QuerySet::new().defer(fields)
}
pub fn only(&self, fields: &[&str]) -> QuerySet<M> {
QuerySet::new().only(fields)
}
pub fn values(&self, fields: &[&str]) -> QuerySet<M> {
QuerySet::new().values(fields)
}
pub fn select_related(&self, fields: &[&str]) -> QuerySet<M> {
QuerySet::new().select_related(fields)
}
pub fn offset(&self, offset: usize) -> QuerySet<M> {
QuerySet::new().offset(offset)
}
pub fn paginate(&self, page: usize, page_size: usize) -> QuerySet<M> {
QuerySet::new().paginate(page, page_size)
}
pub fn prefetch_related(&self, fields: &[&str]) -> QuerySet<M> {
QuerySet::new().prefetch_related(fields)
}
pub fn values_list(&self, fields: &[&str]) -> QuerySet<M> {
QuerySet::new().values_list(fields)
}
pub fn filter_array_overlap(&self, field: &str, values: &[&str]) -> QuerySet<M> {
QuerySet::new().filter_array_overlap(field, values)
}
pub fn filter_array_contains(&self, field: &str, values: &[&str]) -> QuerySet<M> {
QuerySet::new().filter_array_contains(field, values)
}
pub fn filter_jsonb_contains(&self, field: &str, json: &str) -> QuerySet<M> {
QuerySet::new().filter_jsonb_contains(field, json)
}
pub fn filter_jsonb_key_exists(&self, field: &str, key: &str) -> QuerySet<M> {
QuerySet::new().filter_jsonb_key_exists(field, key)
}
pub fn filter_range_contains(&self, field: &str, value: &str) -> QuerySet<M> {
QuerySet::new().filter_range_contains(field, value)
}
pub fn filter_in_subquery<R: super::Model, F>(&self, field: &str, subquery_fn: F) -> QuerySet<M>
where
F: FnOnce(QuerySet<R>) -> QuerySet<R>,
{
QuerySet::new().filter_in_subquery(field, subquery_fn)
}
pub fn filter_not_in_subquery<R: super::Model, F>(
&self,
field: &str,
subquery_fn: F,
) -> QuerySet<M>
where
F: FnOnce(QuerySet<R>) -> QuerySet<R>,
{
QuerySet::new().filter_not_in_subquery(field, subquery_fn)
}
pub fn filter_exists<R: super::Model, F>(&self, subquery_fn: F) -> QuerySet<M>
where
F: FnOnce(QuerySet<R>) -> QuerySet<R>,
{
QuerySet::new().filter_exists(subquery_fn)
}
pub fn filter_not_exists<R: super::Model, F>(&self, subquery_fn: F) -> QuerySet<M>
where
F: FnOnce(QuerySet<R>) -> QuerySet<R>,
{
QuerySet::new().filter_not_exists(subquery_fn)
}
pub fn with_cte(&self, cte: super::cte::CTE) -> QuerySet<M> {
QuerySet::new().with_cte(cte)
}
pub fn full_text_search(&self, field: &str, query: &str) -> QuerySet<M> {
QuerySet::new().full_text_search(field, query)
}
pub fn annotate_subquery<R, F>(&self, name: &str, builder: F) -> QuerySet<M>
where
R: super::Model + 'static,
F: FnOnce(QuerySet<R>) -> QuerySet<R>,
{
QuerySet::new().annotate_subquery(name, builder)
}
pub async fn get_composite(
&self,
pk_values: &std::collections::HashMap<String, super::composite_pk::PkValue>,
) -> reinhardt_core::exception::Result<M>
where
M: Clone + serde::de::DeserializeOwned,
{
QuerySet::new().get_composite(pk_values).await
}
pub async fn create(&self, model: &M) -> reinhardt_core::exception::Result<M> {
let conn = get_connection().await?;
self.create_with_conn(&conn, model).await
}
pub async fn create_with_conn(
&self,
conn: &DatabaseConnection,
model: &M,
) -> reinhardt_core::exception::Result<M> {
let json = serde_json::to_value(model)
.map_err(|e| reinhardt_core::exception::Error::Database(e.to_string()))?;
let obj = json.as_object().ok_or_else(|| {
reinhardt_core::exception::Error::Database("Model must serialize to object".to_string())
})?;
let mut stmt = Query::insert();
stmt.into_table(Alias::new(M::table_name()));
let pk_field = M::primary_key_field();
let (fields, values): (Vec<_>, Vec<_>) = obj
.iter()
.filter(|(k, v)| {
let key = k.as_str();
if key == pk_field {
if v.is_null() {
return false;
}
if let Some(n) = v.as_i64() {
return n != 0;
}
}
if v.is_null()
&& (key == "created_at"
|| key == "updated_at"
|| key.ends_with("_date")
|| key.ends_with("_time")
|| key.ends_with("_at"))
{
return false;
}
true
})
.map(|(k, v)| {
let value = if v.is_null() {
reinhardt_query::value::Value::Int(None)
} else {
Self::json_to_sea_value(v)
};
(Alias::new(k.as_str()), value)
})
.unzip();
stmt.columns(fields);
stmt.values_panic(values);
let all_columns: Vec<_> = obj.keys().map(|k| Alias::new(k.as_str())).collect();
stmt.returning(all_columns);
let (sql, values) = build_insert_sql(&stmt, conn.backend());
let values: Vec<_> = values
.0
.into_iter()
.map(Self::sea_value_to_query_value)
.collect();
let row = conn.query_one(&sql, values).await?;
serde_json::from_value(row.data.clone())
.map_err(|e| reinhardt_core::exception::Error::Database(e.to_string()))
}
fn json_to_sea_value(v: &serde_json::Value) -> reinhardt_query::value::Value {
match v {
serde_json::Value::Null => reinhardt_query::value::Value::Int(None),
serde_json::Value::Bool(b) => reinhardt_query::value::Value::Bool(Some(*b)),
serde_json::Value::Number(n) => {
if let Some(i) = n.as_i64() {
reinhardt_query::value::Value::BigInt(Some(i))
} else if let Some(f) = n.as_f64() {
reinhardt_query::value::Value::Double(Some(f))
} else {
reinhardt_query::value::Value::Int(None)
}
}
serde_json::Value::String(s) => {
if let Ok(uuid) = Uuid::parse_str(s) {
return reinhardt_query::value::Value::Uuid(Some(Box::new(uuid)));
}
if let Ok(dt) = chrono::DateTime::parse_from_rfc3339(s) {
return reinhardt_query::value::Value::ChronoDateTimeUtc(Some(Box::new(
dt.with_timezone(&chrono::Utc),
)));
}
if let Ok(dt) = s.parse::<chrono::DateTime<chrono::Utc>>() {
return reinhardt_query::value::Value::ChronoDateTimeUtc(Some(Box::new(dt)));
}
if let Ok(dt) = s.parse::<chrono::DateTime<chrono::FixedOffset>>() {
return reinhardt_query::value::Value::ChronoDateTimeUtc(Some(Box::new(
dt.with_timezone(&chrono::Utc),
)));
}
reinhardt_query::value::Value::String(Some(Box::new(s.clone())))
}
serde_json::Value::Array(arr) => {
let values: Vec<reinhardt_query::value::Value> =
arr.iter().map(|v| Self::json_to_sea_value(v)).collect();
reinhardt_query::value::Value::Array(
reinhardt_query::value::ArrayType::String,
Some(Box::new(values)),
)
}
serde_json::Value::Object(_obj) => {
reinhardt_query::value::Value::Json(Some(Box::new(v.clone())))
}
}
}
fn sea_value_to_query_value(v: reinhardt_query::value::Value) -> super::connection::QueryValue {
use super::connection::QueryValue;
match v {
reinhardt_query::value::Value::Bool(Some(b)) => QueryValue::Bool(b),
reinhardt_query::value::Value::Bool(None) => QueryValue::Null,
reinhardt_query::value::Value::TinyInt(Some(i)) => QueryValue::Int(i as i64),
reinhardt_query::value::Value::TinyInt(None) => QueryValue::Null,
reinhardt_query::value::Value::SmallInt(Some(i)) => QueryValue::Int(i as i64),
reinhardt_query::value::Value::SmallInt(None) => QueryValue::Null,
reinhardt_query::value::Value::Int(Some(i)) => QueryValue::Int(i as i64),
reinhardt_query::value::Value::Int(None) => QueryValue::Null,
reinhardt_query::value::Value::BigInt(Some(i)) => QueryValue::Int(i),
reinhardt_query::value::Value::BigInt(None) => QueryValue::Null,
reinhardt_query::value::Value::TinyUnsigned(Some(u)) => QueryValue::Int(u as i64),
reinhardt_query::value::Value::TinyUnsigned(None) => QueryValue::Null,
reinhardt_query::value::Value::SmallUnsigned(Some(u)) => QueryValue::Int(u as i64),
reinhardt_query::value::Value::SmallUnsigned(None) => QueryValue::Null,
reinhardt_query::value::Value::Unsigned(Some(u)) => QueryValue::Int(u as i64),
reinhardt_query::value::Value::Unsigned(None) => QueryValue::Null,
reinhardt_query::value::Value::BigUnsigned(Some(u)) => QueryValue::Int(u as i64),
reinhardt_query::value::Value::BigUnsigned(None) => QueryValue::Null,
reinhardt_query::value::Value::Float(Some(f)) => QueryValue::Float(f as f64),
reinhardt_query::value::Value::Float(None) => QueryValue::Null,
reinhardt_query::value::Value::Double(Some(f)) => QueryValue::Float(f),
reinhardt_query::value::Value::Double(None) => QueryValue::Null,
reinhardt_query::value::Value::String(Some(s)) => QueryValue::String((*s).clone()),
reinhardt_query::value::Value::String(None) => QueryValue::Null,
reinhardt_query::value::Value::Bytes(Some(b)) => QueryValue::Bytes((*b).clone()),
reinhardt_query::value::Value::Bytes(None) => QueryValue::Null,
reinhardt_query::value::Value::ChronoDateTime(Some(dt)) => {
QueryValue::Timestamp(dt.and_utc())
}
reinhardt_query::value::Value::ChronoDateTime(None) => QueryValue::Null,
reinhardt_query::value::Value::ChronoDateTimeUtc(Some(dt)) => {
QueryValue::Timestamp(*dt)
}
reinhardt_query::value::Value::ChronoDateTimeUtc(None) => QueryValue::Null,
reinhardt_query::value::Value::Uuid(Some(u)) => QueryValue::Uuid(*u),
reinhardt_query::value::Value::Uuid(None) => QueryValue::Null,
reinhardt_query::value::Value::Json(Some(json)) => QueryValue::String(json.to_string()),
reinhardt_query::value::Value::Json(None) => QueryValue::Null,
_ => QueryValue::Null,
}
}
#[allow(dead_code)]
fn serialize_value(v: &serde_json::Value) -> String {
match v {
serde_json::Value::Null => "NULL".to_string(),
serde_json::Value::Bool(b) => b.to_string().to_uppercase(),
serde_json::Value::Number(n) => n.to_string(),
serde_json::Value::String(s) => {
format!("'{}'", s.replace('\'', "''"))
}
serde_json::Value::Array(arr) => {
let items: Vec<String> = arr.iter().map(Self::serialize_value).collect();
format!("ARRAY[{}]", items.join(", "))
}
serde_json::Value::Object(obj) => {
let json_str = serde_json::to_string(obj).unwrap_or_else(|_| "{}".to_string());
format!("'{}'::jsonb", json_str.replace('\'', "''"))
}
}
}
pub async fn update(&self, model: &M) -> reinhardt_core::exception::Result<M> {
let conn = get_connection().await?;
self.update_with_conn(&conn, model).await
}
pub async fn update_with_conn(
&self,
conn: &DatabaseConnection,
model: &M,
) -> reinhardt_core::exception::Result<M> {
let pk = model.primary_key().ok_or_else(|| {
reinhardt_core::exception::Error::Database("Model must have primary key".to_string())
})?;
let json = serde_json::to_value(model)
.map_err(|e| reinhardt_core::exception::Error::Database(e.to_string()))?;
let obj = json.as_object().ok_or_else(|| {
reinhardt_core::exception::Error::Database("Model must serialize to object".to_string())
})?;
let mut stmt = Query::update();
stmt.table(Alias::new(M::table_name()));
for (k, v) in obj
.iter()
.filter(|(k, _)| k.as_str() != M::primary_key_field())
{
if v.is_null() {
stmt.value_expr(Alias::new(k.as_str()), Expr::cust("NULL"));
} else {
stmt.value(Alias::new(k.as_str()), Self::json_to_sea_value(v));
}
}
let pk_str = pk.to_string();
let pk_value = if let Ok(int_value) = pk_str.parse::<i64>() {
reinhardt_query::value::Value::BigInt(Some(int_value))
} else if let Ok(uuid) = Uuid::parse_str(&pk_str) {
reinhardt_query::value::Value::Uuid(Some(Box::new(uuid)))
} else {
reinhardt_query::value::Value::String(Some(Box::new(pk_str)))
};
stmt.and_where(Expr::col(Alias::new(M::primary_key_field())).eq(pk_value));
let all_columns: Vec<_> = obj.keys().map(|k| Alias::new(k.as_str())).collect();
stmt.returning(all_columns);
let (sql, values) = build_update_sql(&stmt, conn.backend());
let values: Vec<_> = values
.0
.into_iter()
.map(Self::sea_value_to_query_value)
.collect();
let row = conn.query_one(&sql, values).await?;
serde_json::from_value(row.data.clone())
.map_err(|e| reinhardt_core::exception::Error::Database(e.to_string()))
}
pub async fn delete(&self, pk: M::PrimaryKey) -> reinhardt_core::exception::Result<()> {
let conn = get_connection().await?;
self.delete_with_conn(&conn, pk).await
}
fn build_delete_statement(pk: M::PrimaryKey) -> DeleteStatement {
let primary_key_field = M::primary_key_field();
let primary_key_column = M::field_metadata()
.into_iter()
.find(|field| field.name == primary_key_field)
.map(|field| field.db_column_name().to_owned())
.unwrap_or_else(|| primary_key_field.to_owned());
let primary_key_value = M::primary_key_filter_value(pk);
let primary_key_value = QuerySet::<M>::filter_value_to_sea_value(&primary_key_value);
let mut stmt = Query::delete();
stmt.from_table(Alias::new(M::table_name()))
.and_where(Expr::col(Alias::new(primary_key_column)).eq(primary_key_value));
stmt
}
pub async fn delete_with_conn(
&self,
conn: &DatabaseConnection,
pk: M::PrimaryKey,
) -> reinhardt_core::exception::Result<()> {
let stmt = Self::build_delete_statement(pk);
let (sql, values) = build_delete_sql(&stmt, conn.backend());
let values: Vec<_> = values
.0
.into_iter()
.map(Self::sea_value_to_query_value)
.collect();
conn.execute(&sql, values).await?;
Ok(())
}
pub async fn count(&self) -> reinhardt_core::exception::Result<i64> {
let conn = get_connection().await?;
self.count_with_conn(&conn).await
}
pub async fn count_with_conn(
&self,
conn: &DatabaseConnection,
) -> reinhardt_core::exception::Result<i64> {
let stmt = Query::select()
.from(Alias::new(M::table_name()))
.expr_as(Func::count(Expr::asterisk().into()), Alias::new("count"))
.to_owned();
let (sql, values) = build_select_sql(&stmt, conn.backend());
let values: Vec<_> = values
.0
.into_iter()
.map(Self::sea_value_to_query_value)
.collect();
let row = conn.query_one(&sql, values).await?;
row.get::<i64>("count").ok_or_else(|| {
reinhardt_core::exception::Error::Database("Failed to get count".to_string())
})
}
pub fn bulk_create_query(&self, models: &[M]) -> Option<InsertStatement> {
if models.is_empty() {
return None;
}
let json_values: Vec<serde_json::Value> = models
.iter()
.filter_map(|m| serde_json::to_value(m).ok())
.collect();
if json_values.is_empty() {
return None;
}
let first_obj = json_values[0].as_object()?;
let fields: Vec<_> = first_obj.keys().map(|k| Alias::new(k.as_str())).collect();
let mut stmt = Query::insert();
stmt.into_table(Alias::new(M::table_name())).columns(fields);
for val in &json_values {
if let Some(obj) = val.as_object() {
let values: Vec<reinhardt_query::value::Value> = first_obj
.keys()
.map(|field| {
obj.get(field)
.map(|v| {
if v.is_null() {
reinhardt_query::value::Value::Int(None)
} else {
Self::json_to_sea_value(v)
}
})
.unwrap_or(reinhardt_query::value::Value::Int(None))
})
.collect();
stmt.values_panic(values);
}
}
Some(stmt.to_owned())
}
pub fn bulk_create_sql(&self, models: &[M], backend: DatabaseBackend) -> String {
if let Some(stmt) = self.bulk_create_query(models) {
insert_to_string(&stmt, backend)
} else {
String::new()
}
}
pub fn update_queryset(
&self,
queryset: &QuerySet<M>,
updates: &[(&str, &str)],
) -> (String, Vec<String>) {
use crate::orm::query::UpdateValue;
use std::collections::HashMap;
let updates_map: HashMap<String, UpdateValue> = updates
.iter()
.map(|(key, value)| (key.to_string(), UpdateValue::String(value.to_string())))
.collect();
queryset.update_sql(&updates_map)
}
pub fn delete_queryset(&self, queryset: &QuerySet<M>) -> (String, Vec<String>) {
queryset.delete_sql()
}
pub async fn get_or_create(
&self,
lookup_fields: HashMap<String, String>,
defaults: Option<HashMap<String, String>>,
) -> reinhardt_core::exception::Result<(M, bool)> {
let conn = get_connection().await?;
let (select_sql, _) = self.get_or_create_sql(
&lookup_fields,
&defaults.clone().unwrap_or_default(),
conn.backend(),
);
if let Ok(Some(row)) = conn.query_optional(&select_sql, vec![]).await {
let model: M = serde_json::from_value(row.data.clone())
.map_err(|e| reinhardt_core::exception::Error::Database(e.to_string()))?;
return Ok((model, false));
}
let mut all_fields = lookup_fields.clone();
if let Some(defs) = defaults {
all_fields.extend(defs);
}
let fields: Vec<String> = all_fields.keys().cloned().collect();
let values: Vec<String> = all_fields.values().map(|v| format!("'{}'", v)).collect();
let insert_sql = format!(
"INSERT INTO {} ({}) VALUES ({}) RETURNING *",
M::table_name(),
fields.join(", "),
values.join(", ")
);
let row = conn.query_one(&insert_sql, vec![]).await?;
let model: M = serde_json::from_value(row.data.clone())
.map_err(|e| reinhardt_core::exception::Error::Database(e.to_string()))?;
Ok((model, true))
}
pub async fn bulk_create(
&self,
models: Vec<M>,
batch_size: Option<usize>,
ignore_conflicts: bool,
_update_conflicts: bool,
) -> reinhardt_core::exception::Result<Vec<M>> {
if models.is_empty() {
return Ok(vec![]);
}
let conn = get_connection().await?;
let batch_size = batch_size.unwrap_or(models.len());
let mut results = Vec::new();
for chunk in models.chunks(batch_size) {
let json = serde_json::to_value(&chunk[0])
.map_err(|e| reinhardt_core::exception::Error::Database(e.to_string()))?;
let obj = json.as_object().ok_or_else(|| {
reinhardt_core::exception::Error::Database(
"Model must serialize to object".to_string(),
)
})?;
let pk_field = M::primary_key_field();
let field_names: Vec<String> = obj
.iter()
.filter_map(|(k, v)| {
if k == pk_field && v.is_null() {
None
} else {
Some(k.clone())
}
})
.collect();
let value_rows: Vec<Vec<serde_json::Value>> = chunk
.iter()
.map(|model| {
let json = serde_json::to_value(model).unwrap();
let obj = json.as_object().unwrap();
field_names.iter().map(|field| obj[field].clone()).collect()
})
.collect();
let sql = self.bulk_create_sql_detailed(&field_names, &value_rows, ignore_conflicts);
if ignore_conflicts {
conn.execute(&sql, vec![]).await?;
} else {
let sql_with_returning = sql + " RETURNING *";
let rows = conn.query(&sql_with_returning, vec![]).await?;
for row in rows {
let model: M = serde_json::from_value(row.data.clone())
.map_err(|e| reinhardt_core::exception::Error::Database(e.to_string()))?;
results.push(model);
}
}
}
Ok(results)
}
pub async fn bulk_update(
&self,
models: Vec<M>,
fields: Vec<String>,
batch_size: Option<usize>,
) -> reinhardt_core::exception::Result<usize> {
if models.is_empty() || fields.is_empty() {
return Ok(0);
}
let conn = get_connection().await?;
let batch_size = batch_size.unwrap_or(models.len());
let mut total_updated = 0;
for chunk in models.chunks(batch_size) {
let updates: Vec<(M::PrimaryKey, HashMap<String, serde_json::Value>)> = chunk
.iter()
.filter_map(|model| {
let pk = model.primary_key()?.clone();
let json = serde_json::to_value(model).ok()?;
let obj = json.as_object()?;
let mut field_map = HashMap::new();
for field in &fields {
if let Some(val) = obj.get(field) {
field_map.insert(field.clone(), val.clone());
}
}
Some((pk, field_map))
})
.collect();
if !updates.is_empty() {
let sql = self.bulk_update_sql_detailed(&updates, &fields, conn.backend());
let rows_affected = conn.execute(&sql, vec![]).await?;
total_updated += rows_affected as usize;
}
}
Ok(total_updated)
}
pub fn get_or_create_queries(
&self,
lookup_fields: &HashMap<String, String>,
defaults: &HashMap<String, String>,
) -> (SelectStatement, InsertStatement) {
let mut select_stmt = Query::select();
select_stmt
.from(Alias::new(M::table_name()))
.column(ColumnRef::Asterisk);
for (k, v) in lookup_fields.iter() {
select_stmt.and_where(Expr::col(Alias::new(k.as_str())).eq(v.as_str()));
}
let mut insert_fields = lookup_fields.clone();
insert_fields.extend(defaults.clone());
let mut insert_stmt = Query::insert();
insert_stmt.into_table(Alias::new(M::table_name()));
let columns: Vec<_> = insert_fields
.keys()
.map(|k| Alias::new(k.as_str()))
.collect();
let values: Vec<reinhardt_query::prelude::Expr> = insert_fields
.values()
.map(|v| Expr::val(v.clone()))
.collect();
insert_stmt.columns(columns);
insert_stmt.values_panic(values);
(select_stmt.to_owned(), insert_stmt.to_owned())
}
pub fn get_or_create_sql(
&self,
lookup_fields: &HashMap<String, String>,
defaults: &HashMap<String, String>,
backend: DatabaseBackend,
) -> (String, String) {
let (select_stmt, insert_stmt) = self.get_or_create_queries(lookup_fields, defaults);
(
select_to_string(&select_stmt, backend),
insert_to_string(&insert_stmt, backend),
)
}
pub fn bulk_create_sql_detailed(
&self,
field_names: &[String],
value_rows: &[Vec<serde_json::Value>],
ignore_conflicts: bool,
) -> String {
if value_rows.is_empty() {
return String::new();
}
let values_clause: Vec<String> = value_rows
.iter()
.map(|row| {
let values = row
.iter()
.map(|v| match v {
serde_json::Value::Null => "NULL".to_string(),
serde_json::Value::Number(n) => n.to_string(),
serde_json::Value::String(s) => {
format!("'{}'", s.replace("'", "''"))
}
serde_json::Value::Bool(b) => if *b { "TRUE" } else { "FALSE" }.to_string(),
serde_json::Value::Array(_) | serde_json::Value::Object(_) => {
format!("'{}'", v.to_string().replace("'", "''"))
}
})
.collect::<Vec<_>>()
.join(", ");
format!("({})", values)
})
.collect();
let mut sql = format!(
"INSERT INTO {} ({}) VALUES {}",
M::table_name(),
field_names.join(", "),
values_clause.join(", ")
);
if ignore_conflicts {
sql.push_str(" ON CONFLICT DO NOTHING");
}
sql
}
pub fn bulk_update_sql_detailed(
&self,
updates: &[(M::PrimaryKey, HashMap<String, serde_json::Value>)],
fields: &[String],
_backend: DatabaseBackend,
) -> String
where
M::PrimaryKey: std::fmt::Display + Clone,
{
if updates.is_empty() || fields.is_empty() {
return String::new();
}
let table_name = M::table_name();
let mut set_clauses = Vec::new();
for field in fields {
let mut when_clauses = Vec::new();
for (pk, field_map) in updates.iter() {
if let Some(value) = field_map.get(field) {
let val_str = match value {
serde_json::Value::Null => "NULL".to_string(),
serde_json::Value::Bool(b) => b.to_string().to_uppercase(),
serde_json::Value::Number(n) => n.to_string(),
serde_json::Value::String(s) => format!("'{}'", s.replace('\'', "''")),
serde_json::Value::Array(_) | serde_json::Value::Object(_) => {
format!("'{}'", value.to_string().replace('\'', "''"))
}
};
when_clauses.push(format!(
"WHEN \"id\" = '{}' THEN {}",
pk.to_string().replace('\'', "''"),
val_str
));
}
}
if !when_clauses.is_empty() {
set_clauses.push(format!(
"\"{}\" = CASE {} END",
field,
when_clauses.join(" ")
));
}
}
let ids: Vec<String> = updates
.iter()
.map(|(pk, _)| format!("'{}'", pk.to_string().replace('\'', "''")))
.collect();
format!(
"UPDATE \"{}\" SET {} WHERE \"id\" IN ({})",
table_name,
set_clauses.join(", "),
ids.join(", ")
)
}
}
impl<M: Model> Default for Manager<M> {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::{Manager, build_delete_sql};
use crate::orm::FieldSelector;
use crate::orm::Model;
use crate::orm::connection::DatabaseBackend;
use crate::orm::fields::{CharField, Field};
use crate::orm::inspection::FieldInfo;
use crate::orm::query::FilterValue;
use rstest::rstest;
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::fmt;
use uuid::Uuid;
#[derive(Debug, Clone, Serialize, Deserialize)]
struct TestUser {
id: Option<i64>,
name: String,
email: String,
}
impl TestUser {
#[allow(dead_code)]
fn new(name: String, email: String) -> Self {
Self {
id: None,
name,
email,
}
}
}
#[derive(Debug, Clone)]
struct TestUserFields;
impl FieldSelector for TestUserFields {
fn with_alias(self, _alias: &str) -> Self {
self
}
}
impl Model for TestUser {
type PrimaryKey = i64;
type Fields = TestUserFields;
type Objects = Manager<Self>;
fn table_name() -> &'static str {
"test_user"
}
fn primary_key(&self) -> Option<Self::PrimaryKey> {
self.id
}
fn set_primary_key(&mut self, value: Self::PrimaryKey) {
self.id = Some(value);
}
fn primary_key_field() -> &'static str {
"id"
}
fn new_fields() -> Self::Fields {
TestUserFields
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct TestStringUser {
id: String,
}
#[derive(Debug, Clone)]
struct TestStringUserFields;
impl FieldSelector for TestStringUserFields {
fn with_alias(self, _alias: &str) -> Self {
self
}
}
impl Model for TestStringUser {
type PrimaryKey = String;
type Fields = TestStringUserFields;
type Objects = Manager<Self>;
fn table_name() -> &'static str {
"test_string_user"
}
fn primary_key(&self) -> Option<Self::PrimaryKey> {
Some(self.id.clone())
}
fn set_primary_key(&mut self, value: Self::PrimaryKey) {
self.id = value;
}
fn new_fields() -> Self::Fields {
TestStringUserFields
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct TestUuidUser {
id: Uuid,
}
#[derive(Debug, Clone)]
struct TestUuidUserFields;
impl FieldSelector for TestUuidUserFields {
fn with_alias(self, _alias: &str) -> Self {
self
}
}
impl Model for TestUuidUser {
type PrimaryKey = Uuid;
type Fields = TestUuidUserFields;
type Objects = Manager<Self>;
fn table_name() -> &'static str {
"test_uuid_user"
}
fn primary_key(&self) -> Option<Self::PrimaryKey> {
Some(self.id)
}
fn primary_key_filter_value(pk: Self::PrimaryKey) -> FilterValue {
FilterValue::Uuid(pk)
}
fn set_primary_key(&mut self, value: Self::PrimaryKey) {
self.id = value;
}
fn primary_key_field() -> &'static str {
"id"
}
fn new_fields() -> Self::Fields {
TestUuidUserFields
}
}
#[rstest]
fn test_get_preserves_uuid_primary_key_binding() {
let id = Uuid::parse_str("123e4567-e89b-12d3-a456-426614174000")
.expect("UUID literal should be valid");
let query = TestUuidUser::objects().get(id);
assert_eq!(query.filters().len(), 1);
assert!(matches!(&query.filters()[0].value, FilterValue::Uuid(value) if *value == id));
}
#[rstest]
fn test_get_preserves_default_numeric_primary_key_binding() {
let query = TestUser::objects().get(42);
assert_eq!(query.filters().len(), 1);
assert!(matches!(query.filters()[0].value, FilterValue::Integer(42)));
}
#[rstest]
fn test_delete_preserves_default_numeric_primary_key_binding() {
let statement = Manager::<TestUser>::build_delete_statement(42);
let (_sql, values) = build_delete_sql(&statement, DatabaseBackend::Postgres);
assert_eq!(
values.0,
vec![reinhardt_query::value::Value::BigInt(Some(42))]
);
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize)]
struct NumericUserId(i64);
impl fmt::Display for NumericUserId {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
self.0.fmt(formatter)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct NumericNewtypeUser {
id: NumericUserId,
}
impl Model for NumericNewtypeUser {
type PrimaryKey = NumericUserId;
type Fields = TestUserFields;
type Objects = Manager<Self>;
fn table_name() -> &'static str {
"numeric_newtype_user"
}
fn primary_key(&self) -> Option<Self::PrimaryKey> {
Some(self.id)
}
fn set_primary_key(&mut self, value: Self::PrimaryKey) {
self.id = value;
}
fn new_fields() -> Self::Fields {
TestUserFields
}
}
#[rstest]
fn test_manual_numeric_newtype_preserves_numeric_primary_key_binding() {
let query = NumericNewtypeUser::objects().get(NumericUserId(42));
let statement = Manager::<NumericNewtypeUser>::build_delete_statement(NumericUserId(42));
let (_sql, values) = build_delete_sql(&statement, DatabaseBackend::Postgres);
assert_eq!(query.filters().len(), 1);
assert!(matches!(query.filters()[0].value, FilterValue::Integer(42)));
assert_eq!(
values.0,
vec![reinhardt_query::value::Value::BigInt(Some(42))]
);
}
#[rstest]
#[case("01")]
#[case("+1")]
#[case("0001")]
fn test_get_preserves_exact_custom_primary_key_string(#[case] id: &str) {
let query = TestStringUser::objects().get(id.to_owned());
assert_eq!(query.filters().len(), 1);
let FilterValue::String(value) = &query.filters()[0].value else {
panic!("custom primary key should use an exact string binding");
};
assert_eq!(value, id);
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct ExternalId(String);
impl fmt::Display for ExternalId {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.0)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
struct TypedKeyUser {
external_id: ExternalId,
}
#[derive(Debug, Clone)]
struct TypedKeyUserFields;
impl FieldSelector for TypedKeyUserFields {
fn with_alias(self, _alias: &str) -> Self {
self
}
}
impl Model for TypedKeyUser {
type PrimaryKey = ExternalId;
type Fields = TypedKeyUserFields;
type Objects = Manager<Self>;
fn table_name() -> &'static str {
"typed_key_user"
}
fn primary_key(&self) -> Option<Self::PrimaryKey> {
Some(self.external_id.clone())
}
fn primary_key_filter_value(pk: Self::PrimaryKey) -> FilterValue {
FilterValue::String(format!("external:{}", pk.0))
}
fn set_primary_key(&mut self, value: Self::PrimaryKey) {
self.external_id = value;
}
fn primary_key_field() -> &'static str {
"external_id"
}
fn new_fields() -> Self::Fields {
TypedKeyUserFields
}
fn field_metadata() -> Vec<FieldInfo> {
let mut field = CharField::new(64);
field.base.primary_key = true;
field.base.db_column = Some("external_key".to_owned());
field.set_attributes_from_name(Self::primary_key_field());
vec![FieldInfo::from_field(&field)]
}
}
#[rstest]
#[case(
DatabaseBackend::Postgres,
"DELETE FROM \"typed_key_user\" WHERE \"external_key\" = $1"
)]
#[case(
DatabaseBackend::MySql,
"DELETE FROM `typed_key_user` WHERE `external_key` = ?"
)]
#[case(
DatabaseBackend::Sqlite,
"DELETE FROM \"typed_key_user\" WHERE \"external_key\" = ?"
)]
fn delete_uses_primary_key_column_and_typed_binding(
#[case] backend: DatabaseBackend,
#[case] expected_sql: &str,
) {
let primary_key = ExternalId("42".to_owned());
let statement = Manager::<TypedKeyUser>::build_delete_statement(primary_key);
let (sql, values) = build_delete_sql(&statement, backend);
assert_eq!(sql, expected_sql);
assert_eq!(
values.0,
vec![reinhardt_query::value::Value::String(Some(Box::new(
"external:42".to_owned()
)))]
);
}
#[test]
fn test_get_or_create_sql() {
let manager = TestUser::objects();
let mut lookup = HashMap::new();
lookup.insert("email".to_string(), "test@example.com".to_string());
let mut defaults = HashMap::new();
defaults.insert("name".to_string(), "Test User".to_string());
let (select_sql, insert_sql) =
manager.get_or_create_sql(&lookup, &defaults, DatabaseBackend::Postgres);
assert!(select_sql.contains("SELECT") && select_sql.contains("FROM"));
assert!(select_sql.contains("test_user"));
assert!(select_sql.contains("email"));
assert!(select_sql.contains("$1"));
assert!(insert_sql.contains("INSERT"));
assert!(insert_sql.contains("test_user"));
assert!(insert_sql.contains("email"));
assert!(insert_sql.contains("name"));
}
#[test]
fn test_bulk_create_sql() {
use serde_json::json;
let manager = TestUser::objects();
let fields = vec!["name".to_string(), "email".to_string()];
let values = vec![
vec![json!("Alice"), json!("alice@example.com")],
vec![json!("Bob"), json!("bob@example.com")],
];
let sql = manager.bulk_create_sql_detailed(&fields, &values, false);
assert!(sql.contains("INSERT"));
assert!(sql.contains("test_user"));
assert!(sql.contains("name"));
assert!(sql.contains("email"));
assert!(sql.contains("Alice"));
assert!(sql.contains("alice@example.com"));
assert!(sql.contains("Bob"));
assert!(sql.contains("bob@example.com"));
}
#[test]
fn test_bulk_create_sql_with_conflict() {
use serde_json::json;
let manager = TestUser::objects();
let fields = vec!["name".to_string(), "email".to_string()];
let values = vec![vec![json!("Alice"), json!("alice@example.com")]];
let sql = manager.bulk_create_sql_detailed(&fields, &values, true);
assert!(sql.contains("ON CONFLICT DO NOTHING"));
}
#[test]
fn test_bulk_update_sql() {
use serde_json::json;
let manager = TestUser::objects();
let mut updates = Vec::new();
let mut user1_fields = HashMap::new();
user1_fields.insert("name".to_string(), json!("Alice Updated"));
user1_fields.insert("email".to_string(), json!("alice_new@example.com"));
updates.push((1i64, user1_fields));
let mut user2_fields = HashMap::new();
user2_fields.insert("name".to_string(), json!("Bob Updated"));
user2_fields.insert("email".to_string(), json!("bob_new@example.com"));
updates.push((2i64, user2_fields));
let fields = vec!["name".to_string(), "email".to_string()];
let sql = manager.bulk_update_sql_detailed(&updates, &fields, DatabaseBackend::Postgres);
assert!(sql.contains("UPDATE"));
assert!(sql.contains("test_user"));
assert!(sql.contains("SET"));
assert!(sql.contains("name"));
assert!(sql.contains("CASE"));
assert!(sql.contains("email"));
assert!(sql.contains("Alice Updated"));
assert!(sql.contains("Bob Updated"));
assert!(sql.contains("WHERE"));
}
#[test]
fn test_bulk_create_empty() {
use serde_json::Value;
let manager = TestUser::objects();
let fields: Vec<String> = vec![];
let values: Vec<Vec<Value>> = vec![];
let sql = manager.bulk_create_sql_detailed(&fields, &values, false);
assert!(sql.is_empty());
}
#[test]
fn test_bulk_update_empty() {
use serde_json::Value;
let manager = TestUser::objects();
let updates: Vec<(i64, HashMap<String, Value>)> = vec![];
let fields = vec!["name".to_string()];
let sql = manager.bulk_update_sql_detailed(&updates, &fields, DatabaseBackend::Postgres);
assert!(sql.is_empty());
}
#[test]
fn test_manager_new() {
let manager = super::Manager::<TestUser>::new();
let _ = manager;
}
#[test]
fn test_manager_default() {
let manager = super::Manager::<TestUser>::default();
let _ = manager;
}
#[test]
fn test_get_or_create_sql_empty_lookup() {
let manager = TestUser::objects();
let lookup: HashMap<String, String> = HashMap::new();
let defaults: HashMap<String, String> = HashMap::new();
let (select_sql, insert_sql) =
manager.get_or_create_sql(&lookup, &defaults, DatabaseBackend::Postgres);
assert!(select_sql.contains("SELECT") || select_sql.contains("select"));
assert!(insert_sql.contains("INSERT") || insert_sql.contains("insert"));
}
#[test]
fn test_get_or_create_sql_with_multiple_lookups() {
let manager = TestUser::objects();
let mut lookup = HashMap::new();
lookup.insert("email".to_string(), "test@example.com".to_string());
lookup.insert("name".to_string(), "Test User".to_string());
let defaults: HashMap<String, String> = HashMap::new();
let (select_sql, _insert_sql) =
manager.get_or_create_sql(&lookup, &defaults, DatabaseBackend::Postgres);
assert!(select_sql.contains("email"));
assert!(select_sql.contains("name"));
}
#[test]
fn test_bulk_create_sql_single_row() {
use serde_json::json;
let manager = TestUser::objects();
let fields = vec!["name".to_string()];
let values = vec![vec![json!("SingleUser")]];
let sql = manager.bulk_create_sql_detailed(&fields, &values, false);
assert!(sql.contains("INSERT"));
assert!(sql.contains("test_user"));
assert!(sql.contains("SingleUser"));
}
#[test]
fn test_bulk_update_sql_single_field() {
use serde_json::json;
let manager = TestUser::objects();
let mut updates = Vec::new();
let mut user1_fields = HashMap::new();
user1_fields.insert("name".to_string(), json!("Updated Name"));
updates.push((1i64, user1_fields));
let fields = vec!["name".to_string()];
let sql = manager.bulk_update_sql_detailed(&updates, &fields, DatabaseBackend::Postgres);
assert!(sql.contains("UPDATE"));
assert!(sql.contains("name"));
assert!(sql.contains("Updated Name"));
assert!(!sql.contains("email"));
}
#[test]
fn test_json_to_sea_value_string() {
use serde_json::json;
let value = json!("hello");
let sea_value = super::Manager::<TestUser>::json_to_sea_value(&value);
let debug_str = format!("{:?}", sea_value);
assert!(debug_str.contains("hello") || debug_str.contains("String"));
}
#[test]
fn test_json_to_sea_value_integer() {
use serde_json::json;
let value = json!(42);
let sea_value = super::Manager::<TestUser>::json_to_sea_value(&value);
let debug_str = format!("{:?}", sea_value);
assert!(debug_str.contains("42") || debug_str.contains("Int"));
}
#[test]
fn test_json_to_sea_value_float() {
use serde_json::json;
let value = json!(1.5);
let sea_value = super::Manager::<TestUser>::json_to_sea_value(&value);
let debug_str = format!("{:?}", sea_value);
assert!(debug_str.contains("1.5") || debug_str.contains("Double"));
}
#[test]
fn test_json_to_sea_value_bool() {
use serde_json::json;
let value = json!(true);
let sea_value = super::Manager::<TestUser>::json_to_sea_value(&value);
let debug_str = format!("{:?}", sea_value);
assert!(debug_str.contains("true") || debug_str.contains("Bool"));
}
#[test]
fn test_json_to_sea_value_null() {
use serde_json::json;
let value = json!(null);
let sea_value = super::Manager::<TestUser>::json_to_sea_value(&value);
let debug_str = format!("{:?}", sea_value);
assert!(!debug_str.is_empty());
}
#[test]
fn test_json_to_sea_value_array() {
use serde_json::json;
let value = json!([1, 2, 3]);
let sea_value = super::Manager::<TestUser>::json_to_sea_value(&value);
let debug_str = format!("{:?}", sea_value);
assert!(!debug_str.is_empty());
}
#[test]
fn test_json_to_sea_value_object() {
use serde_json::json;
let value = json!({"key": "value"});
let sea_value = super::Manager::<TestUser>::json_to_sea_value(&value);
let debug_str = format!("{:?}", sea_value);
assert!(!debug_str.is_empty());
}
#[test]
fn test_serialize_value_string() {
use serde_json::json;
let value = json!("test_string");
let serialized = super::Manager::<TestUser>::serialize_value(&value);
assert!(serialized.contains("test_string"));
}
#[test]
fn test_serialize_value_number() {
use serde_json::json;
let value = json!(123);
let serialized = super::Manager::<TestUser>::serialize_value(&value);
assert!(serialized.contains("123"));
}
}