feat(rust): port worker pool and pthread wrappers

Replace the common worker pool and debug pthread wrappers with ABI-compatible
Rust implementations. The pool preserves bounded-queue, resize, drain-on-free,
custom-allocator, and no-thread behavior while C supplies only the build-time
multithreading configuration bit.

This moves runtime support without changing C tests or public headers. Codec
and CLI implementations remain C while their dependencies are migrated.

Test Plan:
- cargo fmt --check
- cargo test --all-targets
- cargo clippy --all-targets -- -D warnings
- cargo build --release
- make -C tests poolTests
- ./tests/poolTests
- compile pool and threading shims with -Werror in MT and no-thread modes

Refs: rust/README.md
This commit is contained in:
2026-07-10 20:05:35 +02:00
parent a0f2b2a14c
commit 26b5e202ee
6 changed files with 1037 additions and 549 deletions
+2
View File
@@ -9,5 +9,7 @@ pub mod entropy_common;
pub mod errors;
pub mod fse_decompress;
pub mod mem;
pub mod pool;
pub mod threading;
pub mod xxhash;
pub mod zstd_common;
+733
View File
@@ -0,0 +1,733 @@
use std::cell::UnsafeCell;
use std::mem::{self, MaybeUninit};
use std::os::raw::{c_int, c_void};
use std::ptr;
use std::sync::{Condvar, Mutex};
use std::thread::{self, JoinHandle};
pub type PoolFunction = unsafe extern "C" fn(*mut c_void);
type ZstdAllocFunction = unsafe extern "C" fn(*mut c_void, usize) -> *mut c_void;
type ZstdFreeFunction = unsafe extern "C" fn(*mut c_void, *mut c_void);
/// ABI-compatible representation of `ZSTD_customMem` from `zstd.h`.
#[repr(C)]
#[derive(Clone, Copy)]
pub struct ZstdCustomMem {
custom_alloc: Option<ZstdAllocFunction>,
custom_free: Option<ZstdFreeFunction>,
opaque: *mut c_void,
}
const DEFAULT_CUSTOM_MEM: ZstdCustomMem = ZstdCustomMem {
custom_alloc: None,
custom_free: None,
opaque: ptr::null_mut(),
};
#[derive(Clone, Copy)]
struct Job {
function: PoolFunction,
opaque: *mut c_void,
}
// A job's opaque value is owned by the caller and is required by pool.h to
// remain valid until the callback completes.
unsafe impl Send for Job {}
struct PoolState {
queue_head: usize,
queue_tail: usize,
queue_empty: bool,
num_threads_busy: usize,
thread_limit: usize,
shutdown: bool,
}
struct ThreadStorage {
slots: *mut MaybeUninit<JoinHandle<()>>,
initialized: usize,
}
pub struct PoolCtx {
custom_mem: ZstdCustomMem,
queue: *mut MaybeUninit<Job>,
queue_slots: usize,
state: Mutex<PoolState>,
queue_push_cond: Condvar,
queue_pop_cond: Condvar,
// Only POOL_resize()/POOL_free()/POOL_sizeof() access this field, and each
// does so while holding state. Workers never access thread handles.
threads: UnsafeCell<ThreadStorage>,
}
// All mutable shared state is protected by `state`. The raw queue and handle
// arrays stay allocated until every worker has been joined.
unsafe impl Send for PoolCtx {}
unsafe impl Sync for PoolCtx {}
#[derive(Clone, Copy)]
struct SendPoolPtr(*mut PoolCtx);
unsafe impl Send for SendPoolPtr {}
impl SendPoolPtr {
fn get(self) -> *mut PoolCtx {
self.0
}
}
// This has the same size and stable address as the dummy C context used by the
// original non-multithreaded implementation. It is only compared, never
// dereferenced as a PoolCtx.
static SINGLE_THREADED_POOL: c_int = 0;
fn single_threaded_pool() -> *mut PoolCtx {
ptr::from_ref(&SINGLE_THREADED_POOL).cast_mut().cast()
}
fn is_single_threaded_pool(ctx: *const PoolCtx) -> bool {
ptr::eq(ctx, single_threaded_pool().cast_const())
}
#[cfg(not(test))]
unsafe extern "C" {
fn ZSTD_rust_pool_is_multithreaded() -> c_int;
}
#[cfg(not(test))]
fn multithreading_enabled() -> bool {
// The C shim is compiled with the same preprocessor flags as pool users.
unsafe { ZSTD_rust_pool_is_multithreaded() != 0 }
}
#[cfg(test)]
fn multithreading_enabled() -> bool {
true
}
unsafe fn custom_calloc(size: usize, custom_mem: ZstdCustomMem) -> *mut c_void {
if let Some(alloc) = custom_mem.custom_alloc {
let allocation = unsafe { alloc(custom_mem.opaque, size) };
if !allocation.is_null() {
unsafe { ptr::write_bytes(allocation, 0, size) };
}
allocation
} else {
unsafe { libc::calloc(1, size) }
}
}
unsafe fn custom_free(allocation: *mut c_void, custom_mem: ZstdCustomMem) {
if allocation.is_null() {
return;
}
if let Some(free) = custom_mem.custom_free {
unsafe { free(custom_mem.opaque, allocation) };
} else {
unsafe { libc::free(allocation) };
}
}
fn checked_array_size<T>(len: usize) -> Option<usize> {
mem::size_of::<T>().checked_mul(len)
}
unsafe fn create_threaded_pool(
num_threads: usize,
queue_size: usize,
custom_mem: ZstdCustomMem,
) -> *mut PoolCtx {
if num_threads == 0 {
return ptr::null_mut();
}
let Some(queue_slots) = queue_size.checked_add(1) else {
return ptr::null_mut();
};
let Some(queue_bytes) = checked_array_size::<MaybeUninit<Job>>(queue_slots) else {
return ptr::null_mut();
};
let Some(thread_bytes) = checked_array_size::<MaybeUninit<JoinHandle<()>>>(num_threads) else {
return ptr::null_mut();
};
let ctx_ptr = unsafe { custom_calloc(mem::size_of::<PoolCtx>(), custom_mem) }.cast::<PoolCtx>();
if ctx_ptr.is_null() {
return ptr::null_mut();
}
let queue = unsafe { custom_calloc(queue_bytes, custom_mem) }.cast::<MaybeUninit<Job>>();
if queue.is_null() {
unsafe { custom_free(ctx_ptr.cast(), custom_mem) };
return ptr::null_mut();
}
let thread_slots =
unsafe { custom_calloc(thread_bytes, custom_mem) }.cast::<MaybeUninit<JoinHandle<()>>>();
if thread_slots.is_null() {
unsafe {
custom_free(queue.cast(), custom_mem);
custom_free(ctx_ptr.cast(), custom_mem);
}
return ptr::null_mut();
}
unsafe {
ctx_ptr.write(PoolCtx {
custom_mem,
queue,
queue_slots,
state: Mutex::new(PoolState {
queue_head: 0,
queue_tail: 0,
queue_empty: true,
num_threads_busy: 0,
thread_limit: num_threads,
shutdown: false,
}),
queue_push_cond: Condvar::new(),
queue_pop_cond: Condvar::new(),
threads: UnsafeCell::new(ThreadStorage {
slots: thread_slots,
initialized: 0,
}),
});
}
for thread_id in 0..num_threads {
match spawn_worker(ctx_ptr) {
Ok(handle) => unsafe {
thread_slots.add(thread_id).write(MaybeUninit::new(handle));
(*(*ctx_ptr).threads.get()).initialized = thread_id + 1;
},
Err(_) => {
unsafe {
shutdown_and_join(ctx_ptr);
destroy_context(ctx_ptr);
}
return ptr::null_mut();
}
}
}
ctx_ptr
}
fn spawn_worker(ctx: *mut PoolCtx) -> std::io::Result<JoinHandle<()>> {
let send_ptr = SendPoolPtr(ctx);
thread::Builder::new().spawn(move || worker(send_ptr.get()))
}
fn worker(ctx_ptr: *mut PoolCtx) {
// The creator keeps the context alive until every handle is joined.
let ctx = unsafe { &*ctx_ptr };
loop {
let job = {
let mut state = ctx.state.lock().unwrap_or_else(|error| error.into_inner());
while state.queue_empty || state.num_threads_busy >= state.thread_limit {
if state.shutdown {
return;
}
state = ctx
.queue_pop_cond
.wait(state)
.unwrap_or_else(|error| error.into_inner());
}
let job = unsafe { ctx.queue.add(state.queue_head).read().assume_init() };
state.queue_head = (state.queue_head + 1) % ctx.queue_slots;
state.num_threads_busy += 1;
state.queue_empty = state.queue_head == state.queue_tail;
ctx.queue_push_cond.notify_one();
job
};
unsafe { (job.function)(job.opaque) };
let mut state = ctx.state.lock().unwrap_or_else(|error| error.into_inner());
debug_assert!(state.num_threads_busy > 0);
state.num_threads_busy -= 1;
ctx.queue_push_cond.notify_one();
}
}
unsafe fn shutdown_and_join(ctx_ptr: *mut PoolCtx) {
let ctx = unsafe { &*ctx_ptr };
{
let mut state = ctx.state.lock().unwrap_or_else(|error| error.into_inner());
state.shutdown = true;
}
ctx.queue_push_cond.notify_all();
ctx.queue_pop_cond.notify_all();
let threads = unsafe { &mut *ctx.threads.get() };
for thread_id in 0..threads.initialized {
let handle = unsafe { threads.slots.add(thread_id).read().assume_init() };
let _ = handle.join();
}
threads.initialized = 0;
}
unsafe fn destroy_context(ctx_ptr: *mut PoolCtx) {
let custom_mem = unsafe { (*ctx_ptr).custom_mem };
let queue = unsafe { (*ctx_ptr).queue };
let thread_slots = unsafe { (*(*ctx_ptr).threads.get()).slots };
unsafe { ptr::drop_in_place(ctx_ptr) };
unsafe {
custom_free(queue.cast(), custom_mem);
custom_free(thread_slots.cast(), custom_mem);
custom_free(ctx_ptr.cast(), custom_mem);
}
}
fn queue_is_full(ctx: &PoolCtx, state: &PoolState) -> bool {
if ctx.queue_slots > 1 {
state.queue_head == (state.queue_tail + 1) % ctx.queue_slots
} else {
state.num_threads_busy == state.thread_limit || !state.queue_empty
}
}
unsafe fn add_internal(
ctx: &PoolCtx,
state: &mut PoolState,
function: PoolFunction,
opaque: *mut c_void,
) {
if state.shutdown {
return;
}
unsafe {
ctx.queue
.add(state.queue_tail)
.write(MaybeUninit::new(Job { function, opaque }));
}
state.queue_empty = false;
state.queue_tail = (state.queue_tail + 1) % ctx.queue_slots;
ctx.queue_pop_cond.notify_one();
}
#[no_mangle]
pub extern "C" fn POOL_create(num_threads: usize, queue_size: usize) -> *mut PoolCtx {
unsafe { POOL_create_advanced(num_threads, queue_size, DEFAULT_CUSTOM_MEM) }
}
/// Public ZSTD API alias for `POOL_create` (see `ZSTD_threadPool` in zstd.h).
#[no_mangle]
pub extern "C" fn ZSTD_createThreadPool(num_threads: usize) -> *mut PoolCtx {
POOL_create(num_threads, 0)
}
/// Public ZSTD API alias for `POOL_free`.
///
/// # Safety
/// `pool` must be null or a pointer returned by a pool creation function.
#[no_mangle]
pub unsafe extern "C" fn ZSTD_freeThreadPool(pool: *mut PoolCtx) {
unsafe { POOL_free(pool) }
}
#[no_mangle]
pub unsafe extern "C" fn POOL_create_advanced(
num_threads: usize,
queue_size: usize,
custom_mem: ZstdCustomMem,
) -> *mut PoolCtx {
if !multithreading_enabled() {
return single_threaded_pool();
}
unsafe { create_threaded_pool(num_threads, queue_size, custom_mem) }
}
#[no_mangle]
pub unsafe extern "C" fn POOL_free(ctx: *mut PoolCtx) {
if ctx.is_null() || is_single_threaded_pool(ctx) {
return;
}
unsafe {
shutdown_and_join(ctx);
destroy_context(ctx);
}
}
#[no_mangle]
pub unsafe extern "C" fn POOL_joinJobs(ctx: *mut PoolCtx) {
if ctx.is_null() || is_single_threaded_pool(ctx) {
return;
}
let ctx = unsafe { &*ctx };
let mut state = ctx.state.lock().unwrap_or_else(|error| error.into_inner());
while !state.queue_empty || state.num_threads_busy > 0 {
state = ctx
.queue_push_cond
.wait(state)
.unwrap_or_else(|error| error.into_inner());
}
}
#[no_mangle]
pub unsafe extern "C" fn POOL_resize(ctx: *mut PoolCtx, num_threads: usize) -> c_int {
if is_single_threaded_pool(ctx) {
return 0;
}
if ctx.is_null() || num_threads == 0 {
return 1;
}
let ctx_ref = unsafe { &*ctx };
let mut state = ctx_ref
.state
.lock()
.unwrap_or_else(|error| error.into_inner());
let threads = unsafe { &mut *ctx_ref.threads.get() };
if num_threads <= threads.initialized {
state.thread_limit = num_threads;
drop(state);
ctx_ref.queue_pop_cond.notify_all();
return 0;
}
let Some(thread_bytes) = checked_array_size::<MaybeUninit<JoinHandle<()>>>(num_threads) else {
drop(state);
ctx_ref.queue_pop_cond.notify_all();
return 1;
};
let replacement = unsafe { custom_calloc(thread_bytes, ctx_ref.custom_mem) }
.cast::<MaybeUninit<JoinHandle<()>>>();
if replacement.is_null() {
drop(state);
ctx_ref.queue_pop_cond.notify_all();
return 1;
}
let old_initialized = threads.initialized;
for thread_id in 0..old_initialized {
let handle = unsafe { threads.slots.add(thread_id).read().assume_init() };
unsafe {
replacement.add(thread_id).write(MaybeUninit::new(handle));
}
}
unsafe { custom_free(threads.slots.cast(), ctx_ref.custom_mem) };
threads.slots = replacement;
for thread_id in old_initialized..num_threads {
match spawn_worker(ctx) {
Ok(handle) => {
unsafe {
replacement.add(thread_id).write(MaybeUninit::new(handle));
}
threads.initialized = thread_id + 1;
}
Err(_) => {
drop(state);
ctx_ref.queue_pop_cond.notify_all();
return 1;
}
}
}
state.thread_limit = num_threads;
drop(state);
ctx_ref.queue_pop_cond.notify_all();
0
}
#[no_mangle]
pub unsafe extern "C" fn POOL_sizeof(ctx: *const PoolCtx) -> usize {
if ctx.is_null() {
return 0;
}
if is_single_threaded_pool(ctx) {
return mem::size_of::<c_int>();
}
let ctx = unsafe { &*ctx };
let _state = ctx.state.lock().unwrap_or_else(|error| error.into_inner());
let thread_capacity = unsafe { (*ctx.threads.get()).initialized };
mem::size_of::<PoolCtx>()
.saturating_add(
ctx.queue_slots
.saturating_mul(mem::size_of::<MaybeUninit<Job>>()),
)
.saturating_add(
thread_capacity.saturating_mul(mem::size_of::<MaybeUninit<JoinHandle<()>>>()),
)
}
#[no_mangle]
pub unsafe extern "C" fn POOL_add(ctx: *mut PoolCtx, function: PoolFunction, opaque: *mut c_void) {
if is_single_threaded_pool(ctx) {
unsafe { function(opaque) };
return;
}
if ctx.is_null() {
return;
}
let ctx = unsafe { &*ctx };
let mut state = ctx.state.lock().unwrap_or_else(|error| error.into_inner());
while queue_is_full(ctx, &state) && !state.shutdown {
state = ctx
.queue_push_cond
.wait(state)
.unwrap_or_else(|error| error.into_inner());
}
unsafe { add_internal(ctx, &mut state, function, opaque) };
}
#[no_mangle]
pub unsafe extern "C" fn POOL_tryAdd(
ctx: *mut PoolCtx,
function: PoolFunction,
opaque: *mut c_void,
) -> c_int {
if is_single_threaded_pool(ctx) {
unsafe { function(opaque) };
return 1;
}
if ctx.is_null() {
return 0;
}
let ctx = unsafe { &*ctx };
let mut state = ctx.state.lock().unwrap_or_else(|error| error.into_inner());
if queue_is_full(ctx, &state) {
return 0;
}
unsafe { add_internal(ctx, &mut state, function, opaque) };
1
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::{Duration, Instant};
struct Gate {
started: AtomicUsize,
released: AtomicBool,
}
unsafe extern "C" fn gated_job(opaque: *mut c_void) {
let gate = unsafe { &*(opaque.cast::<Gate>()) };
gate.started.fetch_add(1, Ordering::Release);
while !gate.released.load(Ordering::Acquire) {
thread::yield_now();
}
}
unsafe extern "C" fn increment(opaque: *mut c_void) {
let count = unsafe { &*(opaque.cast::<AtomicUsize>()) };
count.fetch_add(1, Ordering::Relaxed);
}
fn wait_until(mut predicate: impl FnMut() -> bool) {
let deadline = Instant::now() + Duration::from_secs(5);
while !predicate() {
assert!(Instant::now() < deadline, "timed out waiting for pool job");
thread::yield_now();
}
}
#[test]
fn zero_sized_queue_only_accepts_jobs_for_free_workers() {
let ctx = POOL_create(1, 0);
assert!(!ctx.is_null());
let gate = Gate {
started: AtomicUsize::new(0),
released: AtomicBool::new(false),
};
let count = AtomicUsize::new(0);
unsafe { POOL_add(ctx, gated_job, ptr::from_ref(&gate).cast_mut().cast()) };
wait_until(|| gate.started.load(Ordering::Acquire) == 1);
assert_eq!(
unsafe { POOL_tryAdd(ctx, increment, ptr::from_ref(&count).cast_mut().cast()) },
0
);
gate.released.store(true, Ordering::Release);
unsafe { POOL_joinJobs(ctx) };
assert_eq!(
unsafe { POOL_tryAdd(ctx, increment, ptr::from_ref(&count).cast_mut().cast()) },
1
);
unsafe {
POOL_joinJobs(ctx);
POOL_free(ctx);
}
assert_eq!(count.load(Ordering::Relaxed), 1);
}
#[test]
fn bounded_queue_rejects_a_job_when_its_only_slot_is_occupied() {
let ctx = POOL_create(1, 1);
assert!(!ctx.is_null());
let gate = Gate {
started: AtomicUsize::new(0),
released: AtomicBool::new(false),
};
let count = AtomicUsize::new(0);
unsafe { POOL_add(ctx, gated_job, ptr::from_ref(&gate).cast_mut().cast()) };
wait_until(|| gate.started.load(Ordering::Acquire) == 1);
unsafe { POOL_add(ctx, increment, ptr::from_ref(&count).cast_mut().cast()) };
assert_eq!(
unsafe { POOL_tryAdd(ctx, increment, ptr::from_ref(&count).cast_mut().cast()) },
0
);
gate.released.store(true, Ordering::Release);
unsafe { POOL_free(ctx) };
assert_eq!(count.load(Ordering::Relaxed), 1);
}
#[test]
fn add_blocks_until_a_bounded_queue_slot_opens() {
let ctx = POOL_create(1, 1);
assert!(!ctx.is_null());
let gate = Gate {
started: AtomicUsize::new(0),
released: AtomicBool::new(false),
};
let count = std::sync::Arc::new(AtomicUsize::new(0));
let producer_returned = std::sync::Arc::new(AtomicBool::new(false));
unsafe { POOL_add(ctx, gated_job, ptr::from_ref(&gate).cast_mut().cast()) };
wait_until(|| gate.started.load(Ordering::Acquire) == 1);
unsafe {
POOL_add(
ctx,
increment,
std::sync::Arc::as_ptr(&count).cast_mut().cast(),
)
};
let send_ctx = SendPoolPtr(ctx);
let producer_count = std::sync::Arc::clone(&count);
let producer_returned_clone = std::sync::Arc::clone(&producer_returned);
let producer = thread::spawn(move || {
unsafe {
POOL_add(
send_ctx.get(),
increment,
std::sync::Arc::as_ptr(&producer_count).cast_mut().cast(),
)
};
producer_returned_clone.store(true, Ordering::Release);
});
thread::sleep(Duration::from_millis(10));
assert!(!producer_returned.load(Ordering::Acquire));
gate.released.store(true, Ordering::Release);
producer.join().unwrap();
unsafe { POOL_free(ctx) };
assert!(producer_returned.load(Ordering::Acquire));
assert_eq!(count.load(Ordering::Relaxed), 2);
}
#[test]
fn resize_expansion_creates_workers_and_size_tracks_capacity() {
let ctx = POOL_create(1, 3);
assert!(!ctx.is_null());
let original_size = unsafe { POOL_sizeof(ctx) };
assert_eq!(unsafe { POOL_resize(ctx, 3) }, 0);
assert!(unsafe { POOL_sizeof(ctx) } > original_size);
let gate = Gate {
started: AtomicUsize::new(0),
released: AtomicBool::new(false),
};
for _ in 0..3 {
unsafe { POOL_add(ctx, gated_job, ptr::from_ref(&gate).cast_mut().cast()) };
}
wait_until(|| gate.started.load(Ordering::Acquire) == 3);
gate.released.store(true, Ordering::Release);
unsafe { POOL_free(ctx) };
}
#[test]
fn free_drains_queued_jobs_after_reducing_the_limit() {
let ctx = POOL_create(3, 16);
assert!(!ctx.is_null());
let count = AtomicUsize::new(0);
for _ in 0..16 {
unsafe { POOL_add(ctx, increment, ptr::from_ref(&count).cast_mut().cast()) };
}
assert_eq!(unsafe { POOL_resize(ctx, 1) }, 0);
unsafe { POOL_free(ctx) };
assert_eq!(count.load(Ordering::Relaxed), 16);
}
struct AllocStats {
allocations: AtomicUsize,
frees: AtomicUsize,
}
unsafe extern "C" fn counting_alloc(opaque: *mut c_void, size: usize) -> *mut c_void {
let stats = unsafe { &*(opaque.cast::<AllocStats>()) };
stats.allocations.fetch_add(1, Ordering::Relaxed);
unsafe { libc::calloc(1, size) }
}
unsafe extern "C" fn counting_free(opaque: *mut c_void, allocation: *mut c_void) {
let stats = unsafe { &*(opaque.cast::<AllocStats>()) };
stats.frees.fetch_add(1, Ordering::Relaxed);
unsafe { libc::free(allocation) };
}
#[test]
fn advanced_pool_uses_the_custom_allocator_for_owned_storage() {
let stats = AllocStats {
allocations: AtomicUsize::new(0),
frees: AtomicUsize::new(0),
};
let custom_mem = ZstdCustomMem {
custom_alloc: Some(counting_alloc),
custom_free: Some(counting_free),
opaque: ptr::from_ref(&stats).cast_mut().cast(),
};
let ctx = unsafe { POOL_create_advanced(2, 2, custom_mem) };
assert!(!ctx.is_null());
unsafe { POOL_free(ctx) };
assert_eq!(stats.allocations.load(Ordering::Relaxed), 3);
assert_eq!(stats.frees.load(Ordering::Relaxed), 3);
}
#[test]
fn custom_memory_layout_matches_the_c_abi() {
assert_eq!(mem::size_of::<ZstdCustomMem>(), 3 * mem::size_of::<usize>());
assert_eq!(mem::align_of::<ZstdCustomMem>(), mem::align_of::<usize>());
}
#[test]
fn non_threaded_pool_executes_callbacks_synchronously() {
let ctx = single_threaded_pool();
let count = AtomicUsize::new(0);
unsafe {
POOL_add(ctx, increment, ptr::from_ref(&count).cast_mut().cast());
assert_eq!(count.load(Ordering::Relaxed), 1);
assert_eq!(
POOL_tryAdd(ctx, increment, ptr::from_ref(&count).cast_mut().cast()),
1
);
assert_eq!(POOL_resize(ctx, 0), 0);
assert_eq!(POOL_sizeof(ctx), mem::size_of::<c_int>());
POOL_free(ctx);
}
assert_eq!(count.load(Ordering::Relaxed), 2);
}
#[test]
fn rejects_zero_threads_in_threaded_mode() {
assert!(unsafe { create_threaded_pool(0, 1, DEFAULT_CUSTOM_MEM) }.is_null());
assert_eq!(unsafe { POOL_resize(ptr::null_mut(), 1) }, 1);
assert_eq!(unsafe { POOL_sizeof(ptr::null()) }, 0);
}
}
+282
View File
@@ -0,0 +1,282 @@
use std::os::raw::{c_int, c_void};
use std::ptr;
#[cfg(unix)]
use libc::{
pthread_cond_destroy, pthread_cond_init, pthread_cond_t, pthread_condattr_t, pthread_create,
pthread_join, pthread_mutex_destroy, pthread_mutex_init, pthread_mutex_t, pthread_mutexattr_t,
pthread_t,
};
/// Allocating mutex wrapper used by threading.h when `DEBUGLEVEL >= 1`.
#[cfg(unix)]
#[no_mangle]
pub unsafe extern "C" fn ZSTD_pthread_mutex_init(
mutex: *mut *mut pthread_mutex_t,
attr: *const pthread_mutexattr_t,
) -> c_int {
if mutex.is_null() {
return libc::EINVAL;
}
let allocation =
unsafe { libc::malloc(std::mem::size_of::<pthread_mutex_t>()) }.cast::<pthread_mutex_t>();
unsafe { mutex.write(allocation) };
if allocation.is_null() {
return 1;
}
unsafe { pthread_mutex_init(allocation, attr) }
}
/// Destroying mutex wrapper used by threading.h when `DEBUGLEVEL >= 1`.
#[cfg(unix)]
#[no_mangle]
pub unsafe extern "C" fn ZSTD_pthread_mutex_destroy(mutex: *mut *mut pthread_mutex_t) -> c_int {
if mutex.is_null() {
return libc::EINVAL;
}
let allocation = unsafe { mutex.read() };
if allocation.is_null() {
return 0;
}
let result = unsafe { pthread_mutex_destroy(allocation) };
unsafe { libc::free(allocation.cast()) };
result
}
/// Allocating condition-variable wrapper used by threading.h when
/// `DEBUGLEVEL >= 1`.
#[cfg(unix)]
#[no_mangle]
pub unsafe extern "C" fn ZSTD_pthread_cond_init(
cond: *mut *mut pthread_cond_t,
attr: *const pthread_condattr_t,
) -> c_int {
if cond.is_null() {
return libc::EINVAL;
}
let allocation =
unsafe { libc::malloc(std::mem::size_of::<pthread_cond_t>()) }.cast::<pthread_cond_t>();
unsafe { cond.write(allocation) };
if allocation.is_null() {
return 1;
}
unsafe { pthread_cond_init(allocation, attr) }
}
/// Destroying condition-variable wrapper used by threading.h when
/// `DEBUGLEVEL >= 1`.
#[cfg(unix)]
#[no_mangle]
pub unsafe extern "C" fn ZSTD_pthread_cond_destroy(cond: *mut *mut pthread_cond_t) -> c_int {
if cond.is_null() {
return libc::EINVAL;
}
let allocation = unsafe { cond.read() };
if allocation.is_null() {
return 0;
}
let result = unsafe { pthread_cond_destroy(allocation) };
unsafe { libc::free(allocation.cast()) };
result
}
/// POSIX-compatible thread creation entry point. On POSIX this symbol is
/// normally bypassed by threading.h's macro, but keeping it ABI-correct makes
/// the Rust object usable by configurations which reference the wrapper.
#[cfg(unix)]
#[no_mangle]
pub unsafe extern "C" fn ZSTD_pthread_create(
thread: *mut pthread_t,
attr: *const libc::pthread_attr_t,
start_routine: extern "C" fn(*mut c_void) -> *mut c_void,
arg: *mut c_void,
) -> c_int {
if thread.is_null() {
return -1;
}
unsafe { pthread_create(thread, attr, start_routine, arg) }
}
/// POSIX-compatible join which intentionally discards the callback result,
/// matching `ZSTD_pthread_join` in threading.h.
#[cfg(unix)]
#[no_mangle]
pub unsafe extern "C" fn ZSTD_pthread_join(thread: pthread_t) -> c_int {
unsafe { pthread_join(thread, ptr::null_mut()) }
}
#[cfg(windows)]
type Handle = *mut c_void;
#[cfg(windows)]
struct WindowsThreadParams {
start_routine: extern "C" fn(*mut c_void) -> *mut c_void,
arg: *mut c_void,
}
#[cfg(windows)]
unsafe impl Send for WindowsThreadParams {}
#[cfg(windows)]
unsafe extern "system" fn windows_worker(arg: *mut c_void) -> u32 {
let params = unsafe { Box::from_raw(arg.cast::<WindowsThreadParams>()) };
(params.start_routine)(params.arg);
0
}
#[cfg(windows)]
#[link(name = "msvcrt")]
unsafe extern "C" {
fn _beginthreadex(
security: *mut c_void,
stack_size: u32,
start_address: Option<unsafe extern "system" fn(*mut c_void) -> u32>,
arglist: *mut c_void,
initflag: u32,
thread_address: *mut u32,
) -> usize;
fn _errno() -> *mut c_int;
}
#[cfg(windows)]
#[link(name = "kernel32")]
unsafe extern "system" {
fn WaitForSingleObject(handle: Handle, milliseconds: u32) -> u32;
fn CloseHandle(handle: Handle) -> c_int;
fn GetLastError() -> u32;
}
#[cfg(windows)]
#[no_mangle]
pub unsafe extern "C" fn ZSTD_pthread_create(
thread: *mut Handle,
_unused: *const c_void,
start_routine: extern "C" fn(*mut c_void) -> *mut c_void,
arg: *mut c_void,
) -> c_int {
if thread.is_null() {
return -1;
}
unsafe { thread.write(ptr::null_mut()) };
let params = Box::new(WindowsThreadParams { start_routine, arg });
let params = Box::into_raw(params);
let handle = unsafe {
_beginthreadex(
ptr::null_mut(),
0,
Some(windows_worker),
params.cast(),
0,
ptr::null_mut(),
)
};
if handle == 0 {
unsafe { drop(Box::from_raw(params)) };
let errno = unsafe { _errno() };
return if errno.is_null() {
-1
} else {
unsafe { *errno }
};
}
unsafe { thread.write(handle as Handle) };
0
}
#[cfg(windows)]
#[no_mangle]
pub unsafe extern "C" fn ZSTD_pthread_join(thread: Handle) -> c_int {
const INFINITE: u32 = u32::MAX;
const WAIT_OBJECT_0: u32 = 0;
const WAIT_ABANDONED: u32 = 0x80;
if thread.is_null() {
return 0;
}
let result = unsafe { WaitForSingleObject(thread, INFINITE) };
unsafe { CloseHandle(thread) };
match result {
WAIT_OBJECT_0 => 0,
WAIT_ABANDONED => libc::EINVAL,
_ => unsafe { GetLastError() as c_int },
}
}
#[cfg(all(test, unix))]
mod tests {
use super::*;
use std::mem::MaybeUninit;
use std::sync::atomic::{AtomicUsize, Ordering};
static THREAD_CALLS: AtomicUsize = AtomicUsize::new(0);
extern "C" fn thread_main(arg: *mut c_void) -> *mut c_void {
THREAD_CALLS.fetch_add(1, Ordering::Relaxed);
arg
}
#[test]
fn debug_mutex_wrapper_allocates_initializes_and_frees() {
let mut mutex: *mut pthread_mutex_t = ptr::null_mut();
unsafe {
assert_eq!(ZSTD_pthread_mutex_init(&mut mutex, ptr::null()), 0);
assert!(!mutex.is_null());
assert_eq!(libc::pthread_mutex_lock(mutex), 0);
assert_eq!(libc::pthread_mutex_unlock(mutex), 0);
assert_eq!(ZSTD_pthread_mutex_destroy(&mut mutex), 0);
}
}
#[test]
fn debug_condition_wrapper_allocates_initializes_and_frees() {
let mut cond: *mut pthread_cond_t = ptr::null_mut();
unsafe {
assert_eq!(ZSTD_pthread_cond_init(&mut cond, ptr::null()), 0);
assert!(!cond.is_null());
assert_eq!(ZSTD_pthread_cond_destroy(&mut cond), 0);
}
}
#[test]
fn pthread_wrapper_runs_and_joins_callback() {
THREAD_CALLS.store(0, Ordering::Relaxed);
let mut thread = MaybeUninit::<pthread_t>::uninit();
unsafe {
assert_eq!(
ZSTD_pthread_create(
thread.as_mut_ptr(),
ptr::null(),
thread_main,
ptr::null_mut()
),
0
);
assert_eq!(ZSTD_pthread_join(thread.assume_init()), 0);
}
assert_eq!(THREAD_CALLS.load(Ordering::Relaxed), 1);
}
#[test]
fn null_output_pointers_are_rejected() {
unsafe {
assert_eq!(
ZSTD_pthread_mutex_init(ptr::null_mut(), ptr::null()),
libc::EINVAL
);
assert_eq!(
ZSTD_pthread_cond_init(ptr::null_mut(), ptr::null()),
libc::EINVAL
);
assert_eq!(
ZSTD_pthread_create(ptr::null_mut(), ptr::null(), thread_main, ptr::null_mut()),
-1
);
}
}
}