use crate::active_model::{ActiveModel, ActiveModelTrait, ActiveValue};
use crate::model::Model;
use crate::pool::Connection;
use crate::relation_trait::RelationDef;
use crate::value::Value;
use crate::DbError;
const MAX_NESTED_DEPTH: usize = 10;
#[derive(Debug, Clone)]
pub struct ChildEntity {
table: String,
fields: Vec<(String, Value)>,
children: Vec<ChildEntity>,
relation: Option<RelationDef>,
}
impl ChildEntity {
pub fn new(table: impl Into<String>, fields: Vec<(String, Value)>) -> Self {
Self {
table: table.into(),
fields,
children: Vec::new(),
relation: None,
}
}
pub fn from_active<A: ActiveModelTrait>(active: &A) -> Self {
let table = active.table_name().to_string();
let mut fields = Vec::new();
active.for_each_changed(|field, av| {
if let ActiveValue::Set(val) = av {
fields.push((field.to_string(), val.clone()));
}
});
Self {
table,
fields,
children: Vec::new(),
relation: None,
}
}
pub fn with_children(mut self, children: Vec<ChildEntity>) -> Self {
self.children = children;
self
}
pub fn with_relation(mut self, relation: RelationDef) -> Self {
self.relation = Some(relation);
self
}
pub fn table(&self) -> &str {
&self.table
}
pub fn fields(&self) -> &[(String, Value)] {
&self.fields
}
pub fn children(&self) -> &[ChildEntity] {
&self.children
}
pub fn depth(&self) -> usize {
if self.children.is_empty() {
0
} else {
1 + self.children.iter().map(|c| c.depth()).max().unwrap_or(0)
}
}
}
pub struct NestedActiveModel<M: Model> {
parent: ActiveModel<M>,
children: Vec<ChildEntity>,
relation: RelationDef,
cascade_delete: bool,
}
impl<M: Model> NestedActiveModel<M> {
pub fn from_model(parent: ActiveModel<M>, relation: RelationDef) -> Self {
Self {
parent,
children: Vec::new(),
relation,
cascade_delete: false,
}
}
pub fn with_children(mut self, children: Vec<ChildEntity>) -> Self {
self.children = children;
self
}
pub fn cascade_delete(mut self, cascade: bool) -> Self {
self.cascade_delete = cascade;
self
}
pub fn parent(&self) -> &ActiveModel<M> {
&self.parent
}
pub fn parent_mut(&mut self) -> &mut ActiveModel<M> {
&mut self.parent
}
pub fn children(&self) -> &[ChildEntity] {
&self.children
}
pub fn relation(&self) -> &RelationDef {
&self.relation
}
pub fn is_cascade_delete(&self) -> bool {
self.cascade_delete
}
}
#[derive(Debug, Clone)]
pub struct SaveResult {
pub affected_rows: u64,
pub parent_id: Option<Value>,
}
pub async fn nested_save<M: Model>(
conn: &mut dyn Connection,
nested: NestedActiveModel<M>,
) -> Result<SaveResult, DbError>
where
M::PrimaryKey: Into<Value>,
{
for child in &nested.children {
if child.depth() > MAX_NESTED_DEPTH {
return Err(DbError::InvalidInput(format!(
"nested persistence depth exceeds limit ({})",
MAX_NESTED_DEPTH
)));
}
}
conn.begin_transaction().await?;
match do_nested_save(conn, nested).await {
Ok(result) => {
conn.commit().await?;
Ok(result)
}
Err(e) => {
let _ = conn.rollback().await;
Err(e)
}
}
}
async fn do_nested_save<M: Model>(
conn: &mut dyn Connection,
nested: NestedActiveModel<M>,
) -> Result<SaveResult, DbError>
where
M::PrimaryKey: Into<Value>,
{
let table = nested.parent.table_name().to_string();
let mut columns: Vec<String> = Vec::new();
let mut param_values: Vec<Value> = Vec::new();
nested.parent.for_each_changed(|field, av| {
if let ActiveValue::Set(val) = av {
columns.push(field.to_string());
param_values.push(val.clone());
}
});
if columns.is_empty() {
return Err(DbError::QueryError(
"nested_save: no fields set for parent insert".to_string(),
));
}
let placeholders: Vec<String> = (0..columns.len()).map(|_| "?".to_string()).collect();
let parent_sql = format!(
"INSERT INTO {} ({}) VALUES ({})",
table,
columns.join(", "),
placeholders.join(", ")
);
conn.execute_with_params(&parent_sql, ¶m_values).await?;
let parent_id = get_last_insert_id(conn).await?;
let parent_id_value = Value::I64(parent_id);
let mut affected_rows: u64 = 1;
let fk_key = nested.relation.to_key.to_string();
for child in &nested.children {
let child_rows = save_child(conn, child, &fk_key, &parent_id_value).await?;
affected_rows += child_rows;
}
Ok(SaveResult {
affected_rows,
parent_id: Some(parent_id_value),
})
}
fn save_child<'a>(
conn: &'a mut dyn Connection,
child: &'a ChildEntity,
fk_key: &'a str,
parent_id: &'a Value,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
let mut columns: Vec<String> = Vec::new();
let mut params: Vec<Value> = Vec::new();
for (field, value) in &child.fields {
columns.push(field.clone());
params.push(value.clone());
}
if !columns.iter().any(|c| c == fk_key) {
columns.push(fk_key.to_string());
params.push(parent_id.clone());
}
if columns.is_empty() {
return Err(DbError::QueryError(
"nested_save: no fields set for child insert".to_string(),
));
}
let placeholders: Vec<String> = (0..columns.len()).map(|_| "?".to_string()).collect();
let sql = format!(
"INSERT INTO {} ({}) VALUES ({})",
child.table,
columns.join(", "),
placeholders.join(", ")
);
conn.execute_with_params(&sql, ¶ms).await?;
let mut affected_rows: u64 = 1;
if !child.children.is_empty() {
let child_id = get_last_insert_id(conn).await?;
let child_id_value = Value::I64(child_id);
let child_fk = child
.relation
.as_ref()
.map(|r| std::borrow::Borrow::<str>::borrow(&r.to_key))
.unwrap_or("");
for grandchild in &child.children {
let rows = save_child(conn, grandchild, child_fk, &child_id_value).await?;
affected_rows += rows;
}
}
Ok(affected_rows)
})
}
async fn get_last_insert_id(conn: &mut dyn Connection) -> Result<i64, DbError> {
let rows = conn.query("SELECT LAST_INSERT_ID() as id").await?;
if rows.is_empty() {
return Err(DbError::QueryError(
"get_last_insert_id: no result".to_string(),
));
}
let row = &rows[0];
match row.get("id") {
Some(Value::I64(id)) => Ok(*id),
Some(Value::I32(id)) => Ok(*id as i64),
Some(Value::U64(id)) => Ok(*id as i64),
Some(Value::U32(id)) => Ok(*id as i64),
_ => Err(DbError::QueryError(
"get_last_insert_id: unexpected type".to_string(),
)),
}
}
pub async fn nested_delete<M: Model>(
conn: &mut dyn Connection,
nested: &NestedActiveModel<M>,
) -> Result<u64, DbError>
where
M::PrimaryKey: Into<Value>,
{
conn.begin_transaction().await?;
match do_nested_delete(conn, nested).await {
Ok(rows) => {
conn.commit().await?;
Ok(rows)
}
Err(e) => {
let _ = conn.rollback().await;
Err(e)
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum CascadeStrategy {
Restrict,
#[default]
Cascade,
SetNull,
SetDefault,
}
pub async fn nested_delete_with_strategy<M: Model>(
conn: &mut dyn Connection,
nested: &NestedActiveModel<M>,
strategy: CascadeStrategy,
) -> Result<u64, DbError>
where
M::PrimaryKey: Into<Value>,
{
conn.begin_transaction().await?;
let result = match strategy {
CascadeStrategy::Cascade => do_nested_delete(conn, nested).await,
CascadeStrategy::Restrict => do_nested_delete_restrict(conn, nested).await,
CascadeStrategy::SetNull => do_nested_delete_set_null(conn, nested).await,
CascadeStrategy::SetDefault => do_nested_delete_set_default(conn, nested).await,
};
match result {
Ok(rows) => {
conn.commit().await?;
Ok(rows)
}
Err(e) => {
let _ = conn.rollback().await;
Err(e)
}
}
}
async fn do_nested_delete_restrict<M: Model>(
conn: &mut dyn Connection,
nested: &NestedActiveModel<M>,
) -> Result<u64, DbError>
where
M::PrimaryKey: Into<Value>,
{
if !nested.children.is_empty() {
let fk = &nested.relation.to_key;
let table = &nested.relation.to_entity;
let pk_value = nested
.parent
.pk_value()
.ok_or_else(|| DbError::InvalidInput("父实体缺少主键值".to_string()))?;
let count_sql = format!("SELECT COUNT(*) AS cnt FROM {} WHERE {} = ?", table, fk);
let rows = conn.query_with_params(&count_sql, &[pk_value]).await?;
let count = rows
.first()
.and_then(|r| r.get("cnt"))
.and_then(|v| v.as_i64())
.unwrap_or(0);
if count > 0 {
return Err(DbError::InvalidInput(format!(
"存在 {} 个子实体({}),禁止删除",
count, table
)));
}
}
let table = M::table_name();
let pk_col = M::pk_name();
let pk_value = nested
.parent
.pk_value()
.ok_or_else(|| DbError::InvalidInput("父实体缺少主键值".to_string()))?;
let delete_sql = format!("DELETE FROM {} WHERE {} = ?", table, pk_col);
let affected = conn.execute_with_params(&delete_sql, &[pk_value]).await?;
Ok(affected)
}
async fn do_nested_delete_set_null<M: Model>(
conn: &mut dyn Connection,
nested: &NestedActiveModel<M>,
) -> Result<u64, DbError>
where
M::PrimaryKey: Into<Value>,
{
let pk_value = nested
.parent
.pk_value()
.ok_or_else(|| DbError::InvalidInput("父实体缺少主键值".to_string()))?;
let mut total_affected = 0u64;
if !nested.children.is_empty() {
let fk = &nested.relation.to_key;
let table = &nested.relation.to_entity;
let sql = format!("UPDATE {} SET {} = NULL WHERE {} = ?", table, fk, fk);
total_affected += conn
.execute_with_params(&sql, std::slice::from_ref(&pk_value))
.await?;
}
let table = M::table_name();
let pk_col = M::pk_name();
let delete_sql = format!("DELETE FROM {} WHERE {} = ?", table, pk_col);
total_affected += conn.execute_with_params(&delete_sql, &[pk_value]).await?;
Ok(total_affected)
}
async fn do_nested_delete_set_default<M: Model>(
conn: &mut dyn Connection,
nested: &NestedActiveModel<M>,
) -> Result<u64, DbError>
where
M::PrimaryKey: Into<Value>,
{
let pk_value = nested
.parent
.pk_value()
.ok_or_else(|| DbError::InvalidInput("父实体缺少主键值".to_string()))?;
let mut total_affected = 0u64;
if !nested.children.is_empty() {
let fk = &nested.relation.to_key;
let table = &nested.relation.to_entity;
let sql = format!("UPDATE {} SET {} = DEFAULT WHERE {} = ?", table, fk, fk);
total_affected += conn
.execute_with_params(&sql, std::slice::from_ref(&pk_value))
.await?;
}
let table = M::table_name();
let pk_col = M::pk_name();
let delete_sql = format!("DELETE FROM {} WHERE {} = ?", table, pk_col);
total_affected += conn.execute_with_params(&delete_sql, &[pk_value]).await?;
Ok(total_affected)
}
async fn do_nested_delete<M: Model>(
conn: &mut dyn Connection,
nested: &NestedActiveModel<M>,
) -> Result<u64, DbError>
where
M::PrimaryKey: Into<Value>,
{
let parent_pk = nested
.parent
.pk_value()
.ok_or_else(|| DbError::QueryError("nested_delete: parent pk not set".to_string()))?;
let mut affected_rows: u64 = 0;
let relation = &nested.relation;
let fk_key = &relation.to_key;
let child_table = &relation.to_entity;
if !nested.children.is_empty() {
let delete_children_sql = format!("DELETE FROM {} WHERE {} = ?", child_table, fk_key);
let params = vec![parent_pk.clone()];
let rows = conn
.execute_with_params(&delete_children_sql, ¶ms)
.await?;
affected_rows += rows;
}
let parent_table = nested.parent.table_name();
let delete_parent_sql = format!("DELETE FROM {} WHERE id = ?", parent_table);
let params = vec![parent_pk];
let rows = conn
.execute_with_params(&delete_parent_sql, ¶ms)
.await?;
affected_rows += rows;
Ok(affected_rows)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::active_model::ActiveModel;
use crate::model::Model;
use crate::relation_trait::RelationKind;
use crate::Value;
#[derive(Debug, Clone, Default)]
#[allow(dead_code)]
struct User {
id: i64,
name: String,
}
impl Model for User {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"users"
}
fn pk(&self) -> Self::PrimaryKey {
self.id
}
fn set_pk(&mut self, pk: Self::PrimaryKey) {
self.id = pk;
}
}
fn make_relation() -> RelationDef {
RelationDef::new(
"orders",
"users",
"orders",
"id",
"user_id",
RelationKind::HasMany,
)
}
#[test]
fn test_nested_active_model_from_model() {
let user = ActiveModel::from_model(User::default());
let nested = NestedActiveModel::from_model(user, make_relation());
assert_eq!(nested.parent().table_name(), "users");
assert!(!nested.is_cascade_delete());
assert_eq!(nested.children().len(), 0);
}
#[test]
fn test_nested_active_model_with_children() {
let user = ActiveModel::from_model(User::default());
let order = ChildEntity::new("orders", vec![("amount".to_string(), Value::F64(100.0))]);
let nested =
NestedActiveModel::from_model(user, make_relation()).with_children(vec![order]);
assert_eq!(nested.children().len(), 1);
assert_eq!(nested.children()[0].table(), "orders");
}
#[test]
fn test_nested_active_model_cascade_delete() {
let user = ActiveModel::from_model(User::default());
let nested = NestedActiveModel::from_model(user, make_relation()).cascade_delete(true);
assert!(nested.is_cascade_delete());
}
#[test]
fn test_nested_active_model_relation() {
let user = ActiveModel::from_model(User::default());
let nested = NestedActiveModel::from_model(user, make_relation());
assert_eq!(nested.relation().name, "orders");
assert_eq!(nested.relation().to_key, "user_id");
}
#[test]
fn test_child_entity_new() {
let child = ChildEntity::new("orders", vec![("amount".to_string(), Value::F64(50.0))]);
assert_eq!(child.table(), "orders");
assert_eq!(child.fields().len(), 1);
assert_eq!(child.fields()[0].0, "amount");
}
#[test]
fn test_child_entity_with_children() {
let item = ChildEntity::new("order_items", vec![("qty".to_string(), Value::I32(5))]);
let order = ChildEntity::new("orders", vec![("amount".to_string(), Value::F64(100.0))])
.with_children(vec![item]);
assert_eq!(order.children().len(), 1);
assert_eq!(order.depth(), 1);
}
#[test]
fn test_child_entity_depth() {
let leaf = ChildEntity::new("c", vec![]);
assert_eq!(leaf.depth(), 0);
let mid = ChildEntity::new("b", vec![]).with_children(vec![leaf]);
assert_eq!(mid.depth(), 1);
}
#[test]
fn test_depth_limit_constant() {
assert_eq!(MAX_NESTED_DEPTH, 10);
}
#[test]
fn test_save_result() {
let result = SaveResult {
affected_rows: 3,
parent_id: Some(Value::I64(1)),
};
assert_eq!(result.affected_rows, 3);
assert!(result.parent_id.is_some());
}
#[test]
fn test_child_entity_from_active() {
let mut user = ActiveModel::from_model(User::default());
user.set("name", "Alice".into());
let child = ChildEntity::from_active(&user);
assert_eq!(child.table(), "users");
assert_eq!(child.fields().len(), 1);
assert_eq!(child.fields()[0].0, "name");
}
#[test]
fn test_cascade_strategy_default() {
assert_eq!(CascadeStrategy::default(), CascadeStrategy::Cascade);
}
#[test]
fn test_cascade_strategy_variants() {
let strategies = [
CascadeStrategy::Restrict,
CascadeStrategy::Cascade,
CascadeStrategy::SetNull,
CascadeStrategy::SetDefault,
];
assert_eq!(strategies.len(), 4);
assert_ne!(CascadeStrategy::Restrict, CascadeStrategy::Cascade);
assert_ne!(CascadeStrategy::SetNull, CascadeStrategy::SetDefault);
}
}