How to create a single threaded singleton in Rust?

Viewed 941

I'm currently trying to wrap a C library in rust that has a few requirements. The C library can only be run on a single thread, and can only be initialized / cleaned up once on the same thread. I want something something like the following.

extern "C" {
    fn init_lib() -> *mut c_void;
    fn cleanup_lib(ctx: *mut c_void);
}

// This line doesn't work.
static mut CTX: Option<(ThreadId, Rc<Context>)> = None;

struct Context(*mut c_void);

impl Context {
    fn acquire() -> Result<Rc<Context>, Error> {
        // If CTX has a reference on the current thread, clone and return it.
        // Otherwise initialize the library and set CTX.
    }
}

impl Drop for Context {
    fn drop(&mut self) {
        unsafe { cleanup_lib(self.0); }
    }
}

Anyone have a good way to achieve something like this? Every solution I try to come up with involves creating a Mutex / Arc and making the Context type Send and Sync which I don't want as I want it to remain single threaded.

2 Answers

A working solution I came up with was to just implement the reference counting myself, removing the need for Rc entirely.

#![feature(once_cell)]

use std::{error::Error, ffi::c_void, fmt, lazy::SyncLazy, sync::Mutex, thread::ThreadId};

extern "C" {
    fn init_lib() -> *mut c_void;
    fn cleanup_lib(ctx: *mut c_void);
}

#[derive(Debug)]
pub enum ContextError {
    InitOnOtherThread,
}

impl fmt::Display for ContextError {
    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
        match *self {
            ContextError::InitOnOtherThread => {
                write!(f, "Context already initialized on a different thread")
            }
        }
    }
}

impl Error for ContextError {}

struct StaticPtr(*mut c_void);
unsafe impl Send for StaticPtr {}

static CTX: SyncLazy<Mutex<Option<(ThreadId, usize, StaticPtr)>>> =
    SyncLazy::new(|| Mutex::new(None));

pub struct Context(*mut c_void);

impl Context {
    pub fn acquire() -> Result<Context, ContextError> {
        let mut ctx = CTX.lock().unwrap();
        if let Some((id, ref_count, ptr)) = ctx.as_mut() {
            if *id == std::thread::current().id() {
                *ref_count += 1;
                return Ok(Context(ptr.0));
            }

            Err(ContextError::InitOnOtherThread)
        } else {
            let ptr = unsafe { init_lib() };
            *ctx = Some((std::thread::current().id(), 1, StaticPtr(ptr)));
            Ok(Context(ptr))
        }
    }
}

impl Drop for Context {
    fn drop(&mut self) {
        let mut ctx = CTX.lock().unwrap();
        let (_, ref_count, ptr) = ctx.as_mut().unwrap();
        *ref_count -= 1;
        if *ref_count == 0 {
            unsafe {
                cleanup_lib(ptr.0);
            }
            *ctx = None;
        }
    }
}

I think the most 'rustic' way to do this is with std::sync::mpsc::sync_channel and an enum describing library operations.

The only public-facing elements of this module are launch_lib(), the SafeLibRef struct (but not its internals), and the pub fn that are part of the impl SafeLibRef.

Also, this example strongly represents the philosophy that the best way to deal with global state is to not have any.

I have played fast and loose with the Result::unwrap() calls. It would be more responsible to handle error conditions better.

use std::sync::{ atomic::{ AtomicBool, Ordering }, mpsc::{ SyncSender, Receiver, sync_channel } };
use std::ffi::c_void;

extern "C" {
    fn init_lib() -> *mut c_void;
    fn do_op_1(ctx: *mut c_void, a: u16, b: u32, c: u64) -> f64;
    fn do_op_2(ctx: *mut c_void, a: f64) -> bool;
    fn cleanup_lib(ctx: *mut c_void);
}

enum LibOperation {
    Op1(u16,u32,u64,SyncSender<f64>),
    Op2(f64, SyncSender<bool>),
    Terminate(SyncSender<()>),
}

#[derive(Clone)]
pub struct SafeLibRef(SyncSender<LibOperation>);

fn lib_thread(rx: Receiver<LibOperation>) {

    static  LIB_INITIALIZED: AtomicBool = AtomicBool::new(false);
    
    if LIB_INITIALIZED.compare_exchange(false, true, Ordering::SeqCst, Ordering::SeqCst).is_err() {
        panic!("Tried to double-initialize library!");
    }
    
    let libptr = unsafe { init_lib() };
    loop {
        let op = rx.recv();
        if op.is_err() {
            unsafe { cleanup_lib(libptr) };
            break;
        }
        match op.unwrap() {
            LibOperation::Op1(a,b,c,tx_res) => {
                let res: f64 = unsafe { do_op_1(libptr, a, b, c) };
                tx_res.send(res).unwrap();
            },
            LibOperation::Op2(a, tx_res) => {
                let res: bool = unsafe { do_op_2(libptr, a) };
                tx_res.send(res).unwrap();
            }
            LibOperation::Terminate(tx_res) => {
                unsafe { cleanup_lib(libptr) };
                tx_res.send(()).unwrap();
                break;
            }
        }
    }
}

/// This needs to be called no more than once.
/// The resulting SafeLibRef can be cloned and passed around.
pub fn launch_lib() -> SafeLibRef {
    let (tx,rx) = sync_channel(0);
    std::thread::spawn(|| lib_thread(rx));
    SafeLibRef(tx)
}

// This is the interface that most of your code will use
impl SafeLibRef {
    
    pub fn  op_1(&self, a: u16, b: u32, c: u64) -> f64 {
        let (res_tx, res_rx) = sync_channel(1);
        self.0.send(LibOperation::Op1(a, b, c, res_tx)).unwrap();
        res_rx.recv().unwrap()
    }
    pub fn  op_2(&self, a: f64) -> bool {
        let (res_tx, res_rx) = sync_channel(1);
        self.0.send(LibOperation::Op2(a, res_tx)).unwrap();
        res_rx.recv().unwrap()
    }
    pub fn terminate(&self) {
        let (res_tx, res_rx) = sync_channel(1);
        self.0.send(LibOperation::Terminate(res_tx)).unwrap();
        res_rx.recv().unwrap();
    }
}
Related