use std::ffi::{c_char, c_void, CString};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Mutex;
use crate::error::Error;
use crate::ffi::ffi;
use crate::logical_type::LogicalType;
use crate::value::Value;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct PartitionRef {
pub parent_table_id: u64,
pub partition_index: u64,
}
#[derive(Default)]
pub struct Callbacks {
pub locate: Option<Box<dyn Fn(PartitionRef) -> Option<u64> + Send + Sync>>,
pub on_partition_create: Option<Box<dyn Fn(PartitionRef) + Send + Sync>>,
pub on_partition_drop: Option<Box<dyn Fn(PartitionRef) + Send + Sync>>,
pub insert_row: Option<Box<dyn Fn(PartitionRef, Vec<Value>) + Send + Sync>>,
}
pub struct RoutingGuard {
_private: (),
}
struct State {
callbacks: Callbacks,
}
static INSTALLED: AtomicBool = AtomicBool::new(false);
static STATE: Mutex<Option<Box<State>>> = Mutex::new(None);
#[repr(C)]
struct CPartitionRef {
parent_table_id: u64,
partition_index: u64,
}
impl From<CPartitionRef> for PartitionRef {
fn from(r: CPartitionRef) -> Self {
PartitionRef {
parent_table_id: r.parent_table_id,
partition_index: r.partition_index,
}
}
}
#[repr(C)]
struct CHooks {
context: *mut c_void,
locate: Option<extern "C" fn(*mut c_void, CPartitionRef, *mut *mut c_void) -> u8>,
on_partition_create: Option<extern "C" fn(*mut c_void, CPartitionRef, *mut c_void)>,
on_partition_drop: Option<extern "C" fn(*mut c_void, CPartitionRef, *mut c_void)>,
insert_row:
Option<extern "C" fn(*mut c_void, CPartitionRef, *mut c_void, *const *const c_void, usize)>,
}
extern "C" {
fn lbug_partition_routing_install(hooks: *const CHooks) -> libc_int_t;
fn lbug_partition_routing_uninstall();
fn lbug_partition_routing_is_installed() -> u8;
fn lbug_partition_routing_register_schema(
parent_table_id: u64,
prop_names: *const *const c_char,
type_ids: *const u8,
n_props: usize,
) -> libc_int_t;
}
fn with_state<F, R>(context: *mut c_void, f: F) -> R
where
F: FnOnce(&State) -> R,
{
let state = unsafe { &*(context as *const State) };
f(state)
}
fn abort_on_panic(context: &str, result: std::thread::Result<()>) {
if let Err(payload) = result {
let message = payload
.downcast_ref::<String>()
.cloned()
.or_else(|| payload.downcast_ref::<&str>().map(|s| (*s).to_string()))
.unwrap_or_else(|| "unknown panic".to_string());
eprintln!("ladybug routing callback panicked in {context}: {message}; aborting");
std::process::abort();
}
}
extern "C" fn locate_cb(
context: *mut c_void,
pref: CPartitionRef,
handle_out: *mut *mut c_void,
) -> u8 {
let result = std::panic::catch_unwind(|| {
with_state(context, |state| {
state
.callbacks
.locate
.as_ref()
.and_then(|locate| locate(pref.into()))
})
});
match result {
Ok(Some(handle)) => {
unsafe { *handle_out = handle as *mut c_void };
1
}
Ok(None) | Err(_) => 0,
}
}
extern "C" fn create_cb(context: *mut c_void, pref: CPartitionRef, _handle: *mut c_void) {
let result = std::panic::catch_unwind(|| {
with_state(context, |state| {
if let Some(cb) = state.callbacks.on_partition_create.as_ref() {
cb(pref.into());
}
});
});
abort_on_panic("on_partition_create", result);
}
extern "C" fn drop_cb(context: *mut c_void, pref: CPartitionRef, _handle: *mut c_void) {
let result = std::panic::catch_unwind(|| {
with_state(context, |state| {
if let Some(cb) = state.callbacks.on_partition_drop.as_ref() {
cb(pref.into());
}
});
});
abort_on_panic("on_partition_drop", result);
}
extern "C" fn insert_row_cb(
context: *mut c_void,
pref: CPartitionRef,
_handle: *mut c_void,
cells: *const *const c_void,
n_cells: usize,
) {
let result = std::panic::catch_unwind(|| {
with_state(context, |state| {
let Some(cb) = state.callbacks.insert_row.as_ref() else {
return;
};
let mut row = Vec::with_capacity(n_cells);
for i in 0..n_cells {
let cell = unsafe { &*(*cells.add(i)).cast::<ffi::Value>() };
match Value::try_from(cell) {
Ok(value) => row.push(value),
Err(e) => {
eprintln!("ladybug routing: cannot decode routed cell {i}: {e}; aborting");
std::process::abort();
}
}
}
cb(pref.into(), row);
});
});
abort_on_panic("insert_row", result);
}
fn logical_type_id(logical_type: &LogicalType) -> Result<u8, Error> {
match logical_type {
LogicalType::Any => Ok(0),
LogicalType::Bool => Ok(22),
LogicalType::Serial => Ok(13),
LogicalType::Int64 => Ok(23),
LogicalType::Int32 => Ok(24),
LogicalType::Int16 => Ok(25),
LogicalType::Int8 => Ok(26),
LogicalType::UInt64 => Ok(27),
LogicalType::UInt32 => Ok(28),
LogicalType::UInt16 => Ok(29),
LogicalType::UInt8 => Ok(30),
LogicalType::Int128 => Ok(31),
LogicalType::Double => Ok(32),
LogicalType::Float => Ok(33),
LogicalType::Date => Ok(34),
LogicalType::Timestamp => Ok(35),
LogicalType::TimestampTz => Ok(39),
LogicalType::TimestampNs => Ok(38),
LogicalType::TimestampMs => Ok(37),
LogicalType::TimestampSec => Ok(36),
LogicalType::String => Ok(50),
LogicalType::UUID => Ok(59),
LogicalType::Json => Ok(60),
other => Err(Error::FailedQuery(format!(
"partition routing schemas support scalar types only, got {other:?}"
))),
}
}
impl RoutingGuard {
pub fn install(callbacks: Callbacks) -> Result<Self, Error> {
if INSTALLED.swap(true, Ordering::SeqCst) {
return Err(Error::FailedQuery(
"partition routing hooks are already installed".to_string(),
));
}
let mut slot = STATE.lock().unwrap();
let state = Box::new(State { callbacks });
let context: *mut c_void = std::ptr::from_ref(&*state).cast_mut().cast();
let c_hooks = CHooks {
context,
locate: Some(locate_cb),
on_partition_create: Some(create_cb),
on_partition_drop: Some(drop_cb),
insert_row: Some(insert_row_cb),
};
let rc = unsafe { lbug_partition_routing_install(std::ptr::from_ref(&c_hooks)) };
if rc != 0 {
INSTALLED.store(false, Ordering::SeqCst);
return Err(Error::FailedQuery(
"engine rejected partition routing installation".to_string(),
));
}
*slot = Some(state);
Ok(RoutingGuard { _private: () })
}
pub fn register_parent_schema(
&self,
parent_table_id: u64,
columns: Vec<(String, LogicalType)>,
) -> Result<(), Error> {
if columns.is_empty() {
return Err(Error::FailedQuery(
"partition routing schema needs at least one column".to_string(),
));
}
let names: Vec<CString> = columns
.iter()
.map(|(name, _)| {
CString::new(name.as_str()).map_err(|_| {
Error::FailedQuery(format!("property name {name:?} contains a NUL byte"))
})
})
.collect::<Result<_, _>>()?;
let name_ptrs: Vec<*const c_char> = names.iter().map(|n| n.as_ptr()).collect();
let type_ids: Vec<u8> = columns
.iter()
.map(|(_, typ)| logical_type_id(typ))
.collect::<Result<_, _>>()?;
let rc = unsafe {
lbug_partition_routing_register_schema(
parent_table_id,
name_ptrs.as_ptr(),
type_ids.as_ptr(),
columns.len(),
)
};
if rc != 0 {
return Err(Error::FailedQuery(
"engine rejected partition routing schema".to_string(),
));
}
Ok(())
}
pub fn is_installed(&self) -> bool {
unsafe { lbug_partition_routing_is_installed() != 0 }
}
pub fn uninstall(self) {}
}
impl Drop for RoutingGuard {
fn drop(&mut self) {
unsafe { lbug_partition_routing_uninstall() };
INSTALLED.store(false, Ordering::SeqCst);
*STATE.lock().unwrap() = None;
}
}
#[allow(non_camel_case_types)]
type libc_int_t = std::os::raw::c_int;
#[cfg(test)]
mod tests {
use super::*;
use crate::connection::Connection;
use crate::database::{Database, SystemConfig};
use std::collections::HashSet;
use std::sync::{Arc, Mutex};
#[test]
fn remote_partition_round_trip() {
let created: Arc<Mutex<Vec<PartitionRef>>> = Arc::new(Mutex::new(Vec::new()));
let dropped: Arc<Mutex<Vec<PartitionRef>>> = Arc::new(Mutex::new(Vec::new()));
let observed: Arc<Mutex<Vec<(PartitionRef, Vec<Value>)>>> =
Arc::new(Mutex::new(Vec::new()));
let created_cb = created.clone();
let dropped_cb = dropped.clone();
let observed_cb = observed.clone();
let guard = RoutingGuard::install(Callbacks::default()).unwrap();
assert!(
RoutingGuard::install(Callbacks::default()).is_err(),
"double install must fail"
);
drop(guard);
let guard = RoutingGuard::install(Callbacks {
locate: Some(Box::new(|_r: PartitionRef| Some(0xC0FFEE))),
on_partition_create: Some(Box::new(move |r: PartitionRef| {
created_cb.lock().unwrap().push(r);
})),
on_partition_drop: Some(Box::new(move |r: PartitionRef| {
dropped_cb.lock().unwrap().push(r);
})),
insert_row: Some(Box::new(move |r: PartitionRef, row: Vec<Value>| {
observed_cb.lock().unwrap().push((r, row));
})),
})
.unwrap();
assert!(guard.is_installed());
let db_dir = tempfile::tempdir().unwrap();
let db = Database::new(db_dir.path().join("routing"), SystemConfig::default()).unwrap();
let conn = Connection::new(&db).unwrap();
conn.query(
"CREATE NODE TABLE Remote(id INT64, v INT64, PRIMARY KEY(id)) PARTITION BY HASH(v) PARTITIONS 3;",
)
.unwrap();
let parent = created.lock().unwrap()[0].parent_table_id;
guard
.register_parent_schema(
parent,
vec![
("id".to_string(), LogicalType::Int64),
("v".to_string(), LogicalType::Int64),
],
)
.unwrap();
conn.query("CREATE (:Remote {id: 1, v: 10});").unwrap();
conn.query("CREATE (:Remote {id: 2, v: 10});").unwrap();
conn.query("CREATE (:Remote {id: 3, v: 20});").unwrap();
let dir = tempfile::tempdir().unwrap();
let csv = dir.path().join("bulk.csv");
std::fs::write(&csv, "4,30\n5,30\n").unwrap();
conn.query(&format!(
"COPY Remote FROM '{}';",
csv.to_string_lossy().replace('\\', "/")
))
.unwrap();
let result = conn
.query("MATCH (r:Remote) RETURN r.id, r.v ORDER BY r.id;")
.unwrap()
.to_string();
assert_eq!(result, "r.id|r.v\n1|10\n2|10\n3|20\n4|30\n5|30\n");
let seen: HashSet<i64> = observed
.lock()
.unwrap()
.iter()
.map(|(_, row)| match &row[0] {
Value::Int64(id) => *id,
other => panic!("expected Int64 id, got {other:?}"),
})
.collect();
assert_eq!(seen, HashSet::from([1, 2, 3, 4, 5]));
assert_eq!(
created.lock().unwrap().len(),
3,
"one create per HASH partition"
);
conn.query("DROP TABLE Remote;").unwrap();
assert_eq!(dropped.lock().unwrap().len(), 3);
drop(conn);
drop(db);
drop(dir);
db_dir.close().unwrap();
guard.uninstall();
}
}