diff --git a/README.md b/README.md index 110e752..5f127c0 100644 --- a/README.md +++ b/README.md @@ -53,8 +53,11 @@ This support is only available on the `x86_64` and `i686` architectures, however Vendor-specific APIs such as `openblas_get_num_threads()` are not included in header files or exported from the library. See the [public header file](src/libblastrampoline.h) for the most up-to-date documentation on the `libblastrampoline` API. -**Note**: all `lbt_*` functions should be considered thread-unsafe. -Do not attempt to load two BLAS libraries on two different threads at the same time. +### Threading + +By default, all `lbt_*` functions are thread-unsafe; do not reconfigure forwards from multiple threads at once. +Building with `make LBT_THREADSAFE=1` adds a process-global lock around the mutating API (`lbt_forward()`, `lbt_set_forward()`, `lbt_set_forward_by_index()`). +Readers and the BLAS/LAPACK call forwarding itself stay lock-free, so the model remains "configure under the lock, then use". ### Limitations diff --git a/src/Make.inc b/src/Make.inc index c6759df..079ee4e 100644 --- a/src/Make.inc +++ b/src/Make.inc @@ -104,6 +104,18 @@ ifeq ($(OS), Darwin) SYMBOL_TRIMMING := 1 endif +# Optional internal locking around the (otherwise thread-unsafe) library-loading and +# forwarding API. Off by default to preserve historical behavior; enable with +# `make LBT_THREADSAFE=1`. +LBT_THREADSAFE ?= 0 +ifeq ($(LBT_THREADSAFE),1) + LBT_CFLAGS += -DLBT_THREADSAFE + ifneq ($(OS),WINNT) + LBT_CFLAGS += -pthread + LBT_LDFLAGS += -pthread + endif +endif + ifeq ($(VERBOSE),0) ENDCOLOR := "\033[0m" CCCOLOR := "\033[34m" diff --git a/src/libblastrampoline.c b/src/libblastrampoline.c index 39833ab..f0c1993 100644 --- a/src/libblastrampoline.c +++ b/src/libblastrampoline.c @@ -9,6 +9,35 @@ #define DEEPBINDLESS_INTERFACE_ILP64_LOADED 0x02 uint8_t deepbindless_interfaces_loaded = 0x00; +/* + * Optional process-global lock around the state-mutating API (`lbt_forward()`, + * `lbt_set_forward()`, `lbt_set_forward_by_index()`). Historically these are + * documented as thread-unsafe; defining `LBT_THREADSAFE` at compile time wraps + * them in a single lock so that concurrent initializers cannot corrupt the + * forwarding tables. When disabled, `lbt_lock()`/`lbt_unlock()` compile to + * no-ops, preserving the previous behavior exactly. + * + * NB: this guards the *mutators* only. The read-only accessors (e.g. + * `lbt_get_config()`, `lbt_get_forward()`) are intentionally left unlocked, + * so callers must still avoid racing reads against concurrent reconfiguration. + */ +#ifdef LBT_THREADSAFE +#ifdef _OS_WINDOWS_ +static CRITICAL_SECTION lbt_global_lock; +// Initialized in `DllMain()` (DLL_PROCESS_ATTACH) before any forwarding happens. +static inline void lbt_lock(void) { EnterCriticalSection(&lbt_global_lock); } +static inline void lbt_unlock(void) { LeaveCriticalSection(&lbt_global_lock); } +#else +#include +static pthread_mutex_t lbt_global_lock = PTHREAD_MUTEX_INITIALIZER; +static inline void lbt_lock(void) { pthread_mutex_lock(&lbt_global_lock); } +static inline void lbt_unlock(void) { pthread_mutex_unlock(&lbt_global_lock); } +#endif +#else +static inline void lbt_lock(void) { } +static inline void lbt_unlock(void) { } +#endif + int32_t find_symbol_idx(const char * name) { for (int32_t symbol_idx=0; exported_func_names[symbol_idx] != NULL; ++symbol_idx) { @@ -80,9 +109,10 @@ LBT_DLLEXPORT void lbt_set_default_func(const void * addr) { } /* - * Force a forward to a particular value. + * Force a forward to a particular value. Internal, unlocked worker; callers that + * already hold (or intentionally forgo) `lbt_global_lock` use this directly. */ -LBT_DLLEXPORT int32_t lbt_set_forward_by_index(int32_t symbol_idx, const void * addr, int32_t interface, int32_t complex_retstyle, int32_t f2c, int32_t verbose) { +static int32_t set_forward_by_index_impl(int32_t symbol_idx, const void * addr, int32_t interface, int32_t complex_retstyle, int32_t f2c, int32_t verbose) { // Quit out immediately if this is not a interface setting if (interface != LBT_INTERFACE_LP64 && interface != LBT_INTERFACE_ILP64) { return -1; @@ -163,6 +193,13 @@ LBT_DLLEXPORT int32_t lbt_set_forward_by_index(int32_t symbol_idx, const void * return 0; } +LBT_DLLEXPORT int32_t lbt_set_forward_by_index(int32_t symbol_idx, const void * addr, int32_t interface, int32_t complex_retstyle, int32_t f2c, int32_t verbose) { + lbt_lock(); + int32_t ret = set_forward_by_index_impl(symbol_idx, addr, interface, complex_retstyle, f2c, verbose); + lbt_unlock(); + return ret; +} + LBT_DLLEXPORT const void * lbt_get_forward(const char * symbol_name, int32_t interface, int32_t f2c) { // Search symbol list for `symbol_name`` int32_t symbol_idx = find_symbol_idx(symbol_name); @@ -202,22 +239,25 @@ LBT_DLLEXPORT const void * lbt_get_forward(const char * symbol_name, int32_t int } LBT_DLLEXPORT int32_t lbt_set_forward(const char * symbol_name, const void * addr, int32_t interface, int32_t complex_retstyle, int32_t f2c, int32_t verbose) { - // Search symbol list for `symbol_name`, then sub off to `set_forward_by_index()` + // Search symbol list for `symbol_name`, then sub off to `set_forward_by_index_impl()` int32_t symbol_idx = find_symbol_idx(symbol_name); if (symbol_idx == -1) return -1; - int32_t ret = lbt_set_forward_by_index(symbol_idx, addr, interface, complex_retstyle, f2c, verbose); + lbt_lock(); + int32_t ret = set_forward_by_index_impl(symbol_idx, addr, interface, complex_retstyle, f2c, verbose); if (ret == 0) { // Un-mark this symbol as being provided by any of our libraries; // if you use the footgun API, you can keep track of who is providing what. clear_forwarding_mark(symbol_idx, interface); } + lbt_unlock(); return ret; } // Load `libname`, clearing previous mappings if `clear` is set. -LBT_DLLEXPORT int32_t lbt_forward(const char * libname, int32_t clear, int32_t verbose, const char * suffix_hint) { +// Internal, unlocked worker; the public `lbt_forward()` wraps this under `lbt_global_lock`. +static int32_t lbt_forward_impl(const char * libname, int32_t clear, int32_t verbose, const char * suffix_hint) { if (verbose) { printf("Generating forwards to %s (clear: %d, verbose: %d, suffix_hint: '%s')\n", libname, clear, verbose, suffix_hint); } @@ -406,7 +446,7 @@ LBT_DLLEXPORT int32_t lbt_forward(const char * libname, int32_t clear, int32_t v } if (addr != NULL && addr != self_symbol_addr) { - lbt_set_forward_by_index(symbol_idx, addr, interface, complex_retstyle, f2c, verbose); + set_forward_by_index_impl(symbol_idx, addr, interface, complex_retstyle, f2c, verbose); LBT_BITFIELD_SET(forwards, symbol_idx); nforwards++; } @@ -443,6 +483,13 @@ LBT_DLLEXPORT int32_t lbt_forward(const char * libname, int32_t clear, int32_t v return nforwards; } +LBT_DLLEXPORT int32_t lbt_forward(const char * libname, int32_t clear, int32_t verbose, const char * suffix_hint) { + lbt_lock(); + int32_t ret = lbt_forward_impl(libname, clear, verbose, suffix_hint); + lbt_unlock(); + return ret; +} + /* * On windows it's surprisingly difficult to get a handle to ourselves, * and that's because they give it to you in `DllMain()`. ;) @@ -452,6 +499,10 @@ void * _win32_self_handle; BOOL APIENTRY DllMain(HINSTANCE hModule, DWORD code, void *reserved) { if (code == DLL_PROCESS_ATTACH) { _win32_self_handle = (void *)hModule; +#ifdef LBT_THREADSAFE + // Set up the global lock before anyone can call into the forwarding API. + InitializeCriticalSection(&lbt_global_lock); +#endif } else { // We do not want to run our initialization more than once per process. return TRUE; @@ -537,7 +588,9 @@ __attribute__((constructor)) void init(void) { curr_lib_start++; // Load functions from this library, clearing only the first time. - lbt_forward(curr_lib, clear, verbose, suffix_hint); + // We run inside the constructor (single-threaded), so call the unlocked + // worker directly rather than re-entering the public locking wrapper. + lbt_forward_impl(curr_lib, clear, verbose, suffix_hint); clear = 0; } } diff --git a/test/utils.jl b/test/utils.jl index 5a53c44..78cc3a9 100644 --- a/test/utils.jl +++ b/test/utils.jl @@ -95,11 +95,16 @@ function build_libblastrampoline() end cflags_add = "-Werror" * (needs_m32() ? " -m32" : "") + # Build with the optional internal locking compiled in by default, so the + # `LBT_THREADSAFE` code paths (pthread / CRITICAL_SECTION) are exercised across the + # whole test matrix. Set `LBT_THREADSAFE=0` in the environment to instead test the + # plain (no-op lock) build that ships by default. + threadsafe = get(ENV, "LBT_THREADSAFE", "1") dir = mktempdir() srcdir = joinpath(dirname(@__DIR__), "src") global blastrampoline_build_dir = joinpath(dir, "output") run(`$(make) -sC $(pathesc(srcdir)) CFLAGS="$(cflags_add)" ARCH=$(Sys.ARCH) clean`) - run(`$(make) -sC $(pathesc(srcdir)) CFLAGS="$(cflags_add)" ARCH=$(Sys.ARCH) install builddir=$(pathesc(dir))/build prefix=$(pathesc(blastrampoline_build_dir))`) + run(`$(make) -sC $(pathesc(srcdir)) CFLAGS="$(cflags_add)" ARCH=$(Sys.ARCH) LBT_THREADSAFE=$(threadsafe) install builddir=$(pathesc(dir))/build prefix=$(pathesc(blastrampoline_build_dir))`) link_name = blastrampoline_dev_link_name cp(