Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
38 changes: 38 additions & 0 deletions editor-server/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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()
Expand Down
4 changes: 4 additions & 0 deletions editor/src/api.ts
Original file line number Diff line number Diff line change
Expand Up @@ -352,6 +352,10 @@ export const api = {
req<Record<string, unknown>>('GET', `/dataset/frame/${encodeURIComponent(token)}?index=${Math.max(0, Math.floor(index))}`),
trimDatasetEpisode:(token: string, frameIndex: number, side: 'before' | 'after') =>
req<Record<string, unknown>>('POST', '/dataset/trim', { token, frame_index: Math.max(0, Math.floor(frameIndex)), side }, 120000),
publishDatasetReplayFrame:(token: string, frameIndex: number, event: 'play' | 'seek') =>
req<Record<string, unknown>>('POST', '/dataset/replay-event', {
token, frame_index: Math.max(0, Math.floor(frameIndex)), event,
}),
updatePorts:(id: string, patch: Partial<Pick<BnNodeMeta, 'inputs' | 'outputs' | 'input_types' | 'output_types' | 'input_defaults' | 'multi_input_ports'>>) =>
req<BnNodeMeta>('PATCH', `/nodes/${id}/ports`, patch),
updatePortVisibility:(id: string, patch: Pick<BnNodeMeta, 'promoted_inputs' | 'promoted_outputs'>) =>
Expand Down
49 changes: 48 additions & 1 deletion editor/src/components/DatasetBrowserPanel.tsx
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,10 @@ import { api } from '../api'
import { useStore } from '../store'

type AnyRecord = Record<string, any>
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,
Expand Down Expand Up @@ -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<HTMLVideoElement | null>(null)

const catalog = data.portResults?.catalog && typeof data.portResults.catalog === 'object'
Expand All @@ -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<string, unknown> = {}) => {
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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' }} />
<div style={{ minWidth: 0 }}>
<div style={{ color: 'var(--tx2)', fontFamily: 'var(--font-mono)', fontSize: 10, lineHeight: 1.55 }}>
Expand Down
14 changes: 14 additions & 0 deletions editor/src/store.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2777,6 +2777,20 @@ export const useStore = create<Store>((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: {
Expand Down
10 changes: 10 additions & 0 deletions tests/test_editor_media_output.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
45 changes: 45 additions & 0 deletions tests/test_editor_runtime.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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()