#![cfg(feature = "instrument")]
mod entities;
mod helpers;
use std::{
collections::HashMap,
sync::{Arc, Mutex},
};
use entities::order::*;
use es_entity::*;
use helpers::init_pool;
use sqlx::PgPool;
use tracing_subscriber::layer::SubscriberExt;
#[derive(EsRepo, Debug)]
#[es_repo(entity = "Order", delete = "soft")]
pub struct Orders {
pool: PgPool,
#[es_repo(nested)]
items: OrderItems,
}
impl Orders {
pub fn new(pool: PgPool) -> Self {
Self {
pool: pool.clone(),
items: OrderItems::new(pool),
}
}
}
#[derive(EsRepo, Debug)]
#[es_repo(
entity = "OrderItem",
delete = "soft",
columns(order_id(ty = "OrderId", update(persist = false), parent))
)]
pub struct OrderItems {
pool: PgPool,
}
impl OrderItems {
pub fn new(pool: PgPool) -> Self {
Self { pool }
}
}
#[derive(Clone, Default)]
struct SpanCounts(Arc<Mutex<HashMap<String, usize>>>);
impl SpanCounts {
fn count(&self, name: &str) -> usize {
self.0.lock().unwrap().get(name).copied().unwrap_or(0)
}
}
impl<S: tracing::Subscriber> tracing_subscriber::Layer<S> for SpanCounts {
fn on_new_span(
&self,
attrs: &tracing::span::Attributes<'_>,
_id: &tracing::span::Id,
_ctx: tracing_subscriber::layer::Context<'_, S>,
) {
*self
.0
.lock()
.unwrap()
.entry(attrs.metadata().name().to_string())
.or_insert(0) += 1;
}
}
async fn create_order_with_items(
orders: &Orders,
n_items: usize,
) -> anyhow::Result<(OrderId, Vec<OrderItemId>)> {
let order_id = OrderId::new();
let mut order = orders
.create(NewOrderBuilder::default().id(order_id).build().unwrap())
.await?;
let mut item_ids = Vec::new();
for i in 0..n_items {
let item_id = OrderItemId::new();
item_ids.push(item_id);
order.add_item(
NewOrderItemBuilder::default()
.id(item_id)
.order_id(order_id)
.product_name(format!("item-{i}"))
.quantity(1)
.price(9.99)
.build()
.unwrap(),
);
}
orders.update(&mut order).await?;
Ok((order_id, item_ids))
}
#[tokio::test]
async fn update_all_batches_nested_children_across_parents() -> anyhow::Result<()> {
let pool = init_pool().await?;
let orders = Orders::new(pool);
const N_PARENTS: usize = 5;
const M_PERSISTED_ITEMS: usize = 2;
let mut order_ids = Vec::new();
for _ in 0..N_PARENTS {
let (order_id, _items) = create_order_with_items(&orders, M_PERSISTED_ITEMS).await?;
order_ids.push(order_id);
}
let mut loaded = orders.find_all::<Order>(&order_ids).await?;
let mut batch: Vec<Order> = order_ids
.iter()
.map(|id| loaded.remove(id).expect("order was loaded"))
.collect();
for order in batch.iter_mut() {
order
.update_item_quantity("item-0", 42)
.expect("item-0 exists");
order.add_item(
NewOrderItemBuilder::default()
.id(OrderItemId::new())
.order_id(order.id)
.product_name("new-item")
.quantity(1)
.price(1.23)
.build()
.unwrap(),
);
}
let counts = SpanCounts::default();
let subscriber = tracing_subscriber::registry().with(counts.clone());
let _guard = tracing::subscriber::set_default(subscriber);
orders.update_all(&mut batch).await?;
assert_eq!(
counts.count("order_items.update_all_mut"),
1,
"persisted-child updates across all {N_PARENTS} parents should collapse into one \
order_items.update_all_mut call, not one per parent"
);
assert_eq!(
counts.count("order_items.create_all"),
1,
"new-child creates across all {N_PARENTS} parents should collapse into one \
order_items.create_all call, not one per parent"
);
assert_eq!(
counts.count("order_items.update"),
0,
"the batched path should never fall back to the per-child update_in_op"
);
let mut reloaded = orders.find_all::<Order>(&order_ids).await?;
for order_id in &order_ids {
let order = reloaded.remove(order_id).expect("order reloaded");
assert_eq!(order.n_items(), M_PERSISTED_ITEMS + 1, "order {order_id}");
let item0 = order
.find_item_with_name("item-0")
.expect("item-0 still present");
assert_eq!(item0.quantity, 42, "order {order_id}'s item-0 was updated");
assert!(
order.find_item_with_name("new-item").is_some(),
"order {order_id}'s new item was created"
);
}
Ok(())
}