From a1ddbdf7f8f0dd8b2c1babac6e508b94840e0a58 Mon Sep 17 00:00:00 2001 From: Kyle Bayes Date: Fri, 19 Jan 2024 09:02:37 -0800 Subject: [PATCH] Add an implementation of a cache for assets for compilation speedups. PiperOrigin-RevId: 599850003 Change-Id: I83e299e127058a5657d7b262aa98d36451b530f6 --- src/user/user_asset_cache.cc | 327 +++++++++++++++++++++++++++ src/user/user_asset_cache.h | 195 +++++++++++++++++ test/user/user_asset_cache_test.cc | 341 +++++++++++++++++++++++++++++ 3 files changed, 863 insertions(+) create mode 100644 src/user/user_asset_cache.cc create mode 100644 src/user/user_asset_cache.h create mode 100644 test/user/user_asset_cache_test.cc diff --git a/src/user/user_asset_cache.cc b/src/user/user_asset_cache.cc new file mode 100644 index 00000000..4724d748 --- /dev/null +++ b/src/user/user_asset_cache.cc @@ -0,0 +1,327 @@ +// 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. + +#include "user/user_asset_cache.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +// adds a block of data for the asset and returns number of bytes stored +template +std::size_t mjCAsset::Add(const std::string& name, const std::vector& v) { + auto [it, inserted] = blocks_.insert({name, mjCAssetData()}); + if (!inserted) { + return 0; + } + + std::size_t n = v.size() * sizeof(T); + const uint8_t* ptr = reinterpret_cast(v.data()); + mjCAssetData& block = it->second; + + block.bytes = std::make_shared(n); + std::copy(ptr, ptr + n, block.bytes.get()); + block.nbytes = n; + nbytes_ += n; + return n; +} + +template std::size_t mjCAsset::Add(const std::string& name, + const std::vector& v); +template std::size_t mjCAsset::Add(const std::string& name, + const std::vector& v); +template std::size_t mjCAsset::Add(const std::string& name, + const std::vector& v); + + + +// fetches a block of data, sets n to size of data +template +const T* mjCAsset::Get(const std::string& name, std::size_t* n) const { + auto it = blocks_.find(name); + if (it == blocks_.end()) { + *n = 0; + return nullptr; + } + + const mjCAssetData& data = it->second; + + // TODO(kylebayes): This is probably caused by a user bug and needs an + // assertion for debugging purposes + if (data.nbytes % sizeof(T)) { + *n = 0; + return nullptr; + } + + *n = data.nbytes / sizeof(T); + return reinterpret_cast(data.bytes.get()); +} + +template +const int* mjCAsset::Get(const std::string& name, std::size_t* n) const; +template +const float* mjCAsset::Get(const std::string& name, std::size_t* n) const; +template +const double* mjCAsset::Get(const std::string& name, std::size_t* n) const; + + + +// replaces blocks data in asset +void mjCAsset::ReplaceBlocks( + const std::unordered_map& blocks, + std::size_t nbytes) { + blocks_ = blocks; + nbytes_ = nbytes; +} + + + +// makes a copy for user (strip unnecessary items) +mjCAsset mjCAsset::Copy(const mjCAsset& other) { + mjCAsset asset; + asset.id_ = other.Id(); + asset.timestamp_ = other.Timestamp(); + asset.blocks_ = other.blocks_; + asset.nbytes_ = other.nbytes_; + return asset; +} + + + +// sets the total maximum size of the cache in bytes +// low-priority cached assets will be dropped to make the new memory +// requirement +void mjCCache::SetMaxSize(std::size_t size) { + std::lock_guard lock(mutex_); + max_size_ = size; + Trim(); +} + + + +// returns the corresponding timestamp, if the given asset is stored in the cache +const std::string* mjCCache::HasAsset(const std::string& id) { + std::lock_guard lock(mutex_); + auto it = lookup_.find(id); + if (it == lookup_.end()) { + return nullptr; + } + + return &(it->second.Timestamp()); +} + + + +// inserts an asset into the cache, if asset is already in the cache, its data +// is updated only if the timestamps disagree +bool mjCCache::Insert(const mjCAsset& asset) { + std::lock_guard lock(mutex_); + const std::string& id = asset.Id(); + if (asset.References().size() != 1) { + return false; + } + const std::string& filename = *(asset.References().begin()); + auto [it, inserted] = lookup_.insert({id, asset}); + + if (!inserted) { + mjCAsset* asset_ptr = &(it->second); + if (size_ - asset_ptr->BytesCount() + asset.BytesCount() > max_size_) { + return false; + } + models_[filename].insert(asset_ptr); // add it for the model + asset_ptr->AddReference(filename); + if (it->second.Timestamp() == asset.Timestamp()) { + return true; + } + asset_ptr->SetTimestamp(asset.Timestamp()); + size_ = size_ - asset_ptr->BytesCount() + asset.BytesCount(); + asset_ptr->ReplaceBlocks(asset.Blocks(), asset.BytesCount()); + return true; + } else if (size_ + asset.BytesCount() > max_size_) { + return false; + } + + // new asset + mjCAsset* asset_ptr = &(it->second); + asset_ptr->SetInsertNum(insert_num_++); + entries_.insert(asset_ptr); + models_[filename].insert(asset_ptr); + size_ += asset.BytesCount(); + return true; +} + + + +bool mjCCache::Insert(mjCAsset&& asset) { + std::lock_guard lock(mutex_); + const std::string& id = asset.Id(); + if (asset.References().size() != 1) { + return false; + } + const std::string& filename = *(asset.References().begin()); + std::size_t nbytes = asset.BytesCount(); + auto [it, inserted] = lookup_.try_emplace(id, std::move(asset)); + + if (!inserted) { + mjCAsset* asset_ptr = &(it->second); + if (size_ - asset_ptr->BytesCount() + nbytes > max_size_) { + return false; + } + models_[filename].insert(asset_ptr); // add it for the model + asset_ptr->AddReference(std::move(filename)); + if (it->second.Timestamp() == asset.Timestamp()) { + return true; + } + // move data and timestamp over + asset_ptr->SetTimestamp(std::move(asset.timestamp_)); + size_ = size_ - asset_ptr->BytesCount() + nbytes; + asset_ptr->ReplaceBlocks(std::move(asset.blocks_), asset.nbytes_); + return true; + } else if (size_ + nbytes > max_size_) { + return false; + } + + // new asset + mjCAsset* asset_ptr = &(it->second); + asset_ptr->SetInsertNum(insert_num_++); + entries_.insert(asset_ptr); + models_[filename].insert(asset_ptr); + size_ += nbytes; + return true; +} + + + +// returns the asset with the given id, if it exists in the cache +std::optional mjCCache::Get(const std::string& id) { + std::lock_guard lock(mutex_); + auto it = lookup_.find(id); + if (it == lookup_.end()) { + return std::nullopt; + } + + mjCAsset* asset = &(it->second); + asset->IncrementAccess(); + + // update priority queue + entries_.erase(asset); + entries_.insert(asset); + return asset->Copy(*asset); +} + + + +// removes model from the cache along with assets referencing only this model +void mjCCache::RemoveModel(const std::string& filename) { + std::lock_guard lock(mutex_); + for (mjCAsset* asset : models_[filename]) { + asset->RemoveReference(filename); + if (!asset->HasReferences()) { + Delete(asset, filename); + } + } + models_.erase(filename); +} + + + +// Wipes out all internal data for the given model +void mjCCache::Reset(const std::string& filename) { + std::lock_guard lock(mutex_); + for (auto asset : models_[filename]) { + Delete(asset, filename); + } + models_.erase(filename); +} + + + +// Wipes out all internal data +void mjCCache::Reset() { + std::lock_guard lock(mutex_); + entries_.clear(); + lookup_.clear(); + models_.clear(); + size_ = 0; + insert_num_ = 0; +} + + + +std::size_t mjCCache::MaxSize() const { + std::lock_guard lock(mutex_); + return max_size_; +} + + + +std::size_t mjCCache::Size() const { + std::lock_guard lock(mutex_); + return size_; +} + + + +// Deletes a single asset +void mjCCache::DeleteAsset(const std::string& id) { + std::lock_guard lock(mutex_); + auto it = lookup_.find(id); + if (it != lookup_.end()) { + Delete(&(it->second)); + } +} + + + +// Deletes a single asset (internal) +void mjCCache::Delete(mjCAsset* asset) { + size_ -= asset->BytesCount(); + entries_.erase(asset); + for (auto& reference : asset->References()) { + models_[reference].erase(asset); + } + lookup_.erase(asset->Id()); +} + + + +// Deletes a single asset (internal) +void mjCCache::Delete(mjCAsset* asset, const std::string& skip) { + size_ -= asset->BytesCount(); + entries_.erase(asset); + + for (auto& reference : asset->References()) { + if (reference != skip) { + models_[reference].erase(asset); + } + } + lookup_.erase(asset->Id()); +} + + + +// trims out data to meet memory requirements +void mjCCache::Trim() { + while (size_ > max_size_) { + Delete(*entries_.begin()); + } +} diff --git a/src/user/user_asset_cache.h b/src/user/user_asset_cache.h new file mode 100644 index 00000000..e8f0bc7c --- /dev/null +++ b/src/user/user_asset_cache.h @@ -0,0 +1,195 @@ +// 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_USER_ASSET_CACHE_H_ +#define MUJOCO_SRC_USER_ASSET_CACHE_H_ + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +// data associated with an asset +struct mjCAssetData { + std::shared_ptr bytes; // raw serialized bytes of cached data + std::size_t nbytes; // number of bytes stored +}; + +// A class container for a thread-safe asset cache +// +// Each mjCAsset is used to store raw and/or processed data loaded from a +// resource and is defined by a unique ID (usually the full filename of the +// asset). The asset's data can be segregated into blocks for ease of use of +// mix and matching different types of data. Each block is given a unique name +// within the asset for readability. For example, mjCAsset for a mesh may +// include all vertex positions, edges, and the computed volume. +class mjCAsset { + friend class mjCCache; + public: + mjCAsset(std::string filename, std::string id, std::string timestamp) + : id_(std::move(id)), timestamp_(std::move(timestamp)) { + AddReference(filename); + } + + // move and copy constructors + mjCAsset(mjCAsset&& other) = default; + mjCAsset& operator=(mjCAsset&& other) = default; + mjCAsset(const mjCAsset& other) = default; + mjCAsset& operator=(const mjCAsset& other) = default; + + // adds a block of data for the asset and returns number of bytes stored + // loading data into an asset should happen single thread + template std::size_t Add(const std::string& name, + const std::vector& v); + + // fetches a block of data, sets n to size of data + // TODO(kylebayes): The C++ span utility doesn't seem to be supported by + // Google C++ coding standards. For now, we fallback to an C style API. + template + const T* Get(const std::string& name, std::size_t* n) const; + + private: + mjCAsset() = default; + + // helpers for managing models referencing this asset + void AddReference(std::string xml_file) { references_.insert(xml_file); } + void RemoveReference(const std::string& xml_file) { + references_.erase(xml_file); + } + bool HasReferences() const { return !references_.empty(); } + + // replaces data blocks in asset + void ReplaceBlocks( + const std::unordered_map& blocks, + std::size_t nbytes); + + void IncrementAccess() { access_count_++; } + + // makes a copy for user (strip unnecessary references) + static mjCAsset Copy(const mjCAsset& other); + + // setters + void SetInsertNum(std::size_t num) { insert_num_ = num; } + void SetTimestamp(std::string timestamp) { timestamp_ = timestamp; } + + // accessors + const std::string& Id() const { return id_; } + const std::string& Timestamp() const { return timestamp_; } + std::size_t InsertNum() const { return insert_num_; } + std::size_t AccessCount() const { return access_count_; } + std::size_t BytesCount() const { return nbytes_; } + const std::unordered_map& Blocks() const { + return blocks_; + } + const std::set& References() const { return references_; } + + std::string id_; // unique id associated with asset + std::string timestamp_; // opaque timestamp of asset + std::size_t insert_num_; // number when asset was inserted + std::size_t access_count_ = 0; // incremented when getting 0th block + std::size_t nbytes_ = 0; // how many bytes taken up by the asset + + // the actually data of the asset + std::unordered_map blocks_; + + // list of models referencing this asset + std::set references_; +}; + +// the class container for a thread-safe asset cache +class mjCCache { + public: + explicit mjCCache(std::size_t size) : max_size_(size) {} + + // move only + mjCCache(mjCCache&& other) = default; + mjCCache& operator=(mjCCache&& other) = default; + mjCCache(const mjCCache& other) = delete; + mjCCache& operator=(const mjCCache& other) = delete; + + // sets the total maximum size of the cache in bytes + // low-priority cached assets will be dropped to make the new memory + // requirement + void SetMaxSize(std::size_t size); + + // returns the corresponding timestamp, if the given asset is stored in + // the cache + const std::string* HasAsset(const std::string& id); + + // inserts an asset into the cache, if asset is already in the cache, its data + // is updated only if the timestamps disagree + bool Insert(const mjCAsset& asset); + bool Insert(mjCAsset&& asset); + + // returns the asset with the given id, if it exists in the cache + std::optional Get(const std::string& id); + + // deletes the asset from the cache with the given id + void DeleteAsset(const std::string& id); + + // removes model from the cache, assets only referenced by the model will be + // deleted + void RemoveModel(const std::string& filename); + + // Wipes out all assets from the cache for the given model + void Reset(const std::string& filename); + + // Wipes out all internal data + void Reset(); + + // accessors + std::size_t MaxSize() const; + std::size_t Size() const; + + private: + void Delete(mjCAsset* asset); + void Delete(mjCAsset* asset, const std::string& skip); + void Trim(); + + // TODO(kylebayes): We should consider a shared mutex like in + // engine/engine_plugin.cc as some of these methods don't need to be fully + // locked. + mutable std::mutex mutex_; + std::size_t insert_num_ = 0; // a running counter of assets being inserted + std::size_t size_ = 0; // current size of the cache in bytes + std::size_t max_size_ = 0; // max size of the cache in bytes + + // compare function for the priority queue + static constexpr auto compare_ = [](const mjCAsset* e1, + const mjCAsset* e2) { + if (e1->AccessCount() != e2->AccessCount()) { + return e1->AccessCount() < e2->AccessCount(); + } + return e1->InsertNum() < e2->InsertNum(); + }; + + // internal constant look up table for assets + std::unordered_map lookup_; + + // internal priority queue for the cache + std::set entries_; + + // models using the cache along with the assets they reference + std::unordered_map> models_; +}; + +#endif // MUJOCO_SRC_USER_ASSET_CACHE_H_ diff --git a/test/user/user_asset_cache_test.cc b/test/user/user_asset_cache_test.cc new file mode 100644 index 00000000..577380de --- /dev/null +++ b/test/user/user_asset_cache_test.cc @@ -0,0 +1,341 @@ +// 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. + +// Tests for user/user_asset_cache.cc + +#include +#include +#include +#include + +#include +#include +#include "test/fixture.h" +#include "src/user/user_asset_cache.h" + +namespace mujoco { + +using ::testing::ElementsAreArray; +using ::testing::IsNull; +using ::testing::NotNull; +using ::testing::StrEq; + +using AssetCacheTest = MujocoTest; + +namespace { + +constexpr int kMaxSize = 100; // in bytes + +TEST(AssetCacheTest, SizeTest) { + mjCCache cache(kMaxSize); + EXPECT_EQ(cache.Size(), 0); +} + +TEST(AssetCacheTest, HasAssetSuccessTest) { + mjCCache cache(kMaxSize); + + mjCAsset asset("file.xml", "foo.obj", "now"); + cache.Insert(asset); + + EXPECT_THAT(*(cache.HasAsset("foo.obj")), StrEq("now")); +} + +TEST(AssetCacheTest, HasAssetFailureTest) { + mjCCache cache(kMaxSize); + + mjCAsset asset("file.xml", "foo.obj", "now"); + cache.Insert(asset); + + EXPECT_THAT(cache.HasAsset("file2.xml"), nullptr); +} + +TEST(AssetCacheTest, AddSuccessTest) { + mjCCache cache(kMaxSize); + std::vector v1 = {1, 2, 3}; + std::vector v2 = {1.0, 2.0, 3.0}; + + mjCAsset asset("file.xml", "foo.obj", "now"); + std::size_t nbytes1 = asset.Add("v1", v1); + std::size_t nbytes2 = asset.Add("v2", v2); + cache.Insert(asset); + + ASSERT_EQ(nbytes1, 12); + ASSERT_EQ(nbytes2, 24); + ASSERT_EQ(cache.Size(), 36); +} + +TEST(AssetCacheTest, AddFailureTest) { + mjCCache cache(kMaxSize); + std::vector v1 = {1, 2, 3}; + std::vector v2 = {1.0, 2.0, 3.0}; + + mjCAsset asset("file.xml", "foo.obj", "now"); + asset.Add("v1", v1); + std::size_t nbytes = asset.Add("v1", v2); + cache.Insert(asset); + + ASSERT_EQ(nbytes, 0); + ASSERT_EQ(cache.Size(), 12); +} + +TEST(AssetCacheTest, InsertReplaceTest) { + mjCCache cache(kMaxSize); + std::vector v1 = {1, 2, 3}; + std::vector v2 = {1.0, 2.0, 3.0}; + + mjCAsset asset("file.xml", "foo.obj", "now"); + asset.Add("v", v1); + cache.Insert(asset); + + mjCAsset asset2("file.xml", "foo.obj", "nower"); + asset2.Add("v", v2); + bool inserted = cache.Insert(asset2); + EXPECT_TRUE(inserted); + + mjCAsset asset3 = *(cache.Get("foo.obj")); + std::size_t n = 0; + const double* ptr = asset3.Get("v", &n); + std::vector v3 = std::vector(ptr, ptr + n); + EXPECT_THAT(v3, ElementsAreArray(v2)); + + ASSERT_EQ(cache.Size(), 24); +} + +TEST(AssetCacheTest, MoveInsertNewTest) { + mjCCache cache(kMaxSize); + std::vector v1 = {1, 2, 3}; + std::vector v2 = {1.0, 2.0, 3.0}; + + mjCAsset asset("file.xml", "foo.obj", "now"); + std::size_t nbytes1 = asset.Add("v1", v1); + std::size_t nbytes2 = asset.Add("v2", v2); + cache.Insert(std::move(asset)); + + ASSERT_EQ(nbytes1, 12); + ASSERT_EQ(nbytes2, 24); + ASSERT_EQ(cache.Size(), 36); +} + +TEST(AssetCacheTest, MoveInsertReplaceTest) { + mjCCache cache(kMaxSize); + std::vector v1 = {1, 2, 3}; + std::vector v2 = {1.0, 2.0, 3.0}; + + mjCAsset asset("file.xml", "foo.obj", "now"); + asset.Add("v", v1); + cache.Insert(std::move(asset)); + + mjCAsset asset2("file.xml", "foo.obj", "nower"); + asset2.Add("v", v2); + bool inserted = cache.Insert(std::move(asset2)); + EXPECT_TRUE(inserted); + + mjCAsset asset3 = *(cache.Get("foo.obj")); + std::size_t n = 0; + const double* ptr = asset3.Get("v", &n); + std::vector v3 = std::vector(ptr, ptr + n); + EXPECT_THAT(v3, ElementsAreArray(v2)); + + ASSERT_EQ(cache.Size(), 24); +} + +TEST(AssetCacheTest, GetSuccessTest) { + mjCCache cache(kMaxSize); + std::vector v1 = {1, 2, 3}; + std::vector v2 = {1.0, 2.0, 3.0}; + mjCAsset asset("file.xml", "foo.obj", "now"); + asset.Add("v1", v1); + asset.Add("v2", v2); + cache.Insert(asset); + + std::size_t n = 0; + + mjCAsset asset2 = *(cache.Get("foo.obj")); + + const int* ptr1 = asset2.Get("v1", &n); + EXPECT_EQ(n, 3); + std::vector v3 = std::vector(ptr1, ptr1 + n); + EXPECT_THAT(v3, ElementsAreArray(v1)); + + const double* ptr2 = asset2.Get("v2", &n); + EXPECT_EQ(n, 3); + std::vector v4 = std::vector(ptr2, ptr2 + n); + EXPECT_THAT(v4, ElementsAreArray(v3)); +} + +TEST(AssetCacheTest, GetFailueTest) { + mjCCache cache(kMaxSize); + std::vector v = {1, 2, 3}; + mjCAsset asset("file.xml", "foo.obj", "now"); + asset.Add("v", v); + cache.Insert(asset); + + mjCAsset asset2 = *(cache.Get("foo.obj")); + EXPECT_EQ(cache.Get("bar.obj").has_value(), false); + + std::size_t n = 0; + const int* ptr = asset.Get("v2", &n); + EXPECT_THAT(ptr, IsNull()); +} + +// Trim cache based off of access count +TEST(AssetCacheTest, LimitTest1) { + mjCCache cache(kMaxSize); + EXPECT_THAT(cache.MaxSize(), kMaxSize); + std::vector v = {1, 2, 3}; + + mjCAsset asset1 = mjCAsset("file.xml", "foo.obj", "now"); + mjCAsset asset2 = mjCAsset("file.xml", "bar.obj", "now"); + asset1.Add("foo.obj", v); + asset2.Add("bar.obj", v); + cache.Insert(asset1); + cache.Insert(asset2); + + // access asset foo twice, bar one + cache.Get("foo.obj"); + cache.Get("foo.obj"); + cache.Get("bar.obj"); + + // make max size so cache can hold only one asset + cache.SetMaxSize(12); + + // foo should still be in cache + EXPECT_THAT(cache.HasAsset("foo.obj"), NotNull()); + + // bar was accessed less, so is removed + EXPECT_THAT(cache.HasAsset("bar.obj"), IsNull()); +} + +// Trim cache based off of insert order +TEST(AssetCacheTest, LimitTest2) { + mjCCache cache(kMaxSize); + std::vector v = {1, 2, 3}; + mjCAsset asset1("file.xml", "foo.obj", "now"); + mjCAsset asset2("file.xml", "bar.obj", "now"); + asset1.Add("v", v); + asset2.Add("v", v); + + cache.Insert(asset1); + cache.Insert(asset2); + + // get each asset once + mjCAsset asset3 = *(cache.Get("foo.obj")); + mjCAsset asset4 = *(cache.Get("bar.obj")); + + // make max size so cache can hold only one asset + cache.SetMaxSize(12); + + // foo should be gone because it's older + EXPECT_THAT(cache.HasAsset("foo.obj"), IsNull()); + + // bar should still be in cache + EXPECT_THAT(cache.HasAsset("bar.obj"), NotNull()); +} + +TEST(AssetCacheTest, ResetAllTest) { + mjCCache cache(kMaxSize); + mjCAsset asset1("file1.xml", "foo.obj", "now"); + mjCAsset asset2("file2.xml", "bar.obj", "now"); + cache.Insert(asset1); + cache.Insert(asset2); + + EXPECT_THAT(cache.HasAsset("foo.obj"), NotNull()); + EXPECT_THAT(cache.HasAsset("bar.obj"), NotNull()); + + cache.Reset(); + + EXPECT_THAT(cache.HasAsset("foo.obj"), IsNull()); + EXPECT_THAT(cache.HasAsset("bar.obj"), IsNull()); +} + +TEST(AssetCacheTest, ResetModelTest1) { + mjCCache cache(kMaxSize); + mjCAsset asset("file1.xml", "foo.obj", "now"); + mjCAsset asset2("file1.xml", "bar.obj", "now"); + mjCAsset asset3("file2.xml", "bar.obj", "now"); + cache.Insert(asset); + cache.Insert(asset2); + cache.Insert(asset3); + + cache.Reset("file2.xml"); + + EXPECT_THAT(cache.HasAsset("foo.obj"), NotNull()); + EXPECT_THAT(cache.HasAsset("bar.obj"), IsNull()); +} + +TEST(AssetCacheTest, ResetModelTest2) { + mjCCache cache(kMaxSize); + mjCAsset asset("file1.xml", "foo.obj", "now"); + mjCAsset asset2("file2.xml", "foo.obj", "now"); + mjCAsset asset3("file2.xml", "bar.obj", "now"); + cache.Insert(asset); + cache.Insert(asset2); + cache.Insert(asset3); + + cache.Reset("file2.xml"); + + EXPECT_THAT(cache.HasAsset("foo.obj"), IsNull()); + EXPECT_THAT(cache.HasAsset("bar.obj"), IsNull()); +} + +TEST(AssetCacheTest, RemoveModelTest1) { + mjCCache cache(kMaxSize); + mjCAsset asset("file1.xml", "foo.obj", "now"); + mjCAsset asset2("file1.xml", "bar.obj", "now"); + mjCAsset asset3("file2.xml", "bar.obj", "now"); + cache.Insert(asset); + cache.Insert(asset2); + cache.Insert(asset3); + + cache.RemoveModel("file2.xml"); + + EXPECT_THAT(cache.HasAsset("foo.obj"), NotNull()); + EXPECT_THAT(cache.HasAsset("bar.obj"), NotNull()); +} + +TEST(AssetCacheTest, RemoveModelTest2) { + mjCCache cache(kMaxSize); + mjCAsset asset("file1.xml", "foo.obj", "now"); + mjCAsset asset2("file2.xml", "bar.obj", "now"); + cache.Insert(asset); + cache.Insert(asset2); + + cache.Reset("file2.xml"); + + EXPECT_THAT(cache.HasAsset("foo.obj"), NotNull()); + EXPECT_THAT(cache.HasAsset("bar.obj"), IsNull()); +} + +TEST(AssetCacheTest, DeleteAssetSuccessTest) { + mjCCache cache(kMaxSize); + mjCAsset asset("file1.xml", "foo.obj", "now"); + cache.Insert(asset); + + cache.DeleteAsset("foo.obj"); + + EXPECT_THAT(cache.HasAsset("foo.obj"), IsNull()); +} + +TEST(AssetCacheTest, DeleteAssetFailureTest) { + mjCCache cache(kMaxSize); + mjCAsset asset("file.xml", "foo.obj", "now"); + cache.Insert(asset); + + cache.DeleteAsset("bar.obj"); + + EXPECT_THAT(cache.HasAsset("foo.obj"), NotNull()); +} + +} // namespace +} // namespace mujoco