diff --git a/editor-server/server.py b/editor-server/server.py index 901463c..21465d4 100644 --- a/editor-server/server.py +++ b/editor-server/server.py @@ -179,6 +179,11 @@ class DatasetTrimReq(BaseModel): frame_index: int side: str +class DatasetReplayEventReq(BaseModel): + token: str + frame_index: int + event: str + class UpdatePortsReq(BaseModel): inputs: list[str] | None = None outputs: list[str] | None = None @@ -1250,6 +1255,27 @@ def control_node(node_id: str, req: NodeControlReq): meta = _session.node_meta.get(node_id) if meta is None: raise HTTPException(404, "Node not found") + if meta.get("type") == "TrajectorySmoother": + if req.action != "apply": + raise HTTPException(400, "TrajectorySmoother supports the apply control") + control_fn = _runtime_callable("dataset", _RUNTIME_MODULES["dataset"], "apply_configured_smoother") + if control_fn is None: + raise HTTPException(503, "blacknode-dataset smoother runtime is not loaded") + params = dict(meta.get("params") or {}) + try: + outputs = dict(control_fn( + node_id, + str(params.get("method") or "spline"), + float(params.get("strength") if params.get("strength") is not None else 1.0), + preview_source=str(params.get("preview_source") or "action"), + preview_joint=str(params.get("preview_joint") or ""), + )) + except ValueError as exc: + raise HTTPException(409, str(exc)) from exc + for port, value in outputs.items(): + _session.graph._cache[(node_id, port)] = value + _session.graph._dirty.discard(node_id) + return {"ok": True, "node_id": node_id, "outputs": outputs} if meta.get("type") != "EpisodeRecorder": raise HTTPException(400, "This node does not expose direct controls") control_fn = _runtime_callable("dataset", _RUNTIME_MODULES["dataset"], "control_configured_recorder") @@ -2412,6 +2438,18 @@ def dataset_trim(req: DatasetTrimReq): raise HTTPException(500, str(exc)) from exc +@app.post("/api/dataset/replay-event") +@app.post("/dataset/replay-event") +def dataset_replay_event(req: DatasetReplayEventReq): + publish_fn = _runtime_callable("dataset", _RUNTIME_MODULES["dataset"], "publish_replay_event") + if publish_fn is None: + raise HTTPException(503, "blacknode-dataset replay stream runtime is not loaded") + try: + return dict(publish_fn(req.token, req.frame_index, req.event)) + except ValueError as exc: + raise HTTPException(409, str(exc)) from exc + + @app.post("/runtime/stop") def stop_runtime(): _stop_active_cook() diff --git a/editor/src/api.ts b/editor/src/api.ts index af5449b..e76defe 100644 --- a/editor/src/api.ts +++ b/editor/src/api.ts @@ -352,6 +352,10 @@ export const api = { req>('GET', `/dataset/frame/${encodeURIComponent(token)}?index=${Math.max(0, Math.floor(index))}`), trimDatasetEpisode:(token: string, frameIndex: number, side: 'before' | 'after') => req>('POST', '/dataset/trim', { token, frame_index: Math.max(0, Math.floor(frameIndex)), side }, 120000), + publishDatasetReplayFrame:(token: string, frameIndex: number, event: 'play' | 'seek') => + req>('POST', '/dataset/replay-event', { + token, frame_index: Math.max(0, Math.floor(frameIndex)), event, + }), updatePorts:(id: string, patch: Partial>) => req('PATCH', `/nodes/${id}/ports`, patch), updatePortVisibility:(id: string, patch: Pick) => diff --git a/editor/src/components/DatasetBrowserPanel.tsx b/editor/src/components/DatasetBrowserPanel.tsx index faa9d4e..10298d8 100644 --- a/editor/src/components/DatasetBrowserPanel.tsx +++ b/editor/src/components/DatasetBrowserPanel.tsx @@ -3,6 +3,10 @@ import { api } from '../api' import { useStore } from '../store' type AnyRecord = Record +type FrameCallbackVideo = HTMLVideoElement & { + requestVideoFrameCallback?: (callback: (now: number, metadata: { mediaTime: number }) => void) => number + cancelVideoFrameCallback?: (id: number) => void +} const panel: React.CSSProperties = { margin: '7px 9px 3px', padding: 10, borderRadius: 8, @@ -41,6 +45,7 @@ export default function DatasetBrowserPanel({ id, data }: { const [trimPending, setTrimPending] = useState<'before' | 'after' | null>(null) const [trimMessage, setTrimMessage] = useState('') const lastFrame = useRef(-1) + const lastPublishedFrame = useRef(-1) const videoRef = useRef(null) const catalog = data.portResults?.catalog && typeof data.portResults.catalog === 'object' @@ -62,6 +67,7 @@ export default function DatasetBrowserPanel({ id, data }: { setFrame(null) setPlaying(false) lastFrame.current = -1 + lastPublishedFrame.current = -1 }, [token]) const refresh = async (patch: Record = {}) => { @@ -94,6 +100,40 @@ export default function DatasetBrowserPanel({ id, data }: { } } + const publishReplayPosition = (time: number, event: 'play' | 'seek', force = false) => { + if (!token || !fps || totalFrames <= 0) return + const index = Math.min(totalFrames - 1, Math.max(0, Math.floor(time * fps))) + if (!force && index === lastPublishedFrame.current) return + lastPublishedFrame.current = index + void api.publishDatasetReplayFrame(token, index, event).catch(() => { + // A publisher is optional; local dataset replay remains usable by itself. + }) + } + + useEffect(() => { + const player = videoRef.current as FrameCallbackVideo | null + if (!playing || !player || !token || !fps || totalFrames <= 0) return + let cancelled = false + let callbackId: number | null = null + let timerId: number | null = null + const publishFrame = (_now?: number, metadata?: { mediaTime: number }) => { + if (cancelled || player.paused) return + publishReplayPosition(metadata?.mediaTime ?? player.currentTime, 'play') + if (player.requestVideoFrameCallback) callbackId = player.requestVideoFrameCallback(publishFrame) + } + if (player.requestVideoFrameCallback) { + callbackId = player.requestVideoFrameCallback(publishFrame) + } else { + timerId = window.setInterval(() => publishReplayPosition(player.currentTime, 'play'), + Math.max(16, 1000 / fps)) + } + return () => { + cancelled = true + if (callbackId !== null && player.cancelVideoFrameCallback) player.cancelVideoFrameCallback(callbackId) + if (timerId !== null) window.clearInterval(timerId) + } + }, [playing, token, fps, totalFrames]) + const jointNames = frame && Array.isArray(frame.joint_names) ? frame.joint_names.map(String) : [] const leader = (frame?.leader ?? {}) as AnyRecord const observation = (frame?.observation ?? {}) as AnyRecord @@ -251,7 +291,14 @@ export default function DatasetBrowserPanel({ id, data }: { onPlay={() => setPlaying(true)} onPause={() => setPlaying(false)} onEnded={() => setPlaying(false)} onLoadedMetadata={event => void updateReplayFrame(event.currentTarget.currentTime)} onTimeUpdate={event => void updateReplayFrame(event.currentTarget.currentTime)} - onSeeked={event => void updateReplayFrame(event.currentTarget.currentTime)} + onSeeking={event => { + void updateReplayFrame(event.currentTarget.currentTime) + publishReplayPosition(event.currentTarget.currentTime, 'seek') + }} + onSeeked={event => { + void updateReplayFrame(event.currentTarget.currentTime) + publishReplayPosition(event.currentTarget.currentTime, 'seek', true) + }} style={{ width: '100%', maxHeight: 470, background: '#020617', borderRadius: 7, objectFit: 'contain' }} />
diff --git a/editor/src/store.ts b/editor/src/store.ts index 3762868..a34c25b 100644 --- a/editor/src/store.ts +++ b/editor/src/store.ts @@ -2777,6 +2777,20 @@ export const useStore = create((set, get) => ({ ...markActiveTabDirty(s), })) await api.updateParam(id, key, value) + if (node.data.type === 'TrajectorySmoother' + && ['method', 'strength', 'preview_source', 'preview_joint'].includes(key)) { + try { + await get().controlNode(id, 'apply') + } catch (error) { + window.dispatchEvent(new CustomEvent('blacknode:notice', { + detail: { + kind: 'error', + title: 'Smoother update needs an input', + message: error instanceof Error ? error.message : String(error), + }, + })) + } + } if (profileChanged) { window.dispatchEvent(new CustomEvent('blacknode:notice', { detail: { diff --git a/tests/test_editor_media_output.py b/tests/test_editor_media_output.py index e9cc859..2ae515d 100644 --- a/tests/test_editor_media_output.py +++ b/tests/test_editor_media_output.py @@ -89,3 +89,13 @@ def test_dataset_replay_switches_units_and_keeps_canvas_wheel_zoom() -> None: assert "✂ Cut after" in browser assert "The selected frame is kept" in browser assert "trimDatasetEpisode" in api + assert "publishDatasetReplayFrame" in api + assert "requestVideoFrameCallback" in browser + assert "onSeeking=" in browser + + +def test_smoother_parameter_updates_use_direct_control_instead_of_graph_cook() -> None: + store = (ROOT / "editor" / "src" / "store.ts").read_text(encoding="utf-8") + + assert "node.data.type === 'TrajectorySmoother'" in store + assert "await get().controlNode(id, 'apply')" in store diff --git a/tests/test_editor_runtime.py b/tests/test_editor_runtime.py index d0a2a6f..a1025e8 100644 --- a/tests/test_editor_runtime.py +++ b/tests/test_editor_runtime.py @@ -195,6 +195,37 @@ def test_episode_recorder_control_does_not_cook_graph(self): finally: server._session.node_meta.pop("recorder-control-test", None) + def test_trajectory_smoother_control_recomputes_only_smoother(self): + node_id = "smoother-control-test" + server._session.node_meta[node_id] = { + "id": node_id, "type": "TrajectorySmoother", + "params": {"method": "gaussian", "strength": 2.5, + "preview_source": "leader", "preview_joint": "elbow"}, + } + server._session.graph._dirty.add(node_id) + apply = lambda node_id, method, strength, **preview: { + "stream": {"token": "smoothed"}, "preview": "image", + "report": f"{node_id}:{method}:{strength}:{preview['preview_source']}:{preview['preview_joint']}", + } + try: + with ( + patch.object(server, "_runtime_callable", return_value=apply), + patch.object(server, "_prepare_cook") as prepare_cook, + ): + response = TestClient(server.app).post( + f"/nodes/{node_id}/control", json={"action": "apply"}, + ) + self.assertEqual(response.status_code, 200) + self.assertIn("gaussian:2.5:leader:elbow", response.json()["outputs"]["report"]) + self.assertEqual(server._session.graph._cache[(node_id, "stream")], {"token": "smoothed"}) + self.assertNotIn(node_id, server._session.graph._dirty) + prepare_cook.assert_not_called() + finally: + server._session.node_meta.pop(node_id, None) + server._session.graph._dirty.discard(node_id) + for key in [key for key in server._session.graph._cache if key[0] == node_id]: + server._session.graph._cache.pop(key, None) + def test_dataset_media_endpoint_serves_only_runtime_registered_video(self): with tempfile.TemporaryDirectory() as tmp: video = Path(tmp) / "episode.mp4" @@ -255,6 +286,20 @@ def trim(token, frame_index, side): self.assertEqual(response.json()["side"], "before") self.assertEqual(expired.status_code, 409) + def test_dataset_replay_event_endpoint_forwards_browser_playback(self): + publish = lambda token, index, event: { + "ok": True, "token": token, "frame_index": index, "event": event, + "publishers": 1, "subscribers": 1, + } + with patch.object(server, "_runtime_callable", return_value=publish): + response = TestClient(server.app).post( + "/dataset/replay-event", + json={"token": "episode", "frame_index": 12, "event": "seek"}, + ) + self.assertEqual(response.status_code, 200) + self.assertEqual(response.json()["frame_index"], 12) + self.assertEqual(response.json()["event"], "seek") + if __name__ == "__main__": unittest.main()