9 Commits
Author SHA1 Message Date
gokayfem 3f9612774e fix: support Transformers 4.x Grounding DINO loading 2026-08-01 14:00:56 +03:00
gokayfem 7a74f5a079 docs: add citation metadata 2026-08-01 03:18:57 +03:00
Gökay Aydoğan 67344abe6a Merge pull request #167 from octo-patch/octo/20260731-music-generation-tool-recvqgT4IrGB9O
Add MiniMax music generation and cover node
2026-07-31 19:25:50 +03:00
octo-patch e79f316908 Add MiniMax music generation node 2026-07-31 20:25:20 +08:00
Gökay Aydoğan f4bc8b9eef Merge pull request #166 from gokayfem/codex/robotics-vla
feat: add robotics VLA policy toolkit
2026-07-31 01:42:15 +03:00
gokayfem 5779c50b20 test: keep robotics clients optional 2026-07-31 01:37:50 +03:00
gokayfem 48101541a3 feat: add robotics VLA policy toolkit 2026-07-31 01:34:48 +03:00
Gökay Aydoğan fcfdf7b210 Merge pull request #165 from gokayfem/dependabot/github_actions/actions-674967a53d
ci: bump actions/upload-artifact from 4 to 7 in the actions group
2026-07-30 13:33:38 +03:00
dependabot[bot] 56cdd25aa9 ci: bump actions/upload-artifact from 4 to 7 in the actions group
Bumps the actions group with 1 update: [actions/upload-artifact](https://github.com/actions/upload-artifact).


Updates `actions/upload-artifact` from 4 to 7
- [Release notes](https://github.com/actions/upload-artifact/releases)
- [Commits](https://github.com/actions/upload-artifact/compare/v4...v7)

---
updated-dependencies:
- dependency-name: actions/upload-artifact
  dependency-version: '7'
  dependency-type: direct:production
  update-type: version-update:semver-major
  dependency-group: actions
...

Signed-off-by: dependabot[bot] <support@github.com>
2026-07-30 10:21:19 +00:00
20 changed files with 5275 additions and 5 deletions
+1 -1
View File
@@ -85,7 +85,7 @@ jobs:
--cov-report=xml --cov-fail-under=70
- name: Upload coverage report
if: matrix.coverage == true && always()
uses: actions/upload-artifact@v4
uses: actions/upload-artifact@v7
with:
name: coverage-xml
path: coverage.xml
+45
View File
@@ -9,6 +9,50 @@ Versions are published to the [Comfy Registry](https://registry.comfy.org/)
from `pyproject.toml`. A release is only published when `version` changes, so
every user-visible fix needs a version bump.
## [3.5.0] - 2026-07-31
### Added
- A MiniMax music node with fixed global and China endpoints, generation and
cover model selection, regional request fields, URL and hexadecimal response
decoding, and MP3, WAV, and PCM output through the existing audio contract.
### Security
- MiniMax credentials are read only from `MINIMAX_API_KEY`; workflows cannot
supply a key or redirect it to a custom endpoint, and request errors redact
the resolved value before reaching ComfyUI.
## [3.4.0] - 2026-07-31
### Added
- A robotics-safe VLA layer with typed embodiment, observation, and action
contracts; bounded multi-camera history; trajectory inspection and preview;
action-chunk replanning; and explicit bounds, rate, dimension, horizon, and
non-finite-value checks before handoff.
- Native policy clients for OpenPI's WebSocket protocol and NVIDIA Isaac
GR00T's ZeroMQ protocol, plus a portable authenticated HTTP/JPEG protocol for
isolated policy runtimes.
- An isolated current-LeRobot policy server with pre/postprocessor support,
serialized inference, optional idle CPU offload, checkpoint feature metadata,
and environment-only bearer authentication.
- A curated 15-model VLA catalog covering SmolVLA, X-VLA, the OpenPI family,
GR00T N1.7, WALL-OSS, MolmoAct2, VLA-JEPA, LingBot-VA, FastWAM, EO-1, EVO-1,
OpenVLA-OFT, and Octo with explicit readiness and fine-tuning requirements.
- A complete API workflow, setup guide, compatibility matrix, security
guidance, and real-weight SmolVLA validation on an RTX 3090.
### Security
- Workflow JSON never stores robotics API keys. The clients read only
`VLA_POLICY_TOKEN`, `OPENPI_API_KEY`, or `GROOT_API_TOKEN` from the
environment, redact them from errors/reports, reject embedded URL
credentials, and require encrypted transports plus explicit opt-in for
remote endpoints where the upstream protocol supports encryption.
- The included policy server bounds request, camera, history, and response
sizes and never uses pickle across the network.
## [3.3.1] - 2026-07-30
### Fixed
@@ -142,6 +186,7 @@ Tagging began at 3.3.0. Earlier versions link to the commit that declared
them, because retroactively tagging them would run current CI against code
that predates it.
[3.4.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/compare/v3.3.1...v3.4.0
[3.3.1]: https://github.com/gokayfem/ComfyUI_VLM_nodes/compare/v3.3.0...v3.3.1
[3.3.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/releases/tag/v3.3.0
[3.2.0]: https://github.com/gokayfem/ComfyUI_VLM_nodes/commit/44fefcb
+21
View File
@@ -0,0 +1,21 @@
cff-version: 1.2.0
message: "If you use ComfyUI VLM Nodes in your work, please cite it using the metadata below."
type: software
title: "ComfyUI VLM Nodes"
version: "3.5.0"
date-released: 2026-07-31
authors:
- family-names: "Aydoğan"
given-names: "Gökay"
orcid: "https://orcid.org/0000-0002-2343-9433"
abstract: "Production-ready local and API vision-language, structured prompting, audio, and utility nodes for ComfyUI."
keywords:
- ComfyUI
- vision-language models
- multimodal AI
- image understanding
- video understanding
- generative AI
license: Apache-2.0
repository-code: "https://github.com/gokayfem/ComfyUI_VLM_nodes"
url: "https://github.com/gokayfem/ComfyUI_VLM_nodes"
+32
View File
@@ -137,6 +137,38 @@ Authoritative references:
- Model downloads, imports, and package compilation never occur during node
discovery.
## Robotics / VLA policy compatibility
ComfyUI's robotics schemas, safety gate, trajectory tools, and universal HTTP
client run wherever this node pack runs. Policy runtime compatibility is
separate:
| Policy route | ComfyUI client | Policy environment | Practical boundary |
| --- | --- | --- | --- |
| Universal VLA HTTP | Windows, Linux, macOS; CUDA, ROCm, Metal, XPU, CPU | Any host that implements `comfyui-vla-http-v1` | Loopback HTTP or trusted HTTPS; no pickle |
| LeRobot sidecar | Same universal client | Current LeRobot supports Linux, Windows, and macOS; individual policy extras/operators vary | Python/PyTorch live outside ComfyUI; fine-tuned checkpoint required for the target embodiment |
| openpi WebSocket | Lightweight optional client on every ComfyUI platform | Upstream currently tests Ubuntu 22.04 + NVIDIA, inference above 8 GB VRAM | Use WSL/Docker/Linux server; remote transport must be WSS |
| Isaac-GR00T N1.7 ZMQ | Lightweight optional client on every ComfyUI platform | NVIDIA CUDA/Jetson Linux according to upstream deployment matrix | ZMQ has no transport encryption; use a private network/tunnel |
| OpenVLA-OFT | Universal client with a project-specific bridge | Upstream PyTorch/CUDA environment | OFT is the preferred high-frequency multi-image OpenVLA route |
| Octo | Universal client with a project-specific bridge | Isolated JAX environment | Kept as a lightweight research baseline, not the default maintained runtime |
Install only the native client protocols into ComfyUI:
```bash
python -m pip install -r requirements-robotics-client.txt
```
Do not install `lerobot[all]`, openpi, Isaac-GR00T, OpenVLA, or JAX into
ComfyUI's Python. The included LeRobot HTTP sidecar belongs in its own
environment and optionally moves its owned policy to CPU after an idle
interval. It does not flush ComfyUI's accelerator cache.
An embodiment profile is a workflow contract, not a hardware certification.
The supplied profiles are visibly labeled templates. Before real deployment,
replace action bounds/deltas with the trained dataset's semantics and the
manufacturer/controller limits. ComfyUI never opens ROS, serial, CAN, or robot
SDK transports.
Install manually:
```bash
+65
View File
@@ -21,6 +21,71 @@ One checkpoint covers sibling sizes that use the same architecture and loader.
The node does not download every size simply to repeat the same integration
test.
## Robotics VLA pass
Validated on 2026-07-31 through the included isolated LeRobot HTTP policy
server, entirely from WSL and D-drive storage:
- Runtime: Python 3.12.12, LeRobot 0.6.1 from current upstream source,
PyTorch 2.11.0+cu128, and an NVIDIA RTX 3090.
- Checkpoint: `lerobot/smolvla_base` (about 2.5 GiB of D-drive cache), backed
by `HuggingFaceTB/SmolVLM2-500M-Video-Instruct`.
- Real input: local `image (23).png`, a 256x256 outdoor photograph, repeated
across the checkpoint's three declared camera keys with a six-value state
vector and the task “Move the end effector toward the backpack and prepare
to grasp it.”
- Contract: three camera tensors, `observation.state`, the LeRobot
preprocessor, `predict_action_chunk`, the checkpoint postprocessor, bounded
JSON/JPEG transport, action parsing, and the ComfyUI safety layer all ran.
The native checkpoint advertises a 50-step chunk; the server returned four
steps of six actions for this test.
- Five warm requests after one discarded warm-up measured 241.374 ms mean
server inference (242.957 ms median, 234.726–249.980 ms range) and
270.404 ms mean HTTP client time (271.483 ms median,
261.477–282.218 ms range).
- The final raw action chunk was:
```json
[
[0.06258623, -0.11250310, -0.13713294, -0.06168950, -0.00926633, -0.08506130],
[0.15420279, -0.05678255, -0.20159233, 0.06734322, -0.00563951, -0.09575561],
[0.16482556, -0.07453565, -0.17410603, 0.02461835, -0.00256573, 0.15842065],
[0.27048433, -0.09272483, -0.19934477, 0.05491992, 0.05286619, 0.07594281]
]
```
Applying the SO-100/SO-101 template limits from an all-zero previous action
found five per-step delta violations, no bounds violations, and no
non-finite values. `Clamp safely` produced:
```json
[
[0.06258623, -0.1, -0.1, -0.06168950, -0.00926633, -0.08506130],
[0.15420279, -0.05678255, -0.2, 0.03831051, -0.00563951, -0.09575561],
[0.16482556, -0.07453565, -0.17410603, 0.02461835, -0.00256573, 0.05424440],
[0.26482555, -0.09272483, -0.19934477, 0.05491992, 0.05286619, 0.07594281]
]
```
This is an end-to-end loading, preprocessing, inference, transport, parsing,
and safety-contract pass. It is not evidence that a base SmolVLA checkpoint can
control an SO-100 from an arbitrary Internet-style photograph. Actual robot
deployment still requires embodiment-matched fine-tuning, calibrated state and
camera inputs, hardware-certified limits, a deadman/watchdog, collision
handling, and an external emergency stop.
The same real checkpoint was then exercised through ComfyUI's actual local
`POST /prompt` API, not by calling the Python node directly. The graph loaded
and center-cropped the real image to 256x256, constructed the three-camera
checkpoint contract, called the isolated GPU policy, applied the SO-100/SO-101
template safety gate, rendered a 960x480 trajectory preview, and emitted all
three text reports. Final prompt
`86b31f5a-5a8c-4abc-abdb-634e46da5c93` completed successfully: the policy
returned `[4, 6]` actions, the safety gate found three rate violations and
clamped them, `safe_for_handoff` was true under the declared template, and
ComfyUI wrote preview `ComfyUI_temp_icynu_00001_.png`. The reusable acceptance
harness is `tests/manual_robotics_smoke.py`.
## ComfyUI API pass
ComfyUI started from the D-drive WSL installation with all four repaired custom
+84 -2
View File
@@ -1,7 +1,7 @@
# ComfyUI VLM Nodes
Production-oriented vision-language, structured prompting, audio, and utility
nodes for ComfyUI. Version 3.3 supports ComfyUI's selected NVIDIA CUDA, AMD
nodes for ComfyUI. Version 3.4 supports ComfyUI's selected NVIDIA CUDA, AMD
ROCm, Apple Metal, Intel XPU, and CPU device without replacing its PyTorch
build. It removes startup installers and global accelerator cache flushes,
adds real image/video batches and live token streaming, and uses ComfyUI model
@@ -340,6 +340,9 @@ API-format examples are in [`examples/vision`](examples/vision):
The dependency-free text-toolkit example is
[`examples/text_toolkit_api.json`](examples/text_toolkit_api.json).
Robotics policy, safety, and sidecar examples are in
[`examples/robotics`](examples/robotics), including a complete universal
HTTP policy graph.
Upload the named media to ComfyUI's input directory, adjust the filenames and
labels, then submit the JSON object as the `prompt` value to `/prompt`. These
@@ -347,7 +350,7 @@ are API graphs, not frontend workflow-export JSON.
## Node reference
All 78 registered nodes, grouped by their menu category. The **Node ID** is the
All 89 registered nodes, grouped by their menu category. The **Node ID** is the
`class_type` written into workflow and API JSON — search for that string when
you need to find a node you saw on a canvas.
@@ -477,6 +480,26 @@ projector (mmproj).
| Hosted VLM API (Secure) | `HostedVLMAPI` | `STRING`, `STRING`, `INT` |
| Hosted LLM API (Secure) | `PromptGenerateAPI` | `STRING` |
### Robotics / VLA policies
These nodes build and inspect policy observations/actions. They never send
commands to robot hardware. Heavy policy runtimes stay in isolated LeRobot,
openpi, GR00T, OpenVLA/OFT, or JAX environments.
| Node | Node ID | Outputs |
| --- | --- | --- |
| VLA Embodiment Profile | `VLAEmbodimentProfile` | `VLA_EMBODIMENT`, `STRING`, `INT`, `INT` |
| VLA Observation Builder | `VLAObservationBuilder` | `VLA_OBSERVATION`, `STRING`, `INT` |
| VLA Policy — Universal HTTP | `VLAHTTPPolicy` | `VLA_ACTIONS`, `STRING` |
| VLA Policy — OpenPI WebSocket | `VLAOpenPIWebSocketPolicy` | `VLA_ACTIONS`, `STRING` |
| VLA Policy — GR00T N1.7 ZMQ | `VLAGr00tZMQPolicy` | `VLA_ACTIONS`, `STRING` |
| VLA Action Safety Gate | `VLAActionSafety` | `VLA_ACTIONS`, `STRING`, `BOOLEAN` |
| VLA Actions From JSON | `VLAActionsFromJSON` | `VLA_ACTIONS`, `STRING` |
| VLA Action Chunk Replan | `VLAActionChunkReplan` | `VLA_ACTIONS`, `STRING` |
| VLA Action Inspect | `VLAActionInspect` | `STRING`, `STRING`, `INT`, `INT` |
| VLA Trajectory Preview | `VLATrajectoryPreview` | `IMAGE` |
| VLA Model Catalog | `VLAModelCatalog` | `STRING`, `STRING`, `STRING`, `STRING` |
### Text toolkit
Dependency-free string handling, so a VLM response can be shaped without an
@@ -512,9 +535,19 @@ device, backend, and which optional packages are installed.
| --- | --- | --- |
| AudioLDM2 | `AudioLDM2Node` | `*`, `INT`, `AUDIO` |
| Chat Musician | `ChatMusician` | `STRING`, `*`, `INT`, `AUDIO` |
| MiniMax Music | `MiniMaxMusicNode` | `*`, `INT`, `AUDIO` |
| PlayMusic Node | `PlayMusic` | `*` |
| Save Audio | `SaveAudioNode` | — |
MiniMax Music reads `MINIMAX_API_KEY` only from the ComfyUI server
environment. It uses fixed `global_en` and `cn_zh` endpoints, supports music
generation and cover models, decodes URL or hexadecimal responses, and emits
MP3, WAV, or PCM results through the existing waveform and `AUDIO` sockets.
The `aigc_watermark` field is sent only for `cn_zh` requests. See the official
[global](https://platform.minimax.io/docs/api-reference/music-generation) or
[China](https://platform.minimaxi.com/docs/api-reference/music-generation)
music API reference for account and content requirements.
### Legacy model loaders
Kept for existing workflows. New graphs should prefer **Modern VLM**, which
@@ -546,6 +579,35 @@ Current official bitsandbytes wheels are installed automatically only on their
supported OS/architecture combinations. Unsupported machines retain all
non-quantized nodes.
### Robotics / VLA isolated runtimes
The robotics nodes keep policy dependencies outside ComfyUI. The universal
HTTP client works without another package. Native openpi WebSocket and
GR00T ZeroMQ clients use the lightweight optional extra:
```bash
python -m pip install \
-r ComfyUI/custom_nodes/ComfyUI_VLM_nodes/requirements-robotics-client.txt
```
`VLA Model Catalog` covers current SmolVLA, X-VLA, π0/π0-FAST/π0.5,
GR00T N1.7, WALL-OSS, MolmoAct2, VLA-JEPA, LingBot-VA, FastWAM, EO-1,
EVO-1, OpenVLA-OFT, and Octo routes. “Available” means a supported isolated
runtime/checkpoint path; base and architecture-only entries still require
embodiment-specific training and transforms.
Start with SmolVLA for small consumer hardware. The included authenticated
LeRobot sidecar loads one chosen policy, uses its serialized processors,
returns action chunks over bounded JSON/JPEG, keeps it resident for speed,
and can offload it to CPU after an idle timeout. Remote policy URLs require
encrypted transport and explicit opt-in. Tokens are fixed environment
variables (`VLA_POLICY_TOKEN`, `OPENPI_API_KEY`, or `GROOT_API_TOKEN`) and are
never workflow inputs.
See [`examples/robotics/README.md`](examples/robotics/README.md) for D-drive
WSL setup, platform boundaries, current model readiness, observation schemas,
action safety semantics, and the runnable API example.
### Moondream 3 / 3.1 isolated runtime
Moondream's official Photon package pins Pillow below version 11 while
@@ -791,3 +853,23 @@ catalog-only evidence matrix.
Please report reproducible bugs at the
[issue tracker](https://github.com/gokayfem/ComfyUI_VLM_nodes/issues).
<details>
<summary><strong>Cite this project</strong></summary>
If ComfyUI VLM Nodes supports your work, please cite the software. GitHub also
provides ready-to-copy APA and BibTeX entries via **Cite this repository**.
```bibtex
@software{Aydogan_ComfyUI_VLM_Nodes_2026,
author = {Aydoğan, Gökay},
title = {ComfyUI VLM Nodes},
version = {3.5.0},
year = {2026},
url = {https://github.com/gokayfem/ComfyUI_VLM_nodes}
}
```
[ORCID](https://orcid.org/0000-0002-2343-9433) · [Citation metadata](CITATION.cff)
</details>
+36
View File
@@ -39,7 +39,11 @@ restart ComfyUI:
| Mistral | `MISTRAL_API_KEY` |
| Together AI | `TOGETHER_API_KEY` |
| OpenRouter | `OPENROUTER_API_KEY` |
| MiniMax | `MINIMAX_API_KEY` |
| Custom remote endpoint | `CUSTOM_API_KEY` |
| Universal VLA policy server | `VLA_POLICY_TOKEN` |
| openpi WebSocket server | `OPENPI_API_KEY` |
| Isaac-GR00T ZMQ server | `GROOT_API_TOKEN` |
For an interactive POSIX/WSL session, this avoids putting the value in shell
history:
@@ -59,6 +63,38 @@ selected provider's server-side search system and may have separate retention,
regional-availability, and billing terms. Treat it as an explicit data-sharing
choice; do not enable it for content that is outside those terms.
## Robotics policy endpoints
Robotics tokens are also server-side only. Workflow nodes select an endpoint,
but cannot select an arbitrary environment variable or contain the secret
value.
- The universal policy client permits unencrypted HTTP only on loopback.
Remote use requires HTTPS plus `allow_remote=true`; redirects and
environment proxies are disabled.
- The openpi client permits unencrypted WebSocket only on loopback. Remote use
requires WSS plus `allow_remote=true`.
- GR00T's official ZeroMQ protocol has token authentication but no built-in
transport encryption. Keep it on loopback/private infrastructure or place it
inside an authenticated encrypted tunnel. Never expose its port directly to
the public internet.
- Camera payloads are JPEG-compressed and bounded per frame and per request.
Response sizes, camera count, observation history, state/action dimensions,
and action horizons are bounded before use.
- MessagePack ndarray decoders reject object/void dtypes and never deserialize
pickle. The included HTTP sidecar uses bounded JSON instead of LeRobot's
pickle-based asynchronous transport.
- Errors redact the resolved token and authorization-like values. Reports
include only endpoint scheme/host/port, not request headers, full camera
payloads, or state data.
Robot observations may expose people, homes, workplaces, proprietary tasks,
and physical state. Treat them as sensitive even when no API key is present.
The safety node is a data validation gate, not a certified control system.
This package intentionally contains no ROS, serial, CAN, motor, or robot SDK
transport; a separate controller must enforce emergency stop, deadman,
watchdog, collision/workspace, command-age, and manufacturer limits.
## Legacy workflows
Versions before this security update exposed an `api_key` text widget.
+2
View File
@@ -18,6 +18,7 @@ node_list = [
"llavaloader",
"mcllava",
"minicpm",
"minimax_music",
"modern_vlm",
"molmo",
"moondream31",
@@ -26,6 +27,7 @@ node_list = [
"paligemma",
"playmusic",
"qwen2vl",
"robotics",
"sam2",
"sam3_adapter",
"simpletext",
+266
View File
@@ -0,0 +1,266 @@
# Robotics and VLA workflows
The robotics nodes make ComfyUI a policy-development, inspection, and
simulation surface. They do **not** send commands to motors, ROS, CAN, serial,
or a robot SDK.
The boundary is intentional:
```text
camera/state/task
|
v
VLA Observation Builder
|
v
isolated policy server ---> raw action chunk
|
v
VLA Action Safety Gate
|
+-------------+-------------+
| |
v v
inspect / plot / record simulator or your own
supervised controller bridge
```
A real controller bridge must independently enforce a deadman, watchdog,
emergency stop, collision/workspace limits, timestamps, command freshness, and
the manufacturer's limits. A `safe_for_handoff=true` workflow result only
means that the declared ComfyUI profile checks passed.
## Why policy runtimes are isolated
LeRobot, openpi, Isaac-GR00T, OpenVLA/OFT, and Octo use different PyTorch/JAX,
CUDA, Transformers, compiler, and operating-system combinations. Installing
all of those into ComfyUI would replace or constrain the working accelerator
stack and make Windows, macOS, ROCm, and XPU support worse.
The ComfyUI package therefore contains only:
- typed state/action/camera contracts;
- bounded image serialization;
- a dependency-light universal HTTPS/loopback HTTP client;
- exact clients for the official openpi MessagePack WebSocket and GR00T
MessagePack/ZeroMQ protocols;
- action validation, horizon control, inspection, and plotting.
The heavyweight policy stays in its own process, container, WSL distribution,
Linux machine, Mac, or GPU server. This also allows ComfyUI to use AMD ROCm,
Apple Metal, Intel XPU, or CPU while a policy runs on an NVIDIA Linux server.
## Fast path: SmolVLA through the universal sidecar
Use a separate LeRobot environment. Current LeRobot documentation recommends
Python 3.12 and exposes policy-specific extras. On this computer, keep it on
the D drive:
```bash
# WSL
python3.12 -m venv /mnt/d/vla-runtime/lerobot-smolvla
source /mnt/d/vla-runtime/lerobot-smolvla/bin/activate
python -m pip install --upgrade pip
python -m pip install "lerobot[smolvla]"
export VLA_POLICY_TOKEN="$(python -c 'import secrets; print(secrets.token_urlsafe(32))')"
python /mnt/d/ComfyUI_windows_portable/ComfyUI/custom_nodes/ComfyUI_VLM_nodes/examples/robotics/lerobot_policy_server.py \
--policy-type smolvla \
--policy-path YOUR_FINE_TUNED_SMOLVLA_CHECKPOINT \
--device auto \
--actions-per-chunk 16 \
--idle-offload-seconds 300
```
Set the same `VLA_POLICY_TOKEN` in the environment that launches ComfyUI.
Never put it in a workflow. In `VLA Policy — Universal HTTP`, use
`http://127.0.0.1:8787`.
`lerobot/smolvla_base` is a base model. It is a useful fine-tuning starting
point, not a universal zero-shot controller. Use an embodiment-specific
checkpoint whose feature names, action dimensions, state dimensions,
normalization statistics, and camera keys match the workflow.
The sidecar:
- loads only the chosen policy and its serialized pre/post-processors;
- uses `predict_action_chunk` when provided and falls back to `select_action`;
- keeps the model resident by default for low latency;
- can move it to CPU after an idle interval and move it back on demand;
- accepts one request at a time per policy, preventing stateful policy races;
- uses bounded JSON/JPEG rather than pickle;
- never returns tracebacks, environment variables, request data, or
authorization headers.
For a real local API acceptance run, start ComfyUI and the policy sidecar, put
an image in ComfyUI's `input` directory, then run:
```bash
python tests/manual_robotics_smoke.py \
--comfy-url http://127.0.0.1:8188 \
--policy-url http://127.0.0.1:8787 \
--image robot_front.png
```
The script queues the graph through `POST /prompt`, waits on its history entry,
and prints the policy report, safety report, first action, and preview filename.
Install the relevant official LeRobot extra for another policy. Examples are
`lerobot[pi]` for π0/π0.5/π0-FAST and `lerobot[smolvla]` for SmolVLA. Some
newer policy integrations may require installing current LeRobot from source
with their documented extra.
## Native openpi server
Install the small ComfyUI client dependencies:
```bash
python -m pip install -r requirements-robotics-client.txt
```
Run the official openpi policy WebSocket server in its own supported
environment. The upstream runtime is currently tested on Ubuntu 22.04 and an
NVIDIA GPU with more than 8 GB for inference; use WSL/Docker or a remote Linux
server rather than forcing it into a macOS/Windows ComfyUI environment.
Use:
- `Flat keys (DROID / LIBERO)` for observations such as
`observation/image`, `observation/wrist_image`, and `observation/state`;
- `Nested images (ALOHA)` for `state`, an `images` mapping such as
`cam_high`/wrist cameras, and `prompt`.
The workflow supplies key names, but the policy's own transform still defines
the exact shapes and normalization. `OPENPI_API_KEY` is read only from the
ComfyUI server environment. Remote endpoints require WSS and explicit
`allow_remote=true`.
## Native Isaac-GR00T N1.7 server
Install the same lightweight robotics client requirements in ComfyUI. Run the
official GR00T `PolicyServer` beside an embodiment-compatible `Gr00tPolicy`.
The node sends the documented nested contract:
```text
video.<camera> uint8 [batch=1, history, height, width, RGB=3]
state.state float32[batch=1, history, state_dim]
language.task string [batch=1, 1]
```
The official server returns one or more physical-unit action streams with
shape `[batch, horizon, dimension]`. The node flattens those streams while
preserving their named slices. `GROOT_API_TOKEN` remains in the ComfyUI
environment.
GR00T N1.7 currently targets NVIDIA CUDA/Jetson Linux and needs an
embodiment-compatible base or post-trained checkpoint. A ComfyUI client on
Windows, macOS, ROCm, or another machine may call that server over a trusted
network, but remote access must be explicitly enabled. Native GR00T ZMQ does
not encrypt traffic; use a private authenticated network/tunnel. Prefer the
universal HTTPS bridge when transport-layer encryption is required.
## Model catalog: what “available” means
`VLA Model Catalog` distinguishes these states:
| Family | Example checkpoint | Route | Important qualification |
| --- | --- | --- | --- |
| SmolVLA | `lerobot/smolvla_base` | LeRobot HTTP sidecar | 450M and the best small starting point; fine-tune for the robot |
| X-VLA | `lerobot/xvla-base` | LeRobot HTTP sidecar | 0.9B cross-embodiment base; use a matching domain checkpoint |
| π0 | `lerobot/pi0_base` | LeRobot or openpi | Base/fine-tuning model, not a universal drop-in controller |
| π0-FAST | `lerobot/pi0fast-base` | LeRobot or openpi | Faster tokenized action generation |
| π0.5 | `lerobot/pi05_base` | LeRobot or openpi | Open-world generalization; still embodiment-specific |
| GR00T N1.7 | `nvidia/GR00T-N1.7-3B` | GR00T ZMQ or LeRobot | Base has specific zero-shot tags; other robots need post-training |
| X-Square WALL-OSS | `x-square-robot/wall-oss-flow` | LeRobot HTTP sidecar | MoE research model; validate checkpoint terms and embodiment |
| MolmoAct2 | `lerobot/MolmoAct2-SO100_101-LeRobot` | LeRobot HTTP sidecar | Converted SO-100/SO-101 checkpoint |
| VLA-JEPA | `lerobot/VLA-JEPA-Pretrain` | LeRobot HTTP sidecar | DROID pretrain plus LIBERO/SimplerEnv checkpoints |
| LingBot-VA | `lerobot/lingbot_va_base` | LeRobot HTTP sidecar | Prefer its LIBERO-Long/RoboTwin post-train where applicable |
| FastWAM | released LIBERO checkpoint | LeRobot HTTP sidecar | Heavy world-action research runtime |
| EO-1 / EVO-1 | your trained checkpoint | LeRobot HTTP sidecar | Architecture support, not a universal ready-made controller |
| OpenVLA-OFT | compatible OFT fine-tune | dedicated sidecar | OFT is the faster multi-image/high-frequency OpenVLA route |
| Octo small | Octo small 27M | dedicated JAX sidecar | Lightweight legacy research baseline |
The catalog is a verified runtime/checkpoint map, not a promise that a base
checkpoint understands an arbitrary robot. Exact data transforms and
fine-tuning are part of the policy.
## Observation history and real-time use
Connect an `IMAGE` batch to a camera input to represent temporal history. All
camera batches must have the same length, although a one-frame camera may
broadcast. Use:
`Video Slice` → `VLM Adaptive Frame Sampler` or a live capture source →
`VLM Image Pixel Budget` → `VLA Observation Builder`
For closed-loop robotics, do not run an unbounded ComfyUI queue for each motor
tick. Use ComfyUI to prototype and inspect the observation/policy/safety
contract, and use the policy runtime's asynchronous control support for the
actual high-frequency loop. LeRobot supports asynchronous action chunks and
GR00T supports TensorRT deployment; both are better places for timing-critical
execution.
## Action safety semantics
The `VLA Action Safety Gate` checks:
- policy action dimension against the embodiment;
- NaN and infinity;
- minimum and maximum values;
- maximum change per action dimension and control step;
- the requested execution horizon.
Modes:
- `Block unsafe`: raise and stop the workflow on any violation.
- `Clamp safely`: replace non-finite values conservatively, then clamp bounds
and sequential per-step deltas.
- `Hold position on unsafe`: replace the whole chunk with the explicitly
supplied previous/current command.
- `Report only`: preserve the raw trajectory and set
`safe_for_handoff=false`.
For delta-action policies, `previous_action_json` means the previous delta
command, not an absolute joint pose. Define the profile in the same units and
semantics as the policy output.
`VLA Actions From JSON` imports recorded/simulator trajectories without a
network policy. `VLA Action Chunk Replan` blends the unexecuted edge of an old
chunk into a new chunk to reduce discontinuities, then the result should pass
through the safety gate again. This deterministic blend is useful for workflow
experiments but does not replace LeRobot's asynchronous controller or a
policy-specific real-time chunking implementation.
## API example
`vla_http_policy_safety_api.json` is a ComfyUI API prompt graph. Put
`robot_front.png` in `ComfyUI/input`, start a compatible sidecar, then POST:
```json
{"prompt": {"...": "contents of vla_http_policy_safety_api.json"}}
```
It builds an observation, calls the policy, clamps it against the explicit
profile, renders the trajectory, and outputs both inference and safety JSON.
## Security checklist
- Keep all policy tokens in environment variables.
- Leave `allow_remote=false` for local servers.
- Remote universal endpoints must use HTTPS; remote openpi endpoints must use
WSS.
- Never expose GR00T ZMQ directly to an untrusted network.
- Pin checkpoint revisions when reproducibility matters.
- Treat camera images, task language, and robot state as sensitive data.
- Do not connect action JSON directly to hardware without a separate
supervised controller bridge and independent safety system.
Authoritative upstream documentation:
- [LeRobot installation](https://huggingface.co/docs/lerobot/main/en/installation)
- [LeRobot SmolVLA](https://huggingface.co/docs/lerobot/smolvla)
- [LeRobot asynchronous inference](https://huggingface.co/docs/lerobot/async)
- [Physical Intelligence openpi](https://github.com/Physical-Intelligence/openpi)
- [NVIDIA Isaac-GR00T](https://github.com/NVIDIA/Isaac-GR00T)
- [OpenVLA and OFT](https://github.com/openvla/openvla)
- [Octo](https://github.com/octo-models/octo)
+1
View File
@@ -0,0 +1 @@
"""Runnable, dependency-isolated robotics policy bridge examples."""
+434
View File
@@ -0,0 +1,434 @@
#!/usr/bin/env python
"""Isolated LeRobot policy server for the ComfyUI VLA HTTP node.
Run this file in a dedicated environment that contains LeRobot and the
policy-specific dependencies. Do not install LeRobot's full dependency stack
into ComfyUI merely to use this bridge.
"""
from __future__ import annotations
import argparse
import base64
import hmac
import io
import json
import os
import threading
import time
from http import HTTPStatus
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from typing import Any
import numpy as np
import torch
from PIL import Image
MAX_REQUEST_BYTES = 64 * 1024 * 1024
MAX_CAMERAS = 16
MAX_FRAMES_PER_CAMERA = 256
MAX_IMAGE_BYTES = 8 * 1024 * 1024
MAX_IMAGE_PIXELS = 16 * 1024 * 1024
MAX_STATE_DIM = 2_048
MAX_ACTION_DIM = 2_048
MAX_TASK_CHARS = 16_384
def _json_bytes(value: Any) -> bytes:
return json.dumps(
value,
ensure_ascii=False,
allow_nan=False,
separators=(",", ":"),
).encode("utf-8")
def _device(value: str) -> str:
if value != "auto":
return value
if torch.cuda.is_available():
return "cuda"
if hasattr(torch, "xpu") and torch.xpu.is_available():
return "xpu"
if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
return "mps"
return "cpu"
def _decode_image(frame: dict[str, Any]) -> np.ndarray:
if frame.get("encoding") != "base64-jpeg":
raise ValueError("Only base64-jpeg camera frames are supported.")
raw = base64.b64decode(frame["data"], validate=True)
if len(raw) > MAX_IMAGE_BYTES:
raise ValueError("Encoded camera frame exceeds the 8 MiB safety limit.")
with Image.open(io.BytesIO(raw)) as image:
if image.width * image.height > MAX_IMAGE_PIXELS:
raise ValueError("Decoded camera frame exceeds the pixel safety limit.")
return np.asarray(image.convert("RGB"), dtype=np.uint8).copy()
def _decode_observation(payload: dict[str, Any]) -> dict[str, Any]:
if payload.get("schema") != "comfyui-vlm/robot-observation":
raise ValueError("Unsupported observation schema.")
if int(payload.get("version", 0)) != 1:
raise ValueError("Unsupported observation schema version.")
cameras = payload.get("cameras")
if not isinstance(cameras, dict) or not 1 <= len(cameras) <= MAX_CAMERAS:
raise ValueError("cameras must contain between 1 and 16 entries.")
observation: dict[str, Any] = {}
for key, encoded_frames in cameras.items():
key = str(key).strip()
if not key or len(key) > 256 or any(ord(char) < 32 for char in key):
raise ValueError("Camera names must contain 1 to 256 printable characters.")
if not isinstance(encoded_frames, list) or not (
1 <= len(encoded_frames) <= MAX_FRAMES_PER_CAMERA
):
raise ValueError(f"Camera {key!r} has an invalid history.")
# Current LeRobot policy processors accept one current observation.
# ComfyUI may send history for servers/models that use it; this generic
# bridge deliberately selects the latest frame.
array = _decode_image(encoded_frames[-1])
tensor = torch.from_numpy(array).permute(2, 0, 1).to(torch.float32) / 255.0
observation[key] = tensor.unsqueeze(0)
state = np.asarray(payload.get("state"), dtype=np.float32)
if (
state.ndim != 1
or not 1 <= state.size <= MAX_STATE_DIM
or not np.isfinite(state).all()
):
raise ValueError(f"state must contain 1 to {MAX_STATE_DIM} finite values.")
observation["observation.state"] = torch.from_numpy(state).unsqueeze(0)
task = str(payload.get("task", "")).strip()
if not task or len(task) > MAX_TASK_CHARS:
raise ValueError(f"task must contain 1 to {MAX_TASK_CHARS} characters.")
observation["task"] = task
return observation
def _postprocess_chunk(postprocessor, action: torch.Tensor) -> torch.Tensor:
if action.ndim == 1:
action = action.unsqueeze(0)
if action.ndim == 2:
# select_action normally returns [batch, dim].
processed = postprocessor(action)
if processed.ndim == 1:
processed = processed.unsqueeze(0)
return processed.unsqueeze(1) if processed.ndim == 2 else processed
if action.ndim != 3:
raise ValueError(f"Policy returned unsupported action shape {tuple(action.shape)}.")
processed_steps = [postprocessor(action[:, index, :]) for index in range(action.shape[1])]
return torch.stack(processed_steps, dim=1)
def _feature_metadata(features: Any) -> dict[str, dict[str, Any]]:
"""Return the portable part of a LeRobot policy feature contract."""
result: dict[str, dict[str, Any]] = {}
for key, feature in (features or {}).items():
if isinstance(feature, dict):
feature_type = feature.get("type")
shape = feature.get("shape", ())
else:
feature_type = getattr(feature, "type", None)
shape = getattr(feature, "shape", ())
feature_type = getattr(feature_type, "value", feature_type)
dimensions: list[int | str | None] = []
for dimension in shape or ():
if dimension is None:
dimensions.append(None)
continue
try:
dimensions.append(int(dimension))
except (TypeError, ValueError):
dimensions.append(str(dimension))
result[str(key)] = {
"type": str(feature_type) if feature_type is not None else "UNKNOWN",
"shape": dimensions,
}
return result
def _optional_config_int(config: Any, name: str) -> int | None:
value = getattr(config, name, None)
try:
return None if value is None else int(value)
except (TypeError, ValueError):
return None
class PolicyRuntime:
def __init__(
self,
*,
policy_type: str,
policy_path: str,
revision: str | None,
device: str,
actions_per_chunk: int,
idle_offload_seconds: float,
):
self.policy_type = policy_type
self.policy_path = policy_path
self.revision = revision
self.device = _device(device)
self.actions_per_chunk = actions_per_chunk
self.idle_offload_seconds = idle_offload_seconds
self.lock = threading.Lock()
self.policy = None
self.preprocessor = None
self.postprocessor = None
self.resident_device = "unloaded"
self.last_request = 0.0
self.load_seconds = 0.0
self._load()
if idle_offload_seconds > 0 and self.device != "cpu":
threading.Thread(target=self._idle_worker, daemon=True).start()
def _load(self) -> None:
from lerobot.policies import get_policy_class, make_pre_post_processors
started = time.perf_counter()
policy_class = get_policy_class(self.policy_type)
kwargs = {}
if self.revision:
kwargs["revision"] = self.revision
self.policy = policy_class.from_pretrained(self.policy_path, **kwargs)
self.policy.eval()
self.policy.to(self.device)
overrides = {"device": self.device}
self.preprocessor, self.postprocessor = make_pre_post_processors(
self.policy.config,
pretrained_path=self.policy_path,
pretrained_revision=self.revision,
preprocessor_overrides={"device_processor": overrides},
postprocessor_overrides={"device_processor": overrides},
)
self.resident_device = self.device
self.last_request = time.monotonic()
self.load_seconds = time.perf_counter() - started
def _ensure_resident(self) -> None:
if self.resident_device != self.device:
self.policy.to(self.device)
self.resident_device = self.device
def _idle_worker(self) -> None:
interval = min(max(self.idle_offload_seconds / 4, 1.0), 30.0)
while True:
time.sleep(interval)
if time.monotonic() - self.last_request < self.idle_offload_seconds:
continue
if not self.lock.acquire(blocking=False):
continue
try:
if (
self.resident_device != "cpu"
and time.monotonic() - self.last_request >= self.idle_offload_seconds
):
self.policy.to("cpu")
self.resident_device = "cpu"
finally:
self.lock.release()
def metadata(self) -> dict[str, Any]:
config = self.policy.config
return {
"protocol": "comfyui-vla-http-v1",
"policy_type": self.policy_type,
"policy_path": self.policy_path,
"revision": self.revision,
"configured_device": self.device,
"resident_device": self.resident_device,
"actions_per_chunk": self.actions_per_chunk,
"idle_offload_seconds": self.idle_offload_seconds,
"load_seconds": self.load_seconds,
"policy_contract": {
"input_features": _feature_metadata(
getattr(config, "input_features", None)
),
"output_features": _feature_metadata(
getattr(config, "output_features", None)
),
"observation_steps": _optional_config_int(config, "n_obs_steps"),
"native_chunk_size": _optional_config_int(config, "chunk_size"),
"native_action_steps": _optional_config_int(config, "n_action_steps"),
},
}
def infer(self, payload: dict[str, Any]) -> dict[str, Any]:
observation = _decode_observation(payload)
with self.lock:
self._ensure_resident()
started = time.perf_counter()
processed = self.preprocessor(observation)
preprocess_ms = (time.perf_counter() - started) * 1000
started_inference = time.perf_counter()
with torch.inference_mode():
predictor = getattr(self.policy, "predict_action_chunk", None)
if callable(predictor):
action = predictor(processed)
else:
action = self.policy.select_action(processed)
inference_ms = (time.perf_counter() - started_inference) * 1000
started_postprocess = time.perf_counter()
action = _postprocess_chunk(self.postprocessor, action)
if action.ndim == 3:
if action.shape[0] != 1:
raise ValueError("Only policy batch size 1 is supported.")
action = action[0]
elif action.ndim == 1:
action = action.unsqueeze(0)
if action.ndim != 2:
raise ValueError(f"Unexpected final action shape {tuple(action.shape)}.")
action = action[: self.actions_per_chunk].detach().to("cpu", torch.float32)
if not 1 <= int(action.shape[1]) <= MAX_ACTION_DIM:
raise ValueError(
f"Policy action dimension must be in [1, {MAX_ACTION_DIM}]."
)
if not torch.isfinite(action).all():
# Preserve the response for ComfyUI's safety node, but do not
# serialize non-standard JSON numbers.
raise ValueError("Policy returned NaN or infinite action values.")
postprocess_ms = (time.perf_counter() - started_postprocess) * 1000
self.last_request = time.monotonic()
return {
"actions": action.tolist(),
"server_timing": {
"preprocess_ms": preprocess_ms,
"infer_ms": inference_ms,
"postprocess_ms": postprocess_ms,
},
"policy": {
"type": self.policy_type,
"path": self.policy_path,
"device": self.device,
},
}
class PolicyHandler(BaseHTTPRequestHandler):
server_version = "ComfyUI-VLA-Policy/1"
def log_message(self, format_string: str, *args: Any) -> None:
# The request path is safe to log. Headers and bodies may contain
# credentials or camera/state data and are intentionally excluded.
print(f"{self.address_string()} - {format_string % args}")
@property
def runtime(self) -> PolicyRuntime:
return self.server.runtime
@property
def token(self) -> str:
return self.server.token
def _authorized(self) -> bool:
if not self.token:
return True
supplied = self.headers.get("Authorization", "")
expected = f"Bearer {self.token}"
return hmac.compare_digest(supplied, expected)
def _send(self, status: HTTPStatus, value: Any) -> None:
body = _json_bytes(value)
self.send_response(status.value)
self.send_header("Content-Type", "application/json; charset=utf-8")
self.send_header("Content-Length", str(len(body)))
self.send_header("Cache-Control", "no-store")
self.send_header("X-Content-Type-Options", "nosniff")
self.end_headers()
self.wfile.write(body)
def do_GET(self) -> None: # noqa: N802
if self.path not in {"/healthz", "/v1/metadata"}:
self._send(HTTPStatus.NOT_FOUND, {"error": "not_found"})
return
if not self._authorized():
self._send(HTTPStatus.UNAUTHORIZED, {"error": "unauthorized"})
return
if self.path == "/healthz":
self._send(HTTPStatus.OK, {"status": "ok"})
else:
self._send(HTTPStatus.OK, self.runtime.metadata())
def do_POST(self) -> None: # noqa: N802
if self.path != "/v1/infer":
self._send(HTTPStatus.NOT_FOUND, {"error": "not_found"})
return
if not self._authorized():
self._send(HTTPStatus.UNAUTHORIZED, {"error": "unauthorized"})
return
try:
content_length = int(self.headers.get("Content-Length", "0"))
if not 1 <= content_length <= MAX_REQUEST_BYTES:
raise ValueError("Request body size is invalid.")
body = self.rfile.read(content_length)
payload = json.loads(body)
if not isinstance(payload, dict):
raise ValueError("Request body must be a JSON object.")
result = self.runtime.infer(payload)
except (TypeError, ValueError, json.JSONDecodeError) as exc:
self._send(HTTPStatus.BAD_REQUEST, {"error": str(exc)[:1000]})
return
except Exception as exc:
# Do not return tracebacks, request data, environment variables, or
# authorization headers across the network.
self._send(
HTTPStatus.INTERNAL_SERVER_ERROR,
{"error": f"{type(exc).__name__}: {str(exc)[:800]}"},
)
return
self._send(HTTPStatus.OK, result)
def main() -> None:
parser = argparse.ArgumentParser(
description="Serve one LeRobot policy through the ComfyUI VLA HTTP protocol."
)
parser.add_argument("--policy-type", required=True, help="LeRobot policy type, e.g. smolvla")
parser.add_argument("--policy-path", required=True, help="Hub repo id or local checkpoint")
parser.add_argument("--revision", default=None, help="Optional immutable Hub revision")
parser.add_argument("--device", default="auto", help="auto, cuda, mps, xpu, or cpu")
parser.add_argument("--host", default="127.0.0.1")
parser.add_argument("--port", type=int, default=8787)
parser.add_argument("--actions-per-chunk", type=int, default=16)
parser.add_argument(
"--idle-offload-seconds",
type=float,
default=0.0,
help="Move the policy to CPU after this idle period; 0 keeps it resident.",
)
args = parser.parse_args()
if not 1 <= args.port <= 65_535:
parser.error("--port must be in [1, 65535]")
if not 1 <= args.actions_per_chunk <= 4096:
parser.error("--actions-per-chunk must be in [1, 4096]")
if args.idle_offload_seconds < 0:
parser.error("--idle-offload-seconds must be non-negative")
runtime = PolicyRuntime(
policy_type=args.policy_type,
policy_path=args.policy_path,
revision=args.revision,
device=args.device,
actions_per_chunk=args.actions_per_chunk,
idle_offload_seconds=args.idle_offload_seconds,
)
token = os.environ.get("VLA_POLICY_TOKEN", "").strip()
server = ThreadingHTTPServer((args.host, args.port), PolicyHandler)
server.runtime = runtime
server.token = token
print(
f"Policy ready at http://{args.host}:{args.port}/v1/infer "
f"(type={args.policy_type}, device={runtime.device}, auth={'on' if token else 'off'})"
)
try:
server.serve_forever()
except KeyboardInterrupt:
pass
finally:
server.server_close()
if __name__ == "__main__":
main()
@@ -0,0 +1,130 @@
{
"1": {
"class_type": "LoadImage",
"inputs": {
"image": "robot_front.png"
}
},
"2": {
"class_type": "VLAEmbodimentProfile",
"inputs": {
"preset": "Generic 7-DoF joint + gripper",
"control_hz": 20.0,
"state_names_json": "",
"action_names_json": "",
"action_min_json": "",
"action_max_json": "",
"max_delta_json": "",
"camera_names_json": "",
"action_mode_override": ""
}
},
"3": {
"class_type": "VLAObservationBuilder",
"inputs": {
"task": "Pick up the blue cube and place it in the tray.",
"state_json": "[0, 0, 0, 0, 0, 0, 0, 0]",
"primary_image": [
"1",
0
],
"primary_camera": "observation.images.front",
"history_fps": 10.0,
"timestamp": 0.0,
"embodiment": [
"2",
0
]
}
},
"4": {
"class_type": "VLAHTTPPolicy",
"inputs": {
"observation": [
"3",
0
],
"endpoint": "http://127.0.0.1:8787",
"timeout_seconds": 120.0,
"include_history": true,
"allow_remote": false
}
},
"5": {
"class_type": "VLAActionSafety",
"inputs": {
"actions": [
"4",
0
],
"embodiment": [
"2",
0
],
"mode": "Clamp safely",
"execution_horizon": 8,
"previous_action_json": "[0, 0, 0, 0, 0, 0, 0, 0]"
}
},
"6": {
"class_type": "VLATrajectoryPreview",
"inputs": {
"actions": [
"5",
0
],
"width": 960,
"height": 480,
"embodiment": [
"2",
0
]
}
},
"7": {
"class_type": "VLAActionInspect",
"inputs": {
"actions": [
"5",
0
],
"step_index": 0
}
},
"8": {
"class_type": "PreviewImage",
"inputs": {
"images": [
"6",
0
]
}
},
"9": {
"class_type": "ViewText",
"inputs": {
"text": [
"4",
1
]
}
},
"10": {
"class_type": "ViewText",
"inputs": {
"text": [
"5",
1
]
}
},
"11": {
"class_type": "ViewText",
"inputs": {
"text": [
"7",
0
]
}
}
}
+6 -1
View File
@@ -207,7 +207,12 @@ class OpenVocabularyDetector:
processor = transformers.AutoProcessor.from_pretrained(model_path)
model_class = transformers.AutoModelForZeroShotObjectDetection
dtype = torch_dtype(precision)
model = model_class.from_pretrained(model_path, dtype=dtype)
# Transformers 4.x consumes ``torch_dtype``; 5.x renamed it to
# ``dtype``. Passing the 5.x name to 4.x leaks into the model
# constructor and crashes Grounding DINO at runtime.
major = int(str(transformers.__version__).split(".", 1)[0])
dtype_kwargs = {"dtype": dtype} if major >= 5 else {"torch_dtype": dtype}
model = model_class.from_pretrained(model_path, **dtype_kwargs)
model.eval()
self.spec = spec
self.dtype = dtype
+478
View File
@@ -0,0 +1,478 @@
"""MiniMax music generation and cover support with fixed regional routing."""
from __future__ import annotations
import base64
import binascii
import io
import json
import os
from typing import Any
from urllib.parse import urlsplit
import numpy as np
import torch
from .audioldm2 import ANY
from .hosted_api import redact_sensitive
from .runtime import require_module
API_KEY_ENV = "MINIMAX_API_KEY"
REGION_ENDPOINTS = {
"global_en": "https://api.minimax.io/v1/music_generation",
"cn_zh": "https://api.minimaxi.com/v1/music_generation",
}
GENERATION_MODELS = (
"music-3.0",
"music-2.6",
"music-3.0-free",
"music-2.6-free",
)
COVER_MODELS = ("music-cover", "music-cover-free")
MUSIC_MODELS = GENERATION_MODELS + COVER_MODELS
DEFAULT_MODEL = "music-3.0"
REQUEST_FIELDS = frozenset(
{
"model",
"prompt",
"lyrics",
"stream",
"output_format",
"audio_setting",
"lyrics_optimizer",
"is_instrumental",
"audio_url",
"audio_base64",
"cover_feature_id",
}
)
OUTPUT_FORMATS = ("url", "hex")
STREAM_OUTPUT_FORMATS = ("hex",)
AUDIO_FORMATS = ("mp3", "wav", "pcm")
SAMPLE_RATES = (16000, 24000, 32000, 44100)
BITRATES = (32000, 64000, 128000, 256000)
REGIONAL_FIELDS = {"global_en": (), "cn_zh": ("aigc_watermark",)}
STATUS_IN_PROGRESS = 1
STATUS_COMPLETED = 2
MAX_COVER_BYTES = 50 * 1024 * 1024
MAX_AUDIO_BYTES = 128 * 1024 * 1024
def _clean_text(value: object) -> str:
return str(value or "").strip()
def _validate_cover_base64(value: str) -> None:
try:
decoded = base64.b64decode(value, validate=True)
except (binascii.Error, ValueError, TypeError):
raise ValueError("audio_base64 must contain valid base64 data.") from None
if len(decoded) > MAX_COVER_BYTES:
raise ValueError("audio_base64 exceeds the 50 MiB cover input limit.")
def build_music_request(
*,
region: str,
model: str,
prompt: str,
lyrics: str,
stream: bool,
output_format: str,
audio_format: str,
sample_rate: int,
bitrate: int,
lyrics_optimizer: bool,
is_instrumental: bool,
aigc_watermark: bool,
audio_url: str = "",
audio_base64: str = "",
cover_feature_id: str = "",
) -> dict[str, Any]:
"""Validate node inputs and build the documented JSON request body."""
if region not in REGION_ENDPOINTS:
raise ValueError(f"region must be one of {tuple(REGION_ENDPOINTS)}.")
if model not in MUSIC_MODELS:
raise ValueError(f"model must be one of {MUSIC_MODELS}.")
if output_format not in OUTPUT_FORMATS:
raise ValueError(f"output_format must be one of {OUTPUT_FORMATS}.")
if bool(stream) and output_format not in STREAM_OUTPUT_FORMATS:
raise ValueError("Streaming music responses require output_format='hex'.")
if audio_format not in AUDIO_FORMATS:
raise ValueError(f"audio_format must be one of {AUDIO_FORMATS}.")
if int(sample_rate) not in SAMPLE_RATES:
raise ValueError(f"sample_rate must be one of {SAMPLE_RATES}.")
if int(bitrate) not in BITRATES:
raise ValueError(f"bitrate must be one of {BITRATES}.")
clean_prompt = _clean_text(prompt)
clean_lyrics = _clean_text(lyrics)
clean_audio_url = _clean_text(audio_url)
clean_audio_base64 = _clean_text(audio_base64)
clean_cover_feature_id = _clean_text(cover_feature_id)
if len(clean_prompt) > 2000:
raise ValueError("prompt exceeds the 2,000-character music API limit.")
payload: dict[str, Any] = {
"model": model,
"stream": bool(stream),
"output_format": output_format,
"audio_setting": {
"sample_rate": int(sample_rate),
"bitrate": int(bitrate),
"format": audio_format,
},
}
if clean_prompt:
payload["prompt"] = clean_prompt
if clean_lyrics:
payload["lyrics"] = clean_lyrics
if model in COVER_MODELS:
if not 10 <= len(clean_prompt) <= 300:
raise ValueError("Cover generation requires a 10-300 character prompt.")
sources = (clean_audio_url, clean_audio_base64, clean_cover_feature_id)
if sum(bool(value) for value in sources) != 1:
raise ValueError(
"Cover generation requires exactly one of audio_url, "
"audio_base64, or cover_feature_id."
)
if clean_audio_base64:
_validate_cover_base64(clean_audio_base64)
payload["audio_base64"] = clean_audio_base64
elif clean_audio_url:
payload["audio_url"] = clean_audio_url
else:
if not 10 <= len(clean_lyrics) <= 1000:
raise ValueError(
"cover_feature_id requires lyrics between 10 and 1,000 characters."
)
payload["cover_feature_id"] = clean_cover_feature_id
if clean_lyrics and not 10 <= len(clean_lyrics) <= 1000:
raise ValueError("Cover lyrics must be between 10 and 1,000 characters.")
else:
if any((clean_audio_url, clean_audio_base64, clean_cover_feature_id)):
raise ValueError("Cover audio fields require a cover model.")
if len(clean_lyrics) > 3500:
raise ValueError("lyrics exceeds the 3,500-character music API limit.")
if bool(is_instrumental) and not clean_prompt:
raise ValueError("Instrumental generation requires a prompt.")
if not bool(is_instrumental) and not clean_lyrics and not bool(lyrics_optimizer):
raise ValueError(
"Non-instrumental generation requires lyrics or lyrics_optimizer."
)
payload["lyrics_optimizer"] = bool(lyrics_optimizer)
payload["is_instrumental"] = bool(is_instrumental)
if region == "cn_zh":
payload["aigc_watermark"] = bool(aigc_watermark)
return payload
def _response_parts(payload: object) -> tuple[str, int, dict[str, Any]]:
if not isinstance(payload, dict):
raise RuntimeError("MiniMax returned a non-object music response.")
base_response = payload.get("base_resp")
if not isinstance(base_response, dict):
raise RuntimeError("MiniMax returned no base_resp status.")
try:
success_code = int(base_response.get("status_code"))
except (TypeError, ValueError):
raise RuntimeError("MiniMax returned an invalid base_resp status code.") from None
if success_code != 0:
message = _clean_text(base_response.get("status_msg")) or "unknown API error"
raise RuntimeError(f"MiniMax music API error {success_code}: {message}")
data = payload.get("data")
if not isinstance(data, dict):
raise RuntimeError("MiniMax returned no music data object.")
try:
status = int(data.get("status"))
except (TypeError, ValueError):
raise RuntimeError("MiniMax returned an invalid music status.") from None
if status not in {STATUS_IN_PROGRESS, STATUS_COMPLETED}:
raise RuntimeError(f"MiniMax returned unsupported music status {status}.")
audio = data.get("audio", "")
if not isinstance(audio, str):
raise RuntimeError("MiniMax returned a non-string audio value.")
extra_info = payload.get("extra_info")
return audio.strip(), status, extra_info if isinstance(extra_info, dict) else {}
def _stream_audio(response: Any) -> tuple[str, dict[str, Any]]:
audio = ""
extra_info: dict[str, Any] = {}
completed = False
saw_payload = False
for line in response.iter_lines():
raw = line.decode("utf-8") if isinstance(line, bytes) else str(line)
raw = raw.strip()
if raw.startswith("data:"):
raw = raw[5:].strip()
if not raw or raw == "[DONE]" or raw.startswith("event:"):
continue
try:
payload = json.loads(raw)
except json.JSONDecodeError:
raise RuntimeError("MiniMax returned invalid streaming JSON.") from None
chunk, status, metadata = _response_parts(payload)
saw_payload = True
if chunk:
if chunk.startswith(audio):
audio = chunk
elif not audio.startswith(chunk):
audio += chunk
if metadata:
extra_info = metadata
completed = completed or status == STATUS_COMPLETED
if not saw_payload:
raise RuntimeError("MiniMax returned an empty streaming response.")
if not completed:
raise RuntimeError("MiniMax streaming ended before music generation completed.")
if not audio:
raise RuntimeError("MiniMax returned no audio data.")
return audio, extra_info
def _request_audio_value(
client: Any,
endpoint: str,
headers: dict[str, str],
payload: dict[str, Any],
) -> tuple[str, dict[str, Any]]:
if payload["stream"]:
with client.stream("POST", endpoint, headers=headers, json=payload) as response:
response.raise_for_status()
return _stream_audio(response)
response = client.post(endpoint, headers=headers, json=payload)
response.raise_for_status()
audio, status, extra_info = _response_parts(response.json())
if status != STATUS_COMPLETED:
raise RuntimeError(
"MiniMax music generation is still in progress and has no query endpoint."
)
if not audio:
raise RuntimeError("MiniMax returned no audio data.")
return audio, extra_info
def _download_audio(client: Any, url: str) -> bytes:
parsed = urlsplit(url)
if (
parsed.scheme != "https"
or not parsed.hostname
or parsed.username is not None
or parsed.password is not None
):
raise RuntimeError("MiniMax returned an invalid HTTPS audio URL.")
chunks: list[bytes] = []
total = 0
with client.stream("GET", url) as response:
response.raise_for_status()
for chunk in response.iter_bytes():
total += len(chunk)
if total > MAX_AUDIO_BYTES:
raise RuntimeError("MiniMax audio download exceeds 128 MiB.")
chunks.append(chunk)
return b"".join(chunks)
def _audio_bytes(client: Any, value: str, output_format: str) -> bytes:
if output_format == "url":
return _download_audio(client, value)
try:
return bytes.fromhex("".join(value.split()))
except ValueError:
raise RuntimeError("MiniMax returned invalid hexadecimal audio data.") from None
def _metadata_integer(metadata: dict[str, Any], name: str, default: int) -> int:
try:
value = int(metadata.get(name, default))
except (TypeError, ValueError):
return int(default)
return value if value > 0 else int(default)
def _decode_audio(
content: bytes,
audio_format: str,
requested_sample_rate: int,
metadata: dict[str, Any],
) -> tuple[np.ndarray, int]:
if not content:
raise RuntimeError("MiniMax returned an empty audio payload.")
if audio_format == "pcm":
if len(content) % 2:
raise RuntimeError("MiniMax returned an odd-length PCM payload.")
channels = _metadata_integer(metadata, "music_channel", 1)
raw = np.frombuffer(content, dtype="<i2")
if raw.size % channels:
raise RuntimeError("MiniMax PCM samples do not align with the channel count.")
samples = raw.astype(np.float32).reshape(-1, channels) / 32768.0
sample_rate = _metadata_integer(
metadata,
"music_sample_rate",
requested_sample_rate,
)
else:
soundfile = require_module("soundfile", "soundfile>=0.12")
try:
samples, sample_rate = soundfile.read(
io.BytesIO(content),
dtype="float32",
always_2d=True,
)
except Exception as exc:
detail = redact_sensitive(exc)
raise RuntimeError(f"Could not decode MiniMax {audio_format} audio: {detail}") from None
samples = np.asarray(samples, dtype=np.float32)
sample_rate = int(sample_rate)
if samples.ndim != 2 or not samples.size:
raise RuntimeError("MiniMax decoded to an empty audio array.")
if not np.isfinite(samples).all():
raise RuntimeError("MiniMax decoded audio contains non-finite samples.")
return np.ascontiguousarray(samples), int(sample_rate)
class MiniMaxMusicNode:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"region": (tuple(REGION_ENDPOINTS), {"default": "global_en"}),
"model": (MUSIC_MODELS, {"default": DEFAULT_MODEL}),
"prompt": ("STRING", {"default": "", "multiline": True}),
"lyrics": ("STRING", {"default": "", "multiline": True}),
"stream": ("BOOLEAN", {"default": False}),
"output_format": (OUTPUT_FORMATS, {"default": "hex"}),
"audio_format": (AUDIO_FORMATS, {"default": "mp3"}),
"sample_rate": (SAMPLE_RATES, {"default": 44100}),
"bitrate": (BITRATES, {"default": 256000}),
"lyrics_optimizer": ("BOOLEAN", {"default": False}),
"is_instrumental": ("BOOLEAN", {"default": False}),
"aigc_watermark": (
"BOOLEAN",
{
"default": False,
"tooltip": "Sent only to the cn_zh endpoint.",
},
),
},
"optional": {
"audio_url": ("STRING", {"default": ""}),
"audio_base64": ("STRING", {"default": "", "multiline": True}),
"cover_feature_id": ("STRING", {"default": ""}),
"timeout_seconds": (
"FLOAT",
{"default": 600.0, "min": 1.0, "max": 1800.0},
),
"use_system_proxy": ("BOOLEAN", {"default": False}),
},
}
RETURN_NAMES = ("wave_form", "sample_rate", "audio")
RETURN_TYPES = (ANY, "INT", "AUDIO")
OUTPUT_NODE = True
FUNCTION = "generate_music"
CATEGORY = "VLM Nodes/Audio"
DESCRIPTION = (
"Generate music or covers through fixed MiniMax regional endpoints. "
f"The API key is read only from {API_KEY_ENV}."
)
def generate_music(
self,
region,
model,
prompt,
lyrics,
stream,
output_format,
audio_format,
sample_rate,
bitrate,
lyrics_optimizer,
is_instrumental,
aigc_watermark,
audio_url="",
audio_base64="",
cover_feature_id="",
timeout_seconds=600.0,
use_system_proxy=False,
):
payload = build_music_request(
region=region,
model=model,
prompt=prompt,
lyrics=lyrics,
stream=stream,
output_format=output_format,
audio_format=audio_format,
sample_rate=sample_rate,
bitrate=bitrate,
lyrics_optimizer=lyrics_optimizer,
is_instrumental=is_instrumental,
aigc_watermark=aigc_watermark,
audio_url=audio_url,
audio_base64=audio_base64,
cover_feature_id=cover_feature_id,
)
api_key = os.getenv(API_KEY_ENV, "").strip()
if not api_key:
raise ValueError(
f"Set {API_KEY_ENV} in the environment that starts ComfyUI, "
"then restart the server."
)
httpx = require_module("httpx", "httpx>=0.27,<1")
client = httpx.Client(
timeout=max(1.0, min(1800.0, float(timeout_seconds))),
follow_redirects=False,
trust_env=bool(use_system_proxy),
)
headers = {
"Authorization": f"Bearer {api_key}",
"Content-Type": "application/json",
}
try:
value, metadata = _request_audio_value(
client,
REGION_ENDPOINTS[region],
headers,
payload,
)
content = _audio_bytes(client, value, output_format)
samples, actual_rate = _decode_audio(
content,
audio_format,
int(sample_rate),
metadata,
)
legacy = samples[:, 0] if samples.shape[1] == 1 else samples
audio = {
"waveform": torch.from_numpy(samples.T.copy()).unsqueeze(0),
"sample_rate": actual_rate,
}
return (legacy.tolist(), actual_rate, audio)
except Exception as exc:
detail = redact_sensitive(exc, (api_key,))
raise RuntimeError(f"MiniMax music request failed: {detail}") from None
finally:
try:
client.close()
except Exception:
pass
NODE_CLASS_MAPPINGS = {"MiniMaxMusicNode": MiniMaxMusicNode}
NODE_DISPLAY_NAME_MAPPINGS = {"MiniMaxMusicNode": "MiniMax Music"}
__all__ = [
"MiniMaxMusicNode",
"NODE_CLASS_MAPPINGS",
"NODE_DISPLAY_NAME_MAPPINGS",
]
+2457
View File
File diff suppressed because it is too large Load Diff
+10 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "comfyui_vlm_nodes"
version = "3.3.1"
version = "3.5.0"
description = "Production-ready local and API vision-language nodes for ComfyUI"
readme = "README.md"
requires-python = ">=3.10"
@@ -47,6 +47,11 @@ quantization = [
gguf = [
"llama-cpp-python>=0.3.20,<1",
]
robotics-client = [
"msgpack>=1.0.8,<2",
"pyzmq>=26,<28",
"websockets>=14,<17",
]
[project.urls]
Repository = "https://github.com/gokayfem/ComfyUI_VLM_nodes"
@@ -95,6 +100,7 @@ Icon = ""
packages = [
"comfyui_vlm_nodes",
"comfyui_vlm_nodes.examples",
"comfyui_vlm_nodes.examples.robotics",
"comfyui_vlm_nodes.examples.vision",
"comfyui_vlm_nodes.nodes",
"comfyui_vlm_nodes.nodes.joytagger",
@@ -111,6 +117,9 @@ comfyui_vlm_nodes = [
"*.json",
"SECURITY.md",
"examples/*.json",
"examples/robotics/*.py",
"examples/robotics/*.md",
"examples/robotics/*.json",
"examples/vision/*.json",
"requirements*.txt",
]
+5
View File
@@ -0,0 +1,5 @@
# Lightweight native clients only. Heavy VLA policy runtimes stay in a
# separate LeRobot, openpi, GR00T, OpenVLA/OFT, or Octo environment.
msgpack>=1.0.8,<2
pyzmq>=26,<28
websockets>=14,<17
+158
View File
@@ -0,0 +1,158 @@
"""Queue a real robotics policy workflow through ComfyUI's local API.
This is intentionally excluded from pytest: it requires a running ComfyUI
server, a running policy sidecar, a real image in ComfyUI's input directory,
and downloaded policy weights.
"""
from __future__ import annotations
import argparse
import json
import time
import urllib.request
import uuid
def _graph(image: str, policy_endpoint: str) -> dict:
camera_names = [
"observation.images.camera1",
"observation.images.camera2",
"observation.images.camera3",
]
return {
"1": {"class_type": "LoadImage", "inputs": {"image": image}},
"12": {
"class_type": "ImageScale",
"inputs": {
"image": ["1", 0],
"upscale_method": "lanczos",
"width": 256,
"height": 256,
"crop": "center",
},
},
"2": {
"class_type": "VLAEmbodimentProfile",
"inputs": {
"preset": "LeRobot SO-100 / SO-101 template",
"control_hz": 30.0,
"state_names_json": "",
"action_names_json": "",
"action_min_json": "",
"action_max_json": "",
"max_delta_json": "",
"camera_names_json": json.dumps(camera_names),
"action_mode_override": "",
},
},
"3": {
"class_type": "VLAObservationBuilder",
"inputs": {
"task": (
"Move the end effector toward the backpack and prepare to grasp it."
),
"state_json": "[0, 0, 0, 0, 0, 0]",
"primary_image": ["12", 0],
"primary_camera": camera_names[0],
"history_fps": 10.0,
"timestamp": 0.0,
"embodiment": ["2", 0],
"wrist_image": ["12", 0],
"wrist_camera": camera_names[1],
"secondary_image": ["12", 0],
"secondary_camera": camera_names[2],
},
},
"4": {
"class_type": "VLAHTTPPolicy",
"inputs": {
"observation": ["3", 0],
"endpoint": policy_endpoint,
"timeout_seconds": 120.0,
"include_history": True,
"allow_remote": False,
},
},
"5": {
"class_type": "VLAActionSafety",
"inputs": {
"actions": ["4", 0],
"embodiment": ["2", 0],
"mode": "Clamp safely",
"execution_horizon": 4,
"previous_action_json": "[0, 0, 0, 0, 0, 0]",
},
},
"6": {
"class_type": "VLATrajectoryPreview",
"inputs": {
"actions": ["5", 0],
"width": 960,
"height": 480,
"embodiment": ["2", 0],
},
},
"7": {
"class_type": "VLAActionInspect",
"inputs": {"actions": ["5", 0], "step_index": 0},
},
"8": {"class_type": "PreviewImage", "inputs": {"images": ["6", 0]}},
"9": {"class_type": "ViewText", "inputs": {"text": ["4", 1]}},
"10": {"class_type": "ViewText", "inputs": {"text": ["5", 1]}},
"11": {"class_type": "ViewText", "inputs": {"text": ["7", 0]}},
}
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--comfy-url", default="http://127.0.0.1:8188")
parser.add_argument("--policy-url", default="http://127.0.0.1:8787")
parser.add_argument(
"--image",
required=True,
help="Filename relative to the running ComfyUI instance's input directory.",
)
parser.add_argument("--timeout", type=float, default=180.0)
args = parser.parse_args()
base = args.comfy_url.rstrip("/")
body = json.dumps(
{
"prompt": _graph(args.image, args.policy_url),
"client_id": str(uuid.uuid4()),
}
).encode("utf-8")
request = urllib.request.Request(
f"{base}/prompt",
data=body,
headers={"Content-Type": "application/json"},
)
with urllib.request.urlopen(request, timeout=30) as response:
queued = json.load(response)
prompt_id = queued["prompt_id"]
deadline = time.monotonic() + args.timeout
while time.monotonic() < deadline:
with urllib.request.urlopen(
f"{base}/history/{prompt_id}",
timeout=10,
) as response:
history = json.load(response)
if prompt_id not in history:
time.sleep(0.5)
continue
entry = history[prompt_id]
result = {
"prompt_id": prompt_id,
"status": entry.get("status"),
"outputs": entry.get("outputs"),
}
print(json.dumps(result, indent=2))
if entry.get("status", {}).get("status_str") != "success":
raise SystemExit(1)
return
raise SystemExit(f"Timed out waiting for prompt {prompt_id}.")
if __name__ == "__main__":
main()
+291
View File
@@ -0,0 +1,291 @@
from types import SimpleNamespace
import numpy as np
import pytest
from ComfyUI_VLM_nodes.nodes import minimax_music
def generation_request(**overrides):
values = {
"region": "global_en",
"model": "music-3.0",
"prompt": "Reflective acoustic pop",
"lyrics": "[Verse]\nA quiet road under evening light",
"stream": False,
"output_format": "hex",
"audio_format": "mp3",
"sample_rate": 44100,
"bitrate": 256000,
"lyrics_optimizer": False,
"is_instrumental": False,
"aigc_watermark": False,
"audio_url": "",
"audio_base64": "",
"cover_feature_id": "",
}
values.update(overrides)
return minimax_music.build_music_request(**values)
def test_music_contract_matches_current_models_regions_and_formats():
assert minimax_music.REGION_ENDPOINTS == {
"global_en": "https://api.minimax.io/v1/music_generation",
"cn_zh": "https://api.minimaxi.com/v1/music_generation",
}
assert minimax_music.DEFAULT_MODEL == "music-3.0"
assert minimax_music.GENERATION_MODELS == (
"music-3.0",
"music-2.6",
"music-3.0-free",
"music-2.6-free",
)
assert minimax_music.COVER_MODELS == ("music-cover", "music-cover-free")
assert minimax_music.OUTPUT_FORMATS == ("url", "hex")
assert minimax_music.AUDIO_FORMATS == ("mp3", "wav", "pcm")
assert {
"model",
"prompt",
"lyrics",
"stream",
"output_format",
"audio_setting",
"lyrics_optimizer",
"is_instrumental",
"audio_url",
"audio_base64",
"cover_feature_id",
} == minimax_music.REQUEST_FIELDS
assert minimax_music.REGIONAL_FIELDS == {
"global_en": (),
"cn_zh": ("aigc_watermark",),
}
def test_generation_request_covers_generation_and_cn_fields():
request = generation_request(
region="cn_zh",
stream=True,
lyrics="",
lyrics_optimizer=True,
aigc_watermark=True,
audio_format="wav",
sample_rate=32000,
bitrate=128000,
)
assert request == {
"model": "music-3.0",
"prompt": "Reflective acoustic pop",
"stream": True,
"output_format": "hex",
"audio_setting": {
"sample_rate": 32000,
"bitrate": 128000,
"format": "wav",
},
"lyrics_optimizer": True,
"is_instrumental": False,
"aigc_watermark": True,
}
@pytest.mark.parametrize(
("source", "value"),
[
("audio_url", "https://media.example/reference.wav"),
("audio_base64", "dGVzdA=="),
("cover_feature_id", "feature-123"),
],
)
def test_cover_request_supports_each_documented_source(source, value):
overrides = {
"model": "music-cover",
"prompt": "Warm orchestral cover",
"lyrics": "Updated words for the cover",
source: value,
}
request = generation_request(**overrides)
assert request[source] == value
assert "lyrics_optimizer" not in request
assert "is_instrumental" not in request
def test_streaming_requires_hex_and_cover_sources_are_exclusive():
with pytest.raises(ValueError, match="output_format='hex'"):
generation_request(stream=True, output_format="url")
with pytest.raises(ValueError, match="exactly one"):
generation_request(
model="music-cover",
prompt="Warm orchestral cover",
audio_url="https://media.example/reference.wav",
audio_base64="dGVzdA==",
)
def test_stream_response_joins_hex_chunks_and_requires_completion():
class Response:
def iter_lines(self):
return iter(
[
'data: {"data":{"status":1,"audio":"0001"},'
'"base_resp":{"status_code":0}}',
'data: {"data":{"status":2,"audio":"0203"},'
'"extra_info":{"music_sample_rate":32000,"music_channel":2},'
'"base_resp":{"status_code":0}}',
"data: [DONE]",
]
)
audio, metadata = minimax_music._stream_audio(Response())
assert audio == "00010203"
assert metadata == {"music_sample_rate": 32000, "music_channel": 2}
def test_url_and_hex_response_decoding():
class DownloadResponse:
def __enter__(self):
return self
def __exit__(self, *_args):
return False
def raise_for_status(self):
return None
def iter_bytes(self):
return iter((b"ab", b"cd"))
class Client:
def stream(self, method, url):
assert method == "GET"
assert url == "https://media.example/music.wav"
return DownloadResponse()
client = Client()
assert minimax_music._audio_bytes(client, "61626364", "hex") == b"abcd"
assert (
minimax_music._audio_bytes(
client,
"https://media.example/music.wav",
"url",
)
== b"abcd"
)
def test_pcm_decoding_uses_response_sample_rate_and_channel_count():
content = np.array([0, 32767, -32768, 0], dtype="<i2").tobytes()
samples, sample_rate = minimax_music._decode_audio(
content,
"pcm",
44100,
{"music_sample_rate": 32000, "music_channel": 2},
)
assert samples.shape == (2, 2)
assert sample_rate == 32000
assert samples[0, 1] == pytest.approx(32767 / 32768)
def test_node_posts_to_fixed_region_and_returns_comfy_audio(monkeypatch):
captured = {}
class Response:
def raise_for_status(self):
return None
def json(self):
return {
"data": {"status": 2, "audio": "0102"},
"extra_info": {"music_sample_rate": 44100, "music_channel": 2},
"base_resp": {"status_code": 0},
}
class Client:
def __init__(self, **kwargs):
captured["client"] = kwargs
def post(self, endpoint, *, headers, json):
captured["endpoint"] = endpoint
captured["headers"] = headers
captured["request"] = json
return Response()
def close(self):
captured["closed"] = True
def fake_soundfile_read(buffer, **kwargs):
assert buffer.read() == b"\x01\x02"
assert kwargs == {"dtype": "float32", "always_2d": True}
return np.zeros((8, 2), dtype=np.float32), 44100
def fake_require_module(name, *_args):
if name == "httpx":
return SimpleNamespace(Client=Client)
if name == "soundfile":
return SimpleNamespace(read=fake_soundfile_read)
raise AssertionError(name)
monkeypatch.setenv(minimax_music.API_KEY_ENV, "test-key-not-for-production")
monkeypatch.setattr(minimax_music, "require_module", fake_require_module)
result = minimax_music.MiniMaxMusicNode().generate_music(
region="global_en",
model="music-3.0",
prompt="Reflective acoustic pop",
lyrics="[Verse]\nA quiet road under evening light",
stream=False,
output_format="hex",
audio_format="wav",
sample_rate=44100,
bitrate=256000,
lyrics_optimizer=False,
is_instrumental=False,
aigc_watermark=False,
)
assert captured["endpoint"] == minimax_music.REGION_ENDPOINTS["global_en"]
assert captured["headers"]["Authorization"].startswith("Bearer ")
assert captured["client"] == {
"timeout": 600.0,
"follow_redirects": False,
"trust_env": False,
}
assert captured["closed"] is True
assert len(result) == 3
assert result[1] == 44100
assert result[2]["waveform"].shape == (1, 2, 8)
def test_request_failures_redact_the_resolved_key(monkeypatch):
resolved_value = "unit-key"
class Client:
def __init__(self, **_kwargs):
pass
def post(self, *_args, **_kwargs):
raise RuntimeError(f"Authorization: Bearer {resolved_value}")
def close(self):
pass
monkeypatch.setenv(minimax_music.API_KEY_ENV, resolved_value)
monkeypatch.setattr(
minimax_music,
"require_module",
lambda *_args: SimpleNamespace(Client=Client),
)
with pytest.raises(RuntimeError) as captured:
minimax_music.MiniMaxMusicNode().generate_music(
region="global_en",
model="music-3.0",
prompt="Reflective acoustic pop",
lyrics="[Verse]\nA quiet road under evening light",
stream=False,
output_format="hex",
audio_format="mp3",
sample_rate=44100,
bitrate=256000,
lyrics_optimizer=False,
is_instrumental=False,
aigc_watermark=False,
)
assert resolved_value not in str(captured.value)
assert "[REDACTED]" in str(captured.value)
+753
View File
@@ -0,0 +1,753 @@
import json
import re
import threading
from http import HTTPStatus
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from types import SimpleNamespace
import numpy as np
import pytest
import torch
from ComfyUI_VLM_nodes.examples.robotics import lerobot_policy_server
from ComfyUI_VLM_nodes.nodes import robotics
from ComfyUI_VLM_nodes.nodes.robotics import (
EMBODIMENT_SCHEMA,
VLA_MODEL_CATALOG,
RobotActions,
actions_from_json,
actions_from_response,
blend_action_chunks,
build_observation,
make_embodiment,
observation_to_groot_payload,
observation_to_http_payload,
observation_to_openpi_payload,
render_action_preview,
validate_action_trajectory,
validate_policy_url,
validate_ws_url,
validate_zmq_host,
)
def _image(frames=1, height=24, width=32):
return torch.linspace(0, 1, frames * height * width * 3).reshape(
frames, height, width, 3
)
def _joint_profile():
return make_embodiment("Generic 7-DoF joint + gripper")
def _observation(profile=None):
profile = profile or _joint_profile()
return build_observation(
task="Pick up the blue cube.",
state_json=json.dumps([0.0] * profile.state_dim),
primary_image=_image(frames=3),
primary_camera=profile.camera_names[0],
history_fps=15,
timestamp=12.5,
embodiment=profile,
)
def test_embodiment_presets_are_explicit_roundtrippable_contracts():
for preset in robotics.EMBODIMENT_PRESETS:
profile = make_embodiment(preset)
encoded = profile.to_dict()
assert encoded["schema"] == EMBODIMENT_SCHEMA
assert len(encoded["action_names"]) == len(encoded["action_min"])
assert len(encoded["action_names"]) == len(encoded["action_max"])
assert len(encoded["action_names"]) == len(encoded["max_delta_per_step"])
assert robotics.RobotEmbodiment.from_dict(encoded) == profile
assert "Template limits" in profile.notes
def test_embodiment_rejects_invalid_bounds_and_mismatched_overrides():
with pytest.raises(ValueError, match="smaller"):
robotics.RobotEmbodiment(
name="bad",
state_names=("s",),
action_names=("a",),
action_min=(1,),
action_max=(0,),
max_delta=(0.1,),
control_hz=10,
action_mode="absolute",
camera_names=("front",),
)
with pytest.raises(ValueError, match="must contain 8"):
make_embodiment(
"Generic 7-DoF joint + gripper",
action_min=[-1],
)
def test_observation_preserves_history_and_validates_state_and_cameras():
profile = _joint_profile()
observation = _observation(profile)
summary = observation.summary()
assert summary["history_frames"] == 3
assert summary["cameras"][profile.camera_names[0]]["width"] == 32
assert summary["state"] == [0.0] * 8
assert summary["task"] == "Pick up the blue cube."
with pytest.raises(ValueError, match="State dimension"):
build_observation(
task="move",
state_json="[0]",
primary_image=_image(),
primary_camera=profile.camera_names[0],
history_fps=10,
timestamp=0,
embodiment=profile,
)
with pytest.raises(ValueError, match="not declared"):
build_observation(
task="move",
state_json=json.dumps([0] * profile.state_dim),
primary_image=_image(),
primary_camera="unknown",
history_fps=10,
timestamp=0,
embodiment=profile,
)
def test_http_payload_is_bounded_and_contains_no_tensor_details():
payload = observation_to_http_payload(_observation(), include_history=True)
assert payload["schema"] == "comfyui-vlm/robot-observation"
camera_frames = next(iter(payload["cameras"].values()))
assert len(camera_frames) == 3
assert all(item["encoding"] == "base64-jpeg" for item in camera_frames)
assert "device" not in json.dumps(payload)
def test_openpi_payload_supports_flat_and_aloha_shapes():
observation = _observation()
flat = observation_to_openpi_payload(
observation,
layout="Flat keys (DROID / LIBERO)",
state_key="observation/state",
prompt_key="prompt",
)
camera_key = observation.images[0][0]
assert flat[camera_key].shape == (24, 32, 3)
assert flat[camera_key].dtype == np.uint8
assert flat["observation/state"].dtype == np.float32
nested = observation_to_openpi_payload(
observation,
layout="Nested images (ALOHA)",
state_key="state",
prompt_key="prompt",
)
assert nested["images"][camera_key].shape == (3, 24, 32)
def test_openpi_array_codec_roundtrips_and_rejects_object_arrays():
source = np.arange(12, dtype=np.float32).reshape(3, 4)
encoded = robotics._openpi_pack_array(source)
decoded = robotics._openpi_unpack_array(encoded)
assert np.array_equal(decoded, source)
with pytest.raises(ValueError, match="does not support dtype"):
robotics._openpi_pack_array(np.array([object()], dtype=object))
forged = {
b"__ndarray__": True,
b"data": b"",
b"dtype": "|O",
b"shape": (0,),
}
with pytest.raises(ValueError, match="unsafe"):
robotics._openpi_unpack_array(forged)
def test_openpi_client_protocol_and_token_redaction(monkeypatch):
sent = []
class Connection:
responses = [b"metadata", b"actions"]
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def recv(self, timeout):
assert timeout == 5
return self.responses.pop(0)
def send(self, value):
sent.append(value)
connect_calls = []
def connect(uri, **kwargs):
connect_calls.append((uri, kwargs))
return Connection()
fake_ws = SimpleNamespace(connect=connect)
class FakeMsgpack:
@staticmethod
def packb(value, default):
assert value["prompt"] == "Pick up the blue cube."
assert callable(default)
return b"encoded-request"
@staticmethod
def unpackb(value, object_hook):
assert callable(object_hook)
if value == b"metadata":
return {"model": "pi-test"}
return {
"actions": np.zeros((2, 8), dtype=np.float32),
"server_timing": {"infer_ms": 12.5},
}
def fake_require(name, *_args):
return fake_ws if name == "websockets.sync.client" else FakeMsgpack
monkeypatch.setattr(robotics, "require_module", fake_require)
monkeypatch.setenv("OPENPI_API_KEY", "openpi-secret")
actions, report = robotics.call_openpi_policy(
_observation(),
endpoint="ws://127.0.0.1:8000",
timeout_seconds=5,
allow_remote=False,
layout="Flat keys (DROID / LIBERO)",
state_key="observation/state",
prompt_key="prompt",
)
assert actions.values.shape == (2, 8)
assert sent == [b"encoded-request"]
assert connect_calls[0][1]["additional_headers"] == {
"Authorization": "Api-Key openpi-secret"
}
assert report["server_metadata"] == {"model": "pi-test"}
assert "openpi-secret" not in json.dumps(report)
def broken_connect(*_args, **_kwargs):
raise RuntimeError("token=openpi-secret")
fake_ws.connect = broken_connect
with pytest.raises(RuntimeError, match=r"token=\[REDACTED\]") as exc:
robotics.call_openpi_policy(
_observation(),
endpoint="ws://127.0.0.1:8000",
timeout_seconds=5,
allow_remote=False,
layout="Flat keys (DROID / LIBERO)",
state_key="observation/state",
prompt_key="prompt",
)
assert "openpi-secret" not in str(exc.value)
def test_groot_payload_uses_official_nested_batch_time_contract():
observation = _observation()
payload = observation_to_groot_payload(observation)
camera = next(iter(payload["video"].values()))
assert camera.shape == (1, 3, 24, 32, 3)
assert camera.dtype == np.uint8
assert payload["state"]["state"].shape == (1, 3, 8)
assert payload["state"]["state"].dtype == np.float32
assert payload["language"]["task"] == [["Pick up the blue cube."]]
def test_groot_array_codec_and_native_client(monkeypatch):
array = np.arange(6, dtype=np.float32).reshape(2, 3)
encoded = robotics._groot_encode(array)
decoded = robotics._groot_decode(encoded)
assert np.array_equal(decoded, array)
with pytest.raises(TypeError, match="object/void"):
robotics._groot_encode(np.array([object()], dtype=object))
with pytest.raises(ValueError, match="object/void"):
robotics._groot_decode({"nd": True, "kind": "O", "type": "|O"})
sockets = []
class Socket:
def __init__(self):
self.options = []
self.connected = None
self.request = None
self.closed = False
def setsockopt(self, key, value):
self.options.append((key, value))
def connect(self, value):
self.connected = value
def send(self, value):
self.request = value
def recv(self):
return b"response"
def close(self, linger):
assert linger == 0
self.closed = True
class Context:
terminated = False
def socket(self, kind):
assert kind == 1
socket = Socket()
sockets.append(socket)
return socket
def term(self):
self.terminated = True
fake_zmq = SimpleNamespace(
REQ=1,
RCVTIMEO=2,
SNDTIMEO=3,
LINGER=4,
Context=Context,
)
packed_requests = []
class FakeMsgpack:
@staticmethod
def packb(value, default):
packed_requests.append(value)
assert callable(default)
return b"request"
@staticmethod
def unpackb(value, object_hook, raw):
assert value == b"response"
assert callable(object_hook)
assert raw is False
return [
{
"arm": np.zeros((1, 3, 7), dtype=np.float32),
"gripper": np.ones((1, 3, 1), dtype=np.float32),
},
{"server": "ok"},
]
def fake_require(name, *_args):
return fake_zmq if name == "zmq" else FakeMsgpack
monkeypatch.setattr(robotics, "require_module", fake_require)
monkeypatch.setenv("GROOT_API_TOKEN", "groot-secret")
actions, report = robotics.call_groot_policy(
_observation(),
host="127.0.0.1",
port=5555,
timeout_seconds=3,
allow_remote=False,
)
assert actions.values.shape == (3, 8)
assert actions.stream_slices == (("arm", 0, 7), ("gripper", 7, 8))
assert packed_requests[0]["api_token"] == "groot-secret"
assert sockets[0].connected == "tcp://127.0.0.1:5555"
assert sockets[0].closed is True
assert report["policy_info"] == {"server": "ok"}
assert "groot-secret" not in json.dumps(report)
def test_action_response_parses_arrays_and_named_streams():
single = actions_from_response(
{"actions": [[[1, 2], [3, 4]]]},
source="test",
)
assert single.values.shape == (2, 2)
assert single.stream_slices == (("actions", 0, 2),)
streams = actions_from_response(
{
"arm": np.zeros((1, 4, 7), dtype=np.float32),
"gripper": np.ones((1, 4, 1), dtype=np.float32),
"info": {"ignored": True},
},
source="groot",
)
assert streams.values.shape == (4, 8)
assert streams.stream_slices == (("arm", 0, 7), ("gripper", 7, 8))
def test_actions_json_roundtrip_and_chunk_replanning():
profile = _joint_profile()
previous = RobotActions(
torch.zeros((5, 8)),
profile.action_names,
"previous",
)
new = RobotActions(
torch.ones((6, 8)),
profile.action_names,
"new",
)
parsed = actions_from_json(json.dumps(new.to_dict()))
assert torch.equal(parsed.values, new.values)
assert parsed.action_names == new.action_names
replanned, report = blend_action_chunks(
previous,
new,
executed_steps=2,
transition_steps=2,
max_horizon=4,
)
assert replanned.values.shape == (4, 8)
assert torch.allclose(replanned.values[0], torch.full((8,), 1 / 3))
assert torch.allclose(replanned.values[1], torch.full((8,), 2 / 3))
assert torch.equal(replanned.values[2:], torch.ones((2, 8)))
assert report["transition_steps_applied"] == 2
with pytest.raises(ValueError, match="outside"):
blend_action_chunks(
previous,
new,
executed_steps=99,
transition_steps=2,
max_horizon=4,
)
def test_action_safety_clamps_bounds_deltas_and_horizon():
profile = _joint_profile()
values = torch.tensor(
[
[4.0, 0, 0, 0, 0, 0, 0, 2.0],
[-4.0, 0, 0, 0, 0, 0, 0, -2.0],
[0.0, 0, 0, 0, 0, 0, 0, 0.5],
]
)
actions = RobotActions(
values=values,
action_names=profile.action_names,
source="unit",
)
safe, report = validate_action_trajectory(
actions,
profile,
mode="Clamp safely",
execution_horizon=2,
previous_action_json=json.dumps([0.0] * 8),
)
assert safe.horizon == 2
assert torch.all(safe.values <= torch.tensor(profile.action_max))
assert torch.all(safe.values >= torch.tensor(profile.action_min))
limits = torch.tensor(profile.max_delta)
previous = torch.zeros(8)
for step in safe.values:
assert torch.all((step - previous).abs() <= limits + 1.0e-6)
previous = step
assert report["violations"]["total"] > 0
assert report["changed"] is True
assert report["safe_for_handoff"] is True
def test_action_safety_blocks_or_holds_nonfinite_actions():
profile = _joint_profile()
values = torch.zeros((2, 8))
values[0, 2] = float("nan")
actions = RobotActions(values, profile.action_names, "unit")
with pytest.raises(ValueError, match="blocked"):
validate_action_trajectory(
actions,
profile,
mode="Block unsafe",
execution_horizon=2,
)
held, report = validate_action_trajectory(
actions,
profile,
mode="Hold position on unsafe",
execution_horizon=2,
previous_action_json=json.dumps([0.25] * 8),
)
assert torch.allclose(held.values, torch.full((2, 8), 0.25))
assert report["safe_for_handoff"] is True
with pytest.raises(ValueError, match="requires previous_action_json"):
validate_action_trajectory(
actions,
profile,
mode="Hold position on unsafe",
execution_horizon=2,
)
def test_action_json_rejects_nonfinite_before_serialization():
actions = RobotActions(
torch.tensor([[float("inf")]]),
("action",),
"unit",
)
with pytest.raises(ValueError, match="NaN or infinity"):
actions.to_dict()
def test_policy_endpoint_security_defaults():
assert (
validate_policy_url("http://127.0.0.1:8787", allow_remote=False)
== "http://127.0.0.1:8787/v1/infer"
)
assert validate_policy_url(
"https://policy.example/v1/infer",
allow_remote=True,
) == "https://policy.example/v1/infer"
with pytest.raises(ValueError, match="HTTPS"):
validate_policy_url("http://policy.example", allow_remote=True)
with pytest.raises(ValueError, match="disabled"):
validate_policy_url("https://policy.example", allow_remote=False)
with pytest.raises(ValueError, match="embedded credentials"):
validate_policy_url("https://secret@example.com", allow_remote=True)
assert validate_ws_url("127.0.0.1:8000", allow_remote=False) == "ws://127.0.0.1:8000"
with pytest.raises(ValueError, match="WSS"):
validate_ws_url("ws://policy.example", allow_remote=True)
assert validate_zmq_host("localhost", allow_remote=False) == "localhost"
with pytest.raises(ValueError, match="disabled"):
validate_zmq_host("policy.example", allow_remote=False)
class _PolicyHandler(BaseHTTPRequestHandler):
observed_auth = None
observed_payload = None
def log_message(self, *_args):
pass
def do_POST(self): # noqa: N802
type(self).observed_auth = self.headers.get("Authorization")
length = int(self.headers["Content-Length"])
type(self).observed_payload = json.loads(self.rfile.read(length))
body = json.dumps(
{
"actions": [[0.0] * 8, [0.1] * 8],
"action_names": [f"a{index}" for index in range(8)],
}
).encode()
self.send_response(HTTPStatus.OK)
self.send_header("Content-Type", "application/json")
self.send_header("Content-Length", str(len(body)))
self.end_headers()
self.wfile.write(body)
def test_http_policy_real_loopback_request_uses_env_token_without_leaking(monkeypatch):
server = ThreadingHTTPServer(("127.0.0.1", 0), _PolicyHandler)
thread = threading.Thread(target=server.serve_forever, daemon=True)
thread.start()
monkeypatch.setenv("VLA_POLICY_TOKEN", "top-secret-value")
try:
actions, report = robotics.call_http_policy(
_observation(),
endpoint=f"http://127.0.0.1:{server.server_port}",
timeout_seconds=5,
allow_remote=False,
include_history=False,
)
finally:
server.shutdown()
server.server_close()
thread.join(timeout=5)
assert actions.values.shape == (2, 8)
assert _PolicyHandler.observed_auth == "Bearer top-secret-value"
assert _PolicyHandler.observed_payload["task"] == "Pick up the blue cube."
assert report["authenticated"] is True
assert "top-secret-value" not in json.dumps(report)
assert "top-secret-value" not in repr(actions)
def test_http_policy_error_redacts_server_token(monkeypatch):
class BrokenClient:
def __init__(self, **_kwargs):
pass
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def stream(self, *_args, **_kwargs):
raise RuntimeError("Authorization: Bearer secret-http-token")
monkeypatch.setenv("VLA_POLICY_TOKEN", "secret-http-token")
monkeypatch.setattr(
robotics,
"require_module",
lambda *_args: SimpleNamespace(Client=BrokenClient),
)
with pytest.raises(RuntimeError) as exc:
robotics.call_http_policy(
_observation(),
endpoint="http://127.0.0.1:8787",
timeout_seconds=5,
allow_remote=False,
include_history=False,
)
assert "secret-http-token" not in str(exc.value)
assert "[REDACTED]" in str(exc.value)
def test_http_policy_rejects_declared_oversized_response(monkeypatch):
class Response:
headers = {
"content-length": str(robotics.MAX_HTTP_RESPONSE_BYTES + 1),
}
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def raise_for_status(self):
pass
def iter_bytes(self):
raise AssertionError("Oversized response body must not be read.")
class Client:
def __init__(self, **_kwargs):
pass
def __enter__(self):
return self
def __exit__(self, *_args):
pass
def stream(self, *_args, **_kwargs):
return Response()
monkeypatch.setattr(
robotics,
"require_module",
lambda *_args: SimpleNamespace(Client=Client),
)
with pytest.raises(RuntimeError, match="32 MiB"):
robotics.call_http_policy(
_observation(),
endpoint="http://127.0.0.1:8787",
timeout_seconds=5,
allow_remote=False,
include_history=False,
)
def test_catalog_has_unique_official_entries_and_clear_readiness():
assert 12 <= len(VLA_MODEL_CATALOG) <= 30
checkpoints = []
for label, info in VLA_MODEL_CATALOG.items():
assert label == info.label
assert info.official_url.startswith("https://")
assert info.backend
assert info.status
if info.checkpoint:
checkpoints.append(info.checkpoint)
assert len(checkpoints) == len(set(checkpoints))
assert any(info.family == "SmolVLA" for info in VLA_MODEL_CATALOG.values())
assert any(info.family == "Isaac GR00T N1.7" for info in VLA_MODEL_CATALOG.values())
assert any(info.family == "OpenVLA-OFT" for info in VLA_MODEL_CATALOG.values())
def test_trajectory_preview_is_a_comfy_image():
profile = _joint_profile()
actions = RobotActions(
torch.linspace(-0.5, 0.5, 5 * 8).reshape(5, 8),
profile.action_names,
"preview",
)
preview = render_action_preview(actions, embodiment=profile, width=640, height=320)
assert preview.shape == (1, 320, 640, 3)
assert preview.dtype == torch.float32
assert 0 <= float(preview.min()) <= float(preview.max()) <= 1
def test_robotics_nodes_are_registered_and_have_safe_categories():
expected = {
"VLAEmbodimentProfile",
"VLAObservationBuilder",
"VLAHTTPPolicy",
"VLAOpenPIWebSocketPolicy",
"VLAGr00tZMQPolicy",
"VLAActionSafety",
"VLAActionsFromJSON",
"VLAActionChunkReplan",
"VLAActionInspect",
"VLATrajectoryPreview",
"VLAModelCatalog",
}
assert expected == set(robotics.NODE_CLASS_MAPPINGS)
for node in robotics.NODE_CLASS_MAPPINGS.values():
assert node.CATEGORY.startswith("VLM Nodes/Robotics")
assert "forceInput" not in repr(node.INPUT_TYPES())
def test_lerobot_sidecar_is_packaged_and_avoids_pickle_transport():
root = Path(robotics.__file__).parents[1]
server = root / "examples" / "robotics" / "lerobot_policy_server.py"
source = server.read_text(encoding="utf-8")
compile(source, str(server), "exec")
assert "pickle.loads" not in source
assert "VLA_POLICY_TOKEN" in source
assert "ThreadingHTTPServer" in source
def test_lerobot_sidecar_exposes_checkpoint_feature_contract(monkeypatch):
visual = SimpleNamespace(
type=SimpleNamespace(value="VISUAL"),
shape=(3, 256, 256),
)
action = {"type": "ACTION", "shape": (6,)}
assert lerobot_policy_server._feature_metadata(
{"observation.images.camera1": visual, "action": action}
) == {
"observation.images.camera1": {
"type": "VISUAL",
"shape": [3, 256, 256],
},
"action": {"type": "ACTION", "shape": [6]},
}
assert lerobot_policy_server._optional_config_int(
SimpleNamespace(chunk_size=50), "chunk_size"
) == 50
monkeypatch.setattr(
lerobot_policy_server,
"_decode_image",
lambda _frame: np.zeros((1, 1, 3), dtype=np.uint8),
)
oversized_task = {
"schema": "comfyui-vlm/robot-observation",
"version": 1,
"cameras": {
"observation.images.front": [
{"encoding": "base64-jpeg", "data": "unused"}
]
},
"state": [0.0],
"task": "x" * (lerobot_policy_server.MAX_TASK_CHARS + 1),
}
with pytest.raises(ValueError, match="task must contain"):
lerobot_policy_server._decode_observation(oversized_task)
def test_robotics_client_requirements_match_optional_extra():
root = Path(robotics.__file__).parents[1]
pyproject = (root / "pyproject.toml").read_text(encoding="utf-8")
match = re.search(r"^robotics-client\s*=\s*\[(.*?)^\]", pyproject, re.M | re.S)
assert match is not None
optional_extra = set(re.findall(r'"([^"]+)"', match.group(1)))
requirement_file = {
line.strip()
for line in (root / "requirements-robotics-client.txt")
.read_text(encoding="utf-8")
.splitlines()
if line.strip() and not line.lstrip().startswith("#")
}
assert optional_extra == requirement_file