use nodedb_physical::physical_plan::PhysicalPlan;
use nodedb_types::DatabaseId;
use crate::control::router::vshard::VShardRouter;
use crate::control::state::SharedState;
use crate::types::VShardId;
fn owning_core(router: &VShardRouter, database_id: DatabaseId, collection: &str) -> Option<usize> {
router.resolve(VShardId::from_collection_in_database(
database_id,
collection,
))
}
fn cores_diverge(
router: &VShardRouter,
database_id: DatabaseId,
coll_a: &str,
coll_b: &str,
) -> bool {
match (
owning_core(router, database_id, coll_a),
owning_core(router, database_id, coll_b),
) {
(Some(a), Some(b)) => a != b,
_ => true,
}
}
pub(crate) fn cross_collection_cores_diverge(
state: &SharedState,
database_id: DatabaseId,
coll_a: &str,
coll_b: &str,
) -> bool {
let dispatcher = match state.dispatcher.lock() {
Ok(d) => d,
Err(poisoned) => poisoned.into_inner(),
};
cores_diverge(dispatcher.router(), database_id, coll_a, coll_b)
}
pub(crate) fn ensure_cross_collection_colocated(
state: &SharedState,
database_id: DatabaseId,
op: &'static str,
source: &str,
target: &str,
) -> crate::Result<()> {
if cross_collection_cores_diverge(state, database_id, source, target) {
return Err(crate::Error::CrossCollectionNotColocated {
op,
source_collection: source.to_string(),
target_collection: target.to_string(),
});
}
Ok(())
}
pub(crate) fn guard_cross_collection_write(
state: &SharedState,
database_id: DatabaseId,
plan: &PhysicalPlan,
) -> crate::Result<()> {
use nodedb_physical::physical_plan::KvOp;
match plan {
PhysicalPlan::Kv(KvOp::TransferItem {
source_collection,
dest_collection,
..
}) => ensure_cross_collection_colocated(
state,
database_id,
"TRANSFER",
source_collection,
dest_collection,
),
_ => Ok(()),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn collection_on_core(router: &VShardRouter, core: usize) -> String {
for i in 0u64.. {
let name = format!("col_{i}");
if owning_core(router, DatabaseId::DEFAULT, &name) == Some(core) {
return name;
}
}
unreachable!("router covers all cores")
}
#[test]
fn single_core_never_diverges() {
let router = VShardRouter::round_robin(1);
for i in 0u64..64 {
for j in 0u64..64 {
let a = format!("col_{i}");
let b = format!("col_{j}");
assert!(
!cores_diverge(&router, DatabaseId::DEFAULT, &a, &b),
"single-core node must treat all collections as co-resident ({a}, {b})"
);
}
}
}
#[test]
fn same_core_does_not_diverge() {
let router = VShardRouter::round_robin(4);
let a = collection_on_core(&router, 2);
assert!(!cores_diverge(&router, DatabaseId::DEFAULT, &a, &a));
}
#[test]
fn different_cores_diverge() {
let router = VShardRouter::round_robin(4);
let a = collection_on_core(&router, 0);
let b = collection_on_core(&router, 1);
assert_ne!(a, b);
assert!(
cores_diverge(&router, DatabaseId::DEFAULT, &a, &b),
"collections on cores 0 and 1 must be reported as divergent"
);
}
}