use std::any::Any;
use std::collections::{HashMap, hash_map::Entry};
use std::ffi::{CStr, CString, c_char};
use std::ptr;
use std::sync::{Arc, RwLock};
pub struct DiContainer {
services: RwLock<HashMap<String, Arc<dyn Any + Send + Sync>>>,
}
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,
}
#[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();
LAST_ERROR.with(|e| {
*e.borrow_mut() = Some(sanitized);
});
}
#[unsafe(no_mangle)]
pub extern "C" fn di_container_new() -> *mut DiContainer {
let container = Box::new(DiContainer {
services: RwLock::new(HashMap::new()),
});
Box::into_raw(container)
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_container_free(container: *mut DiContainer) {
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 {
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().clone();
let child = Box::new(DiContainer {
services: RwLock::new(services),
});
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 {
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();
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 {
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_resolve(
container: *mut DiContainer,
type_name: *const c_char,
) -> DiResult {
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();
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 {
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();
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 {
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();
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 {
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 {
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 {
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) {
if !service.is_null() {
drop(unsafe { Box::from_raw(service) });
}
}
#[unsafe(no_mangle)]
pub extern "C" fn di_error_message() -> *mut c_char {
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() {
LAST_ERROR.with(|e| {
*e.borrow_mut() = None;
});
}
#[unsafe(no_mangle)]
pub unsafe extern "C" fn di_string_free(s: *mut c_char) {
if !s.is_null() {
drop(unsafe { CString::from_raw(s) });
}
}
#[unsafe(no_mangle)]
pub extern "C" fn di_version() -> *const c_char {
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 {
if container.is_null() {
return -1;
}
let container = unsafe { &*container };
container.services.read().unwrap().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) };
}
#[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_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);
});
}
}