use crate::error::DiError;
use std::any::Any;
use std::collections::{HashMap, hash_map::Entry};
use std::ffi::{CStr, CString, c_char};
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::ptr;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, PoisonError, RwLock};
pub struct DiContainer {
services: RwLock<HashMap<String, Arc<dyn Any + Send + Sync>>>,
locked: AtomicBool,
}
pub struct DiService {
type_name: String,
data: Vec<u8>,
}
#[repr(C)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DiErrorCode {
Ok = 0,
NotFound = 1,
InvalidArgument = 2,
AlreadyRegistered = 3,
InternalError = 4,
SerializationError = 5,
Locked = 6,
}
impl From<&DiError> for DiErrorCode {
fn from(err: &DiError) -> Self {
match err {
DiError::NotFound { .. } => Self::NotFound,
DiError::AlreadyRegistered { .. } => Self::AlreadyRegistered,
DiError::Locked => Self::Locked,
DiError::CircularDependency { .. }
| DiError::CreationFailed { .. }
| DiError::ParentDropped
| DiError::Internal(_) => Self::InternalError,
}
}
}
#[repr(C)]
pub struct DiResult {
pub code: DiErrorCode,
pub service: *mut DiService,
}
thread_local! {
static LAST_ERROR: std::cell::RefCell<Option<String>> = const { std::cell::RefCell::new(None) };
}
fn set_last_error(msg: impl Into<String>) {
let sanitized: String = msg
.into()
.chars()
.map(|c| if c.is_control() { ' ' } else { c })
.collect();
let _ = LAST_ERROR.try_with(|e| {
*e.borrow_mut() = Some(sanitized);
});
}
fn ffi_guard<T>(default: T, f: impl FnOnce() -> T) -> T {
match catch_unwind(AssertUnwindSafe(f)) {
Ok(value) => value,
Err(payload) => {
let msg = if let Some(s) = payload.downcast_ref::<&str>() {
format!("internal panic: {s}")
} else if let Some(s) = payload.downcast_ref::<String>() {
format!("internal panic: {s}")
} else {
String::from("internal panic")
};
set_last_error(msg);
default
}
}
}
#[cfg(test)]
fn panicking_entry_str() -> DiErrorCode {
ffi_guard(DiErrorCode::InternalError, || panic!("test panic: boom"))
}
#[cfg(test)]
fn panicking_entry_string() -> DiErrorCode {
ffi_guard(DiErrorCode::InternalError, || {
let detail = String::from("boom-string");
panic!("test panic: {detail}")
})
}
#[unsafe(no_mangle)]
pub extern "C" fn di_container_new() -> *mut DiContainer {
ffi_guard(ptr::null_mut(), || {
let container = Box::new(DiContainer {
services: RwLock::new(HashMap::new()),
locked: AtomicBool::new(false),
});
Box::into_raw(container)
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_container_free(container: *mut DiContainer) {
ffi_guard((), || {
if !container.is_null() {
drop(unsafe { Box::from_raw(container) });
}
});
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_container_scope(container: *mut DiContainer) -> *mut DiContainer {
ffi_guard(ptr::null_mut(), || {
if container.is_null() {
set_last_error("Container pointer is null");
return ptr::null_mut();
}
let parent = unsafe { &*container };
let services = parent
.services
.read()
.unwrap_or_else(PoisonError::into_inner)
.clone();
let child = Box::new(DiContainer {
services: RwLock::new(services),
locked: AtomicBool::new(false),
});
Box::into_raw(child)
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_register_singleton(
container: *mut DiContainer,
type_name: *const c_char,
data: *const u8,
data_len: usize,
) -> DiErrorCode {
ffi_guard(DiErrorCode::InternalError, || {
if container.is_null() {
set_last_error("Container pointer is null");
return DiErrorCode::InvalidArgument;
}
if type_name.is_null() {
set_last_error("Type name is null");
return DiErrorCode::InvalidArgument;
}
let type_name_str = if let Ok(s) = unsafe { CStr::from_ptr(type_name) }.to_str() {
s.to_string()
} else {
set_last_error("Type name is not valid UTF-8");
return DiErrorCode::InvalidArgument;
};
if data.is_null() && data_len > 0 {
set_last_error("Data pointer is null but length is non-zero");
return DiErrorCode::InvalidArgument;
}
let data_vec = if data_len > 0 {
unsafe { std::slice::from_raw_parts(data, data_len) }.to_vec()
} else {
Vec::new()
};
let container = unsafe { &*container };
let mut services = container
.services
.write()
.unwrap_or_else(PoisonError::into_inner);
if container.locked.load(Ordering::Acquire) {
set_last_error("Container is locked - cannot register new services");
return DiErrorCode::Locked;
}
match services.entry(type_name_str) {
Entry::Occupied(entry) => {
set_last_error(format!("Service '{}' is already registered", entry.key()));
DiErrorCode::AlreadyRegistered
}
Entry::Vacant(entry) => {
let service_data: Arc<dyn Any + Send + Sync> = Arc::new(data_vec);
entry.insert(service_data);
DiErrorCode::Ok
}
}
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_register_singleton_json(
container: *mut DiContainer,
type_name: *const c_char,
json_data: *const c_char,
) -> DiErrorCode {
ffi_guard(DiErrorCode::InternalError, || {
if json_data.is_null() {
set_last_error("JSON data is null");
return DiErrorCode::InvalidArgument;
}
let json_str = match unsafe { CStr::from_ptr(json_data) }.to_str() {
Ok(s) => s,
Err(_) => {
set_last_error("JSON data is not valid UTF-8");
return DiErrorCode::InvalidArgument;
}
};
let json_bytes = json_str.as_bytes();
unsafe {
di_register_singleton(container, type_name, json_bytes.as_ptr(), json_bytes.len())
}
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_remove(
container: *mut DiContainer,
type_name: *const c_char,
) -> DiErrorCode {
ffi_guard(DiErrorCode::InternalError, || {
if container.is_null() {
set_last_error("Container pointer is null");
return DiErrorCode::InvalidArgument;
}
if type_name.is_null() {
set_last_error("Type name is null");
return DiErrorCode::InvalidArgument;
}
let type_name_str = match unsafe { CStr::from_ptr(type_name) }.to_str() {
Ok(s) => s,
Err(_) => {
set_last_error("Type name is not valid UTF-8");
return DiErrorCode::InvalidArgument;
}
};
let container = unsafe { &*container };
let removed = container
.services
.write()
.unwrap_or_else(PoisonError::into_inner)
.remove(type_name_str);
if removed.is_some() {
DiErrorCode::Ok
} else {
set_last_error(format!("Service '{type_name_str}' not found"));
DiErrorCode::NotFound
}
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_clear(container: *mut DiContainer) -> DiErrorCode {
ffi_guard(DiErrorCode::InternalError, || {
if container.is_null() {
set_last_error("Container pointer is null");
return DiErrorCode::InvalidArgument;
}
let container = unsafe { &*container };
container
.services
.write()
.unwrap_or_else(PoisonError::into_inner)
.clear();
DiErrorCode::Ok
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_lock(container: *mut DiContainer) {
ffi_guard((), || {
if container.is_null() {
set_last_error("Container pointer is null");
return;
}
let container = unsafe { &*container };
let guard = container
.services
.write()
.unwrap_or_else(PoisonError::into_inner);
container.locked.store(true, Ordering::Release);
drop(guard);
});
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_is_locked(container: *const DiContainer) -> i32 {
ffi_guard(-1, || {
if container.is_null() {
return -1;
}
i32::from(unsafe { &*container }.locked.load(Ordering::Acquire))
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_resolve(
container: *mut DiContainer,
type_name: *const c_char,
) -> DiResult {
let internal_error = DiResult {
code: DiErrorCode::InternalError,
service: ptr::null_mut(),
};
ffi_guard(internal_error, || {
if container.is_null() {
set_last_error("Container pointer is null");
return DiResult {
code: DiErrorCode::InvalidArgument,
service: ptr::null_mut(),
};
}
if type_name.is_null() {
set_last_error("Type name is null");
return DiResult {
code: DiErrorCode::InvalidArgument,
service: ptr::null_mut(),
};
}
let type_name_str = match unsafe { CStr::from_ptr(type_name) }.to_str() {
Ok(s) => s.to_string(),
Err(_) => {
set_last_error("Type name is not valid UTF-8");
return DiResult {
code: DiErrorCode::InvalidArgument,
service: ptr::null_mut(),
};
}
};
let container = unsafe { &*container };
let services = container
.services
.read()
.unwrap_or_else(PoisonError::into_inner);
match services.get(&type_name_str) {
Some(service_arc) => {
if let Some(data) = service_arc.downcast_ref::<Vec<u8>>() {
let service = Box::new(DiService {
type_name: type_name_str,
data: data.clone(),
});
DiResult {
code: DiErrorCode::Ok,
service: Box::into_raw(service),
}
} else {
set_last_error("Internal error: service data type mismatch");
DiResult {
code: DiErrorCode::InternalError,
service: ptr::null_mut(),
}
}
}
None => {
set_last_error(format!("Service '{}' not found", type_name_str));
DiResult {
code: DiErrorCode::NotFound,
service: ptr::null_mut(),
}
}
}
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_resolve_json(
container: *mut DiContainer,
type_name: *const c_char,
) -> *mut c_char {
ffi_guard(ptr::null_mut(), || {
if container.is_null() {
set_last_error("Container pointer is null");
return ptr::null_mut();
}
if type_name.is_null() {
set_last_error("Type name is null");
return ptr::null_mut();
}
let type_name_str = match unsafe { CStr::from_ptr(type_name) }.to_str() {
Ok(s) => s.to_string(),
Err(_) => {
set_last_error("Type name is not valid UTF-8");
return ptr::null_mut();
}
};
let container = unsafe { &*container };
let services = container
.services
.read()
.unwrap_or_else(PoisonError::into_inner);
match services.get(&type_name_str) {
Some(service_arc) => {
if let Some(data) = service_arc.downcast_ref::<Vec<u8>>() {
match std::str::from_utf8(data) {
Ok(json_str) => match CString::new(json_str) {
Ok(cstr) => cstr.into_raw(),
Err(_) => {
set_last_error("JSON string contains null bytes");
ptr::null_mut()
}
},
Err(_) => {
set_last_error("Service data is not valid UTF-8");
ptr::null_mut()
}
}
} else {
set_last_error("Internal error: service data type mismatch");
ptr::null_mut()
}
}
None => {
set_last_error(format!("Service '{}' not found", type_name_str));
ptr::null_mut()
}
}
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_contains(container: *mut DiContainer, type_name: *const c_char) -> i32 {
ffi_guard(-1, || {
if container.is_null() || type_name.is_null() {
return -1;
}
let type_name_str = match unsafe { CStr::from_ptr(type_name) }.to_str() {
Ok(s) => s,
Err(_) => return -1,
};
let container = unsafe { &*container };
let services = container
.services
.read()
.unwrap_or_else(PoisonError::into_inner);
if services.contains_key(type_name_str) {
1
} else {
0
}
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_service_data(service: *const DiService) -> *const u8 {
ffi_guard(ptr::null(), || {
if service.is_null() {
return ptr::null();
}
unsafe { &*service }.data.as_ptr()
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_service_data_len(service: *const DiService) -> usize {
ffi_guard(0, || {
if service.is_null() {
return 0;
}
unsafe { &*service }.data.len()
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_service_type_name(service: *const DiService) -> *const c_char {
ffi_guard(ptr::null(), || {
if service.is_null() {
return ptr::null();
}
let service = unsafe { &*service };
match CString::new(service.type_name.clone()) {
Ok(cstr) => cstr.into_raw(),
Err(_) => ptr::null(),
}
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_service_free(service: *mut DiService) {
ffi_guard((), || {
if !service.is_null() {
drop(unsafe { Box::from_raw(service) });
}
});
}
#[unsafe(no_mangle)]
pub extern "C" fn di_error_message() -> *mut c_char {
ffi_guard(ptr::null_mut(), || {
LAST_ERROR.with(|e| {
let error = e.borrow();
match &*error {
Some(msg) => match CString::new(msg.as_str()) {
Ok(cstr) => cstr.into_raw(),
Err(_) => ptr::null_mut(),
},
None => ptr::null_mut(),
}
})
})
}
#[unsafe(no_mangle)]
pub extern "C" fn di_error_clear() {
ffi_guard((), || {
LAST_ERROR.with(|e| {
*e.borrow_mut() = None;
});
});
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_string_free(s: *mut c_char) {
ffi_guard((), || {
if !s.is_null() {
drop(unsafe { CString::from_raw(s) });
}
});
}
#[unsafe(no_mangle)]
pub extern "C" fn di_version() -> *const c_char {
ffi_guard(ptr::null(), || {
static VERSION: &[u8] = concat!(env!("CARGO_PKG_VERSION"), "\0").as_bytes();
VERSION.as_ptr() as *const c_char
})
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_service_count(container: *const DiContainer) -> i64 {
ffi_guard(-1, || {
if container.is_null() {
return -1;
}
let container = unsafe { &*container };
let services = container
.services
.read()
.unwrap_or_else(PoisonError::into_inner);
services.len() as i64
})
}
#[cfg(test)]
mod tests {
use super::*;
fn with_container(f: impl FnOnce(*mut DiContainer)) {
let container = di_container_new();
assert!(!container.is_null());
f(container);
unsafe { di_container_free(container) };
}
fn last_error_string() -> String {
let error = di_error_message();
assert!(!error.is_null(), "expected a last error message");
let message = unsafe { CStr::from_ptr(error) }
.to_str()
.unwrap()
.to_owned();
unsafe { di_string_free(error) };
message
}
fn with_silent_panic_hook<T>(f: impl FnOnce() -> T) -> T {
let previous = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {}));
let result = f();
std::panic::set_hook(previous);
result
}
#[test]
fn test_ffi_guard_catches_str_panic() {
di_error_clear();
let code = with_silent_panic_hook(panicking_entry_str);
assert_eq!(code, DiErrorCode::InternalError);
let msg = last_error_string();
assert!(msg.contains("internal panic"), "got: {msg}");
assert!(msg.contains("boom"), "got: {msg}");
}
#[test]
fn test_ffi_guard_catches_string_panic() {
di_error_clear();
let code = with_silent_panic_hook(panicking_entry_string);
assert_eq!(code, DiErrorCode::InternalError);
let msg = last_error_string();
assert!(msg.contains("internal panic"), "got: {msg}");
assert!(msg.contains("boom-string"), "got: {msg}");
}
#[test]
fn test_every_entry_point_routes_through_ffi_guard() {
const KNOWN_ENTRY_POINTS: usize = 21;
let source = include_str!("ffi.rs");
let lines: Vec<&str> = source.lines().collect();
let marker = concat!("extern ", '"', "C", '"', " fn ");
let mut scanned = 0_usize;
for (start, line) in lines.iter().enumerate() {
let Some(after_marker) = line.split(marker).nth(1) else {
continue;
};
let is_test_only = lines[..start]
.iter()
.rev()
.take_while(|l| l.starts_with('#') || l.starts_with("//"))
.any(|l| l.contains("#[cfg(test)]"));
if is_test_only {
continue;
}
let name = after_marker.split(['(', '<', ' ']).next().unwrap_or("?");
let Some(offset) = lines[start + 1..].iter().position(|l| *l == "}") else {
panic!("{name}: no closing brace at column zero - heuristic broke");
};
let body = lines[start..=start + 1 + offset].join("\n");
assert!(
body.contains("ffi_guard("),
"{name} does not route through ffi_guard; a panic inside it \
would unwind across the FFI boundary"
);
scanned += 1;
}
assert!(
scanned >= KNOWN_ENTRY_POINTS,
"only {scanned} entry points scanned (expected at least \
{KNOWN_ENTRY_POINTS}) - the source heuristic has drifted"
);
}
#[test]
fn test_poisoned_lock_stays_usable() {
struct SendPtr(*mut DiContainer);
unsafe impl Send for SendPtr {}
unsafe impl Sync for SendPtr {}
let container = std::sync::Arc::new(SendPtr(di_container_new()));
assert!(!container.0.is_null());
let joined = {
let container = std::sync::Arc::clone(&container);
with_silent_panic_hook(move || {
std::thread::spawn(move || {
let _guard = unsafe { &*container.0 }.services.write().unwrap();
panic!("deliberate panic while holding the write guard");
})
.join()
})
};
assert!(joined.is_err(), "the poisoning thread must have panicked");
let poisoned = unsafe { &*container.0 }.services.is_poisoned();
assert!(poisoned, "the services lock must now be poisoned");
let name = CString::new("Poisoned").unwrap();
let data = b"data";
unsafe {
assert_eq!(
di_register_singleton(container.0, name.as_ptr(), data.as_ptr(), data.len()),
DiErrorCode::Ok,
"registration must recover from lock poisoning"
);
assert_eq!(di_contains(container.0, name.as_ptr()), 1);
assert_eq!(di_service_count(container.0), 1);
let resolved = di_resolve(container.0, name.as_ptr());
assert_eq!(resolved.code, DiErrorCode::Ok);
di_service_free(resolved.service);
let json = di_resolve_json(container.0, name.as_ptr());
assert!(!json.is_null(), "JSON resolve must recover from poisoning");
di_string_free(json);
let child = di_container_scope(container.0);
assert!(!child.is_null(), "scoping must recover from poisoning");
assert_eq!(di_contains(child, name.as_ptr()), 1);
di_container_free(child);
di_lock(container.0);
assert_eq!(di_is_locked(container.0), 1);
assert_eq!(di_remove(container.0, name.as_ptr()), DiErrorCode::Ok);
assert_eq!(di_clear(container.0), DiErrorCode::Ok);
assert_eq!(di_service_count(container.0), 0);
di_container_free(container.0);
}
}
#[test]
fn test_container_lifecycle() {
unsafe {
let container = di_container_new();
assert!(!container.is_null());
di_container_free(container);
}
}
#[test]
fn test_register_and_resolve() {
with_container(|container| unsafe {
let type_name = CString::new("TestService").unwrap();
let data = b"hello world";
let result =
di_register_singleton(container, type_name.as_ptr(), data.as_ptr(), data.len());
assert_eq!(result, DiErrorCode::Ok);
let resolve_result = di_resolve(container, type_name.as_ptr());
assert_eq!(resolve_result.code, DiErrorCode::Ok);
assert!(!resolve_result.service.is_null());
let service = resolve_result.service;
assert_eq!(di_service_data_len(service), 11);
let data_ptr = di_service_data(service);
let resolved_data = std::slice::from_raw_parts(data_ptr, 11);
assert_eq!(resolved_data, b"hello world");
di_service_free(service);
});
}
#[test]
fn test_not_found() {
with_container(|container| unsafe {
let type_name = CString::new("NonExistent").unwrap();
let result = di_resolve(container, type_name.as_ptr());
assert_eq!(result.code, DiErrorCode::NotFound);
assert!(result.service.is_null());
});
}
#[test]
fn test_contains() {
with_container(|container| unsafe {
let type_name = CString::new("TestService").unwrap();
assert_eq!(di_contains(container, type_name.as_ptr()), 0);
let data = b"test";
di_register_singleton(container, type_name.as_ptr(), data.as_ptr(), data.len());
assert_eq!(di_contains(container, type_name.as_ptr()), 1);
});
}
#[test]
fn test_scope() {
unsafe {
let parent = di_container_new();
let type_name = CString::new("ParentService").unwrap();
let data = b"parent";
di_register_singleton(parent, type_name.as_ptr(), data.as_ptr(), data.len());
let child = di_container_scope(parent);
assert!(!child.is_null());
assert_eq!(di_contains(child, type_name.as_ptr()), 1);
di_container_free(child);
di_container_free(parent);
}
}
#[test]
fn test_duplicate_registration() {
with_container(|container| unsafe {
let type_name = CString::new("Duplicate").unwrap();
let data = b"data";
let first =
di_register_singleton(container, type_name.as_ptr(), data.as_ptr(), data.len());
assert_eq!(first, DiErrorCode::Ok);
let second =
di_register_singleton(container, type_name.as_ptr(), data.as_ptr(), data.len());
assert_eq!(second, DiErrorCode::AlreadyRegistered);
});
}
#[test]
fn test_concurrent_duplicate_registration() {
struct SendPtr(*mut DiContainer);
unsafe impl Send for SendPtr {}
unsafe impl Sync for SendPtr {}
const THREADS: usize = 8;
const ROUNDS: usize = 100;
let container = SendPtr(di_container_new());
let container = std::sync::Arc::new(container);
for round in 0..ROUNDS {
let barrier = std::sync::Arc::new(std::sync::Barrier::new(THREADS));
let name = std::sync::Arc::new(CString::new(format!("Svc{round}")).unwrap());
let handles: Vec<_> = (0..THREADS)
.map(|_| {
let container = std::sync::Arc::clone(&container);
let barrier = std::sync::Arc::clone(&barrier);
let name = std::sync::Arc::clone(&name);
std::thread::spawn(move || {
let data = b"race";
barrier.wait();
unsafe {
di_register_singleton(
container.0,
name.as_ptr(),
data.as_ptr(),
data.len(),
)
}
})
})
.collect();
let results: Vec<_> = handles.into_iter().map(|h| h.join().unwrap()).collect();
let ok = results.iter().filter(|r| **r == DiErrorCode::Ok).count();
let dup = results
.iter()
.filter(|r| **r == DiErrorCode::AlreadyRegistered)
.count();
assert_eq!(ok, 1, "round {round}: exactly one registration must win");
assert_eq!(
dup,
THREADS - 1,
"round {round}: the rest must see AlreadyRegistered"
);
}
unsafe { di_container_free(container.0) };
}
#[test]
fn test_concurrent_lock_vs_registration() {
struct SendPtr(*mut DiContainer);
unsafe impl Send for SendPtr {}
unsafe impl Sync for SendPtr {}
const THREADS: usize = 8;
let container = std::sync::Arc::new(SendPtr(di_container_new()));
let barrier = std::sync::Arc::new(std::sync::Barrier::new(THREADS + 1));
let locker = {
let container = std::sync::Arc::clone(&container);
let barrier = std::sync::Arc::clone(&barrier);
std::thread::spawn(move || {
barrier.wait();
unsafe { di_lock(container.0) };
})
};
let handles: Vec<_> = (0..THREADS)
.map(|i| {
let container = std::sync::Arc::clone(&container);
let barrier = std::sync::Arc::clone(&barrier);
std::thread::spawn(move || {
let name = CString::new(format!("Racer{i}")).unwrap();
let data = b"race";
barrier.wait();
let code = unsafe {
di_register_singleton(container.0, name.as_ptr(), data.as_ptr(), data.len())
};
(i, code)
})
})
.collect();
locker.join().unwrap();
let results: Vec<(usize, DiErrorCode)> =
handles.into_iter().map(|h| h.join().unwrap()).collect();
let mut expected_count = 0_i64;
for (i, code) in &results {
let name = CString::new(format!("Racer{i}")).unwrap();
let present = unsafe { di_contains(container.0, name.as_ptr()) };
match code {
DiErrorCode::Ok => {
expected_count += 1;
assert_eq!(present, 1, "Racer{i}: Ok registration was lost");
}
DiErrorCode::Locked => {
assert_eq!(present, 0, "Racer{i}: Locked registration landed anyway");
}
other => panic!("Racer{i}: must be Ok or Locked, got {other:?}"),
}
}
let count = unsafe { di_service_count(container.0) };
assert_eq!(count, expected_count, "count must match the Ok results");
unsafe {
assert_eq!(di_is_locked(container.0), 1);
let data = b"race";
for i in 0..THREADS {
let name = CString::new(format!("Racer{i}")).unwrap();
let result =
di_register_singleton(container.0, name.as_ptr(), data.as_ptr(), data.len());
assert_eq!(result, DiErrorCode::Locked);
}
let fresh = CString::new("LateRacer").unwrap();
let result =
di_register_singleton(container.0, fresh.as_ptr(), data.as_ptr(), data.len());
assert_eq!(result, DiErrorCode::Locked);
di_container_free(container.0);
}
}
#[test]
fn test_json_register_and_resolve() {
with_container(|container| unsafe {
let type_name = CString::new("JsonService").unwrap();
let json = CString::new("{\"name\":\"test\"}").unwrap();
let result = di_register_singleton_json(container, type_name.as_ptr(), json.as_ptr());
assert_eq!(result, DiErrorCode::Ok);
let resolved = di_resolve_json(container, type_name.as_ptr());
assert!(!resolved.is_null());
assert_eq!(
CStr::from_ptr(resolved).to_str().unwrap(),
"{\"name\":\"test\"}"
);
di_string_free(resolved);
});
}
#[test]
fn test_resolve_json_not_found() {
with_container(|container| unsafe {
let type_name = CString::new("Missing").unwrap();
di_error_clear();
let resolved = di_resolve_json(container, type_name.as_ptr());
assert!(resolved.is_null());
let error = di_error_message();
assert!(!error.is_null());
assert!(
CStr::from_ptr(error)
.to_str()
.unwrap()
.contains("not found")
);
di_string_free(error);
});
}
#[test]
fn test_service_type_name() {
with_container(|container| unsafe {
let type_name = CString::new("NamedService").unwrap();
let data = b"data";
di_register_singleton(container, type_name.as_ptr(), data.as_ptr(), data.len());
let result = di_resolve(container, type_name.as_ptr());
assert_eq!(result.code, DiErrorCode::Ok);
let name = di_service_type_name(result.service);
assert!(!name.is_null());
assert_eq!(CStr::from_ptr(name).to_str().unwrap(), "NamedService");
di_string_free(name.cast_mut());
di_service_free(result.service);
});
}
#[test]
fn test_error_message_set_and_clear() {
with_container(|container| unsafe {
let type_name = CString::new("Missing").unwrap();
let result = di_resolve(container, type_name.as_ptr());
assert_eq!(result.code, DiErrorCode::NotFound);
let error = di_error_message();
assert!(!error.is_null());
di_string_free(error);
di_error_clear();
assert!(di_error_message().is_null());
});
}
#[test]
fn test_version() {
unsafe {
let version = di_version();
assert!(!version.is_null());
assert_eq!(
CStr::from_ptr(version).to_str().unwrap(),
env!("CARGO_PKG_VERSION")
);
}
}
#[test]
fn test_service_count() {
with_container(|container| unsafe {
assert_eq!(di_service_count(container), 0);
let first = CString::new("First").unwrap();
let second = CString::new("Second").unwrap();
let data = b"data";
di_register_singleton(container, first.as_ptr(), data.as_ptr(), data.len());
assert_eq!(di_service_count(container), 1);
di_register_singleton(container, second.as_ptr(), data.as_ptr(), data.len());
assert_eq!(di_service_count(container), 2);
});
}
#[test]
fn test_remove_round_trip() {
with_container(|container| unsafe {
let type_name = CString::new("Removable").unwrap();
let data = b"data";
di_register_singleton(container, type_name.as_ptr(), data.as_ptr(), data.len());
assert_eq!(di_contains(container, type_name.as_ptr()), 1);
assert_eq!(di_remove(container, type_name.as_ptr()), DiErrorCode::Ok);
assert_eq!(di_contains(container, type_name.as_ptr()), 0);
assert_eq!(
di_remove(container, type_name.as_ptr()),
DiErrorCode::NotFound
);
let msg = last_error_string();
assert!(msg.contains("Removable"), "got: {msg}");
});
}
#[test]
fn test_clear() {
with_container(|container| unsafe {
let first = CString::new("First").unwrap();
let second = CString::new("Second").unwrap();
let data = b"data";
di_register_singleton(container, first.as_ptr(), data.as_ptr(), data.len());
di_register_singleton(container, second.as_ptr(), data.as_ptr(), data.len());
assert_eq!(di_service_count(container), 2);
assert_eq!(di_clear(container), DiErrorCode::Ok);
assert_eq!(di_service_count(container), 0);
});
}
#[test]
fn test_lock_blocks_registration_but_allows_remove_and_clear() {
with_container(|container| unsafe {
let existing = CString::new("Existing").unwrap();
let blocked = CString::new("Blocked").unwrap();
let data = b"data";
di_register_singleton(container, existing.as_ptr(), data.as_ptr(), data.len());
di_lock(container);
let result =
di_register_singleton(container, blocked.as_ptr(), data.as_ptr(), data.len());
assert_eq!(result, DiErrorCode::Locked);
let msg = last_error_string();
assert!(msg.contains("locked"), "got: {msg}");
let json = CString::new("{}").unwrap();
let json_result =
di_register_singleton_json(container, blocked.as_ptr(), json.as_ptr());
assert_eq!(json_result, DiErrorCode::Locked);
assert_eq!(di_remove(container, existing.as_ptr()), DiErrorCode::Ok);
assert_eq!(di_clear(container), DiErrorCode::Ok);
});
}
#[test]
fn test_is_locked_transitions() {
with_container(|container| unsafe {
assert_eq!(di_is_locked(container), 0);
di_lock(container);
assert_eq!(di_is_locked(container), 1);
});
}
#[test]
fn test_null_container_sentinels() {
unsafe {
assert_eq!(di_is_locked(std::ptr::null_mut()), -1);
let name = CString::new("Anything").unwrap();
assert_eq!(di_contains(std::ptr::null_mut(), name.as_ptr()), -1);
}
}
#[test]
fn test_error_code_from_core_error() {
assert_eq!(
DiErrorCode::from(&DiError::not_found::<u32>()),
DiErrorCode::NotFound
);
assert_eq!(
DiErrorCode::from(&DiError::already_registered::<u32>()),
DiErrorCode::AlreadyRegistered
);
assert_eq!(DiErrorCode::from(&DiError::Locked), DiErrorCode::Locked);
assert_eq!(
DiErrorCode::from(&DiError::circular::<u32>()),
DiErrorCode::InternalError
);
assert_eq!(
DiErrorCode::from(&DiError::Internal(String::from("boom"))),
DiErrorCode::InternalError
);
}
}