diff --git a/src/engine/CMakeLists.txt b/src/engine/CMakeLists.txt index 875b97ec..0f56e84a 100644 --- a/src/engine/CMakeLists.txt +++ b/src/engine/CMakeLists.txt @@ -37,6 +37,7 @@ set(MUJOCO_ENGINE_SRCS engine_derivative_fd.h engine_forward.c engine_forward.h + engine_global_table.h engine_inverse.c engine_inverse.h engine_island.c diff --git a/src/engine/engine_global_table.h b/src/engine/engine_global_table.h new file mode 100644 index 00000000..d619b62f --- /dev/null +++ b/src/engine/engine_global_table.h @@ -0,0 +1,282 @@ +// Copyright 2024 DeepMind Technologies Limited +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#ifndef MUJOCO_SRC_ENGINE_ENGINE_GLOBAL_TABLE_H_ +#define MUJOCO_SRC_ENGINE_ENGINE_GLOBAL_TABLE_H_ + +#include +#include +#include +#include +#include +#include +#include +#include + +#include "engine/engine_util_errmem.h" + +namespace mujoco { +static constexpr int kCacheLineBytes = 256; + +static inline bool CaseInsensitiveEqual(std::string_view s1, std::string_view s2) { + if (s1.length() != s2.length()) { + return false; + } + auto len = s1.length(); + for (decltype(len) i = 0; i < len; ++i) { + if (std::tolower(s1[i]) != std::tolower(s2[i])) { + return false; + } + } + return true; +} + +// A table intended for use as global storage for extension objects such as plugins, implemented as +// a linked list of array "blocks". This is a compromise that maintains a good degree of memory +// locality while not invalidating existing pointers when growing the table. It is expected that for +// most users, the number of objects loaded will be small enough to fit in the initial block, and +// so the global table will behave like an array. Since pointers are never invalidated, we do not +// need to apply a read lock on the global table when resolving an element. +template +struct alignas(kCacheLineBytes) TableBlock { + static constexpr int kBlockSize = 15; + + TableBlock() : objects{}, next(nullptr) { + static_assert( + sizeof(TableBlock) / kCacheLineBytes == + sizeof(TableBlock::objects) / kCacheLineBytes + + (sizeof(TableBlock::objects) % kCacheLineBytes > 0), + "TableBlock::next doesn't fit in the same cache line as the end of TableBlock::objects"); + } + + T objects[kBlockSize]; + TableBlock* next; +}; + + +using Mutex = std::mutex; + +class ReentrantWriteLock { + public: + ReentrantWriteLock(Mutex& mutex) : mutex_(mutex) { + if (LockCountOnCurrentThread() == 0) { + mutex_.lock(); + } + ++LockCountOnCurrentThread(); + } + + ~ReentrantWriteLock() { + if (--LockCountOnCurrentThread() == 0) { + mutex_.unlock(); + } + } + + private: + Mutex& mutex_; + + static int& LockCountOnCurrentThread() noexcept { + thread_local int counter = 0; + return counter; + } +}; + +template +class GlobalTable { + public: + static GlobalTable& GetSingleton() { + static_assert(std::is_trivially_destructible_v>); + static GlobalTable global; + return global; + } + + // Each extension object type T must implement these. + using ErrorMessage = char[512]; + static const char* HumanReadableTypeName(); + static std::string_view ObjectKey(const T&); + static bool ObjectEqual(const T&, const T&); + static bool CopyObject(T& dst, const T& src, ErrorMessage& err); + + int count() { + return count_.load(std::memory_order_acquire); + } + + ReentrantWriteLock LockExclusively() { + return ReentrantWriteLock(mutex()); + } + + int AppendIfUnique(const T& obj) { + ErrorMessage err = "\0"; + + // ========= ATTENTION! ======================================================================== + // Do not handle objects with nontrivial destructors outside of this lambda. + // Do not call mju_error inside this lambda. + int slot = [&]() { + auto lock = LockExclusively(); + + int count = count_.load(std::memory_order_acquire); + int local_idx = 0; + TableBlock* block = &first_block_; + + // check if a non-identical object has already been registered + for (int i = 0; i < count; ++i, ++local_idx) { + if (local_idx == TableBlock::kBlockSize) { + local_idx = 0; + block = block->next; + } + const T& existing = block->objects[local_idx]; + if (CaseInsensitiveEqual(ObjectKey(obj), ObjectKey(existing))) { + if (!ObjectEqual(obj, existing)) { + std::snprintf(err, sizeof(err), "%s '%s' is already registered", + HumanReadableTypeName(), std::string(ObjectKey(obj)).c_str()); + return -1; + } else { + return i; + } + } + } + + // allocate a new block if the last allocated block is full + if (local_idx == TableBlock::kBlockSize) { + local_idx = 0; + block->next = new(std::nothrow) TableBlock; + if (!block->next) { + std::snprintf(err, sizeof(err), "failed to allocate memory for a new %s table block", + HumanReadableTypeName()); + return -1; + } + block = block->next; + } + + // copy the new object into the table + if (!CopyObject(block->objects[local_idx], obj, err)) { + return -1; + } + + // increment the global count with a release memory barrier + count_.store(count + 1, std::memory_order_release); + + return count; + }(); + // ========= ATTENTION! ======================================================================== + // End of safe lambda, do not handle objects with non-trivial destructors beyond this point. + + // registration failed, throw an mju_error + if (slot < 0) { + err[sizeof(err) - 1] = '\0'; + mju_error("%s", err); + } + + return slot; + } + + // look up by slot number, assuming that count() has already been called + const T* GetAtSlotUnsafe(int slot, int nslot) { + if (slot < 0 || slot >= nslot) { + return nullptr; + } + + TableBlock* block = &first_block_; + + // iterate over blocks in the global table until the local index is less than the block size + int local_idx = slot; + while (local_idx >= TableBlock::kBlockSize) { + local_idx -= TableBlock::kBlockSize; + block = block->next; + if (!block) { + return nullptr; + } + } + + // local_idx is now a valid index into the current block + T* obj = &(block->objects[local_idx]); + + // check if obj has been initialized + if (obj && ObjectKey(*obj).empty()) { + return nullptr; + } + + return obj; + } + + // look up by key, assuming that count() has already been called + const T* GetByKeyUnsafe(std::string_view key, int* slot, int nslot) { + if (slot) *slot = -1; + + if (key.empty()) { + return nullptr; + } + + TableBlock* block = &first_block_; + int found_slot = 0; + while (block) { + for (int i = 0; + i < TableBlock::kBlockSize && found_slot < nslot; + ++i, ++found_slot) { + const T& obj = block->objects[i]; + + // reached an uninitialized object, which means that iterated beyond the object count + // this should never happen if `count` was actually returned by count() + std::string_view candidate_key = ObjectKey(obj); + if (candidate_key.empty()) { + return nullptr; + } + + // check if key matches the query + if (CaseInsensitiveEqual(candidate_key, key)) { + if (slot) *slot = found_slot; + return &obj; + } + } + + block = block->next; + } + + return nullptr; + } + + const T* GetAtSlot(int slot) { + // count() uses memory_order_acquire which acts as a barrier that guarantees that all + // objects up to `count` have been completely inserted + return GetAtSlotUnsafe(slot, count()); + } + + const T* GetByKey(std::string_view key, int* slot) { + // count() uses memory_order_acquire which acts as a barrier that guarantees that all + // objects up to `count` have been completely inserted + return GetByKeyUnsafe(key, slot, count()); + } + + private: + GlobalTable() { + new(mutex_) Mutex; + } + + Mutex& mutex() { + return *std::launder(reinterpret_cast(&mutex_)); + } + + TableBlock first_block_; + std::atomic_int count_; + + // A mutex whose destructor is never run. + // When a C++ program terminates, the destructors for function static objects and globals will be + // executed by whichever thread started that termination but there is no guarantee that other + // threads have terminated. In other words, a static object may be accessed by another thread + // after it is deleted. We avoid destruction issues by never running the destructor. + alignas(Mutex) unsigned char mutex_[sizeof(Mutex)]; +}; + +} // namespace mujoco + +#endif // MUJOCO_SRC_ENGINE_ENGINE_GLOBAL_TABLE_H_ diff --git a/src/engine/engine_plugin.cc b/src/engine/engine_plugin.cc index 50b93d61..93ea4233 100644 --- a/src/engine/engine_plugin.cc +++ b/src/engine/engine_plugin.cc @@ -20,20 +20,13 @@ #include "engine/engine_plugin.h" -#include #include -#include -#include #include #include #include -#include #include -#include #include -#include -#include -#include +#include extern "C" { #if defined(_WIN32) || defined(__CYGWIN__) @@ -45,6 +38,7 @@ extern "C" { } #include +#include "engine/engine_global_table.h" #include "engine/engine_util_errmem.h" // set default plugin definition @@ -53,139 +47,11 @@ void mjp_defaultPlugin(mjpPlugin* plugin) { } namespace { +using mujoco::GlobalTable; + constexpr int kMaxNameLength = 1024; constexpr int kMaxAttributes = 255; -constexpr int kCacheLine = 256; - -// A table of registered plugins, implemented as a linked list of array "blocks". -// This is a compromise that maintains a good degree of memory locality while not invalidating -// existing pointers when growing the table. It is expected that for most users, the number of -// plugins loaded into a program will be small enough to fit in the initial block, and so the global -// table will behave like an array. Since pointers are never invalidated, we do not need to apply a -// read lock on the global table when resolving a plugin. -template -struct alignas(kCacheLine) PluginTable { - static constexpr int kBlockSize = 15; - - PluginTable() { - for (int i = 0; i < kBlockSize; ++i) { - std::memset(&plugins[i], 0, sizeof(plugins[i])); - } - } - - T plugins[kBlockSize]; - PluginTable* next = nullptr; -}; - -static_assert( - sizeof(PluginTable) / kCacheLine == - sizeof(PluginTable::plugins) / kCacheLine - + (sizeof(PluginTable::plugins) % kCacheLine > 0), - "PluginTable::next doesn't fit in the same cache line as the end of PluginTable::plugins"); - -using Mutex = std::shared_mutex; - -class ReentrantWriteLock { - public: - ReentrantWriteLock(Mutex& mutex) : mutex_(mutex) { - if (LockCountOnCurrentThread() == 0) { - mutex_.lock(); - } - ++LockCountOnCurrentThread(); - } - - ~ReentrantWriteLock() { - if (--LockCountOnCurrentThread() == 0) { - mutex_.unlock(); - } - } - - private: - Mutex& mutex_; - - static int& LockCountOnCurrentThread() noexcept { - thread_local int counter = 0; - return counter; - } -}; - -template -class Global { - public: - Global() { - new(mutex_) Mutex; - } - PluginTable& table() { - return table_; - } - std::atomic_int& count() { - return count_; - } - Mutex& mutex() { - return *std::launder(reinterpret_cast(&mutex_)); - } - - ReentrantWriteLock lock_mutex_exclusively() { - return ReentrantWriteLock(mutex()); - } - - private: - PluginTable table_; - std::atomic_int count_; - - // A mutex whose destructor is never run. - // When a C++ program terminates, the destructors for function static objects and globals will be - // executed by whichever thread started that termination but there is no guarantee that other - // threads have terminated. In other words, a static object may be accessed by another thread - // after it is deleted. We avoid destruction issues by never running the destructor. - alignas(Mutex) unsigned char mutex_[sizeof(Mutex)]; -}; - -template -Global& GetGlobal() { - static Global global; - static_assert(std::is_trivially_destructible_v); - return global; -} - -// allocate new block for the global table -template -PluginTable* AddNewTableBlock(PluginTable* table) { - char err[512]; - err[0] = '\0'; - table->next = new(std::nothrow) PluginTable; - if (!table->next) { - std::snprintf(err, sizeof(err), "failed to allocate memory for the global plugin table"); - return nullptr; - } - return table->next; -} - -// look up a plugin by slot number -template -const T* GetAtSlot(int slot, int nslot) { - if (slot < 0 || slot >= nslot) { - return nullptr; - } - - Global& global = GetGlobal(); - PluginTable* table = &global.table(); - - // iterate over blocks in the global table until the local index is less than the block size - int local_idx = slot; - while (local_idx >= PluginTable::kBlockSize) { - local_idx -= PluginTable::kBlockSize; - table = table->next; - if (!table) { - return nullptr; - } - } - - // local_idx is now a valid index into the current block - return &(table->plugins[local_idx]); -} - // return the length of a null-terminated string, or -1 if it is not terminated after kMaxNameLength int strklen(const char* s) { for (int i = 0; i < kMaxNameLength; ++i) { @@ -211,68 +77,6 @@ std::unique_ptr CopyName(const char* s) { return out; } -// check if two plugins are identical -bool PluginsAreIdentical(const mjpPlugin& plugin1, const mjpPlugin& plugin2) { - if (plugin1.name && !plugin2.name) { - return false; - } - if (plugin2.name && !plugin1.name) { - return false; - } - if (plugin1.name && plugin2.name && - std::strncmp(plugin1.name, plugin2.name, kMaxNameLength)) { - return false; - } - - if (plugin1.nattribute != plugin2.nattribute) { - return false; - } - for (int i = 0; i < plugin1.nattribute; ++i) { - if (plugin1.attributes[i] && !plugin2.attributes[i]) { - return false; - } - if (plugin2.attributes[i] && !plugin1.attributes[i]) { - return false; - } - if (plugin1.attributes[i] && plugin2.attributes[i] && - std::strncmp(plugin1.attributes[i], plugin2.attributes[i], - kMaxNameLength)) { - return false; - } - } - - const char* ptr1 = reinterpret_cast(&plugin1.attributes) + - sizeof(plugin1.attributes); - const char* ptr2 = reinterpret_cast(&plugin2.attributes) + - sizeof(plugin2.attributes); - std::size_t remaining_size = - sizeof(mjpPlugin) - (ptr1 - reinterpret_cast(&plugin1)); - return !std::memcmp(ptr1, ptr2, remaining_size); -} - -// does case insensitive comparison -bool PrefixesAreIdentical(const char* p1, const char* p2) { - int i = 0; - for (; p1[i] != '\0'; i++) { - if (std::tolower(p1[i]) != std::tolower(p2[i])) { - return false; - } - } - - return p2[i] == '\0'; -} - -// check if two resource providers are identical -bool ResourceProvidersAreIdentical(const mjpResourceProvider* p1, const mjpResourceProvider* p2) { - return (PrefixesAreIdentical(p1->prefix, p2->prefix) && - p1->open == p2->open && - p1->read == p2->read && - p1->close == p2->close && - p1->getdir == p2->getdir && - p1->modified == p2->modified && - p1->data == p2->data); -} - // check if prefix is a valid URI scheme format bool IsValidURISchemeFormat(const char* prefix) { int len; @@ -299,8 +103,168 @@ bool IsValidURISchemeFormat(const char* prefix) { return true; } +// seek the nth config attrib of a plugin instance by counting null terminators +const char* PluginAttrSeek(const mjModel* m, int plugin_id, int attrib_id) { + const char* ptr = m->plugin_attr + m->plugin_attradr[plugin_id]; + for (int i = 0; i < attrib_id; ++i) { + while (*ptr) { + ++ptr; + } + ++ptr; + } + return ptr; +} } // namespace +template <> +const char* GlobalTable::HumanReadableTypeName() { + return "plugin"; +} + +template <> +std::string_view GlobalTable::ObjectKey(const mjpPlugin& plugin) { + return std::string_view(plugin.name, strklen(plugin.name)); +} + +// check if two plugins are identical +template <> +bool GlobalTable::ObjectEqual(const mjpPlugin& plugin1, const mjpPlugin& plugin2) { + if (plugin1.name && !plugin2.name) { + return false; + } + if (plugin2.name && !plugin1.name) { + return false; + } + if (plugin1.name && plugin2.name && + std::strncmp(plugin1.name, plugin2.name, kMaxNameLength)) { + return false; + } + + if (plugin1.nattribute != plugin2.nattribute) { + return false; + } + for (int i = 0; i < plugin1.nattribute; ++i) { + if (plugin1.attributes[i] && !plugin2.attributes[i]) { + return false; + } + if (plugin1.attributes[i] && !plugin2.attributes[i]) { + return false; + } + if (plugin1.attributes[i] && plugin2.attributes[i] && + std::strncmp(plugin1.attributes[i], plugin2.attributes[i], + kMaxNameLength)) { + return false; + } + } + + const char* ptr1 = reinterpret_cast(&plugin1.attributes) + + sizeof(plugin1.attributes); + const char* ptr2 = reinterpret_cast(&plugin2.attributes) + + sizeof(plugin2.attributes); + std::size_t remaining_size = + sizeof(mjpPlugin) - (ptr1 - reinterpret_cast(&plugin1)); + return !std::memcmp(ptr1, ptr2, remaining_size); +} + +template <> +bool GlobalTable::CopyObject(mjpPlugin& dst, const mjpPlugin& src, ErrorMessage& err) { + // check and copy the plugin name + std::unique_ptr name = CopyName(src.name); + if (!name) { + if (strklen(src.name) == -1) { + std::snprintf(err, sizeof(err), + "plugin->name length exceeds the maximum limit of %d", kMaxNameLength); + } else { + std::snprintf(err, sizeof(err), "failed to allocate memory for plugin name"); + } + return false; + } + + // check and copy plugin attributes + std::unique_ptr[]> attributes_list; + if (src.nattribute) { + attributes_list.reset(new(std::nothrow) std::unique_ptr[src.nattribute]); + if (!attributes_list) { + std::snprintf(err, sizeof(err), "failed to allocate memory for plugin attribute list"); + return false; + } + for (int i = 0; i < src.nattribute; ++i) { + std::unique_ptr attr = CopyName(src.attributes[i]); + if (!attr) { + if (strklen(src.attributes[i]) == -1) { + std::snprintf( + err, sizeof(err), + "length of plugin attribute %d exceeds the maximum limit of %d", i, kMaxAttributes); + } else { + std::snprintf(err, sizeof(err), "failed to allocate memory for plugin attribute %d", i); + } + return false; + } + attributes_list[i].swap(attr); + } + } + + // release the attribute names from unique_ptr into a plain array + const char** attributes = nullptr; + if (src.nattribute) { + attributes = new(std::nothrow) const char*[src.nattribute]; + if (!attributes) { + std::snprintf(err, sizeof(err), "failed to allocate memory for plugin attribute array"); + return -1; + } + for (int i = 0; i < src.nattribute; ++i) { + attributes[i] = attributes_list[i].release(); + } + } + + dst = src; + dst.name = name.release(); + dst.attributes = attributes; + + return true; +} + +template <> +const char* GlobalTable::HumanReadableTypeName() { + return "resource provider"; +} + +template <> +std::string_view GlobalTable::ObjectKey(const mjpResourceProvider& plugin) { + return std::string_view(plugin.prefix, strklen(plugin.prefix)); +} + +// check if two resource providers are identical +template <> +bool GlobalTable::ObjectEqual(const mjpResourceProvider& p1, const mjpResourceProvider& p2) { + return (CaseInsensitiveEqual(p1.prefix, p2.prefix) && + p1.open == p2.open && + p1.read == p2.read && + p1.close == p2.close && + p1.getdir == p2.getdir && + p1.modified == p2.modified && + p1.data == p2.data); +} + +template <> +bool GlobalTable::CopyObject(mjpResourceProvider& dst, const mjpResourceProvider& src, ErrorMessage& err) { + // copy prefix + std::unique_ptr prefix = CopyName(src.prefix); + if (!prefix) { + if (strklen(src.prefix) == -1) { + std::snprintf(err, sizeof(err), + "provider->prefix length exceeds the maximum limit of %d", kMaxNameLength); + } else { + std::snprintf(err, sizeof(err), "failed to allocate memory for resource provider prefix"); + } + return false; + } + + dst = src; + dst.prefix = prefix.release(); + return true; +} + // globally register a plugin (thread-safe), return new slot id int mjp_registerPlugin(const mjpPlugin* plugin) { if (!plugin->name) { @@ -314,196 +278,34 @@ int mjp_registerPlugin(const mjpPlugin* plugin) { kMaxAttributes); } - char err[512]; - err[0] = '\0'; - - // ========= ATTENTION! ========================================================================== - // Do not handle objects with nontrivial destructors outside of this lambda. - // Do not call mju_error inside this lambda. - int slot = [&]() -> int { - // check and copy the plugin name - std::unique_ptr name = CopyName(plugin->name); - if (!name) { - if (strklen(plugin->name) == -1) { - std::snprintf(err, sizeof(err), - "plugin->name length exceeds the maximum limit of %d", kMaxNameLength); - } else { - std::snprintf(err, sizeof(err), "failed to allocate memory for plugin name"); - } - return -1; - } - - // check and copy plugin attributes - std::vector> attributes_vec; - if (plugin->nattribute) { - attributes_vec.reserve(plugin->nattribute); - for (int i = 0; i < plugin->nattribute; ++i) { - std::unique_ptr attr = CopyName(plugin->attributes[i]); - if (!attr) { - if (strklen(plugin->attributes[i]) == -1) { - std::snprintf( - err, sizeof(err), - "plugin->attributes[%d] exceeds the maximum limit of %d", i, kMaxAttributes); - } else { - std::snprintf(err, sizeof(err), "failed to allocate memory for plugin attribute"); - } - return -1; - } - attributes_vec.emplace_back(std::move(attr)); - } - } - - Global& global = GetGlobal(); - auto lock = global.lock_mutex_exclusively(); - - int count = global.count().load(std::memory_order_acquire); - int local_idx = 0; - PluginTable* table = &global.table(); - - // check if a non-identical plugin with the same name has already been registered - for (int i = 0; i < count; ++i, ++local_idx) { - if (local_idx == PluginTable::kBlockSize) { - local_idx = 0; - table = table->next; - } - mjpPlugin& existing = table->plugins[local_idx]; - if (std::strcmp(plugin->name, existing.name) == 0) { - if (PluginsAreIdentical(*plugin, existing)) { - return i; - } else { - std::snprintf(err, sizeof(err), "plugin '%s' is already registered", plugin->name); - return -1; - } - } - } - - // allocate a new block of PluginTable if the last allocated block is full - if (local_idx == PluginTable::kBlockSize) { - local_idx = 0; - table = AddNewTableBlock(table); - if (!table) { - return -1; - } - } - - // release the attribute names from unique_ptr into a plain array - const char** attributes = nullptr; - if (plugin->nattribute) { - attributes = new(std::nothrow) const char*[plugin->nattribute]; - if (!attributes) { - std::snprintf(err, sizeof(err), "failed to allocate memory for plugin attribute array"); - return -1; - } - for (int i = 0; i < plugin->nattribute; ++i) { - attributes[i] = attributes_vec[i].release(); - } - } - - // all checked passed, actually register the plugin into the global table - mjpPlugin& registered_plugin = table->plugins[local_idx]; - registered_plugin = *plugin; - registered_plugin.name = name.release(); - registered_plugin.attributes = attributes; - - // increment the global plugin count with a release memory barrier - global.count().store(count + 1, std::memory_order_release); - - return count; - }(); - - // ========= ATTENTION! ========================================================================== - // End of safe lambda, do not handle objects with non-trivial destructors beyond this point. - - // plugin registration failed, throw an mju_error - if (slot < 0) { - err[sizeof(err) - 1] = '\0'; - mju_error("%s", err); - } - - return slot; + return GlobalTable::GetSingleton().AppendIfUnique(*plugin); } // look up plugin by slot number, assuming that mjp_pluginCount has already been called const mjpPlugin* mjp_getPluginAtSlotUnsafe(int slot, int nslot) { - const mjpPlugin* plugin = GetAtSlot(slot, nslot); - if (!plugin || !plugin->name) { - return nullptr; - } - return plugin; + return GlobalTable::GetSingleton().GetAtSlotUnsafe(slot, nslot); } // look up plugin by name, assuming that mjp_pluginCount has already been called const mjpPlugin* mjp_getPluginUnsafe(const char* name, int* slot, int nslot) { - if (slot) *slot = -1; - - if (!name || !name[0]) { - return nullptr; - } - - Global& plugins = GetGlobal(); - PluginTable* table = &plugins.table(); - int found_slot = 0; - while (table) { - for (int i = 0; - i < PluginTable::kBlockSize && found_slot < nslot; - ++i, ++found_slot) { - const mjpPlugin& plugin = table->plugins[i]; - - // reached an uninitialized plugin, which means that iterated beyond the plugin count - // this should never happen if `count` was actually returned by mjp_pluginCount - if (!plugin.name) { - return nullptr; - } - if (std::strcmp(plugin.name, name) == 0) { - if (slot) *slot = found_slot; - return &plugin; - } - } - table = table->next; - } - - return nullptr; + return GlobalTable::GetSingleton().GetByKeyUnsafe(name, slot, nslot); } // return the number of globally registered plugins int mjp_pluginCount() { - return GetGlobal().count().load(std::memory_order_acquire); + return GlobalTable::GetSingleton().count(); } // look up a plugin by slot number const mjpPlugin* mjp_getPluginAtSlot(int slot) { - const int count = mjp_pluginCount(); - - // mjp_pluginCount uses memory_order_acquire which acts as a barrier that guarantees that all - // plugins up to `count` have been completely inserted - return mjp_getPluginAtSlotUnsafe(slot, count); + return GlobalTable::GetSingleton().GetAtSlot(slot); } // look up a plugin by name, optionally also get its registered slot number const mjpPlugin* mjp_getPlugin(const char* name, int* slot) { - const int count = mjp_pluginCount(); - int found_slot = -1; - const mjpPlugin* plugin = mjp_getPluginUnsafe(name, &found_slot, count); - - if (slot) *slot = found_slot; - return plugin; + return GlobalTable::GetSingleton().GetByKey(name, slot); } -namespace { -// seek the nth config attrib of a plugin instance by counting null terminators -const char* PluginAttrSeek(const mjModel* m, int plugin_id, int attrib_id) { - const char* ptr = m->plugin_attr + m->plugin_attradr[plugin_id]; - for (int i = 0; i < attrib_id; ++i) { - while (*ptr) { - ++ptr; - } - ++ptr; - } - return ptr; -} -} // namespace - - // return a config attribute of a plugin instance // NULL: invalid plugin instance ID or attribute name const char* mj_getPluginConfig(const mjModel* m, int plugin_id, const char* attrib) { @@ -544,95 +346,16 @@ int mjp_registerResourceProvider(const mjpResourceProvider* provider) { return -1; } - char err[512]; - err[0] = '\0'; - - // ========= ATTENTION! ========================================================================== - // Do not handle objects with nontrivial destructors outside of this lambda. - // Do not call mju_error inside this lambda. - int slot = [&]() -> int { - std::unique_ptr prefix; - - // copy prefix - prefix = CopyName(provider->prefix); - if (!prefix) { - if (strklen(provider->prefix) == -1) { - std::snprintf(err, sizeof(err), - "provider->prefix length exceeds the maximum limit of %d", kMaxNameLength); - } else { - std::snprintf(err, sizeof(err), "failed to allocate memory for resource provider prefix"); - } - return -1; - } - - Global& global = GetGlobal(); - auto lock = global.lock_mutex_exclusively(); - int count = global.count().load(std::memory_order_acquire); - int local_idx = 0; - PluginTable* table = &global.table(); - - // check if a non-identical provider with the same name has already been registered - for (int i = 0; i < count; ++i, ++local_idx) { - if (local_idx == PluginTable::kBlockSize) { - local_idx = 0; - table = table->next; - } - mjpResourceProvider& existing = table->plugins[local_idx]; - - if (existing.prefix != nullptr) { - // if identical then return slot number - if (PrefixesAreIdentical(provider->prefix, existing.prefix)) { - if (ResourceProvidersAreIdentical(provider, &existing)) { - return i; - } else { - std::snprintf(err, sizeof(err), - "a resource provider with prefix '%s' cannot be registered", - provider->prefix); - return -1; - } - } - } - } - - // allocate a new block of PluginTable if the last allocated block is full - if (local_idx == PluginTable::kBlockSize) { - local_idx = 0; - table = AddNewTableBlock(table); - if (!table) { - return -1; - } - } - - // all checks passed, register the plugin into the global table - mjpResourceProvider& registered_provider = table->plugins[local_idx]; - registered_provider = *provider; - registered_provider.prefix = prefix.release(); - - // increment the global plugin count with a release memory barrier - global.count().store(count + 1, std::memory_order_release); - return count; - }(); - - // ========= ATTENTION! ========================================================================== - // End of safe lambda, do not handle objects with non-trivial destructors beyond this point. - - // plugin registration failed, throw a warning - if (slot < 0) { - err[sizeof(err) - 1] = '\0'; - mju_warning("%s", err); - } - - return slot+1; + return GlobalTable::GetSingleton().AppendIfUnique(*provider) + 1; } // return the number of globally registered resource providers int mjp_resourceProviderCount() { - return GetGlobal().count().load(std::memory_order_acquire); + return GlobalTable::GetSingleton().count(); } // look up a resource provider that matches its prefix against the given resource scheme const mjpResourceProvider* mjp_getResourceProvider(const char* resource_name) { - const int count = mjp_resourceProviderCount(); if (!resource_name || !resource_name[0]) { return nullptr; } @@ -650,41 +373,13 @@ const mjpResourceProvider* mjp_getResourceProvider(const char* resource_name) { return nullptr; } - Global& global = GetGlobal(); - auto lock = global.lock_mutex_exclusively(); - PluginTable* table = &global.table(); - int found_slot = 0; - - while (table) { - for (int i = 0; - i < PluginTable::kBlockSize && found_slot < count; - ++i, ++found_slot) { - const mjpResourceProvider& provider = table->plugins[i]; - const char *prefix = provider.prefix; - - if (prefix != nullptr && - PrefixesAreIdentical(prefix, file_prefix.c_str())) { - return &provider; - } - } - table = table->next; - } - - return nullptr; + return GlobalTable::GetSingleton().GetByKey(file_prefix.c_str(), nullptr); } // look up a resource provider by slot number const mjpResourceProvider* mjp_getResourceProviderAtSlot(int slot) { - // mjp_resourceProviderCount uses memory_order_acquire which acts as a barrier - // that guarantees that all providers up to `count` have been completely inserted - const int count = mjp_resourceProviderCount(); - // shift slot to be zero-indexed - const mjpResourceProvider* provider = GetAtSlot(slot - 1, count); - if (!provider || provider->prefix[0] == '\0') { - return nullptr; - } - return provider; + return GlobalTable::GetSingleton().GetAtSlot(slot - 1); } // load plugins from a dynamic library @@ -704,9 +399,9 @@ void mj_loadAllPluginLibraries(const char* directory, int nplugin_before; int nplugin_after; - Global& global = GetGlobal(); + auto& plugin_table = GlobalTable::GetSingleton(); { - auto lock = global.lock_mutex_exclusively(); + auto lock = plugin_table.LockExclusively(); nplugin_before = mjp_pluginCount(); mj_loadPluginLibrary(dso_path.c_str()); nplugin_after = mjp_pluginCount();