Refactor VFS logic out of ResourceProvider plugin code.

PiperOrigin-RevId: 561302536
Change-Id: I952f9e76e85a7e3a111e0338f307378a6ff58197
This commit is contained in:
Kyle Bayes
2023-08-30 04:46:32 -07:00
committed by Copybara-Service
parent 215433f5f7
commit f86e8b449f
22 changed files with 310 additions and 445 deletions
+6 -24
View File
@@ -661,7 +661,7 @@ void mj_saveModel(const mjModel* m, const char* filename, void* buffer, int buff
// load model from binary MJB resource
static mjModel* _mj_loadModel(const char* filename, int vfs_provider) {
mjModel* mj_loadModel(const char* filename, const mjVFS* vfs) {
int header[NHEADER] = {0};
int expected_header[NHEADER] = {ID, sizeof(mjtNum), getnint(), getnsize(), getnptr()};
int ints[256];
@@ -670,8 +670,11 @@ static mjModel* _mj_loadModel(const char* filename, int vfs_provider) {
mjModel *m = 0;
mjResource* r = NULL;
if((r = mju_openResource(filename, vfs_provider)) == NULL) {
return NULL;
// first try vfs, otherwise try a provider or OS filesystem
if ((r = mju_openVfsResource(filename, vfs)) == NULL) {
if ((r = mju_openResource(filename)) == NULL) {
return NULL;
}
}
const void* buffer = NULL;
@@ -785,27 +788,6 @@ static mjModel* _mj_loadModel(const char* filename, int vfs_provider) {
// load model from binary MJB file
// if vfs is not NULL, look up file in vfs before reading from disk
mjModel* mj_loadModel(const char* filename, const mjVFS* vfs) {
if (vfs == NULL) {
return _mj_loadModel(filename, 0);
}
int index = mj_registerVfsProvider(vfs);
if (index < 1) {
mjERROR("could not allocate memory");
return NULL;
}
mjModel* model = _mj_loadModel(filename, index);
mjp_unregisterResourceProvider(index);
return model;
}
// de-allocate mjModel
void mj_deleteModel(mjModel* m) {
if (m) {
+16 -102
View File
@@ -24,6 +24,7 @@
#include <cctype>
#include <cstddef>
#include <cstdlib>
#include <cstdio>
#include <cstring>
#include <memory>
#include <mutex>
@@ -57,9 +58,6 @@ constexpr int kMaxAttributes = 255;
constexpr int kCacheLine = 256;
// vfs prefix
constexpr const char* kVfsPrefix = mjVFS_PREFIX;
// 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
@@ -270,6 +268,7 @@ bool ResourceProvidersAreIdentical(const mjpResourceProvider* p1, const mjpResou
p1->open == p2->open &&
p1->read == p2->read &&
p1->close == p2->close &&
p1->getdir == p2->getdir &&
p1->data == p2->data);
}
@@ -532,18 +531,7 @@ void mjp_defaultResourceProvider(mjpResourceProvider* provider) {
// globally register a resource provider (thread-safe), return new slot id
int mjp_registerResourceProvider(const mjpResourceProvider* provider) {
// check against reserved prefixes
if (PrefixesAreIdentical(kVfsPrefix, provider->prefix)) {
mju_warning("provider->prefix is '%s' which is reserved", provider->prefix);
return -1;
}
return mjp_registerResourceProviderInternal(provider);
}
// internal version of mjp_registerResourceProvider without prechecks on reserved prefixes
int mjp_registerResourceProviderInternal(const mjpResourceProvider* provider) {
// check if prefix is valid URI scheme format
// check if prefix is valid URI scheme format
if (!IsValidURISchemeFormat(provider->prefix)) {
mju_warning("provider->prefix is '%s' which is not a valid URI scheme format",
provider->prefix);
@@ -562,35 +550,25 @@ int mjp_registerResourceProviderInternal(const mjpResourceProvider* provider) {
// Do not handle objects with nontrivial destructors outside of this lambda.
// Do not call mju_error inside this lambda.
int slot = [&]() -> int {
bool vfs_provider = false;
std::unique_ptr<char[]> prefix;
// check if this is a VFS provider
if (PrefixesAreIdentical(mjVFS_PREFIX, provider->prefix)) {
vfs_provider = true;
}
// copy prefix
if (!vfs_provider) {
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;
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<mjpResourceProvider>& global = GetGlobal<mjpResourceProvider>();
auto lock = global.lock_mutex_exclusively();
int count = global.count().load(std::memory_order_acquire);
int local_idx = 0;
int free_idx = -1, free_local_idx = -1;
PluginTable<mjpResourceProvider>* table = &global.table();
PluginTable<mjpResourceProvider>* free_table = nullptr;
// check if a non-identical provider with the same name has already been registered
for (int i = 0; i < count; ++i, ++local_idx) {
@@ -600,17 +578,7 @@ int mjp_registerResourceProviderInternal(const mjpResourceProvider* provider) {
}
mjpResourceProvider& existing = table->plugins[local_idx];
// VFS providers can safely go in open slots
if (vfs_provider && existing.prefix == nullptr && free_idx == -1) {
free_table = table;
free_idx = i;
free_local_idx = local_idx;
// can skip the rest
break;
}
if (!vfs_provider && existing.prefix != nullptr) {
if (existing.prefix != nullptr) {
// if identical then return slot number
if (PrefixesAreIdentical(provider->prefix, existing.prefix)) {
if (ResourceProvidersAreIdentical(provider, &existing)) {
@@ -626,7 +594,7 @@ int mjp_registerResourceProviderInternal(const mjpResourceProvider* provider) {
}
// allocate a new block of PluginTable if the last allocated block is full
if (free_local_idx == -1 && local_idx == PluginTable<mjpResourceProvider>::kBlockSize) {
if (local_idx == PluginTable<mjpResourceProvider>::kBlockSize) {
local_idx = 0;
table = AddNewTableBlock<mjpResourceProvider>(table);
if (!table) {
@@ -636,19 +604,12 @@ int mjp_registerResourceProviderInternal(const mjpResourceProvider* provider) {
// all checks passed, register the plugin into the global table
mjpResourceProvider& registered_provider = table->plugins[local_idx];
if (free_local_idx != -1) {
registered_provider = free_table->plugins[free_local_idx];
}
registered_provider = *provider;
registered_provider.prefix = (!vfs_provider) ? prefix.release() : kVfsPrefix;
registered_provider.prefix = prefix.release();
// increment the global plugin count with a release memory barrier
if (free_idx == -1) {
free_idx = count;
global.count().store(count + 1, std::memory_order_release);
}
return free_idx;
global.count().store(count + 1, std::memory_order_release);
return count;
}();
// ========= ATTENTION! ==========================================================================
@@ -663,46 +624,6 @@ int mjp_registerResourceProviderInternal(const mjpResourceProvider* provider) {
return slot+1;
}
// globally unregister resource provider (thread-safe)
// only used for VFS resource providers
void mjp_unregisterResourceProvider(int slot) {
// shift slot to zero-index
slot--;
if (slot < 0) {
return;
}
// get global table, acquire lock
Global<mjpResourceProvider>& global = GetGlobal<mjpResourceProvider>();
auto lock = global.lock_mutex_exclusively();
int count = global.count().load(std::memory_order_acquire);
if (slot >= count) {
return;
}
PluginTable<mjpResourceProvider>* 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<mjpResourceProvider>::kBlockSize) {
local_idx -= PluginTable<mjpResourceProvider>::kBlockSize;
table = table->next;
if (!table) {
return;
}
}
// local_idx is now a valid index into the current block
mjpResourceProvider& provider = table->plugins[local_idx];
// no-op for anything other than VFS resource providers
if (provider.prefix == kVfsPrefix) {
provider.prefix = nullptr;
}
}
// return the number of globally registered resource providers
int mjp_resourceProviderCount() {
return GetGlobal<mjpResourceProvider>().count().load(std::memory_order_acquire);
@@ -728,12 +649,6 @@ const mjpResourceProvider* mjp_getResourceProvider(const char* resource_name) {
return nullptr;
}
// since multiple VFS resource providers can be registered with the same
// prefix, it doesn't make sense to try to match against them
if (PrefixesAreIdentical(kVfsPrefix, file_prefix.c_str())) {
return nullptr;
}
Global<mjpResourceProvider>& global = GetGlobal<mjpResourceProvider>();
auto lock = global.lock_mutex_exclusively();
PluginTable<mjpResourceProvider>* table = &global.table();
@@ -743,7 +658,6 @@ const mjpResourceProvider* mjp_getResourceProvider(const char* resource_name) {
for (int i = 0;
i < PluginTable<mjpPlugin>::kBlockSize && found_slot < count;
++i, ++found_slot) {
const mjpResourceProvider& provider = table->plugins[i];
const char *prefix = provider.prefix;
-3
View File
@@ -31,9 +31,6 @@ MJAPI int mjp_registerPlugin(const mjpPlugin* plugin);
// globally register a resource provider (thread-safe), return new slot id
MJAPI int mjp_registerResourceProvider(const mjpResourceProvider* provider);
// globally unregister a resource provider (thread-safe)
MJAPI void mjp_unregisterResourceProvider(int slot);
// return the number of globally registered plugins
MJAPI int mjp_pluginCount();
+39 -82
View File
@@ -16,40 +16,23 @@
#include <limits.h>
#include <stddef.h>
#include <stdint.h>
#include <stdio.h>
#include <string.h>
#include <mujoco/mjmodel.h>
#include <mujoco/mjplugin.h>
#include "engine/engine_plugin.h"
#include "engine/engine_util_errmem.h"
// file buffer used internally for the OS filesystem
typedef struct {
void* buffer;
int nbuffer;
uint8_t* buffer; // raw bytes from file
size_t nbuffer; // size of buffer in bytes
} file_buffer;
// helper function to fill data from resource provider into provider
static void fillResource(const mjpResourceProvider* provider, mjResource* resource) {
if (provider == NULL) {
resource->read = NULL;
resource->close = NULL;
resource->getdir = NULL;
resource->provider_data = NULL;
} else {
resource->read = provider->read;
resource->close = provider->close;
resource->getdir = provider->getdir;
resource->provider_data = provider->data;
}
}
// open the given resource; if the name doesn't have a prefix matching with a
// resource provider, then the default_provider is used
// if default_provider non-positive, then the OS filesystem is used
mjResource* mju_openResource(const char* name, int default_provider) {
// resource provider, then the OS filesystem is used
mjResource* mju_openResource(const char* name) {
mjResource* resource = (mjResource*) mju_malloc(sizeof(mjResource));
const mjpResourceProvider* provider = NULL;
if (resource == NULL) {
@@ -57,19 +40,22 @@ mjResource* mju_openResource(const char* name, int default_provider) {
return NULL;
}
// clear out resource
memset(resource, 0, sizeof(mjResource));
// copy name
resource->name = mju_malloc(sizeof(char) * (strlen(name) + 1));
if (resource->name == NULL) {
mju_free(resource);
mju_closeResource(resource);
mjERROR("could not allocate memory");
return NULL;
}
strcpy(resource->name, name);
memcpy(resource->name, name, sizeof(char) * (strlen(name) + 1));
// find provider based off prefix of name
provider = mjp_getResourceProvider(name);
if (provider != NULL) {
fillResource(provider, resource);
resource->provider = provider;
if (provider->open(resource)) {
return resource;
}
@@ -77,48 +63,19 @@ mjResource* mju_openResource(const char* name, int default_provider) {
mju_warning("mju_openResource: could not open resource '%s' "
"using a resource provider matching prefix '%s'",
name, provider->prefix);
mju_free(resource->name);
mju_free(resource);
return NULL;
}
// fallback to default provider
if (default_provider > 0) {
provider = mjp_getResourceProviderAtSlot(default_provider);
if (provider == NULL) {
mju_warning("mju_openResource: unknown resource provider at slot %d",
default_provider);
mju_free(resource->name);
mju_free(resource);
return NULL;
}
fillResource(provider, resource);
if (provider->open(resource)) {
return resource;
}
mju_warning("mju_openResource: could not open resource '%s' "
"with default provider at slot %d",
name, default_provider);
mju_free(resource->name);
mju_free(resource);
mju_closeResource(resource);
return NULL;
}
// lastly fallback to OS filesystem
else {
fillResource(NULL, resource);
resource->data = mju_malloc(sizeof(file_buffer));
file_buffer* fb = (file_buffer*) resource->data;
fb->buffer = mju_fileToMemory(name, &(fb->nbuffer));
if (fb->buffer == NULL) {
mju_warning("mju_openResource: unknown file '%s'", name);
mju_free(fb);
mju_free(resource->name);
mju_free(resource);
return NULL;
}
resource->provider = NULL;
resource->data = mju_malloc(sizeof(file_buffer));
file_buffer* fb = (file_buffer*) resource->data;
fb->buffer = mju_fileToMemory(name, &(fb->nbuffer));
if (fb->buffer == NULL) {
mju_warning("mju_openResource: unknown file '%s'", name);
mju_closeResource(resource);
return NULL;
}
return resource;
}
@@ -131,20 +88,20 @@ void mju_closeResource(mjResource* resource) {
return;
}
// use the resource provider to close resource
if (resource->close) {
resource->close(resource);
}
// if provider is NULL, then OS filesystem is used
else {
// use the resource provider close callback
if (resource->provider && resource->provider->close) {
resource->provider->close(resource);
} else {
// clear OS filesystem if present
file_buffer* fb = (file_buffer*) resource->data;
mju_free(fb->buffer);
mju_free(fb);
if (fb) {
if (fb->buffer) mju_free(fb->buffer);
mju_free(fb);
}
}
// free name and resource
mju_free(resource->name);
// free resource
if (resource->name) mju_free(resource->name);
mju_free(resource);
}
@@ -157,12 +114,12 @@ int mju_readResource(mjResource* resource, const void** buffer) {
return 0;
}
if (resource->read) {
return resource->read(resource, buffer);
if (resource->provider) {
return resource->provider->read(resource, buffer);
}
// if provider read callback is NULL, then OS filesystem is used
// if provider is NULL, then OS filesystem is used
const file_buffer* fb = (file_buffer*) resource->data;
*buffer = fb->buffer;
return fb->nbuffer;
@@ -175,14 +132,14 @@ void mju_getResourceDir(mjResource* resource, const char** dir, int* ndir) {
*dir = NULL;
*ndir = 0;
if (!resource) {
if (resource == NULL) {
return;
}
// provider is not OS filesystem
if (resource->read) {
if (resource->getdir) {
resource->getdir(resource, dir, ndir);
if (resource->provider) {
if (resource->provider->getdir) {
resource->provider->getdir(resource, dir, ndir);
}
} else {
*dir = resource->name;
@@ -211,7 +168,7 @@ int mju_dirnamelen(const char* path) {
// read file into memory buffer (allocated here with mju_malloc)
void* mju_fileToMemory(const char* filename, int* filesize) {
void* mju_fileToMemory(const char* filename, size_t* filesize) {
// open file
*filesize = 0;
FILE* fp = fopen(filename, "rb");
+9 -4
View File
@@ -15,6 +15,8 @@
#ifndef MUJOCO_SRC_ENGINE_ENGINE_RESOURCE_H_
#define MUJOCO_SRC_ENGINE_ENGINE_RESOURCE_H_
#include <stddef.h>
#include <mujoco/mjexport.h>
#include "engine/engine_plugin.h"
@@ -23,9 +25,8 @@ extern "C" {
#endif
// open the given resource; if the name doesn't have a prefix matching with a
// resource provider, then the default_provider is used
// if default_provider non-positive, then the OS filesystem is used
MJAPI mjResource* mju_openResource(const char* name, int default_provider);
// resource provider, then the OS filesystem is used
MJAPI mjResource* mju_openResource(const char* name);
// close the given resource; no-op if resource is NULL
MJAPI void mju_closeResource(mjResource* resource);
@@ -37,11 +38,15 @@ MJAPI int mju_readResource(mjResource* resource, const void** buffer);
// sets for a resource with a name partitioned as {dir}{filename}, the dir and ndir pointers
MJAPI void mju_getResourceDir(mjResource* resource, const char** dir, int* ndir);
// Returns > 0 if resource has been modified since last read, 0 if not, and < 0
// if inconclusive
MJAPI int mju_isModifiedResource(const mjResource* resource);
// get the length of the dirname portion of a given path
int mju_dirnamelen(const char* path);
// read file into memory buffer (allocated here with mju_malloc)
void* mju_fileToMemory(const char* filename, int* filesize);
void* mju_fileToMemory(const char* filename, size_t* filesize);
#ifdef __cplusplus
}
+49 -16
View File
@@ -14,10 +14,10 @@
#include "engine/engine_vfs.h"
#include <stddef.h>
#include <string.h>
#include "engine/engine_array_safety.h"
#include "engine/engine_plugin.h"
#include "engine/engine_resource.h"
#include "engine/engine_util_errmem.h"
#include "engine/engine_util_misc.h"
@@ -90,7 +90,7 @@ int mj_addFileVFS(mjVFS* vfs, const char* directory, const char* filename) {
mjSTRNCPY(vfs->filename[vfs->nfile], newname);
// allocate and read
int filesize = 0;
size_t filesize = 0;
vfs->filedata[vfs->nfile] = mju_fileToMemory(fullname, &filesize);
if (!vfs->filedata[vfs->nfile]) {
return -1;
@@ -211,11 +211,11 @@ void mj_deleteVFS(mjVFS* vfs) {
// open callback for the VFS resource provider
static int vfs_open_callback(mjResource* resource) {
if (!resource || !resource->provider_data || !resource->name) {
if (!resource || !resource->name || !resource->data) {
return 0;
}
const mjVFS* vfs = (const mjVFS*) resource->provider_data;
const mjVFS* vfs = (const mjVFS*) resource->data;
return mj_findFileVFS(vfs, resource->name) >= 0;
}
@@ -223,12 +223,12 @@ static int vfs_open_callback(mjResource* resource) {
// read callback for the VFS resource provider
static int vfs_read_callback(mjResource* resource, const void** buffer) {
if (!resource || !resource->provider_data) {
if (!resource || !resource->name || !resource->data) {
*buffer = NULL;
return -1;
}
const mjVFS* vfs = (const mjVFS*) resource->provider_data;
const mjVFS* vfs = (const mjVFS*) resource->data;
int i = mj_findFileVFS(vfs, resource->name);
if (i < 0) {
*buffer = NULL;
@@ -260,16 +260,49 @@ static void vfs_getdir_callback(mjResource* resource, const char** dir, int* ndi
// registers a VFS resource provider; returns the index of the provider
int mj_registerVfsProvider(const mjVFS* vfs) {
mjpResourceProvider provider = {
.prefix = mjVFS_PREFIX,
.open = &vfs_open_callback,
.read = &vfs_read_callback,
.close = &vfs_close_callback,
.getdir = &vfs_getdir_callback,
.data = (void*) vfs
// open VFS resource
mjResource* mju_openVfsResource(const char* name, const mjVFS* vfs) {
if (vfs == NULL) {
return NULL;
}
// VFS provider
static struct mjpResourceProvider provider = {
.prefix = NULL,
.data = NULL,
.open = &vfs_open_callback,
.read = &vfs_read_callback,
.close = &vfs_close_callback,
.getdir = &vfs_getdir_callback,
};
return mjp_registerResourceProviderInternal(&provider);
// create resource
mjResource* resource = (mjResource*) mju_malloc(sizeof(mjResource));
if (resource == NULL) {
mjERROR("could not allocate memory");
return NULL;
}
// clear out resource
memset(resource, 0, sizeof(mjResource));
// copy name
resource->name = mju_malloc(sizeof(char) * (strlen(name) + 1));
if (resource->name == NULL) {
mju_closeResource(resource);
mjERROR("could not allocate memory");
return NULL;
}
memcpy(resource->name, name, sizeof(char) * (strlen(name) + 1));
resource->data = (void*) vfs;
// open resource
resource->provider = &provider;
if (provider.open(resource)) {
return resource;
}
// not found in VFS
mju_closeResource(resource);
return NULL;
}
+3 -2
View File
@@ -17,6 +17,7 @@
#include <mujoco/mjexport.h>
#include <mujoco/mjmodel.h>
#include <mujoco/mjplugin.h>
#ifdef __cplusplus
extern "C" {
@@ -40,8 +41,8 @@ MJAPI int mj_deleteFileVFS(mjVFS* vfs, const char* filename);
// delete all files from VFS
MJAPI void mj_deleteVFS(mjVFS* vfs);
// registers a VFS resource provider; returns the index of the provider
MJAPI int mj_registerVfsProvider(const mjVFS* vfs);
// open VFS resource
MJAPI mjResource* mju_openVfsResource(const char* name, const mjVFS* vfs);
#ifdef __cplusplus
}