diff --git a/src/experimental/platform/CMakeLists.txt b/src/experimental/platform/CMakeLists.txt index 6429db47..2cdc5f4e 100644 --- a/src/experimental/platform/CMakeLists.txt +++ b/src/experimental/platform/CMakeLists.txt @@ -50,6 +50,7 @@ target_sources(${MUJOCO_PLATFORM_TARGET_NAME} interaction.h model_holder.cc model_holder.h + object_launcher_plugin.cc picture_gui.h picture_gui.cc plugin.cc diff --git a/src/experimental/platform/object_launcher_plugin.cc b/src/experimental/platform/object_launcher_plugin.cc new file mode 100644 index 00000000..0ba23894 --- /dev/null +++ b/src/experimental/platform/object_launcher_plugin.cc @@ -0,0 +1,214 @@ +// Copyright 2026 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 +#include +#include + +#include +#include +#include "experimental/platform/imgui_widgets.h" +#include "experimental/platform/plugin.h" + +namespace mujoco::studio { + +class ObjectLauncher { + public: + ObjectLauncher() : rng_(std::random_device{}()) {} + + void UpdateGui() { + using platform::ImGui_Input; + + ImGui::Checkbox("Enable Key Binding (Ctrl+Shift+Enter)", &enabled_); + ImGui_Input("Size", &size_, {0.01f, 1.0f, 0.01, 0.1}); + ImGui_Input("Speed", &speed_, {0.01f, 100.0f, 0.1, 1.0}); + ImGui_Input("Mass", &mass_, {0.01f, 100.0f, 0.01, 0.1}); + ImGui_Input("Life", &lifetime_, {0.0f, 60.0f, 0.1, 1.0}); + + int shape = type_ == mjGEOM_BOX ? 0 : 1; + const char* names[] = {"Box", "Sphere"}; + ImGui::Combo("Shape", &shape, names, 2); + type_ = shape == 0 ? mjGEOM_BOX : mjGEOM_SPHERE; + + if (ImGui::Button("Launch", ImVec2(-1.0f, 0.0f))) { + active_ = true; + } + + if (ImGui::Button("Clear")) { + for (auto& object : objects_) { + object.expiration = -1; + } + } + } + + void HandleKeyboardEvent() { if (enabled_) active_ = true; } + + bool UpdateSpecPreCompile(mjSpec* spec, const mjModel* model, + const mjData* data, const mjvCamera* camera) { + // Remove expired objects. + auto it = std::remove_if( + objects_.begin(), objects_.end(), [&](const ObjectInfo& o) { + const bool expired = + o.body_id >= 0 && o.expiration != 0 && o.expiration < data->time; + if (expired) { + mjsBody* body = mjs_findBody(spec, o.name.c_str()); + if (body) { + mjs_delete(spec, body->element); + } + } + return expired; + }); + if (it != objects_.end()) { + objects_.erase(it, objects_.end()); + return true; + } + + if (!active_) return false; + active_ = false; + + mjsBody* world = mjs_findBody(spec, "world"); + if (!world) return false; + mjsBody* body = mjs_addBody(world, nullptr); + if (!body) return false; + mjsJoint* joint = mjs_addJoint(body, nullptr); + if (!joint) return false; + mjsGeom* geom = mjs_addGeom(body, nullptr); + if (!geom) return false; + + ObjectInfo& object = objects_.emplace_back(); + object.name = "projectile" + std::to_string(counter_++);; + object.expiration = data->time + lifetime_; + + mjtNum pos[3]; + mjtNum dir[3]; + mjtNum up[3]; + mjv_cameraFrame(pos, dir, up, nullptr, data, camera); + mjs_setName(body->element, object.name.c_str()); + + joint->type = mjJNT_FREE; + body->mass = mass_; + geom->type = type_; + geom->size[0] = size_; + geom->size[1] = size_; + geom->size[2] = size_; + // Slightly in front of the camera. + body->pos[0] = pos[0] + (dir[0] * 0.1); + body->pos[1] = pos[1] + (dir[1] * 0.1); + body->pos[2] = pos[2] + (dir[2] * 0.1); + // Randomize the color. + geom->rgba[0] = std::uniform_real_distribution(0.3f, 1.0f)(rng_); + geom->rgba[1] = std::uniform_real_distribution(0.3f, 1.0f)(rng_); + geom->rgba[2] = std::uniform_real_distribution(0.3f, 1.0f)(rng_); + geom->rgba[3] = 1.0; + // Launch it slightly upwards to get a nice arc. + launch_vel_[0] = (dir[0] * speed_) + up[0]; + launch_vel_[1] = (dir[1] * speed_) + up[1]; + launch_vel_[2] = (dir[2] * speed_) + up[2]; + return true; + } + + void UpdateSpecPostCompile(const mjSpec* spec, const mjModel* model, + mjData* data) { + if (objects_.empty()) { + return; + } + + ObjectInfo& object = objects_.back(); + if (object.launched) { + return; + } + object.launched = true; + + const int body_id = mj_name2id(model, mjOBJ_BODY, object.name.c_str()); + if (body_id < 0) { + return; + } + int joint_id = model->body_jntadr[body_id]; + if (joint_id < 0 || model->jnt_type[joint_id] != mjJNT_FREE) { + return; + } + int qvel_addr = model->jnt_dofadr[joint_id]; + if (qvel_addr < 0) { + return; + } + object.body_id = body_id; + data->qvel[qvel_addr + 0] = launch_vel_[0]; + data->qvel[qvel_addr + 1] = launch_vel_[1]; + data->qvel[qvel_addr + 2] = launch_vel_[2]; + } + + private: + struct ObjectInfo { + std::string name; + int body_id = -1; + mjtNum expiration = 0; + bool launched = false; + }; + + std::mt19937 rng_; + int counter_ = 0; + bool enabled_ = false; + bool active_ = false; + mjtNum size_ = 0.13365; + mjtNum speed_ = 10.0; + mjtNum mass_ = 10.0; + mjtNum lifetime_ = 5.0; + mjtGeom type_ = mjGEOM_BOX; + mjtNum launch_vel_[3] = {0, 0, 0}; + std::vector objects_; +}; + +} // namespace mujoco::studio + +mjPLUGIN_LIB_INIT { + using mujoco::studio::ObjectLauncher; + + static ObjectLauncher plugin; + + mujoco::platform::GuiPlugin gui; + gui.data = &plugin; + gui.name = "ObjectLauncher"; + gui.update = [](mujoco::platform::GuiPlugin* self) { + auto* plugin = static_cast(self->data); + plugin->UpdateGui(); + }; + mujoco::platform::RegisterPlugin(gui); + + mujoco::platform::KeyHandlerPlugin key_handler; + key_handler.data = &plugin; + key_handler.name = "ObjectLauncher"; + key_handler.key_chord = ImGuiKey_Enter | ImGuiMod_Ctrl | ImGuiMod_Shift; + key_handler.on_key_pressed = [](mujoco::platform::KeyHandlerPlugin* self) { + auto* plugin = static_cast(self->data); + plugin->HandleKeyboardEvent(); + }; + mujoco::platform::RegisterPlugin(key_handler); + + mujoco::platform::SpecEditorPlugin spec_editor; + spec_editor.data = &plugin; + spec_editor.name = "ObjectLauncher"; + spec_editor.pre_compile = [](mujoco::platform::SpecEditorPlugin* self, + mjSpec* spec, const mjModel* model, + const mjData* data, const mjvCamera* camera) { + auto* plugin = static_cast(self->data); + return plugin->UpdateSpecPreCompile(spec, model, data, camera); + }; + spec_editor.post_compile = [](mujoco::platform::SpecEditorPlugin* self, + const mjSpec* spec, const mjModel* model, + mjData* data) { + auto* plugin = static_cast(self->data); + return plugin->UpdateSpecPostCompile(spec, model, data); + }; + mujoco::platform::RegisterPlugin(spec_editor); +} diff --git a/src/experimental/platform/plugin.cc b/src/experimental/platform/plugin.cc index fb790f63..3ddb6e1b 100644 --- a/src/experimental/platform/plugin.cc +++ b/src/experimental/platform/plugin.cc @@ -22,6 +22,8 @@ using GuiPlugin = mujoco::platform::GuiPlugin; using ModelPlugin = mujoco::platform::ModelPlugin; +using KeyHandlerPlugin = mujoco::platform::KeyHandlerPlugin; +using SpecEditorPlugin = mujoco::platform::SpecEditorPlugin; namespace mujoco::platform { @@ -71,3 +73,5 @@ void ForEachPlugin(const std::function& fn) { MUJOCO_SPECIALIZE_PLUGIN(GuiPlugin, "gui plugin"); MUJOCO_SPECIALIZE_PLUGIN(ModelPlugin, "model plugin"); +MUJOCO_SPECIALIZE_PLUGIN(KeyHandlerPlugin, "key handler plugin"); +MUJOCO_SPECIALIZE_PLUGIN(SpecEditorPlugin, "spec editor plugin"); diff --git a/src/experimental/platform/plugin.h b/src/experimental/platform/plugin.h index f10fd641..5e47fae9 100644 --- a/src/experimental/platform/plugin.h +++ b/src/experimental/platform/plugin.h @@ -82,6 +82,45 @@ struct ModelPlugin final { void* data = nullptr; }; +// Plugin for handling custom keyboard events. +struct KeyHandlerPlugin final { + using OnKeyPressedFn = void (*)(KeyHandlerPlugin* self); + + // The name of the plugin; must be unique. + const char* name = ""; + + // The ImGui key codes for the key combination that triggers the plugin. + int key_chord = 0; + + // The function to be called when the above key combination is pressed. + OnKeyPressedFn on_key_pressed = nullptr; + + // Optional data pointer. + void* data = nullptr; +}; + +// Plugin for editing the mjSpec. +struct SpecEditorPlugin final { + using PreCompileFn = bool (*)(SpecEditorPlugin* self, mjSpec* spec, + const mjModel* model, const mjData* data, + const mjvCamera* camera); + using PostCompileFn = void (*)(SpecEditorPlugin* self, const mjSpec* spec, + const mjModel* model, mjData* data); + + // The name of the plugin; must be unique. + const char* name = ""; + + // Callback that edits the spec. If it returns true, then the spec will be + // recompiled and `post_compile` will be called with the result. + PreCompileFn pre_compile = nullptr; + + // Callback that is called after the spec has been recompiled. + PostCompileFn post_compile = nullptr; + + // Optional data pointer. + void* data = nullptr; +}; + } // namespace mujoco::platform #endif // MUJOCO_SRC_EXPERIMENTAL_PLATFORM_PLUGIN_H_ diff --git a/src/experimental/studio/app.cc b/src/experimental/studio/app.cc index 6da86d2f..352e8c8a 100644 --- a/src/experimental/studio/app.cc +++ b/src/experimental/studio/app.cc @@ -25,7 +25,6 @@ #include #include #include -#include #include #include #include @@ -108,8 +107,7 @@ static constexpr std::array kPercentRealTime = { }; // clang-format on -App::App(Config config) - : rng_(std::random_device()()), ini_path_(std::move(config.ini_path)) { +App::App(Config config) : ini_path_(std::move(config.ini_path)) { platform::Window::Config window_config; window_config.renderer_backend = platform::Renderer::GetBackend(); window_config.offscreen_mode = config.offscreen_mode; @@ -222,10 +220,6 @@ void App::OnModelLoaded(std::string filename, ModelKind model_kind) { SetSpeedIndex(i); } } -if (spec_op_) { - spec_op_(); - spec_op_ = nullptr; - } platform::ForEachPlugin([&](auto* plugin) { if (plugin->post_model_loaded) { @@ -391,6 +385,23 @@ void App::ProcessPendingLoads() { } } + if (spec_op_) { + spec_op_(); + spec_op_ = nullptr; + } + + // Allow plugins to edit the spec as well. + platform::ForEachPlugin([&](auto* plugin) { + if (plugin->pre_compile) { + if (plugin->pre_compile(plugin, spec(), model(), data(), &camera_)) { + Recompile(); + if (plugin->post_compile) { + plugin->post_compile(plugin, spec(), model(), data()); + } + }; + } + }); + // Check plugins to see if we need to load a new model. platform::ForEachPlugin([&](auto* plugin) { if (plugin->get_model_to_load) { @@ -715,60 +726,65 @@ void App::HandleKeyboardEvents() { ToggleFlag(vis_options_.geomgroup[4]); } else if (ImGui_IsChordJustPressed(ImGuiKey_5)) { ToggleFlag(vis_options_.geomgroup[5]); - } else if (has_model()) { - if (ImGui_IsChordJustPressed(ImGuiKey_Escape)) { - ui_.camera_idx = - platform::SetCamera(model(), &camera_, platform::kTumbleCameraIdx); - } else if (ImGui_IsChordJustPressed(ImGuiKey_LeftBracket)) { - ui_.camera_idx = - platform::SetCamera(model(), &camera_, ui_.camera_idx - 1); - } else if (ImGui_IsChordJustPressed(ImGuiKey_RightBracket)) { - ui_.camera_idx = - platform::SetCamera(model(), &camera_, ui_.camera_idx + 1); + } else if (has_model() && ImGui_IsChordJustPressed(ImGuiKey_Escape)) { + ui_.camera_idx = + platform::SetCamera(model(), &camera_, platform::kTumbleCameraIdx); + } else if (has_model() && ImGui_IsChordJustPressed(ImGuiKey_LeftBracket)) { + ui_.camera_idx = platform::SetCamera(model(), &camera_, ui_.camera_idx - 1); + } else if (has_model() && ImGui_IsChordJustPressed(ImGuiKey_RightBracket)) { + ui_.camera_idx = platform::SetCamera(model(), &camera_, ui_.camera_idx + 1); + // WASD camera controls for free camera. + } else if (is_freecam_wasd && + (ImGui::IsKeyDown(ImGuiKey_W) || ImGui::IsKeyDown(ImGuiKey_S) || + ImGui::IsKeyDown(ImGuiKey_A) || ImGui::IsKeyDown(ImGuiKey_D) || + ImGui::IsKeyDown(ImGuiKey_Q) || ImGui::IsKeyDown(ImGuiKey_E))) { + bool moved = false; + + // Move (dolly) forward/backward using W and S keys. + if (ImGui::IsKeyDown(ImGuiKey_W)) { + MoveCamera(platform::CameraMotion::TRUCK_DOLLY, 0, tmp_.cam_speed); + moved = true; + } else if (ImGui::IsKeyDown(ImGuiKey_S)) { + MoveCamera(platform::CameraMotion::TRUCK_DOLLY, 0, -tmp_.cam_speed); + moved = true; } - // WASD camera controls for free camera. - if (is_freecam_wasd) { - bool moved = false; + // Strafe (truck) left/right using A and D keys. + if (ImGui::IsKeyDown(ImGuiKey_A)) { + MoveCamera(platform::CameraMotion::TRUCK_DOLLY, -tmp_.cam_speed, 0); + moved = true; + } else if (ImGui::IsKeyDown(ImGuiKey_D)) { + MoveCamera(platform::CameraMotion::TRUCK_DOLLY, tmp_.cam_speed, 0); + moved = true; + } - // Move (dolly) forward/backward using W and S keys. - if (ImGui::IsKeyDown(ImGuiKey_W)) { - MoveCamera(platform::CameraMotion::TRUCK_DOLLY, 0, tmp_.cam_speed); - moved = true; - } else if (ImGui::IsKeyDown(ImGuiKey_S)) { - MoveCamera(platform::CameraMotion::TRUCK_DOLLY, 0, -tmp_.cam_speed); - moved = true; + // Move (pedestal) up/down using Q and E keys. + if (ImGui::IsKeyDown(ImGuiKey_Q)) { + MoveCamera(platform::CameraMotion::TRUCK_PEDESTAL, 0, tmp_.cam_speed); + moved = true; + } else if (ImGui::IsKeyDown(ImGuiKey_E)) { + MoveCamera(platform::CameraMotion::TRUCK_PEDESTAL, 0, -tmp_.cam_speed); + moved = true; + } + + if (moved) { + tmp_.cam_speed += 0.001f; + + const float max_speed = ImGui::GetIO().KeyShift ? 0.1 : 0.01f; + if (tmp_.cam_speed > max_speed) { + tmp_.cam_speed = max_speed; } - - // Strafe (truck) left/right using A and D keys. - if (ImGui::IsKeyDown(ImGuiKey_A)) { - MoveCamera(platform::CameraMotion::TRUCK_DOLLY, -tmp_.cam_speed, 0); - moved = true; - } else if (ImGui::IsKeyDown(ImGuiKey_D)) { - MoveCamera(platform::CameraMotion::TRUCK_DOLLY, tmp_.cam_speed, 0); - moved = true; - } - - // Move (pedestal) up/down using Q and E keys. - if (ImGui::IsKeyDown(ImGuiKey_Q)) { - MoveCamera(platform::CameraMotion::TRUCK_PEDESTAL, 0, tmp_.cam_speed); - moved = true; - } else if (ImGui::IsKeyDown(ImGuiKey_E)) { - MoveCamera(platform::CameraMotion::TRUCK_PEDESTAL, 0, -tmp_.cam_speed); - moved = true; - } - - if (moved) { - tmp_.cam_speed += 0.001f; - - const float max_speed = ImGui::GetIO().KeyShift ? 0.1 : 0.01f; - if (tmp_.cam_speed > max_speed) { - tmp_.cam_speed = max_speed; + } else { + tmp_.cam_speed = 0.001f; + } + } else { + platform::ForEachPlugin([&](auto* plugin) { + if (plugin->key_chord && plugin->on_key_pressed) { + if (ImGui_IsChordJustPressed(plugin->key_chord)) { + plugin->on_key_pressed(plugin); } - } else { - tmp_.cam_speed = 0.001f; } - } + }); } } diff --git a/src/experimental/studio/app.h b/src/experimental/studio/app.h index 170aadd6..d2cce3a5 100644 --- a/src/experimental/studio/app.h +++ b/src/experimental/studio/app.h @@ -19,7 +19,6 @@ #include #include #include -#include #include #include #include @@ -224,8 +223,6 @@ class App { bool has_model() const { return model_holder_ && model_holder_->model(); } bool has_data() const { return model_holder_ && model_holder_->data(); } - std::mt19937 rng_; - std::string ini_path_; std::string model_name_; // Used if model_kind_ is kModelFromBuffer. std::string model_path_;