Refactor the projectile launcher into a plugin.

Add a simple GUI for it that allows users to enable/disable
the key binding as well as control some parameters.

PiperOrigin-RevId: 872816645
Change-Id: I8abb2803e17b4f6195d447220c436d4701b24db8
This commit is contained in:
Haroon Qureshi
2026-02-20 03:18:55 -08:00
committed by Copybara-Service
parent 28ad603f6b
commit 0083cf7fa8
6 changed files with 329 additions and 58 deletions
+1
View File
@@ -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
@@ -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 <random>
#include <string>
#include <vector>
#include <imgui.h>
#include <mujoco/mujoco.h>
#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<float>(0.3f, 1.0f)(rng_);
geom->rgba[1] = std::uniform_real_distribution<float>(0.3f, 1.0f)(rng_);
geom->rgba[2] = std::uniform_real_distribution<float>(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<ObjectInfo> 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<mujoco::studio::ObjectLauncher*>(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<mujoco::studio::ObjectLauncher*>(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<mujoco::studio::ObjectLauncher*>(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<mujoco::studio::ObjectLauncher*>(self->data);
return plugin->UpdateSpecPostCompile(spec, model, data);
};
mujoco::platform::RegisterPlugin(spec_editor);
}
+4
View File
@@ -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<void(T*)>& 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");
+39
View File
@@ -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_
+71 -55
View File
@@ -25,7 +25,6 @@
#include <filesystem>
#include <functional>
#include <memory>
#include <random>
#include <span>
#include <string>
#include <string_view>
@@ -108,8 +107,7 @@ static constexpr std::array<const char*, 31> 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<platform::ModelPlugin>([&](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<platform::SpecEditorPlugin>([&](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<platform::ModelPlugin>([&](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<platform::KeyHandlerPlugin>([&](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;
}
}
});
}
}
-3
View File
@@ -19,7 +19,6 @@
#include <functional>
#include <memory>
#include <optional>
#include <random>
#include <span>
#include <string>
#include <string_view>
@@ -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_;