Add VFS handling logic to mj_openResource.

PiperOrigin-RevId: 653124909
Change-Id: I4d031eb0bfb0070f8493753d2aeaac60ac283b85
This commit is contained in:
Kyle Bayes
2024-07-17 00:40:20 -07:00
committed by Copybara-Service
parent 7f46e936f6
commit 2cf352c2e6
13 changed files with 244 additions and 247 deletions
+105 -132
View File
@@ -18,7 +18,6 @@
#include <sys/types.h>
#include <sys/stat.h>
#include <climits>
#include <cstddef>
#include <cstdint>
#include <cstdio>
@@ -35,12 +34,16 @@
#endif
#include <mujoco/mjplugin.h>
#include <mujoco/mujoco.h>
#include "engine/engine_plugin.h"
#include "engine/engine_util_errmem.h"
#include "engine/engine_util_misc.h"
#include "user/user_util.h"
#include "user/user_vfs.h"
namespace {
using mujoco::user::FileToMemory;
// file buffer used internally for the OS filesystem
struct FileSpec {
bool is_read; // set to nonzero if buffer was read into
@@ -48,11 +51,74 @@ struct FileSpec {
time_t mtime; // last modified time
};
// callback for opening a resource, returns zero on failure
int FileOpen(mjResource* resource, char* error, size_t nerror) {
resource->provider = nullptr;
resource->data = new FileSpec;
const char* name = resource->name;
FileSpec* spec = (FileSpec*) resource->data;
spec->is_read = false;
struct stat file_stat;
if (stat(name, &file_stat) == 0) {
memcpy(&spec->mtime, &file_stat.st_mtime, sizeof(time_t));
} else {
if (error) {
snprintf(error, nerror, "Error opening file '%s': %s", name,
strerror(errno));
}
return 0;
}
mju_encodeBase64(resource->timestamp, (uint8_t*) &spec->mtime,
sizeof(time_t));
return 1;
}
// OS filesystem read callback
int FileRead(mjResource* resource, const void** buffer) {
FileSpec* spec = (FileSpec*) resource->data;
// only read once from file
if (!spec->is_read) {
spec->buffer = FileToMemory(resource->name);
spec->is_read = true;
}
*buffer = spec->buffer.data();
return spec->buffer.size();
}
// OS filesystem close callback
void FileClose(mjResource* resource) {
FileSpec* spec = (FileSpec*) resource->data;
if (spec) delete spec;
}
// OS filesystem getdir callback
void FileGetDir(mjResource* resource, const char** dir, int* ndir) {
*dir = resource->name;
*ndir = mjuu_dirnamelen(resource->name);
}
// OS filesystem modified callback
int FileModified(const mjResource* resource, const char*timestamp) {
if (mju_isValidBase64(timestamp) != sizeof(time_t)) {
return 1; // error (assume modified)
}
time_t time1, time2;
mju_decodeBase64((uint8_t*) &time1, timestamp);
time2 = ((FileSpec*) resource->data)->mtime;
double diff = difftime(time2, time1);
if (diff < 0) return -1;
if (diff > 0) return 1;
return 0;
}
} // namespace
// open the given resource; if the name doesn't have a prefix matching with a
// resource provider, then the OS filesystem is used
mjResource* mju_openResource(const char* name, char* error, size_t nerror) {
mjResource* mju_openResource(const char* name, const mjVFS* vfs,
char* error, size_t nerror) {
// no error so far
if (error) {
error[0] = '\0';
@@ -60,7 +126,10 @@ mjResource* mju_openResource(const char* name, char* error, size_t nerror) {
mjResource* resource = (mjResource*) mju_malloc(sizeof(mjResource));
if (resource == nullptr) {
mjERROR("could not allocate memory");
if (error) {
strncpy(error, "could not allocate memory", nerror);
error[nerror - 1] = '\0';
}
return nullptr;
}
@@ -70,16 +139,30 @@ mjResource* mju_openResource(const char* name, char* error, size_t nerror) {
// copy name
resource->name = (char*) mju_malloc(sizeof(char) * (strlen(name) + 1));
if (resource->name == nullptr) {
if (error) {
strncpy(error, "could not allocate memory", nerror);
error[nerror - 1] = '\0';
}
mju_closeResource(resource);
mjERROR("could not allocate memory");
return nullptr;
}
memcpy(resource->name, name, sizeof(char) * (strlen(name) + 1));
// first priority is to check the VFS
if (vfs != nullptr) {
const mjpResourceProvider* provider = GetVfsResourceProvider();
resource->data = (void*) vfs;
resource->provider = provider;
if (provider->open(resource)) {
return resource;
}
}
// find provider based off prefix of name
const mjpResourceProvider* provider = mjp_getResourceProvider(name);
if (provider != nullptr) {
resource->provider = provider;
resource->data = nullptr;
if (provider->open(resource)) {
return resource;
}
@@ -95,24 +178,11 @@ mjResource* mju_openResource(const char* name, char* error, size_t nerror) {
}
// lastly fallback to OS filesystem
resource->provider = nullptr;
resource->data = new FileSpec;
FileSpec* spec = (FileSpec*) resource->data;
spec->is_read = false;
struct stat file_stat;
if (stat(name, &file_stat) == 0) {
memcpy(&spec->mtime, &file_stat.st_mtime, sizeof(time_t));
} else {
if (error) {
snprintf(error, nerror, "Error opening file '%s': %s", name,
strerror(errno));
}
mju_closeResource(resource);
return nullptr;
if (FileOpen(resource, error, nerror)) {
return resource;
}
mju_encodeBase64(resource->timestamp, (uint8_t*) &spec->mtime,
sizeof(time_t));
return resource;
mju_closeResource(resource);
return nullptr;
}
@@ -124,14 +194,11 @@ void mju_closeResource(mjResource* resource) {
}
// use the resource provider close callback
if (resource->provider && resource->provider->close) {
resource->provider->close(resource);
const mjpResourceProvider* provider = resource->provider;
if (provider) {
if (provider->close) provider->close(resource);
} else {
// clear OS filesystem if present
FileSpec* spec = (FileSpec*) resource->data;
if (spec) {
delete spec;
}
FileClose(resource); // clear OS filesystem if present
}
// free resource
@@ -149,36 +216,26 @@ int mju_readResource(mjResource* resource, const void** buffer) {
}
// if provider is NULL, then OS filesystem is used
FileSpec* spec = (FileSpec*) resource->data;
// only read once from file
if (!spec->is_read) {
spec->buffer = mju_fileToMemory(resource->name);
spec->is_read = true;
}
*buffer = spec->buffer.data();
return spec->buffer.size();
return FileRead(resource, buffer);
}
// get directory path of resource
void mju_getResourceDir(mjResource* resource, const char** dir, int* ndir) {
*dir = NULL;
*dir = nullptr;
*ndir = 0;
if (resource == NULL) {
if (resource == nullptr) {
return;
}
// provider is not OS filesystem
if (resource->provider) {
if (resource->provider->getdir) {
resource->provider->getdir(resource, dir, ndir);
}
const mjpResourceProvider* provider = resource->provider;
if (provider) {
if (provider->getdir) provider->getdir(resource, dir, ndir);
} else {
*dir = resource->name;
*ndir = mju_dirnamelen(resource->name);
// fallback to OS filesystem
FileGetDir(resource, dir, ndir);
}
}
@@ -197,89 +254,5 @@ int mju_isModifiedResource(const mjResource* resource, const char* timestamp) {
}
// fallback to OS filesystem
if (mju_isValidBase64(timestamp) != sizeof(time_t)) {
return 1; // error (assume modified)
}
time_t time1, time2;
mju_decodeBase64((uint8_t*) &time1, timestamp);
time2 = ((FileSpec*) resource->data)->mtime;
double diff = difftime(time2, time1);
if (diff < 0) return -1;
if (diff > 0) return 1;
return 0;
}
// get the length of the dirname portion of a given path
int mju_dirnamelen(const char* path) {
if (!path) {
return 0;
}
int pos = -1;
for (int i = 0; path[i]; ++i) {
if (path[i] == '/' || path[i] == '\\') {
pos = i;
}
}
return pos + 1;
}
// read file into memory buffer (allocated here with mju_malloc)
std::vector<uint8_t> mju_fileToMemory(const char* filename) {
FILE* fp = fopen(filename, "rb");
if (!fp) {
return {};
}
// find size
if (fseek(fp, 0, SEEK_END) != 0) {
fclose(fp);
mju_warning("Failed to calculate size for '%s'", filename);
return {};
}
// ensure file size fits in int
long long_filesize = ftell(fp); // NOLINT(runtime/int)
if (long_filesize > INT_MAX) {
fclose(fp);
mju_warning("File size over 2GB is not supported. File: '%s'", filename);
return {};
} else if (long_filesize < 0) {
fclose(fp);
mju_warning("Failed to calculate size for '%s'", filename);
return {};
}
std::vector<uint8_t> buffer(long_filesize);
// go back to start of file
if (fseek(fp, 0, SEEK_SET) != 0) {
fclose(fp);
mju_warning("Read error while reading '%s'", filename);
return {};
}
// allocate and read
std::size_t bytes_read = fread(buffer.data(), 1, buffer.size(), fp);
// check that read data matches file size
if (bytes_read != buffer.size()) { // SHOULD NOT OCCUR
if (ferror(fp)) {
fclose(fp);
mju_warning("Read error while reading '%s'", filename);
return {};
} else if (feof(fp)) {
buffer.resize(bytes_read);
}
}
// close file, return contents
fclose(fp);
return buffer;
return FileModified(resource, timestamp);
}