diff --git a/flashdreams/flashdreams/serving/realtime/__init__.py b/flashdreams/flashdreams/serving/realtime/__init__.py
new file mode 100644
index 000000000..ae762506d
--- /dev/null
+++ b/flashdreams/flashdreams/serving/realtime/__init__.py
@@ -0,0 +1,4 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+"""Transport-neutral realtime serving primitives."""
diff --git a/flashdreams/flashdreams/serving/realtime/frame_bus.py b/flashdreams/flashdreams/serving/realtime/frame_bus.py
new file mode 100644
index 000000000..56b94289e
--- /dev/null
+++ b/flashdreams/flashdreams/serving/realtime/frame_bus.py
@@ -0,0 +1,98 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+"""Single-slot latest-frame publishing for realtime transports."""
+
+from __future__ import annotations
+
+import time
+from dataclasses import dataclass
+from threading import Condition
+from typing import Generic, TypeVar
+
+FrameT = TypeVar("FrameT")
+
+
+@dataclass(frozen=True, slots=True)
+class PublishedFrame(Generic[FrameT]):
+ """Frame payload plus its monotonically increasing publication count."""
+
+ payload: FrameT
+ count: int
+
+
+class LatestFrameBus(Generic[FrameT]):
+ """Thread-safe single-slot bus for latest-frame transports.
+
+ The bus stores only the newest frame. Waiters can block until a frame
+ newer than ``last_seen_count`` is published, or until ``close`` wakes them.
+ """
+
+ def __init__(self) -> None:
+ self._condition = Condition()
+ self._latest: FrameT | None = None
+ self._frame_count = 0
+ self._closed = False
+
+ @property
+ def closed(self) -> bool:
+ with self._condition:
+ return self._closed
+
+ @property
+ def frame_count(self) -> int:
+ with self._condition:
+ return self._frame_count
+
+ def publish(self, frame: FrameT) -> int:
+ """Publish ``frame`` and return its publication count."""
+ with self._condition:
+ if self._closed:
+ raise RuntimeError("Cannot publish to a closed LatestFrameBus.")
+ self._latest = frame
+ self._frame_count += 1
+ self._condition.notify_all()
+ return self._frame_count
+
+ def latest(self) -> PublishedFrame[FrameT] | None:
+ """Return the latest frame without blocking, if one has been published."""
+ with self._condition:
+ if self._latest is None:
+ return None
+ return PublishedFrame(payload=self._latest, count=self._frame_count)
+
+ def wait_for_frame(
+ self,
+ *,
+ last_seen_count: int = 0,
+ timeout_s: float | None = None,
+ ) -> PublishedFrame[FrameT] | None:
+ """Block until a newer frame is available or the bus closes.
+
+ Returns ``None`` on timeout or when closed before a newer frame is
+ published.
+ """
+ if last_seen_count < 0:
+ raise ValueError("last_seen_count must be >= 0")
+ if timeout_s is not None and timeout_s < 0:
+ raise ValueError("timeout_s must be >= 0")
+
+ deadline = None if timeout_s is None else time.monotonic() + timeout_s
+ with self._condition:
+ while self._latest is None or self._frame_count <= last_seen_count:
+ if self._closed:
+ return None
+ if deadline is None:
+ self._condition.wait()
+ continue
+ remaining_s = deadline - time.monotonic()
+ if remaining_s <= 0:
+ return None
+ self._condition.wait(timeout=remaining_s)
+ return PublishedFrame(payload=self._latest, count=self._frame_count)
+
+ def close(self) -> None:
+ """Wake all waiters and reject future publications."""
+ with self._condition:
+ self._closed = True
+ self._condition.notify_all()
diff --git a/flashdreams/flashdreams/serving/realtime/input.py b/flashdreams/flashdreams/serving/realtime/input.py
new file mode 100644
index 000000000..1b488bb61
--- /dev/null
+++ b/flashdreams/flashdreams/serving/realtime/input.py
@@ -0,0 +1,385 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+"""Input state and sparse-control helpers for realtime serving."""
+
+from __future__ import annotations
+
+from collections import deque
+from dataclasses import dataclass, field
+from typing import Literal
+
+import numpy as np
+
+DEFAULT_SUPPORTED_KEYS = frozenset({"w", "a", "s", "d", "q", "e", "i", "k", "j", "l"})
+DRIVING_SUPPORTED_KEYS = frozenset(
+ {"w", "a", "s", "d", "up", "down", "left", "right", "space"}
+)
+WSAD_SUPPORTED_KEYS = frozenset({"w", "a", "s", "d"})
+KEY_ALIASES = {
+ "arrowup": "w",
+ "arrowleft": "a",
+ "arrowdown": "s",
+ "arrowright": "d",
+}
+
+
+@dataclass(frozen=True, slots=True)
+class ResetRequest:
+ """Transport-neutral request to reset the realtime rollout."""
+
+ reason: str | None = None
+ request_id: str | None = None
+
+
+@dataclass(frozen=True, slots=True)
+class PromptRequest:
+ """Transport-neutral prompt update request."""
+
+ prompt: str
+ negative_prompt: str | None = None
+ request_id: str | None = None
+
+
+@dataclass(frozen=True, slots=True)
+class ImageRequest:
+ """Transport-neutral image update request."""
+
+ data: bytes
+ content_type: str
+ request_id: str | None = None
+
+
+@dataclass(frozen=True, slots=True)
+class SparseInputSnapshot:
+ """Sparse input state sampled at a realtime loop boundary."""
+
+ timestamp_s: float
+ pressed_keys: frozenset[str] = field(default_factory=frozenset)
+ effective_keys: frozenset[str] = field(default_factory=frozenset)
+ reset: ResetRequest | None = None
+ prompt: PromptRequest | None = None
+ image: ImageRequest | None = None
+
+
+def normalize_key(key: str) -> str:
+ normalized = key.strip().lower()
+ return KEY_ALIASES.get(normalized, normalized)
+
+
+@dataclass(slots=True)
+class KeyboardState:
+ pressed_keys: set[str] = field(default_factory=set)
+ supported_keys: frozenset[str] = DEFAULT_SUPPORTED_KEYS
+ _press_order: dict[str, int] = field(default_factory=dict)
+ _press_counter: int = 0
+
+ def apply_event(self, *, event: str, key: str) -> bool:
+ normalized_key = normalize_key(key)
+ if normalized_key not in self.supported_keys:
+ return False
+
+ normalized_event = event.strip().lower()
+ if normalized_event == "keydown":
+ self.pressed_keys.add(normalized_key)
+ self._press_counter += 1
+ self._press_order[normalized_key] = self._press_counter
+ return True
+ if normalized_event == "keyup":
+ self.pressed_keys.discard(normalized_key)
+ self._press_order.pop(normalized_key, None)
+ return True
+ return False
+
+ def snapshot(self) -> frozenset[str]:
+ return frozenset(self.pressed_keys)
+
+ def sparse_snapshot(self, *, timestamp_s: float) -> SparseInputSnapshot:
+ return SparseInputSnapshot(
+ timestamp_s=timestamp_s,
+ pressed_keys=self.snapshot(),
+ effective_keys=self.resolved_effective_keys(),
+ )
+
+ def _latest_pressed(self, keys: tuple[str, ...]) -> str | None:
+ latest_key: str | None = None
+ latest_idx = -1
+ for key in keys:
+ if key not in self.pressed_keys:
+ continue
+ idx = self._press_order.get(key, -1)
+ if idx >= latest_idx:
+ latest_idx = idx
+ latest_key = key
+ return latest_key
+
+ def resolved_effective_keys(self) -> frozenset[str]:
+ effective: set[str] = set()
+ for key in (
+ self._latest_pressed(("w", "s")),
+ self._latest_pressed(("a", "d", "j", "l")),
+ self._latest_pressed(("q", "e")),
+ self._latest_pressed(("i", "k")),
+ ):
+ if key is not None:
+ effective.add(key)
+ return frozenset(key for key in effective if key in self.supported_keys)
+
+
+PoseSegment = tuple[float, float, frozenset[str]]
+
+
+class KeyboardResampler:
+ """Resample sparse keydown/keyup edges into a chunk timeline."""
+
+ def __init__(
+ self,
+ *,
+ fps: int,
+ start_v: float = 0.0,
+ supported_keys: frozenset[str] = DEFAULT_SUPPORTED_KEYS,
+ ) -> None:
+ if fps <= 0:
+ raise ValueError("fps must be > 0")
+ self._fps = fps
+ self._dt = 1.0 / fps
+ self._supported_keys = supported_keys
+ self.next_chunk_start_v = start_v
+ self._event_log: deque[tuple[float, dict[str, str]]] = deque()
+ self._carried_state = KeyboardState(supported_keys=supported_keys)
+
+ @property
+ def fps(self) -> int:
+ return self._fps
+
+ @property
+ def dt(self) -> float:
+ return self._dt
+
+ def on_edge(self, *, arrival_t: float, event: str, key: str) -> None:
+ self._event_log.append((arrival_t, {"event": event, "key": key}))
+
+ def sample_chunk(self, num_frames: int) -> tuple[list[PoseSegment], list[float]]:
+ if num_frames < 1:
+ raise ValueError("num_frames must be >= 1")
+
+ chunk_start_v = self.next_chunk_start_v
+ chunk_end_v = chunk_start_v + num_frames * self._dt
+
+ while self._event_log and self._event_log[0][0] < chunk_start_v:
+ _, payload = self._event_log.popleft()
+ self._carried_state.apply_event(**payload)
+
+ segments: list[PoseSegment] = []
+ prev_t = chunk_start_v
+ prev_state = self._carried_state.resolved_effective_keys()
+ while self._event_log and self._event_log[0][0] <= chunk_end_v:
+ event_t, payload = self._event_log.popleft()
+ if event_t > prev_t:
+ segments.append((prev_t, event_t, prev_state))
+ self._carried_state.apply_event(**payload)
+ prev_state = self._carried_state.resolved_effective_keys()
+ prev_t = event_t
+ if prev_t < chunk_end_v:
+ segments.append((prev_t, chunk_end_v, prev_state))
+ elif not segments:
+ segments.append((chunk_start_v, chunk_end_v, prev_state))
+
+ frame_times = [chunk_start_v + (i + 1) * self._dt for i in range(num_frames)]
+ self.next_chunk_start_v = chunk_end_v
+ return segments, frame_times
+
+ def reset(self, *, start_v: float) -> None:
+ self._event_log.clear()
+ self._carried_state = KeyboardState(supported_keys=self._supported_keys)
+ self.next_chunk_start_v = start_v
+
+ def event_log_size(self) -> int:
+ return len(self._event_log)
+
+
+def _rotation_matrix(axis: str, angle_rad: float) -> np.ndarray:
+ cos_t = np.float32(np.cos(angle_rad))
+ sin_t = np.float32(np.sin(angle_rad))
+ if axis == "x":
+ return np.array(
+ [
+ [1.0, 0.0, 0.0],
+ [0.0, cos_t, -sin_t],
+ [0.0, sin_t, cos_t],
+ ],
+ dtype=np.float32,
+ )
+ if axis == "y":
+ return np.array(
+ [
+ [cos_t, 0.0, sin_t],
+ [0.0, 1.0, 0.0],
+ [-sin_t, 0.0, cos_t],
+ ],
+ dtype=np.float32,
+ )
+ if axis == "z":
+ return np.array(
+ [
+ [cos_t, -sin_t, 0.0],
+ [sin_t, cos_t, 0.0],
+ [0.0, 0.0, 1.0],
+ ],
+ dtype=np.float32,
+ )
+ return np.eye(3, dtype=np.float32)
+
+
+@dataclass(slots=True)
+class CameraPoseIntegrator:
+ """Integrate a piecewise-constant keyboard timeline into a camera trajectory."""
+
+ move_speed_per_s: float = 0.8
+ rotate_speed_rad_per_s: float = float(np.deg2rad(32.0))
+ pitch_limit_rad: float = float(np.deg2rad(85.0))
+ coordinate_system: Literal["RDF", "FLU"] = "RDF"
+ _current_pose: np.ndarray = field(
+ default_factory=lambda: np.eye(4, dtype=np.float32),
+ )
+ _current_pitch: float = 0.0
+
+ def __post_init__(self) -> None:
+ if self.coordinate_system not in {"RDF", "FLU"}:
+ raise ValueError(
+ "coordinate_system must be 'RDF' (right-down-forward) "
+ "or 'FLU' (forward-left-up)"
+ )
+
+ def reset(self, pose: np.ndarray | None = None) -> None:
+ if pose is None:
+ self._current_pose = np.eye(4, dtype=np.float32)
+ self._current_pitch = 0.0
+ return
+ if pose.shape != (4, 4):
+ raise ValueError(f"Expected pose shape (4, 4), got {pose.shape}")
+ self._current_pose = pose.astype(np.float32, copy=True)
+ if self.coordinate_system == "FLU":
+ self._current_pitch = float(np.arcsin(np.clip(pose[2, 0], -1.0, 1.0)))
+ else:
+ self._current_pitch = float(np.arctan2(pose[2, 1], pose[1, 1]))
+
+ def current_pose(self) -> np.ndarray:
+ return self._current_pose.copy()
+
+ def _advance(self, *, state: frozenset[str], duration: float) -> None:
+ if duration <= 0:
+ return
+
+ yaw_rate = 0.0
+ if self.coordinate_system == "FLU":
+ if "a" in state or "j" in state:
+ yaw_rate += self.rotate_speed_rad_per_s
+ if "d" in state or "l" in state:
+ yaw_rate -= self.rotate_speed_rad_per_s
+ else:
+ if "a" in state or "j" in state:
+ yaw_rate -= self.rotate_speed_rad_per_s
+ if "d" in state or "l" in state:
+ yaw_rate += self.rotate_speed_rad_per_s
+ pitch_rate = 0.0
+ if "i" in state:
+ pitch_rate += self.rotate_speed_rad_per_s
+ if "k" in state:
+ pitch_rate -= self.rotate_speed_rad_per_s
+
+ yaw_delta = yaw_rate * duration
+ pitch_delta = pitch_rate * duration
+
+ new_pitch = self._current_pitch + pitch_delta
+ if -self.pitch_limit_rad <= new_pitch <= self.pitch_limit_rad:
+ self._current_pitch = new_pitch
+ else:
+ pitch_delta = 0.0
+
+ rot = self._current_pose[:3, :3]
+ trans = self._current_pose[:3, 3]
+ if self.coordinate_system == "FLU":
+ rot_pitch = _rotation_matrix("y", -pitch_delta)
+ rot_yaw = _rotation_matrix("z", yaw_delta)
+ else:
+ rot_pitch = _rotation_matrix("x", pitch_delta)
+ rot_yaw = _rotation_matrix("y", yaw_delta)
+ rot_new = rot_yaw @ rot @ rot_pitch
+
+ forward_rate = 0.0
+ if "w" in state:
+ forward_rate += self.move_speed_per_s
+ if "s" in state:
+ forward_rate -= self.move_speed_per_s
+ right_rate = 0.0
+ if "e" in state:
+ right_rate += self.move_speed_per_s
+ if "q" in state:
+ right_rate -= self.move_speed_per_s
+
+ if self.coordinate_system == "FLU":
+ vec_forward = rot_new[:, 0]
+ vec_right = -rot_new[:, 1]
+ forward_flat = np.array(
+ [vec_forward[0], vec_forward[1], 0.0], dtype=np.float32
+ )
+ right_flat = np.array([vec_right[0], vec_right[1], 0.0], dtype=np.float32)
+ else:
+ vec_right = rot_new[:, 0]
+ vec_forward = rot_new[:, 2]
+ forward_flat = np.array(
+ [vec_forward[0], 0.0, vec_forward[2]], dtype=np.float32
+ )
+ right_flat = np.array([vec_right[0], 0.0, vec_right[2]], dtype=np.float32)
+ forward_norm = np.linalg.norm(forward_flat)
+ right_norm = np.linalg.norm(right_flat)
+ if forward_norm > 0:
+ forward_flat /= forward_norm
+ if right_norm > 0:
+ right_flat /= right_norm
+
+ move_vec = forward_flat * (forward_rate * duration) + right_flat * (
+ right_rate * duration
+ )
+ self._current_pose = np.eye(4, dtype=np.float32)
+ self._current_pose[:3, :3] = rot_new
+ self._current_pose[:3, 3] = trans + move_vec
+
+ def integrate_chunk(
+ self,
+ *,
+ segments: list[PoseSegment],
+ frame_times: list[float],
+ ) -> np.ndarray:
+ if not segments:
+ raise ValueError("segments must be non-empty")
+ if not frame_times:
+ raise ValueError("frame_times must be non-empty")
+ chunk_start = segments[0][0]
+ chunk_end = segments[-1][1]
+ if any(
+ frame_times[i] >= frame_times[i + 1] for i in range(len(frame_times) - 1)
+ ):
+ raise ValueError("frame_times must be strictly increasing")
+ if frame_times[0] < chunk_start - 1e-9 or frame_times[-1] > chunk_end + 1e-9:
+ raise ValueError(
+ "frame_times must lie within the chunk window "
+ f"[{chunk_start}, {chunk_end}]"
+ )
+
+ poses: list[np.ndarray] = []
+ cur_t = chunk_start
+ ft_idx = 0
+ for _, seg_end, seg_state in segments:
+ while ft_idx < len(frame_times) and frame_times[ft_idx] <= seg_end:
+ target_t = frame_times[ft_idx]
+ self._advance(state=seg_state, duration=target_t - cur_t)
+ cur_t = target_t
+ poses.append(self._current_pose.copy())
+ ft_idx += 1
+ if seg_end > cur_t:
+ self._advance(state=seg_state, duration=seg_end - cur_t)
+ cur_t = seg_end
+
+ return np.stack(poses, axis=0).astype(np.float32)
diff --git a/flashdreams/flashdreams/serving/realtime/media.py b/flashdreams/flashdreams/serving/realtime/media.py
new file mode 100644
index 000000000..266e7ac52
--- /dev/null
+++ b/flashdreams/flashdreams/serving/realtime/media.py
@@ -0,0 +1,189 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+"""Transport-neutral realtime media conversion helpers."""
+
+from __future__ import annotations
+
+import importlib
+import io
+from typing import TYPE_CHECKING, Literal
+
+import numpy as np
+
+if TYPE_CHECKING:
+ import torch
+
+FrameLayout = Literal["hwc", "chw", "thwc", "tchw", "bvtchw"]
+ValueRange = Literal["minus_one_one", "zero_one", "uint8"]
+
+
+def _as_numpy(value: object, *, sync_device: bool) -> np.ndarray:
+ if isinstance(value, np.ndarray):
+ return value
+
+ try:
+ import torch
+ except ModuleNotFoundError as exc:
+ raise TypeError("Expected a numpy array or torch tensor.") from exc
+
+ if not isinstance(value, torch.Tensor):
+ raise TypeError("Expected a numpy array or torch tensor.")
+
+ tensor = value.detach()
+ if tensor.device.type != "cpu":
+ if not sync_device:
+ raise ValueError(
+ "Converting a non-CPU tensor requires sync_device=True so the "
+ "host/device synchronization is explicit at the call site."
+ )
+ tensor = tensor.cpu()
+ if tensor.is_floating_point():
+ tensor = tensor.float()
+ return tensor.numpy()
+
+
+def _scale_rgb(array: np.ndarray, *, value_range: ValueRange) -> np.ndarray:
+ if value_range not in ("minus_one_one", "zero_one", "uint8"):
+ raise ValueError(f"Unsupported value_range={value_range!r}.")
+ if array.dtype == np.uint8:
+ if value_range != "uint8":
+ raise ValueError("uint8 inputs require value_range='uint8'.")
+ return np.ascontiguousarray(array)
+
+ values = array.astype(np.float32, copy=False)
+ if value_range == "minus_one_one":
+ values = (values + 1.0) / 2.0 * 255.0
+ elif value_range == "zero_one":
+ values = values * 255.0
+ return np.ascontiguousarray(values.clip(0, 255).astype(np.uint8))
+
+
+def rgb_frame_to_uint8(
+ frame: object,
+ *,
+ layout: Literal["hwc", "chw"] = "hwc",
+ value_range: ValueRange = "uint8",
+ sync_device: bool = True,
+) -> np.ndarray:
+ """Convert one RGB frame to contiguous ``HWC`` uint8 host memory."""
+ array = _as_numpy(frame, sync_device=sync_device)
+ if layout == "hwc":
+ if array.ndim != 3 or array.shape[-1] != 3:
+ raise ValueError(
+ f"Expected HWC RGB frame with shape [H, W, 3], got {array.shape}"
+ )
+ return _scale_rgb(array, value_range=value_range)
+ if layout == "chw":
+ if array.ndim != 3 or array.shape[0] != 3:
+ raise ValueError(
+ f"Expected CHW RGB frame with shape [3, H, W], got {array.shape}"
+ )
+ return _scale_rgb(np.transpose(array, (1, 2, 0)), value_range=value_range)
+ raise ValueError(f"Unsupported layout={layout!r}.")
+
+
+def rgb_array_to_uint8_frames(
+ data: object,
+ *,
+ layout: FrameLayout,
+ value_range: ValueRange = "minus_one_one",
+ sync_device: bool = True,
+) -> list[np.ndarray]:
+ """Convert a tensor/array video chunk to ``HWC`` uint8 RGB frames."""
+ array = _as_numpy(data, sync_device=sync_device)
+ if layout == "hwc" or layout == "chw":
+ return [
+ rgb_frame_to_uint8(
+ array,
+ layout=layout,
+ value_range=value_range,
+ sync_device=sync_device,
+ )
+ ]
+ if layout == "thwc":
+ if array.ndim != 4 or array.shape[-1] != 3:
+ raise ValueError(
+ f"Expected THWC RGB chunk with shape [T, H, W, 3], got {array.shape}"
+ )
+ frames = array
+ elif layout == "tchw":
+ if array.ndim != 4 or array.shape[1] != 3:
+ raise ValueError(
+ f"Expected TCHW RGB chunk with shape [T, 3, H, W], got {array.shape}"
+ )
+ frames = np.transpose(array, (0, 2, 3, 1))
+ elif layout == "bvtchw":
+ if (
+ array.ndim != 6
+ or array.shape[0] != 1
+ or array.shape[1] != 1
+ or array.shape[3] != 3
+ ):
+ raise ValueError(
+ "Expected single-batch single-view video chunk "
+ f"[1, 1, T, 3, H, W], got {array.shape}"
+ )
+ frames = np.transpose(array[0, 0], (0, 2, 3, 1))
+ else:
+ raise ValueError(f"Unsupported layout={layout!r}.")
+
+ return [_scale_rgb(frame, value_range=value_range) for frame in frames]
+
+
+def tensor_chunk_to_rgb_frames(
+ video_chunk: torch.Tensor,
+ *,
+ sync_device: bool = True,
+) -> list[np.ndarray]:
+ """Convert common model output tensor layouts to RGB uint8 frames."""
+ value_range: ValueRange = (
+ "minus_one_one" if video_chunk.is_floating_point() else "uint8"
+ )
+ if video_chunk.ndim == 4:
+ return rgb_array_to_uint8_frames(
+ video_chunk,
+ layout="tchw",
+ value_range=value_range,
+ sync_device=sync_device,
+ )
+ if video_chunk.ndim == 6:
+ return rgb_array_to_uint8_frames(
+ video_chunk,
+ layout="bvtchw",
+ value_range=value_range,
+ sync_device=sync_device,
+ )
+ raise ValueError(
+ "Expected video chunk [T, C, H, W] or [1, 1, T, 3, H, W], "
+ f"got {tuple(video_chunk.shape)}"
+ )
+
+
+def encode_rgb_frame_to_jpeg(
+ frame: object,
+ *,
+ quality: int = 85,
+ layout: Literal["hwc", "chw"] = "hwc",
+ value_range: ValueRange = "uint8",
+ sync_device: bool = True,
+) -> bytes:
+ """JPEG-encode one RGB frame.
+
+ Pillow is imported lazily so importing realtime serving modules does not
+ force image-codec dependencies into non-JPEG transports.
+ """
+ if not 1 <= quality <= 100:
+ raise ValueError("quality must be between 1 and 100")
+
+ image_module = importlib.import_module("PIL.Image")
+ image_from_array = getattr(image_module, "fromarray")
+ rgb_uint8 = rgb_frame_to_uint8(
+ frame,
+ layout=layout,
+ value_range=value_range,
+ sync_device=sync_device,
+ )
+ output = io.BytesIO()
+ image_from_array(rgb_uint8).save(output, format="JPEG", quality=quality)
+ return output.getvalue()
diff --git a/flashdreams/flashdreams/serving/realtime/mjpeg.py b/flashdreams/flashdreams/serving/realtime/mjpeg.py
new file mode 100644
index 000000000..b504581cc
--- /dev/null
+++ b/flashdreams/flashdreams/serving/realtime/mjpeg.py
@@ -0,0 +1,115 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+"""Transport-neutral MJPEG streaming helpers."""
+
+from __future__ import annotations
+
+from collections.abc import Callable
+from http import HTTPStatus
+from typing import Protocol
+
+from flashdreams.serving.realtime.frame_bus import LatestFrameBus
+
+DEFAULT_MJPEG_BOUNDARY = "frame"
+WaitForMjpegFrame = Callable[[int], tuple[bytes, int] | None]
+
+
+class StopEvent(Protocol):
+ def is_set(self) -> bool: ...
+
+
+class MjpegResponse(Protocol):
+ def send_response(self, code: int | HTTPStatus) -> None: ...
+
+ def send_header(self, keyword: str, value: str) -> None: ...
+
+ def end_headers(self) -> None: ...
+
+
+class MjpegWriter(Protocol):
+ def write(self, data: bytes, /) -> object: ...
+
+ def flush(self) -> object: ...
+
+
+def mjpeg_content_type(boundary: str = DEFAULT_MJPEG_BOUNDARY) -> str:
+ return f"multipart/x-mixed-replace; boundary={boundary}"
+
+
+def format_mjpeg_part(jpeg: bytes, *, boundary: str = DEFAULT_MJPEG_BOUNDARY) -> bytes:
+ """Format one JPEG payload as a multipart MJPEG response part."""
+ return (
+ (
+ f"--{boundary}\r\n"
+ "Content-Type: image/jpeg\r\n"
+ f"Content-Length: {len(jpeg)}\r\n\r\n"
+ ).encode("ascii")
+ + jpeg
+ + b"\r\n"
+ )
+
+
+def send_mjpeg_response_headers(
+ response: MjpegResponse, *, boundary: str = DEFAULT_MJPEG_BOUNDARY
+) -> None:
+ """Send generic no-cache multipart MJPEG response headers."""
+ response.send_response(HTTPStatus.OK)
+ response.send_header(
+ "Cache-Control", "no-store, no-cache, must-revalidate, max-age=0"
+ )
+ response.send_header("Pragma", "no-cache")
+ response.send_header("Content-Type", mjpeg_content_type(boundary))
+ response.end_headers()
+
+
+def write_mjpeg_stream(
+ writer: MjpegWriter,
+ wait_for_frame: WaitForMjpegFrame,
+ *,
+ boundary: str = DEFAULT_MJPEG_BOUNDARY,
+) -> None:
+ """Write MJPEG parts until ``wait_for_frame`` returns ``None``."""
+ last_seen = 0
+ try:
+ while True:
+ result = wait_for_frame(last_seen)
+ if result is None:
+ return
+ jpeg, last_seen = result
+ writer.write(format_mjpeg_part(jpeg, boundary=boundary))
+ writer.flush()
+ except (BrokenPipeError, ConnectionResetError):
+ return
+
+
+def publish_latest_jpeg(
+ bus: LatestFrameBus[bytes], jpeg: bytes, *, stop_event: StopEvent
+) -> None:
+ """Publish ``jpeg`` unless shutdown is already in progress."""
+ if stop_event.is_set():
+ return
+ try:
+ bus.publish(jpeg)
+ except RuntimeError:
+ if not stop_event.is_set():
+ raise
+
+
+def wait_for_latest_jpeg(
+ bus: LatestFrameBus[bytes],
+ *,
+ last_seen_count: int,
+ stop_event: StopEvent,
+ poll_timeout_s: float = 1.0,
+) -> tuple[bytes, int] | None:
+ """Wait for a JPEG newer than ``last_seen_count`` or return ``None`` on close."""
+ while not stop_event.is_set():
+ frame = bus.wait_for_frame(
+ last_seen_count=last_seen_count, timeout_s=poll_timeout_s
+ )
+ if frame is not None:
+ return frame.payload, frame.count
+ if bus.closed:
+ return None
+ return None
diff --git a/flashdreams/flashdreams/serving/realtime/presenter.py b/flashdreams/flashdreams/serving/realtime/presenter.py
new file mode 100644
index 000000000..088cb8485
--- /dev/null
+++ b/flashdreams/flashdreams/serving/realtime/presenter.py
@@ -0,0 +1,181 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+"""Transport-neutral realtime presentation queue and pacing helpers."""
+
+from __future__ import annotations
+
+import queue
+import time
+from collections import deque
+from collections.abc import Callable, Iterator
+from dataclasses import dataclass
+from typing import Generic, Protocol, TypeVar
+
+import numpy as np
+
+from flashdreams.serving.realtime.media import rgb_frame_to_uint8
+
+FrameT = TypeVar("FrameT")
+
+
+class SupportsGetNowait(Protocol[FrameT]):
+ def get_nowait(self) -> FrameT: ...
+
+
+@dataclass(frozen=True, slots=True)
+class QueueDrainResult(Generic[FrameT]):
+ """Summary of frames drained from a producer queue into a presentation queue."""
+
+ pulled: int
+ accepted: int
+ skipped: int
+ dropped: tuple[FrameT, ...]
+
+
+@dataclass(frozen=True, slots=True)
+class PresentationWait:
+ """Sleep interval consumed while waiting for the next presentation slot."""
+
+ begin_time: float
+ end_time: float
+
+ @property
+ def duration_s(self) -> float:
+ return self.end_time - self.begin_time
+
+
+class PresentationQueue(Generic[FrameT]):
+ """Small FIFO queue for frames ready to present.
+
+ ``capacity`` bounds queued frames by dropping the oldest ready frame when a
+ newer frame arrives. That policy keeps realtime outputs moving forward
+ under backpressure instead of building unbounded latency.
+ """
+
+ def __init__(self, capacity: int | None = None) -> None:
+ if capacity is not None and capacity <= 0:
+ raise ValueError("capacity must be greater than 0")
+ self._capacity = capacity
+ self._frames: deque[FrameT] = deque()
+
+ @property
+ def capacity(self) -> int | None:
+ return self._capacity
+
+ def __len__(self) -> int:
+ return len(self._frames)
+
+ def __bool__(self) -> bool:
+ return bool(self._frames)
+
+ def __iter__(self) -> Iterator[FrameT]:
+ return iter(self._frames)
+
+ def append(self, frame: FrameT) -> FrameT | None:
+ """Append ``frame`` and return any frame dropped by capacity pressure."""
+ dropped = None
+ if self._capacity is not None and len(self._frames) >= self._capacity:
+ dropped = self._frames.popleft()
+ self._frames.append(frame)
+ return dropped
+
+ def pop_ready(self) -> FrameT | None:
+ """Return the oldest ready frame, or ``None`` if no frame is ready."""
+ if not self._frames:
+ return None
+ return self._frames.popleft()
+
+ def popleft(self) -> FrameT:
+ """Return the oldest ready frame, raising ``IndexError`` if empty."""
+ return self._frames.popleft()
+
+ def clear(self) -> int:
+ """Drop all queued frames and return the number removed."""
+ count = len(self._frames)
+ self._frames.clear()
+ return count
+
+ def drain_nowait(
+ self,
+ source: SupportsGetNowait[FrameT],
+ *,
+ include: Callable[[FrameT], bool] | None = None,
+ prepare: Callable[[FrameT], None] | None = None,
+ ) -> QueueDrainResult[FrameT]:
+ """Drain currently available frames from ``source`` without blocking."""
+ pulled = 0
+ accepted = 0
+ skipped = 0
+ accepted_frames: list[FrameT] = []
+ while True:
+ try:
+ frame = source.get_nowait()
+ except queue.Empty:
+ break
+ pulled += 1
+ if include is not None and not include(frame):
+ skipped += 1
+ continue
+ accepted_frames.append(frame)
+ accepted += 1
+
+ if self._capacity is None:
+ retained_entries = [(False, frame) for frame in self._frames] + [
+ (True, frame) for frame in accepted_frames
+ ]
+ dropped: list[FrameT] = []
+ else:
+ entries = [(False, frame) for frame in self._frames] + [
+ (True, frame) for frame in accepted_frames
+ ]
+ overflow = max(0, len(entries) - self._capacity)
+ dropped = [frame for _, frame in entries[:overflow]]
+ retained_entries = entries[overflow:]
+
+ if prepare is not None:
+ for is_new, frame in retained_entries:
+ if is_new:
+ prepare(frame)
+ self._frames.clear()
+ self._frames.extend(frame for _, frame in retained_entries)
+
+ return QueueDrainResult(
+ pulled=pulled,
+ accepted=accepted,
+ skipped=skipped,
+ dropped=tuple(dropped),
+ )
+
+
+def wait_until_present_time(
+ present_time: float,
+ *,
+ poll_timeout_s: float,
+ clock: Callable[[], float] = time.perf_counter,
+ sleep: Callable[[float], None] = time.sleep,
+) -> PresentationWait | None:
+ """Sleep until ``present_time`` is due, capped by ``poll_timeout_s``.
+
+ Returns the consumed sleep interval for tracing, or ``None`` when the
+ target presentation time is already due.
+ """
+ if poll_timeout_s < 0:
+ raise ValueError("poll_timeout_s must be >= 0")
+ now = clock()
+ if now >= present_time:
+ return None
+ begin_time = now
+ sleep(min(poll_timeout_s, max(0.0, present_time - now)))
+ return PresentationWait(begin_time=begin_time, end_time=clock())
+
+
+def materialize_rgb_host_uint8(frame: object) -> np.ndarray:
+ """Materialize a frame-like object to contiguous ``(H, W, 3)`` uint8 RGB."""
+ to_numpy = getattr(frame, "to_numpy", None)
+ if callable(to_numpy):
+ frame = to_numpy()
+ array = np.asarray(frame)
+ if array.ndim == 3 and array.shape[-1] > 3:
+ array = array[..., :3]
+ return rgb_frame_to_uint8(array, value_range="uint8")
diff --git a/flashdreams/flashdreams/serving/realtime/timing.py b/flashdreams/flashdreams/serving/realtime/timing.py
new file mode 100644
index 000000000..9891cfcdc
--- /dev/null
+++ b/flashdreams/flashdreams/serving/realtime/timing.py
@@ -0,0 +1,546 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+"""Transport-neutral realtime timing records and trace helpers."""
+
+from __future__ import annotations
+
+from collections import defaultdict, deque
+from collections.abc import Iterable, Iterator, Mapping
+from dataclasses import dataclass
+from threading import Lock
+from typing import Protocol
+
+TraceComponentValue = str | int | float | bool | None
+
+
+def trace_time_ns(seconds: float) -> int:
+ return int(seconds * 1_000_000_000)
+
+
+def event_dependencies(*events: int | None) -> list[int]:
+ return [event for event in events if event is not None]
+
+
+def _duration_ms(start_time: float | None, end_time: float | None) -> float | None:
+ if start_time is None or end_time is None:
+ return None
+ return (end_time - start_time) * 1000.0
+
+
+@dataclass
+class FrameTimes:
+ frame_index: int
+ intended_present_time: float
+ image_ready_time: float | None = None
+ sample_display_pose_time: float | None = None
+ present_time: float | None = None
+
+ def input_to_present_ms(self, *, input_sample_time: float) -> float | None:
+ return _duration_ms(input_sample_time, self.present_time)
+
+ def present_jitter_ms(self) -> float | None:
+ return _duration_ms(self.intended_present_time, self.present_time)
+
+
+@dataclass(frozen=True)
+class ChunkPrediction:
+ """Predicted timestamps for a chunk's pipeline stages."""
+
+ first_present: float
+
+ @classmethod
+ def create(
+ cls, *, request_time: float, frame_interval_s: float
+ ) -> "ChunkPrediction":
+ return cls(first_present=request_time + frame_interval_s)
+
+
+@dataclass
+class ChunkTimes:
+ """Mutable timing record that travels with one realtime chunk."""
+
+ chunk_index: int
+ input_sample_time: float
+ request_time: float
+ request_poses_ready_time: float
+ frames: list[FrameTimes]
+ prediction: ChunkPrediction | None = None
+ chunk_render_start_time: float | None = None
+ chunk_ready_time: float | None = None
+
+ @classmethod
+ def create(
+ cls,
+ chunk_index: int,
+ input_sample_time: float,
+ request_time: float,
+ request_poses_ready_time: float,
+ intended_present_times: list[float],
+ prediction: ChunkPrediction | None = None,
+ ) -> "ChunkTimes":
+ frames = [
+ FrameTimes(frame_index=index, intended_present_time=time_value)
+ for index, time_value in enumerate(intended_present_times)
+ ]
+ return cls(
+ chunk_index=chunk_index,
+ input_sample_time=input_sample_time,
+ request_time=request_time,
+ request_poses_ready_time=request_poses_ready_time,
+ frames=frames,
+ prediction=prediction,
+ )
+
+ def stage_durations_ms(self) -> dict[str, float]:
+ """Return available milestone-derived durations for this chunk."""
+ return chunk_stage_durations_ms(self)
+
+
+@dataclass(frozen=True)
+class VideoModelTimings:
+ """Observable backend timing milestones for one video-model chunk."""
+
+ condition_start_time: float
+ condition_ready_time: float
+ model_start_time: float
+ model_ready_time: float
+ merge_start_time: float
+ merge_ready_time: float
+ decode_start_time: float | None = None
+ decode_ready_time: float | None = None
+ cache_update_start_time: float | None = None
+ cache_update_ready_time: float | None = None
+
+ def stage_durations_ms(self) -> dict[str, float]:
+ durations: dict[str, float] = {}
+ _add_duration(
+ durations,
+ "condition",
+ self.condition_start_time,
+ self.condition_ready_time,
+ )
+ _add_duration(durations, "model", self.model_start_time, self.model_ready_time)
+ _add_duration(
+ durations,
+ "cache_update",
+ self.cache_update_start_time,
+ self.cache_update_ready_time,
+ )
+ _add_duration(
+ durations,
+ "decode",
+ self.decode_start_time,
+ self.decode_ready_time,
+ )
+ _add_duration(durations, "merge", self.merge_start_time, self.merge_ready_time)
+ _add_duration(
+ durations, "total", self.condition_start_time, self.merge_ready_time
+ )
+ return durations
+
+
+class ChunkHistory:
+ def __init__(self, capacity: int) -> None:
+ self._deque: deque[ChunkTimes] = deque(maxlen=capacity)
+
+ def append(self, chunk: ChunkTimes) -> None:
+ self._deque.append(chunk)
+
+ def __iter__(self) -> Iterator[ChunkTimes]:
+ return iter(self._deque)
+
+ def __len__(self) -> int:
+ return len(self._deque)
+
+ def summary(self) -> "RecentTimingSummary":
+ return summarize_chunk_history(self._deque)
+
+
+class TraceSink(Protocol):
+ def add_thread(self, name: str) -> int: ...
+
+ def add_instant(
+ self,
+ name: str,
+ *,
+ thread: int,
+ time_ns: int,
+ depends_on: list[int] | None = None,
+ **components: TraceComponentValue,
+ ) -> int: ...
+
+ def add_range(
+ self,
+ name: str,
+ *,
+ thread: int,
+ begin_ns: int,
+ end_ns: int,
+ depends_on: list[int] | None = None,
+ **components: TraceComponentValue,
+ ) -> int: ...
+
+
+@dataclass(frozen=True)
+class TraceContext:
+ sink: TraceSink
+ main_thread: int
+ worker_thread: int
+ lock: Lock
+
+ @classmethod
+ def create(cls, sink: TraceSink) -> "TraceContext":
+ return cls(
+ sink=sink,
+ main_thread=sink.add_thread("main"),
+ worker_thread=sink.add_thread("pipeline-worker"),
+ lock=Lock(),
+ )
+
+ def add_instant(
+ self,
+ name: str,
+ *,
+ thread: int,
+ time_ns: int,
+ depends_on: list[int] | None = None,
+ **components: TraceComponentValue,
+ ) -> int:
+ with self.lock:
+ return self.sink.add_instant(
+ name,
+ thread=thread,
+ time_ns=time_ns,
+ depends_on=depends_on,
+ **components,
+ )
+
+ def add_range(
+ self,
+ name: str,
+ *,
+ thread: int,
+ begin_ns: int,
+ end_ns: int,
+ depends_on: list[int] | None = None,
+ **components: TraceComponentValue,
+ ) -> int:
+ with self.lock:
+ return self.sink.add_range(
+ name,
+ thread=thread,
+ begin_ns=begin_ns,
+ end_ns=end_ns,
+ depends_on=depends_on,
+ **components,
+ )
+
+
+@dataclass(frozen=True)
+class VideoModelTraceEvents:
+ condition_event_id: int
+ model_event_id: int
+ merge_event_id: int
+ cache_update_event_id: int | None = None
+ decode_event_id: int | None = None
+
+ @property
+ def final_event_id(self) -> int:
+ return self.merge_event_id
+
+
+def emit_video_model_timing_ranges(
+ trace_context: TraceContext,
+ *,
+ timings: VideoModelTimings,
+ thread: int,
+ depends_on: list[int] | None = None,
+ chunk_index: int,
+) -> VideoModelTraceEvents:
+ """Emit standard trace ranges for backend-visible video-model stages."""
+ condition_event = trace_context.add_range(
+ "condition_raster",
+ thread=thread,
+ begin_ns=trace_time_ns(timings.condition_start_time),
+ end_ns=trace_time_ns(timings.condition_ready_time),
+ depends_on=depends_on,
+ chunk_index=chunk_index,
+ )
+ model_event = trace_context.add_range(
+ "model_generate",
+ thread=thread,
+ begin_ns=trace_time_ns(timings.model_start_time),
+ end_ns=trace_time_ns(timings.model_ready_time),
+ depends_on=event_dependencies(condition_event),
+ chunk_index=chunk_index,
+ )
+ cache_update_event = _add_optional_trace_range(
+ trace_context,
+ "cache_update",
+ thread=thread,
+ begin_time=timings.cache_update_start_time,
+ end_time=timings.cache_update_ready_time,
+ depends_on=event_dependencies(model_event),
+ chunk_index=chunk_index,
+ )
+ decode_event = _add_optional_trace_range(
+ trace_context,
+ "decode",
+ thread=thread,
+ begin_time=timings.decode_start_time,
+ end_time=timings.decode_ready_time,
+ depends_on=event_dependencies(
+ cache_update_event if cache_update_event is not None else model_event
+ ),
+ chunk_index=chunk_index,
+ )
+ last_event = decode_event
+ if last_event is None:
+ last_event = cache_update_event
+ if last_event is None:
+ last_event = model_event
+ merge_event = trace_context.add_range(
+ "frame_merge",
+ thread=thread,
+ begin_ns=trace_time_ns(timings.merge_start_time),
+ end_ns=trace_time_ns(timings.merge_ready_time),
+ depends_on=event_dependencies(last_event),
+ chunk_index=chunk_index,
+ )
+ return VideoModelTraceEvents(
+ condition_event_id=condition_event,
+ model_event_id=model_event,
+ cache_update_event_id=cache_update_event,
+ decode_event_id=decode_event,
+ merge_event_id=merge_event,
+ )
+
+
+def chunk_stage_durations_ms(chunk: ChunkTimes) -> dict[str, float]:
+ durations: dict[str, float] = {}
+ _add_duration(
+ durations, "input_to_request", chunk.input_sample_time, chunk.request_time
+ )
+ _add_duration(
+ durations,
+ "request_to_poses_ready",
+ chunk.request_time,
+ chunk.request_poses_ready_time,
+ )
+ _add_duration(
+ durations,
+ "queue_wait",
+ chunk.request_poses_ready_time,
+ chunk.chunk_render_start_time,
+ )
+ _add_duration(
+ durations,
+ "chunk_render",
+ chunk.chunk_render_start_time,
+ chunk.chunk_ready_time,
+ )
+ if chunk.frames:
+ first_frame = chunk.frames[0]
+ _add_duration(
+ durations,
+ "chunk_ready_to_first_image",
+ chunk.chunk_ready_time,
+ first_frame.image_ready_time,
+ )
+ _add_duration(
+ durations,
+ "first_image_to_present",
+ first_frame.image_ready_time,
+ first_frame.present_time,
+ )
+ _add_duration(
+ durations,
+ "input_to_first_present",
+ chunk.input_sample_time,
+ first_frame.present_time,
+ )
+ return durations
+
+
+@dataclass(frozen=True)
+class StageDurationSummary:
+ count: int
+ avg_ms: float
+ min_ms: float
+ max_ms: float
+ median_ms: float
+ p90_ms: float
+
+
+@dataclass(frozen=True)
+class RecentTimingSummary:
+ chunk_count: int
+ stages: dict[str, StageDurationSummary]
+
+
+def summarize_stage_durations(
+ samples: Iterable[Mapping[str, float]],
+) -> dict[str, StageDurationSummary]:
+ grouped: dict[str, list[float]] = defaultdict(list)
+ for sample in samples:
+ for stage_name, duration_ms in sample.items():
+ grouped[stage_name].append(float(duration_ms))
+ return {
+ stage_name: _summarize_values(values)
+ for stage_name, values in sorted(grouped.items())
+ }
+
+
+def summarize_chunk_history(chunks: Iterable[ChunkTimes]) -> RecentTimingSummary:
+ chunk_list = list(chunks)
+ return RecentTimingSummary(
+ chunk_count=len(chunk_list),
+ stages=summarize_stage_durations(
+ chunk.stage_durations_ms() for chunk in chunk_list
+ ),
+ )
+
+
+class RollingChunkTimingSummary:
+ def __init__(self, capacity: int) -> None:
+ self._chunks: deque[ChunkTimes] = deque(maxlen=capacity)
+
+ def append(self, chunk: ChunkTimes) -> None:
+ self._chunks.append(chunk)
+
+ def reset(self) -> None:
+ self._chunks.clear()
+
+ def summary(self) -> RecentTimingSummary:
+ return summarize_chunk_history(self._chunks)
+
+
+@dataclass(frozen=True)
+class InputToPresentSummary:
+ window_s: float
+ samples: int
+ wall_present_fps: float
+ avg_raw_control_to_present_ms: float
+ avg_adj_control_to_present_ms: float
+
+ def log_message(self) -> str:
+ return (
+ "[profile] e2e "
+ f"wall_present_fps={self.wall_present_fps:.1f} "
+ f"avg_adj_control_to_present_ms={self.avg_adj_control_to_present_ms:.2f} "
+ f"avg_raw_control_to_present_ms={self.avg_raw_control_to_present_ms:.2f} "
+ f"samples={self.samples}"
+ )
+
+
+class InputToPresentProfileWindow:
+ """Rolling wall-clock input-to-present summary window."""
+
+ def __init__(self, *, interval_s: float = 2.0) -> None:
+ self.interval_s = interval_s
+ self.reset()
+
+ def reset(self, *, interval_s: float | None = None) -> None:
+ if interval_s is not None:
+ self.interval_s = interval_s
+ self._sum_raw_ms = 0.0
+ self._sum_adj_ms = 0.0
+ self._count = 0
+ self._window_start: float | None = None
+
+ def record(
+ self,
+ *,
+ present_time: float,
+ input_sample_time: float,
+ frame_index: int,
+ frame_interval_s: float,
+ ) -> InputToPresentSummary | None:
+ raw_ms = (present_time - input_sample_time) * 1000.0
+ scheduled_ms = frame_index * (frame_interval_s * 1000.0)
+ adj_ms = raw_ms - scheduled_ms
+ self._sum_raw_ms += raw_ms
+ self._sum_adj_ms += adj_ms
+ self._count += 1
+ if self._window_start is None:
+ self._window_start = present_time
+
+ window_s = present_time - self._window_start
+ if window_s < self.interval_s:
+ return None
+
+ samples = self._count
+ summary = InputToPresentSummary(
+ window_s=window_s,
+ samples=samples,
+ wall_present_fps=float(samples) / window_s if window_s > 1e-9 else 0.0,
+ avg_raw_control_to_present_ms=self._sum_raw_ms / float(samples),
+ avg_adj_control_to_present_ms=self._sum_adj_ms / float(samples),
+ )
+ self._sum_raw_ms = 0.0
+ self._sum_adj_ms = 0.0
+ self._count = 0
+ self._window_start = present_time
+ return summary
+
+
+def _add_duration(
+ durations: dict[str, float],
+ name: str,
+ start_time: float | None,
+ end_time: float | None,
+) -> None:
+ duration = _duration_ms(start_time, end_time)
+ if duration is not None:
+ durations[name] = duration
+
+
+def _add_optional_trace_range(
+ trace_context: TraceContext,
+ name: str,
+ *,
+ thread: int,
+ begin_time: float | None,
+ end_time: float | None,
+ depends_on: list[int] | None,
+ chunk_index: int,
+) -> int | None:
+ if begin_time is None or end_time is None:
+ return None
+ return trace_context.add_range(
+ name,
+ thread=thread,
+ begin_ns=trace_time_ns(begin_time),
+ end_ns=trace_time_ns(end_time),
+ depends_on=depends_on,
+ chunk_index=chunk_index,
+ )
+
+
+def _summarize_values(values: list[float]) -> StageDurationSummary:
+ ordered = sorted(values)
+ count = len(ordered)
+ if count <= 0:
+ raise ValueError("Cannot summarize an empty stage duration list.")
+ return StageDurationSummary(
+ count=count,
+ avg_ms=sum(ordered) / float(count),
+ min_ms=ordered[0],
+ max_ms=ordered[-1],
+ median_ms=_percentile_sorted(ordered, 0.5),
+ p90_ms=_percentile_sorted(ordered, 0.9),
+ )
+
+
+def _percentile_sorted(values: list[float], percentile: float) -> float:
+ if not values:
+ raise ValueError("Cannot compute percentile of an empty list.")
+ if len(values) == 1:
+ return values[0]
+ clamped = min(1.0, max(0.0, percentile))
+ index = clamped * (len(values) - 1)
+ lower_index = int(index)
+ upper_index = min(lower_index + 1, len(values) - 1)
+ fraction = index - lower_index
+ return values[lower_index] * (1.0 - fraction) + values[upper_index] * fraction
diff --git a/flashdreams/flashdreams/serving/webrtc/controls.py b/flashdreams/flashdreams/serving/webrtc/controls.py
index 3743a0344..7f1d8461e 100644
--- a/flashdreams/flashdreams/serving/webrtc/controls.py
+++ b/flashdreams/flashdreams/serving/webrtc/controls.py
@@ -1,337 +1,36 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
-"""Keyboard state, sparse-edge event resampling, and camera pose integration."""
+"""WebRTC adapters for shared realtime input helpers."""
from __future__ import annotations
-from collections import deque
-from dataclasses import dataclass, field
-from typing import Literal
-
-import numpy as np
-
-DEFAULT_SUPPORTED_KEYS = frozenset({"w", "a", "s", "d", "q", "e", "i", "k", "j", "l"})
-WSAD_SUPPORTED_KEYS = frozenset({"w", "a", "s", "d"})
-KEY_ALIASES = {
- "arrowup": "w",
- "arrowleft": "a",
- "arrowdown": "s",
- "arrowright": "d",
-}
-
-
-def normalize_key(key: str) -> str:
- normalized = key.strip().lower()
- return KEY_ALIASES.get(normalized, normalized)
-
-
-@dataclass(slots=True)
-class KeyboardState:
- pressed_keys: set[str] = field(default_factory=set)
- supported_keys: frozenset[str] = DEFAULT_SUPPORTED_KEYS
- _press_order: dict[str, int] = field(default_factory=dict)
- _press_counter: int = 0
-
- def apply_event(self, *, event: str, key: str) -> bool:
- normalized_key = normalize_key(key)
- if normalized_key not in self.supported_keys:
- return False
-
- normalized_event = event.strip().lower()
- if normalized_event == "keydown":
- self.pressed_keys.add(normalized_key)
- self._press_counter += 1
- self._press_order[normalized_key] = self._press_counter
- return True
- if normalized_event == "keyup":
- self.pressed_keys.discard(normalized_key)
- self._press_order.pop(normalized_key, None)
- return True
- return False
-
- def snapshot(self) -> frozenset[str]:
- return frozenset(self.pressed_keys)
-
- def _latest_pressed(self, keys: tuple[str, ...]) -> str | None:
- latest_key: str | None = None
- latest_idx = -1
- for key in keys:
- if key not in self.pressed_keys:
- continue
- idx = self._press_order.get(key, -1)
- if idx >= latest_idx:
- latest_idx = idx
- latest_key = key
- return latest_key
-
- def resolved_effective_keys(self) -> frozenset[str]:
- effective: set[str] = set()
- for key in (
- self._latest_pressed(("w", "s")),
- self._latest_pressed(("a", "d", "j", "l")),
- self._latest_pressed(("q", "e")),
- self._latest_pressed(("i", "k")),
- ):
- if key is not None:
- effective.add(key)
- return frozenset(key for key in effective if key in self.supported_keys)
-
-
-PoseSegment = tuple[float, float, frozenset[str]]
-
-
-class KeyboardResampler:
- """Resample sparse keydown/keyup edges into a chunk timeline."""
-
- def __init__(
- self,
- *,
- fps: int,
- start_v: float = 0.0,
- supported_keys: frozenset[str] = DEFAULT_SUPPORTED_KEYS,
- ) -> None:
- if fps <= 0:
- raise ValueError("fps must be > 0")
- self._fps = fps
- self._dt = 1.0 / fps
- self._supported_keys = supported_keys
- self.next_chunk_start_v = start_v
- self._event_log: deque[tuple[float, dict[str, str]]] = deque()
- self._carried_state = KeyboardState(supported_keys=supported_keys)
-
- @property
- def fps(self) -> int:
- return self._fps
-
- @property
- def dt(self) -> float:
- return self._dt
-
- def on_edge(self, *, arrival_t: float, event: str, key: str) -> None:
- self._event_log.append((arrival_t, {"event": event, "key": key}))
-
- def sample_chunk(self, num_frames: int) -> tuple[list[PoseSegment], list[float]]:
- if num_frames < 1:
- raise ValueError("num_frames must be >= 1")
-
- chunk_start_v = self.next_chunk_start_v
- chunk_end_v = chunk_start_v + num_frames * self._dt
-
- while self._event_log and self._event_log[0][0] < chunk_start_v:
- _, payload = self._event_log.popleft()
- self._carried_state.apply_event(**payload)
-
- segments: list[PoseSegment] = []
- prev_t = chunk_start_v
- prev_state = self._carried_state.resolved_effective_keys()
- while self._event_log and self._event_log[0][0] <= chunk_end_v:
- event_t, payload = self._event_log.popleft()
- if event_t > prev_t:
- segments.append((prev_t, event_t, prev_state))
- self._carried_state.apply_event(**payload)
- prev_state = self._carried_state.resolved_effective_keys()
- prev_t = event_t
- if prev_t < chunk_end_v:
- segments.append((prev_t, chunk_end_v, prev_state))
- elif not segments:
- segments.append((chunk_start_v, chunk_end_v, prev_state))
-
- frame_times = [chunk_start_v + (i + 1) * self._dt for i in range(num_frames)]
- self.next_chunk_start_v = chunk_end_v
- return segments, frame_times
-
- def reset(self, *, start_v: float) -> None:
- self._event_log.clear()
- self._carried_state = KeyboardState(supported_keys=self._supported_keys)
- self.next_chunk_start_v = start_v
-
- def event_log_size(self) -> int:
- return len(self._event_log)
-
-
-def _rotation_matrix(axis: str, angle_rad: float) -> np.ndarray:
- cos_t = np.float32(np.cos(angle_rad))
- sin_t = np.float32(np.sin(angle_rad))
- if axis == "x":
- return np.array(
- [
- [1.0, 0.0, 0.0],
- [0.0, cos_t, -sin_t],
- [0.0, sin_t, cos_t],
- ],
- dtype=np.float32,
- )
- if axis == "y":
- return np.array(
- [
- [cos_t, 0.0, sin_t],
- [0.0, 1.0, 0.0],
- [-sin_t, 0.0, cos_t],
- ],
- dtype=np.float32,
- )
- if axis == "z":
- return np.array(
- [
- [cos_t, -sin_t, 0.0],
- [sin_t, cos_t, 0.0],
- [0.0, 0.0, 1.0],
- ],
- dtype=np.float32,
- )
- return np.eye(3, dtype=np.float32)
-
-
-@dataclass(slots=True)
-class CameraPoseIntegrator:
- """Integrate a piecewise-constant keyboard timeline into a camera trajectory."""
-
- move_speed_per_s: float = 0.8
- rotate_speed_rad_per_s: float = float(np.deg2rad(32.0))
- pitch_limit_rad: float = float(np.deg2rad(85.0))
- coordinate_system: Literal["RDF", "FLU"] = "RDF"
- _current_pose: np.ndarray = field(
- default_factory=lambda: np.eye(4, dtype=np.float32),
- )
- _current_pitch: float = 0.0
-
- def __post_init__(self) -> None:
- if self.coordinate_system not in {"RDF", "FLU"}:
- raise ValueError(
- "coordinate_system must be 'RDF' (right-down-forward) "
- "or 'FLU' (forward-left-up)"
- )
-
- def reset(self, pose: np.ndarray | None = None) -> None:
- if pose is None:
- self._current_pose = np.eye(4, dtype=np.float32)
- self._current_pitch = 0.0
- return
- if pose.shape != (4, 4):
- raise ValueError(f"Expected pose shape (4, 4), got {pose.shape}")
- self._current_pose = pose.astype(np.float32, copy=True)
- if self.coordinate_system == "FLU":
- self._current_pitch = float(np.arcsin(np.clip(pose[2, 0], -1.0, 1.0)))
- else:
- self._current_pitch = float(np.arctan2(pose[2, 1], pose[1, 1]))
-
- def current_pose(self) -> np.ndarray:
- return self._current_pose.copy()
-
- def _advance(self, *, state: frozenset[str], duration: float) -> None:
- if duration <= 0:
- return
-
- yaw_rate = 0.0
- if self.coordinate_system == "FLU":
- if "a" in state or "j" in state:
- yaw_rate += self.rotate_speed_rad_per_s
- if "d" in state or "l" in state:
- yaw_rate -= self.rotate_speed_rad_per_s
- else:
- if "a" in state or "j" in state:
- yaw_rate -= self.rotate_speed_rad_per_s
- if "d" in state or "l" in state:
- yaw_rate += self.rotate_speed_rad_per_s
- pitch_rate = 0.0
- if "i" in state:
- pitch_rate += self.rotate_speed_rad_per_s
- if "k" in state:
- pitch_rate -= self.rotate_speed_rad_per_s
-
- yaw_delta = yaw_rate * duration
- pitch_delta = pitch_rate * duration
-
- new_pitch = self._current_pitch + pitch_delta
- if -self.pitch_limit_rad <= new_pitch <= self.pitch_limit_rad:
- self._current_pitch = new_pitch
- else:
- pitch_delta = 0.0
-
- rot = self._current_pose[:3, :3]
- trans = self._current_pose[:3, 3]
- if self.coordinate_system == "FLU":
- rot_pitch = _rotation_matrix("y", -pitch_delta)
- rot_yaw = _rotation_matrix("z", yaw_delta)
- else:
- rot_pitch = _rotation_matrix("x", pitch_delta)
- rot_yaw = _rotation_matrix("y", yaw_delta)
- rot_new = rot_yaw @ rot @ rot_pitch
-
- forward_rate = 0.0
- if "w" in state:
- forward_rate += self.move_speed_per_s
- if "s" in state:
- forward_rate -= self.move_speed_per_s
- right_rate = 0.0
- if "e" in state:
- right_rate += self.move_speed_per_s
- if "q" in state:
- right_rate -= self.move_speed_per_s
-
- if self.coordinate_system == "FLU":
- vec_forward = rot_new[:, 0]
- vec_right = -rot_new[:, 1]
- forward_flat = np.array(
- [vec_forward[0], vec_forward[1], 0.0], dtype=np.float32
- )
- right_flat = np.array([vec_right[0], vec_right[1], 0.0], dtype=np.float32)
- else:
- vec_right = rot_new[:, 0]
- vec_forward = rot_new[:, 2]
- forward_flat = np.array(
- [vec_forward[0], 0.0, vec_forward[2]], dtype=np.float32
- )
- right_flat = np.array([vec_right[0], 0.0, vec_right[2]], dtype=np.float32)
- forward_norm = np.linalg.norm(forward_flat)
- right_norm = np.linalg.norm(right_flat)
- if forward_norm > 0:
- forward_flat /= forward_norm
- if right_norm > 0:
- right_flat /= right_norm
-
- move_vec = forward_flat * (forward_rate * duration) + right_flat * (
- right_rate * duration
- )
- self._current_pose = np.eye(4, dtype=np.float32)
- self._current_pose[:3, :3] = rot_new
- self._current_pose[:3, 3] = trans + move_vec
-
- def integrate_chunk(
- self,
- *,
- segments: list[PoseSegment],
- frame_times: list[float],
- ) -> np.ndarray:
- if not segments:
- raise ValueError("segments must be non-empty")
- if not frame_times:
- raise ValueError("frame_times must be non-empty")
- chunk_start = segments[0][0]
- chunk_end = segments[-1][1]
- if any(
- frame_times[i] >= frame_times[i + 1] for i in range(len(frame_times) - 1)
- ):
- raise ValueError("frame_times must be strictly increasing")
- if frame_times[0] < chunk_start - 1e-9 or frame_times[-1] > chunk_end + 1e-9:
- raise ValueError(
- "frame_times must lie within the chunk window "
- f"[{chunk_start}, {chunk_end}]"
- )
-
- poses: list[np.ndarray] = []
- cur_t = chunk_start
- ft_idx = 0
- for _, seg_end, seg_state in segments:
- while ft_idx < len(frame_times) and frame_times[ft_idx] <= seg_end:
- target_t = frame_times[ft_idx]
- self._advance(state=seg_state, duration=target_t - cur_t)
- cur_t = target_t
- poses.append(self._current_pose.copy())
- ft_idx += 1
- if seg_end > cur_t:
- self._advance(state=seg_state, duration=seg_end - cur_t)
- cur_t = seg_end
-
- return np.stack(poses, axis=0).astype(np.float32)
+from flashdreams.serving.realtime.input import (
+ DEFAULT_SUPPORTED_KEYS,
+ KEY_ALIASES,
+ WSAD_SUPPORTED_KEYS,
+ CameraPoseIntegrator,
+ ImageRequest,
+ KeyboardResampler,
+ KeyboardState,
+ PoseSegment,
+ PromptRequest,
+ ResetRequest,
+ SparseInputSnapshot,
+ normalize_key,
+)
+
+__all__ = [
+ "DEFAULT_SUPPORTED_KEYS",
+ "KEY_ALIASES",
+ "WSAD_SUPPORTED_KEYS",
+ "CameraPoseIntegrator",
+ "ImageRequest",
+ "KeyboardResampler",
+ "KeyboardState",
+ "PoseSegment",
+ "PromptRequest",
+ "ResetRequest",
+ "SparseInputSnapshot",
+ "normalize_key",
+]
diff --git a/flashdreams/flashdreams/serving/webrtc/manager.py b/flashdreams/flashdreams/serving/webrtc/manager.py
index 5e97813f9..1a17c8b6b 100644
--- a/flashdreams/flashdreams/serving/webrtc/manager.py
+++ b/flashdreams/flashdreams/serving/webrtc/manager.py
@@ -19,7 +19,7 @@
from aiortc import RTCConfiguration, RTCPeerConnection, RTCSessionDescription
from loguru import logger
-from flashdreams.serving.webrtc.controls import KeyboardResampler
+from flashdreams.serving.realtime.input import KeyboardResampler
from flashdreams.serving.webrtc.media import BufferedVideoTrack
from flashdreams.serving.webrtc.server import SessionBusyError
from flashdreams.serving.webrtc.warmup import (
diff --git a/flashdreams/flashdreams/serving/webrtc/media.py b/flashdreams/flashdreams/serving/webrtc/media.py
index df5925e69..e438c3b1d 100644
--- a/flashdreams/flashdreams/serving/webrtc/media.py
+++ b/flashdreams/flashdreams/serving/webrtc/media.py
@@ -6,41 +6,25 @@
import asyncio
from collections.abc import Callable
from fractions import Fraction
+from typing import TYPE_CHECKING
import numpy as np
-import torch
from aiortc import MediaStreamTrack
from aiortc.mediastreams import MediaStreamError
from av import VideoFrame
from loguru import logger
+from flashdreams.serving.realtime.media import tensor_chunk_to_rgb_frames
+
+if TYPE_CHECKING:
+ import torch
+
_STALL_THRESHOLD_MS = 1.0
_PACING_LAG_LOG_MS = 5.0
-def tensor_chunk_to_rgb_frames(video_chunk: torch.Tensor) -> list[np.ndarray]:
- """Convert common model output tensor layouts to RGB uint8 frames."""
- if video_chunk.ndim == 4:
- frames = video_chunk.float().permute(0, 2, 3, 1).numpy()
- frames = ((frames + 1.0) / 2.0 * 255.0).clip(0, 255).astype(np.uint8)
- return [np.ascontiguousarray(frame) for frame in frames]
- if video_chunk.ndim == 6:
- if video_chunk.shape[0] != 1 or video_chunk.shape[1] != 1:
- raise ValueError(
- "Expected single-batch single-view video chunk [1, 1, T, 3, H, W], "
- f"got {tuple(video_chunk.shape)}"
- )
- chunk = video_chunk[0, 0]
- if chunk.dtype == torch.uint8:
- frames = chunk.permute(0, 2, 3, 1).cpu().numpy()
- else:
- frames = chunk.float().permute(0, 2, 3, 1).cpu().numpy()
- frames = ((frames + 1.0) / 2.0 * 255.0).clip(0, 255).astype(np.uint8)
- return [np.ascontiguousarray(frame) for frame in frames]
- raise ValueError(
- "Expected video chunk [T, C, H, W] or [1, 1, T, 3, H, W], "
- f"got {tuple(video_chunk.shape)}"
- )
+def _default_frame_converter(video_chunk: torch.Tensor) -> list[np.ndarray]:
+ return tensor_chunk_to_rgb_frames(video_chunk, sync_device=True)
class BufferedVideoTrack(MediaStreamTrack):
@@ -66,7 +50,7 @@ def __init__(
self._next_deadline_s: float | None = None
self._pts = 0
self._maxsize = maxsize
- self._frame_converter = frame_converter or tensor_chunk_to_rgb_frames
+ self._frame_converter = frame_converter or _default_frame_converter
self._frames: asyncio.Queue[np.ndarray | None] = asyncio.Queue(maxsize=maxsize)
self._closed = False
diff --git a/flashdreams/tests/test_realtime_mjpeg.py b/flashdreams/tests/test_realtime_mjpeg.py
new file mode 100644
index 000000000..0730f5748
--- /dev/null
+++ b/flashdreams/tests/test_realtime_mjpeg.py
@@ -0,0 +1,144 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+import threading
+from http import HTTPStatus
+
+import pytest
+
+from flashdreams.serving.realtime.frame_bus import LatestFrameBus
+from flashdreams.serving.realtime.mjpeg import (
+ format_mjpeg_part,
+ mjpeg_content_type,
+ publish_latest_jpeg,
+ send_mjpeg_response_headers,
+ wait_for_latest_jpeg,
+ write_mjpeg_stream,
+)
+
+pytestmark = pytest.mark.ci_cpu
+
+
+class _Response:
+ def __init__(self) -> None:
+ self.status: int | HTTPStatus | None = None
+ self.headers: list[tuple[str, str]] = []
+ self.ended = False
+
+ def send_response(self, code: int | HTTPStatus) -> None:
+ self.status = code
+
+ def send_header(self, keyword: str, value: str) -> None:
+ self.headers.append((keyword, value))
+
+ def end_headers(self) -> None:
+ self.ended = True
+
+
+class _Writer:
+ def __init__(self) -> None:
+ self.parts: list[bytes] = []
+ self.flushes = 0
+
+ def write(self, data: bytes) -> object:
+ self.parts.append(data)
+ return len(data)
+
+ def flush(self) -> object:
+ self.flushes += 1
+ return None
+
+
+def test_format_mjpeg_part_includes_boundary_headers_and_payload() -> None:
+ part = format_mjpeg_part(b"\xff\xd8jpeg\xff\xd9", boundary="test-boundary")
+
+ assert part == (
+ b"--test-boundary\r\n"
+ b"Content-Type: image/jpeg\r\n"
+ b"Content-Length: 8\r\n\r\n"
+ b"\xff\xd8jpeg\xff\xd9\r\n"
+ )
+
+
+def test_send_mjpeg_response_headers_uses_no_cache_multipart_type() -> None:
+ response = _Response()
+
+ send_mjpeg_response_headers(response, boundary="test-boundary")
+
+ assert response.status == HTTPStatus.OK
+ assert ("Pragma", "no-cache") in response.headers
+ assert (
+ "Content-Type",
+ "multipart/x-mixed-replace; boundary=test-boundary",
+ ) in response.headers
+ assert response.ended
+
+
+def test_mjpeg_content_type_uses_boundary() -> None:
+ assert mjpeg_content_type("custom") == "multipart/x-mixed-replace; boundary=custom"
+
+
+def test_write_mjpeg_stream_writes_until_waiter_returns_none() -> None:
+ frames = [(b"first", 1), (b"second", 2)]
+ seen_counts: list[int] = []
+
+ def wait_for_frame(last_seen_count: int) -> tuple[bytes, int] | None:
+ seen_counts.append(last_seen_count)
+ if not frames:
+ return None
+ return frames.pop(0)
+
+ writer = _Writer()
+
+ write_mjpeg_stream(writer, wait_for_frame, boundary="test")
+
+ assert seen_counts == [0, 1, 2]
+ assert writer.parts == [
+ format_mjpeg_part(b"first", boundary="test"),
+ format_mjpeg_part(b"second", boundary="test"),
+ ]
+ assert writer.flushes == 2
+
+
+def test_publish_latest_jpeg_ignores_stopped_publish() -> None:
+ bus = LatestFrameBus[bytes]()
+ stop_event = threading.Event()
+ stop_event.set()
+
+ publish_latest_jpeg(bus, b"jpeg", stop_event=stop_event)
+
+ assert bus.latest() is None
+
+
+def test_wait_for_latest_jpeg_returns_frame_and_count() -> None:
+ bus = LatestFrameBus[bytes]()
+ stop_event = threading.Event()
+ bus.publish(b"old")
+ bus.publish(b"new")
+
+ frame = wait_for_latest_jpeg(
+ bus,
+ last_seen_count=1,
+ stop_event=stop_event,
+ poll_timeout_s=0.01,
+ )
+
+ assert frame == (b"new", 2)
+
+
+def test_wait_for_latest_jpeg_returns_none_after_bus_close() -> None:
+ bus = LatestFrameBus[bytes]()
+ stop_event = threading.Event()
+ bus.publish(b"old")
+ bus.close()
+
+ frame = wait_for_latest_jpeg(
+ bus,
+ last_seen_count=1,
+ stop_event=stop_event,
+ poll_timeout_s=0.01,
+ )
+
+ assert frame is None
diff --git a/flashdreams/tests/test_realtime_presenter.py b/flashdreams/tests/test_realtime_presenter.py
new file mode 100644
index 000000000..c7f568f63
--- /dev/null
+++ b/flashdreams/tests/test_realtime_presenter.py
@@ -0,0 +1,139 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+import queue
+
+import numpy as np
+import pytest
+
+from flashdreams.serving.realtime.presenter import (
+ PresentationQueue,
+ materialize_rgb_host_uint8,
+ wait_until_present_time,
+)
+
+pytestmark = pytest.mark.ci_cpu
+
+
+def test_presentation_queue_drops_oldest_frame_at_capacity() -> None:
+ frames = PresentationQueue[str](capacity=2)
+
+ assert frames.append("first") is None
+ assert frames.append("second") is None
+ assert frames.append("third") == "first"
+
+ assert list(frames) == ["second", "third"]
+
+
+def test_presentation_queue_drains_available_frames_with_filter_and_prepare() -> None:
+ source: queue.Queue[int] = queue.Queue()
+ for value in [1, 2, 3, 4]:
+ source.put(value)
+ prepared: list[int] = []
+ ready = PresentationQueue[int](capacity=2)
+
+ result = ready.drain_nowait(
+ source,
+ include=lambda value: value % 2 == 0,
+ prepare=prepared.append,
+ )
+
+ assert result.pulled == 4
+ assert result.accepted == 2
+ assert result.skipped == 2
+ assert result.dropped == ()
+ assert prepared == [2, 4]
+ assert list(ready) == [2, 4]
+
+
+def test_presentation_queue_reports_drops_during_drain() -> None:
+ source: queue.Queue[str] = queue.Queue()
+ for value in ["first", "second", "third"]:
+ source.put(value)
+ ready = PresentationQueue[str](capacity=1)
+
+ result = ready.drain_nowait(source)
+
+ assert result.accepted == 3
+ assert result.dropped == ("first", "second")
+ assert ready.pop_ready() == "third"
+ assert ready.pop_ready() is None
+
+
+def test_presentation_queue_prepares_only_frames_retained_after_capacity() -> None:
+ source: queue.Queue[str] = queue.Queue()
+ for value in ["first", "second", "third"]:
+ source.put(value)
+ prepared: list[str] = []
+ ready = PresentationQueue[str](capacity=1)
+
+ result = ready.drain_nowait(source, prepare=prepared.append)
+
+ assert result.accepted == 3
+ assert result.dropped == ("first", "second")
+ assert prepared == ["third"]
+ assert list(ready) == ["third"]
+
+
+def test_presentation_queue_clear_flushes_ready_frames() -> None:
+ ready = PresentationQueue[int]()
+ ready.append(1)
+ ready.append(2)
+
+ assert ready.clear() == 2
+ assert len(ready) == 0
+
+
+def test_presentation_queue_popleft_raises_when_empty() -> None:
+ ready = PresentationQueue[int]()
+
+ with pytest.raises(IndexError):
+ ready.popleft()
+
+
+def test_wait_until_present_time_sleeps_until_poll_timeout() -> None:
+ times = iter([1.0, 1.25])
+ sleeps: list[float] = []
+
+ wait = wait_until_present_time(
+ 2.0,
+ poll_timeout_s=0.25,
+ clock=lambda: next(times),
+ sleep=sleeps.append,
+ )
+
+ assert wait is not None
+ assert wait.begin_time == 1.0
+ assert wait.end_time == 1.25
+ assert wait.duration_s == pytest.approx(0.25)
+ assert sleeps == [0.25]
+
+
+def test_wait_until_present_time_returns_none_when_due() -> None:
+ wait = wait_until_present_time(
+ 1.0,
+ poll_timeout_s=0.25,
+ clock=lambda: 1.1,
+ sleep=lambda _seconds: None,
+ )
+
+ assert wait is None
+
+
+def test_materialize_rgb_host_uint8_strips_alpha_from_lazy_frame() -> None:
+ class LazyFrame:
+ def to_numpy(self) -> np.ndarray:
+ return np.array(
+ [[[1, 2, 3, 255], [4, 5, 6, 255]]],
+ dtype=np.uint8,
+ )
+
+ frame = materialize_rgb_host_uint8(LazyFrame())
+
+ assert frame.flags.c_contiguous
+ np.testing.assert_array_equal(
+ frame,
+ np.array([[[1, 2, 3], [4, 5, 6]]], dtype=np.uint8),
+ )
diff --git a/flashdreams/tests/test_realtime_serving.py b/flashdreams/tests/test_realtime_serving.py
new file mode 100644
index 000000000..5c89d6c3f
--- /dev/null
+++ b/flashdreams/tests/test_realtime_serving.py
@@ -0,0 +1,212 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+import threading
+
+import numpy as np
+import pytest
+import torch
+
+from flashdreams.serving.realtime.frame_bus import LatestFrameBus
+from flashdreams.serving.realtime.input import (
+ ImageRequest,
+ KeyboardState,
+ PromptRequest,
+ ResetRequest,
+ SparseInputSnapshot,
+)
+from flashdreams.serving.realtime.media import (
+ encode_rgb_frame_to_jpeg,
+ rgb_array_to_uint8_frames,
+ tensor_chunk_to_rgb_frames,
+)
+
+pytestmark = pytest.mark.ci_cpu
+
+
+def test_keyboard_state_builds_sparse_snapshot_with_effective_keys() -> None:
+ state = KeyboardState()
+
+ assert state.apply_event(event="keydown", key="ArrowLeft")
+ assert state.apply_event(event="keydown", key="d")
+
+ snapshot = state.sparse_snapshot(timestamp_s=12.5)
+
+ assert snapshot == SparseInputSnapshot(
+ timestamp_s=12.5,
+ pressed_keys=frozenset({"a", "d"}),
+ effective_keys=frozenset({"d"}),
+ )
+
+
+def test_input_request_containers_are_transport_neutral() -> None:
+ snapshot = SparseInputSnapshot(
+ timestamp_s=1.0,
+ reset=ResetRequest(reason="new_session", request_id="reset-1"),
+ prompt=PromptRequest(prompt="drive forward", request_id="prompt-1"),
+ image=ImageRequest(
+ data=b"image-bytes",
+ content_type="image/jpeg",
+ request_id="image-1",
+ ),
+ )
+
+ assert snapshot.reset is not None
+ assert snapshot.reset.reason == "new_session"
+ assert snapshot.prompt is not None
+ assert snapshot.prompt.prompt == "drive forward"
+ assert snapshot.image is not None
+ assert snapshot.image.content_type == "image/jpeg"
+
+
+def test_latest_frame_bus_publishes_single_slot_frames() -> None:
+ bus = LatestFrameBus[bytes]()
+
+ assert bus.latest() is None
+ first_count = bus.publish(b"first")
+ second_count = bus.publish(b"second")
+
+ assert first_count == 1
+ assert second_count == 2
+ latest = bus.latest()
+ assert latest is not None
+ assert latest.payload == b"second"
+ assert latest.count == 2
+ waited = bus.wait_for_frame(last_seen_count=1, timeout_s=0.01)
+ assert waited is not None
+ assert waited.payload == b"second"
+ assert waited.count == 2
+
+
+def test_latest_frame_bus_close_wakes_waiters() -> None:
+ bus = LatestFrameBus[bytes]()
+ waiter_blocking = threading.Event()
+ results: list[object] = []
+ original_wait = bus._condition.wait
+
+ def wait_after_ready(timeout: float | None = None) -> bool:
+ waiter_blocking.set()
+ return original_wait(timeout=timeout)
+
+ setattr(bus._condition, "wait", wait_after_ready)
+
+ def wait_for_frame() -> None:
+ results.append(bus.wait_for_frame(last_seen_count=0, timeout_s=5.0))
+
+ thread = threading.Thread(target=wait_for_frame)
+ thread.start()
+ assert waiter_blocking.wait(timeout=1.0)
+
+ bus.close()
+ thread.join(timeout=1.0)
+
+ assert not thread.is_alive()
+ assert results == [None]
+ with pytest.raises(RuntimeError, match="closed LatestFrameBus"):
+ bus.publish(b"late")
+
+
+def test_realtime_media_matches_legacy_tensor_chunk_pixels() -> None:
+ chunk = torch.tensor(
+ [
+ [
+ [[-1.0, 0.0], [1.0, 2.0]],
+ [[-2.0, 0.5], [0.0, 1.0]],
+ [[1.0, -1.0], [0.0, 0.0]],
+ ],
+ ],
+ dtype=torch.float32,
+ )
+
+ shared_frames = tensor_chunk_to_rgb_frames(chunk)
+
+ assert len(shared_frames) == 1
+ np.testing.assert_array_equal(
+ shared_frames[0],
+ np.array(
+ [[[0, 0, 255], [127, 191, 0]], [[255, 127, 127], [255, 255, 127]]],
+ dtype=np.uint8,
+ ),
+ )
+
+
+def test_realtime_media_supports_omnidreams_uint8_layout() -> None:
+ chunk = torch.zeros((1, 1, 2, 3, 4, 5), dtype=torch.uint8)
+ chunk[0, 0, 1, 0] = 255
+
+ frames = tensor_chunk_to_rgb_frames(chunk)
+
+ assert len(frames) == 2
+ assert frames[0].shape == (4, 5, 3)
+ assert frames[0].dtype == np.uint8
+ assert frames[1][0, 0, 0] == 255
+
+
+@pytest.mark.parametrize(
+ "chunk",
+ [
+ torch.full((1, 3, 2, 2), -1.0, dtype=torch.bfloat16),
+ torch.full((1, 1, 1, 3, 2, 2), -1.0, dtype=torch.bfloat16),
+ ],
+)
+def test_realtime_media_promotes_bfloat16_tensor_chunks(
+ chunk: torch.Tensor,
+) -> None:
+ frames = tensor_chunk_to_rgb_frames(chunk)
+
+ assert len(frames) == 1
+ assert frames[0].shape == (2, 2, 3)
+ assert frames[0].dtype == np.uint8
+ assert frames[0].max() == 0
+
+
+def test_realtime_media_rejects_non_rgb_bvtchw_layout() -> None:
+ chunk = torch.zeros((1, 1, 2, 4, 5, 6), dtype=torch.uint8)
+
+ with pytest.raises(ValueError, match=r"\[1, 1, T, 3, H, W\]"):
+ tensor_chunk_to_rgb_frames(chunk)
+
+
+def test_realtime_media_rejects_scaled_value_range_for_uint8() -> None:
+ chunk = np.zeros((1, 2, 2, 3), dtype=np.uint8)
+
+ with pytest.raises(ValueError, match="uint8 inputs require value_range='uint8'"):
+ rgb_array_to_uint8_frames(
+ chunk,
+ layout="thwc",
+ value_range="minus_one_one",
+ )
+
+
+def test_realtime_media_converts_array_chunks() -> None:
+ chunk = np.array(
+ [
+ [[[0.0, 0.5, 1.0], [1.0, 0.0, 0.5]]],
+ [[[0.25, 0.75, 1.0], [0.0, 0.0, 0.0]]],
+ ],
+ dtype=np.float32,
+ )
+
+ frames = rgb_array_to_uint8_frames(
+ chunk,
+ layout="thwc",
+ value_range="zero_one",
+ )
+
+ assert len(frames) == 2
+ np.testing.assert_array_equal(
+ frames[0],
+ np.array([[[0, 127, 255], [255, 0, 127]]], dtype=np.uint8),
+ )
+
+
+def test_realtime_media_encodes_jpeg_bytes() -> None:
+ pytest.importorskip("PIL.Image")
+ frame = np.full((4, 5, 3), 127, dtype=np.uint8)
+
+ jpeg = encode_rgb_frame_to_jpeg(frame, quality=80)
+
+ assert jpeg.startswith(b"\xff\xd8")
+ assert jpeg.endswith(b"\xff\xd9")
diff --git a/flashdreams/tests/test_realtime_timing.py b/flashdreams/tests/test_realtime_timing.py
new file mode 100644
index 000000000..623b31afd
--- /dev/null
+++ b/flashdreams/tests/test_realtime_timing.py
@@ -0,0 +1,326 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+from collections.abc import Iterator
+from dataclasses import dataclass
+
+import pytest
+
+from flashdreams.serving.realtime.timing import (
+ ChunkHistory,
+ ChunkTimes,
+ InputToPresentProfileWindow,
+ TraceComponentValue,
+ TraceContext,
+ VideoModelTimings,
+ chunk_stage_durations_ms,
+ emit_video_model_timing_ranges,
+ summarize_chunk_history,
+)
+
+pytestmark = pytest.mark.ci_cpu
+
+
+@dataclass(frozen=True)
+class _TraceEvent:
+ name: str
+ depends_on: list[int]
+ components: dict[str, TraceComponentValue]
+
+
+class _RecordingTraceSink:
+ def __init__(self) -> None:
+ self.threads: list[str] = []
+ self.events: list[_TraceEvent] = []
+
+ def add_thread(self, name: str) -> int:
+ self.threads.append(name)
+ return len(self.threads) - 1
+
+ def add_instant(
+ self,
+ name: str,
+ *,
+ thread: int,
+ time_ns: int,
+ depends_on: list[int] | None = None,
+ **components: TraceComponentValue,
+ ) -> int:
+ del thread, time_ns
+ return self._append_event(name, depends_on, components)
+
+ def add_range(
+ self,
+ name: str,
+ *,
+ thread: int,
+ begin_ns: int,
+ end_ns: int,
+ depends_on: list[int] | None = None,
+ **components: TraceComponentValue,
+ ) -> int:
+ del thread
+ event_components = dict(components)
+ event_components["begin_ns"] = begin_ns
+ event_components["end_ns"] = end_ns
+ return self._append_event(name, depends_on, event_components)
+
+ def _append_event(
+ self,
+ name: str,
+ depends_on: list[int] | None,
+ components: dict[str, TraceComponentValue],
+ ) -> int:
+ self.events.append(
+ _TraceEvent(
+ name=name,
+ depends_on=[] if depends_on is None else depends_on,
+ components=components,
+ )
+ )
+ return len(self.events) - 1
+
+
+class _NamedIdTraceSink(_RecordingTraceSink):
+ def __init__(self, event_ids: dict[str, int]) -> None:
+ super().__init__()
+ self._event_ids = event_ids
+
+ def _append_event(
+ self,
+ name: str,
+ depends_on: list[int] | None,
+ components: dict[str, TraceComponentValue],
+ ) -> int:
+ super()._append_event(name, depends_on, components)
+ return self._event_ids[name]
+
+
+def _chunk_times() -> ChunkTimes:
+ chunk = ChunkTimes.create(
+ chunk_index=7,
+ input_sample_time=1.0,
+ request_time=1.1,
+ request_poses_ready_time=1.2,
+ intended_present_times=[1.5, 1.6],
+ )
+ chunk.chunk_render_start_time = 1.25
+ chunk.chunk_ready_time = 1.45
+ chunk.frames[0].image_ready_time = 1.47
+ chunk.frames[0].sample_display_pose_time = 1.49
+ chunk.frames[0].present_time = 1.52
+ return chunk
+
+
+def test_chunk_times_create_allocates_frame_times() -> None:
+ chunk = ChunkTimes.create(
+ chunk_index=0,
+ input_sample_time=1.0,
+ request_time=1.0,
+ request_poses_ready_time=1.001,
+ intended_present_times=[1.5, 1.6, 1.7],
+ )
+
+ assert [frame.frame_index for frame in chunk.frames] == [0, 1, 2]
+ assert [frame.intended_present_time for frame in chunk.frames] == [1.5, 1.6, 1.7]
+
+
+def test_chunk_history_iter_returns_iterator() -> None:
+ first = _chunk_times()
+ second = _chunk_times()
+ history = ChunkHistory(capacity=2)
+ history.append(first)
+ history.append(second)
+
+ iterator = iter(history)
+
+ assert isinstance(iterator, Iterator)
+ assert next(iterator) is first
+ assert next(iterator) is second
+
+
+def test_chunk_stage_durations_are_derived_from_milestones() -> None:
+ durations = chunk_stage_durations_ms(_chunk_times())
+
+ assert durations["input_to_request"] == pytest.approx(100.0)
+ assert durations["request_to_poses_ready"] == pytest.approx(100.0)
+ assert durations["queue_wait"] == pytest.approx(50.0)
+ assert durations["chunk_render"] == pytest.approx(200.0)
+ assert durations["chunk_ready_to_first_image"] == pytest.approx(20.0)
+ assert durations["first_image_to_present"] == pytest.approx(50.0)
+ assert durations["input_to_first_present"] == pytest.approx(520.0)
+
+
+def test_summarize_chunk_history_builds_stage_statistics() -> None:
+ first = _chunk_times()
+ second = _chunk_times()
+ second.chunk_render_start_time = 2.0
+ second.chunk_ready_time = 2.4
+
+ summary = summarize_chunk_history([first, second])
+
+ assert summary.chunk_count == 2
+ assert summary.stages["chunk_render"].count == 2
+ assert summary.stages["chunk_render"].avg_ms == pytest.approx(300.0)
+ assert summary.stages["chunk_render"].p90_ms == pytest.approx(380.0)
+
+
+def test_video_model_timings_include_optional_decode_and_cache_durations() -> None:
+ timings = VideoModelTimings(
+ condition_start_time=1.0,
+ condition_ready_time=1.1,
+ model_start_time=1.1,
+ model_ready_time=1.6,
+ merge_start_time=1.65,
+ merge_ready_time=1.7,
+ cache_update_start_time=1.2,
+ cache_update_ready_time=1.3,
+ decode_start_time=1.4,
+ decode_ready_time=1.55,
+ )
+
+ durations = timings.stage_durations_ms()
+
+ assert durations["condition"] == pytest.approx(100.0)
+ assert durations["model"] == pytest.approx(500.0)
+ assert durations["cache_update"] == pytest.approx(100.0)
+ assert durations["decode"] == pytest.approx(150.0)
+ assert durations["merge"] == pytest.approx(50.0)
+ assert durations["total"] == pytest.approx(700.0)
+
+
+def test_emit_video_model_timing_ranges_adds_optional_subranges() -> None:
+ sink = _RecordingTraceSink()
+ trace_context = TraceContext.create(sink)
+ timings = VideoModelTimings(
+ condition_start_time=1.0,
+ condition_ready_time=1.1,
+ model_start_time=1.1,
+ model_ready_time=1.6,
+ merge_start_time=1.65,
+ merge_ready_time=1.7,
+ cache_update_start_time=1.2,
+ cache_update_ready_time=1.3,
+ decode_start_time=1.4,
+ decode_ready_time=1.55,
+ )
+
+ events = emit_video_model_timing_ranges(
+ trace_context,
+ timings=timings,
+ thread=trace_context.worker_thread,
+ depends_on=[123],
+ chunk_index=9,
+ )
+
+ names = [event.name for event in sink.events]
+ assert names == [
+ "condition_raster",
+ "model_generate",
+ "cache_update",
+ "decode",
+ "frame_merge",
+ ]
+ assert sink.events[0].depends_on == [123]
+ assert sink.events[2].depends_on == [events.model_event_id]
+ assert sink.events[3].depends_on == [events.cache_update_event_id]
+ assert sink.events[4].depends_on == [events.decode_event_id]
+ assert events.final_event_id == events.merge_event_id
+
+
+def test_emit_video_model_timing_ranges_keeps_zero_valued_event_ids() -> None:
+ sink = _NamedIdTraceSink(
+ {
+ "condition_raster": 10,
+ "model_generate": 11,
+ "cache_update": 0,
+ "decode": 12,
+ "frame_merge": 13,
+ }
+ )
+ trace_context = TraceContext.create(sink)
+ timings = VideoModelTimings(
+ condition_start_time=1.0,
+ condition_ready_time=1.1,
+ model_start_time=1.1,
+ model_ready_time=1.6,
+ merge_start_time=1.65,
+ merge_ready_time=1.7,
+ cache_update_start_time=1.2,
+ cache_update_ready_time=1.3,
+ decode_start_time=1.4,
+ decode_ready_time=1.55,
+ )
+
+ events = emit_video_model_timing_ranges(
+ trace_context,
+ timings=timings,
+ thread=trace_context.worker_thread,
+ chunk_index=9,
+ )
+
+ assert events.cache_update_event_id == 0
+ assert sink.events[3].depends_on == [0]
+
+
+def test_emit_video_model_timing_ranges_uses_zero_decode_event_for_merge() -> None:
+ sink = _NamedIdTraceSink(
+ {
+ "condition_raster": 10,
+ "model_generate": 11,
+ "cache_update": 12,
+ "decode": 0,
+ "frame_merge": 13,
+ }
+ )
+ trace_context = TraceContext.create(sink)
+ timings = VideoModelTimings(
+ condition_start_time=1.0,
+ condition_ready_time=1.1,
+ model_start_time=1.1,
+ model_ready_time=1.6,
+ merge_start_time=1.65,
+ merge_ready_time=1.7,
+ cache_update_start_time=1.2,
+ cache_update_ready_time=1.3,
+ decode_start_time=1.4,
+ decode_ready_time=1.55,
+ )
+
+ events = emit_video_model_timing_ranges(
+ trace_context,
+ timings=timings,
+ thread=trace_context.worker_thread,
+ chunk_index=9,
+ )
+
+ assert events.decode_event_id == 0
+ assert sink.events[4].depends_on == [0]
+
+
+def test_input_to_present_profile_window_returns_summary_on_interval() -> None:
+ window = InputToPresentProfileWindow(interval_s=0.25)
+
+ assert (
+ window.record(
+ present_time=1.0,
+ input_sample_time=0.9,
+ frame_index=0,
+ frame_interval_s=0.1,
+ )
+ is None
+ )
+ summary = window.record(
+ present_time=1.3,
+ input_sample_time=0.9,
+ frame_index=1,
+ frame_interval_s=0.1,
+ )
+
+ assert summary is not None
+ assert summary.samples == 2
+ assert summary.avg_raw_control_to_present_ms == pytest.approx(250.0)
+ assert summary.avg_adj_control_to_present_ms == pytest.approx(200.0)
+ assert "avg_raw_control_to_present_ms=250.00" in summary.log_message()
diff --git a/integrations/flashvsr/flashvsr/grpc/streaming_view.py b/integrations/flashvsr/flashvsr/grpc/streaming_view.py
index eec54f48b..0aceab1f9 100644
--- a/integrations/flashvsr/flashvsr/grpc/streaming_view.py
+++ b/integrations/flashvsr/flashvsr/grpc/streaming_view.py
@@ -27,6 +27,8 @@
import torch
from loguru import logger
+from flashdreams.serving.realtime.media import encode_rgb_frame_to_jpeg
+
DEFAULT_VIEWER_CHUNK_QUEUE_DEPTH = 8
DEFAULT_VIEWER_JPEG_QUALITY = 90
DEFAULT_VIEWER_JPEG_BACKEND = "auto"
@@ -58,10 +60,15 @@ def _load_pillow_image():
def _encode_jpeg_rgb(frame: np.ndarray, quality: int) -> bytes:
"""Encode one uint8 RGB frame to JPEG bytes."""
- image = _load_pillow_image().fromarray(np.ascontiguousarray(frame))
- buf = io.BytesIO()
- image.save(buf, format="JPEG", quality=quality)
- return buf.getvalue()
+ try:
+ return encode_rgb_frame_to_jpeg(frame, quality=quality, value_range="uint8")
+ except ModuleNotFoundError as exc:
+ if exc.name not in ("PIL", "PIL.Image"):
+ raise
+ raise RuntimeError(
+ "JPEG support requires Pillow. Install the flashvsr integration "
+ "dependencies before using JPEG input or the browser viewer."
+ ) from exc
def _load_torchvision_encode_jpeg():
diff --git a/integrations/flashvsr/flashvsr/grpc/uplift_client.py b/integrations/flashvsr/flashvsr/grpc/uplift_client.py
index 59d399f8d..3a4cbc7ea 100644
--- a/integrations/flashvsr/flashvsr/grpc/uplift_client.py
+++ b/integrations/flashvsr/flashvsr/grpc/uplift_client.py
@@ -33,7 +33,6 @@
"""
import argparse
-import io
import sys
import time
import uuid
@@ -44,6 +43,7 @@
import mediapy as media
import numpy as np
+from flashdreams.serving.realtime.media import encode_rgb_frame_to_jpeg
from flashvsr.grpc.protos import flashvsr_pb2 as pb2
from flashvsr.grpc.protos import flashvsr_pb2_grpc as pb2_grpc
@@ -106,21 +106,17 @@ def video_to_rgb_bytes(frames_np: np.ndarray) -> bytes:
def encode_jpeg_frames(frames_np: np.ndarray, quality: int) -> list[bytes]:
"""uint8 [T,H,W,3] → one JPEG byte string per frame."""
try:
- from PIL import Image
- except ImportError as exc:
+ return [
+ encode_rgb_frame_to_jpeg(frame, quality=quality, value_range="uint8")
+ for frame in frames_np
+ ]
+ except ModuleNotFoundError as exc:
+ if exc.name not in ("PIL", "PIL.Image"):
+ raise
raise RuntimeError(
"JPEG input requires Pillow in the client environment"
) from exc
- encoded = []
- for frame in frames_np:
- buf = io.BytesIO()
- Image.fromarray(np.ascontiguousarray(frame)).save(
- buf, format="JPEG", quality=quality
- )
- encoded.append(buf.getvalue())
- return encoded
-
def build_chunk_request(
*,
diff --git a/integrations/flashvsr/tests/test_grpc_client.py b/integrations/flashvsr/tests/test_grpc_client.py
index 501db5b23..d6a8994ea 100644
--- a/integrations/flashvsr/tests/test_grpc_client.py
+++ b/integrations/flashvsr/tests/test_grpc_client.py
@@ -17,9 +17,13 @@
import numpy as np
import pytest
-from flashvsr.grpc import uplift_server as grpc_server
from flashvsr.grpc.protos import flashvsr_pb2 as pb2
-from flashvsr.grpc.uplift_client import build_chunk_request, build_chunks
+from flashvsr.grpc.streaming_view import _encode_jpeg_rgb
+from flashvsr.grpc.uplift_client import (
+ build_chunk_request,
+ build_chunks,
+ encode_jpeg_frames,
+)
pytestmark = pytest.mark.ci_cpu
@@ -55,7 +59,30 @@ def test_build_chunk_request_raw_display_only() -> None:
assert request.display_only
+def test_encode_jpeg_frames_returns_jpeg_payloads() -> None:
+ pytest.importorskip("PIL.Image")
+ frames = np.zeros((2, 4, 6, 3), dtype=np.uint8)
+
+ encoded = encode_jpeg_frames(frames, quality=90)
+
+ assert len(encoded) == 2
+ assert all(frame.startswith(b"\xff\xd8") for frame in encoded)
+ assert all(frame.endswith(b"\xff\xd9") for frame in encoded)
+
+
+def test_streaming_view_cpu_jpeg_encoder_returns_jpeg_payload() -> None:
+ pytest.importorskip("PIL.Image")
+ frame = np.zeros((4, 6, 3), dtype=np.uint8)
+
+ encoded = _encode_jpeg_rgb(frame, quality=90)
+
+ assert encoded.startswith(b"\xff\xd8")
+ assert encoded.endswith(b"\xff\xd9")
+
+
def test_attention_mode_auto_uses_sparse(monkeypatch: pytest.MonkeyPatch) -> None:
+ from flashvsr.grpc import uplift_server as grpc_server
+
monkeypatch.setattr(grpc_server, "_sparse_attention_available", lambda: True)
assert grpc_server._resolve_attention_mode("auto") == "sparse"
@@ -66,6 +93,8 @@ def test_attention_mode_auto_uses_sparse(monkeypatch: pytest.MonkeyPatch) -> Non
def test_attention_mode_auto_falls_back_to_full_when_sparse_unavailable(
monkeypatch: pytest.MonkeyPatch,
) -> None:
+ from flashvsr.grpc import uplift_server as grpc_server
+
monkeypatch.setattr(
grpc_server,
"_sparse_attention_available",
diff --git a/integrations/omnidreams/omnidreams/grpc/utils.py b/integrations/omnidreams/omnidreams/grpc/utils.py
index bb850016c..7bbce05cc 100644
--- a/integrations/omnidreams/omnidreams/grpc/utils.py
+++ b/integrations/omnidreams/omnidreams/grpc/utils.py
@@ -48,6 +48,8 @@
from scipy.spatial.transform import Rotation, Slerp
from torch import Tensor
+from flashdreams.serving.realtime.media import encode_rgb_frame_to_jpeg
+
def decode_image(
image_bytes: bytes,
@@ -99,12 +101,12 @@ def encode_image(
Returns:
Encoded image bytes.
"""
+ if format.upper() == "JPEG":
+ return encode_rgb_frame_to_jpeg(image_np, quality=quality, value_range="uint8")
+
img = Image.fromarray(image_np)
buf = io.BytesIO()
- if format.upper() == "JPEG":
- img.save(buf, format="JPEG", quality=quality)
- else:
- img.save(buf, format=format)
+ img.save(buf, format=format)
return buf.getvalue()
diff --git a/integrations/omnidreams/omnidreams/interactive_drive/app.py b/integrations/omnidreams/omnidreams/interactive_drive/app.py
index 8ac0b0ff9..0fccebabb 100644
--- a/integrations/omnidreams/omnidreams/interactive_drive/app.py
+++ b/integrations/omnidreams/omnidreams/interactive_drive/app.py
@@ -21,7 +21,6 @@
PresenterBackend,
run_main_loop,
)
-from omnidreams.interactive_drive.runtime.timing import TraceContext, TraceSink
from omnidreams.interactive_drive.scene_loader import (
load_scene_bundle,
reseed_scene_bundle,
@@ -42,6 +41,8 @@
from omnidreams.interactive_drive.video_model.chunk_pipeline import ChunkPipeline
from omnidreams.interactive_drive.video_model.local import LocalVideoModelAdapter
+from flashdreams.serving.realtime.timing import TraceContext, TraceSink
+
# Cadence for the event-pump loop that keeps the presenter alive while a
# scene parses on a background thread. ~60 Hz keeps input latency low and
# the loading indicator smooth without burning a core.
diff --git a/integrations/omnidreams/omnidreams/interactive_drive/cli.py b/integrations/omnidreams/omnidreams/interactive_drive/cli.py
index f4149bbd5..2313aabe4 100644
--- a/integrations/omnidreams/omnidreams/interactive_drive/cli.py
+++ b/integrations/omnidreams/omnidreams/interactive_drive/cli.py
@@ -21,11 +21,12 @@
WorldModelProfileConfig,
)
from omnidreams.interactive_drive.log import configure_logging
-from omnidreams.interactive_drive.runtime.timing import TraceSink
from omnidreams.interactive_drive.synthetic_scene import build_synthetic_scene_to_temp
from omnidreams.interactive_drive.world_model.manifest import load_world_model_manifest
from omnidreams.scenes import local_scene_archive_path
+from flashdreams.serving.realtime.timing import TraceSink
+
# Package root (from this file's location) so packaged-asset defaults below
# resolve relative to the install, not the user's cwd. Bundled configs live at
# ``interactive_drive/configs/``; scene USDZs are staged into
diff --git a/integrations/omnidreams/omnidreams/interactive_drive/input/keyboard.py b/integrations/omnidreams/omnidreams/interactive_drive/input/keyboard.py
index 18c610875..a8083a659 100644
--- a/integrations/omnidreams/omnidreams/interactive_drive/input/keyboard.py
+++ b/integrations/omnidreams/omnidreams/interactive_drive/input/keyboard.py
@@ -11,6 +11,14 @@
VehicleState,
)
+from flashdreams.serving.realtime.input import (
+ DRIVING_SUPPORTED_KEYS,
+ normalize_key,
+)
+from flashdreams.serving.realtime.input import (
+ KeyboardState as RealtimeKeyboardState,
+)
+
class KeyboardState:
"""Owns live keyboard state plus the runtime UI affordances the loop reads.
@@ -26,7 +34,7 @@ class KeyboardState:
def __init__(self) -> None:
self._lock = threading.Lock()
- self._pressed: set[str] = set()
+ self._keyboard = RealtimeKeyboardState(supported_keys=DRIVING_SUPPORTED_KEYS)
self._view_mode = "rgb"
self._drive_command: DriverCommand | None = None
self._reset_pending = False
@@ -44,10 +52,10 @@ def __init__(self) -> None:
def set_key(self, name: str, down: bool) -> None:
with self._lock:
- if down:
- self._pressed.add(name)
- else:
- self._pressed.discard(name)
+ self._keyboard.apply_event(
+ event="keydown" if down else "keyup",
+ key=name,
+ )
def set_view_mode(self, mode: str) -> None:
with self._lock:
@@ -113,7 +121,7 @@ def view_mode(self) -> str:
def command(self) -> DriverCommand:
with self._lock:
drive_command = self._drive_command
- pressed = set(self._pressed)
+ pressed = set(self._keyboard.snapshot())
if drive_command is not None:
if "space" in pressed:
return DriverCommand(
@@ -130,18 +138,19 @@ def command(self) -> DriverCommand:
def command_from_snapshot(snapshot: ControlSnapshot) -> DriverCommand:
- throttle = 1.0 if {"w", "up"} & snapshot.pressed else 0.0
- brake = 1.0 if {"s", "down"} & snapshot.pressed else 0.0
+ pressed = {normalize_key(key) for key in snapshot.pressed}
+ throttle = 1.0 if {"w", "up"} & pressed else 0.0
+ brake = 1.0 if {"s", "down"} & pressed else 0.0
steer = 0.0
- if {"a", "left"} & snapshot.pressed:
+ if {"a", "left"} & pressed:
steer += 1.0
- if {"d", "right"} & snapshot.pressed:
+ if {"d", "right"} & pressed:
steer -= 1.0
return DriverCommand(
throttle=throttle,
brake=brake,
steer=steer,
- stop="space" in snapshot.pressed,
+ stop="space" in pressed,
)
diff --git a/integrations/omnidreams/omnidreams/interactive_drive/runtime/loop.py b/integrations/omnidreams/omnidreams/interactive_drive/runtime/loop.py
index 581b07709..e0e309e3d 100644
--- a/integrations/omnidreams/omnidreams/interactive_drive/runtime/loop.py
+++ b/integrations/omnidreams/omnidreams/interactive_drive/runtime/loop.py
@@ -2,9 +2,7 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
import os
-import queue
import time
-from collections import deque
from collections.abc import Callable
from dataclasses import dataclass, replace
from typing import Protocol
@@ -12,32 +10,35 @@
from loguru import logger
from omnidreams.interactive_drive.input.backend import InputBackend
from omnidreams.interactive_drive.runtime.runtime_controls import RuntimeControls
-from omnidreams.interactive_drive.runtime.timing import (
+from omnidreams.interactive_drive.simulation.backend import SimulationBackend
+from omnidreams.interactive_drive.types import DriverCommand, PresentedFrame
+from omnidreams.interactive_drive.video_model.chunk_pipeline import (
+ ChunkPipeline,
+ ChunkRequest,
+ QueuedFrame,
+)
+
+from flashdreams.serving.realtime.presenter import (
+ PresentationQueue,
+ wait_until_present_time,
+)
+from flashdreams.serving.realtime.timing import (
ChunkHistory,
ChunkPrediction,
ChunkTimes,
+ InputToPresentProfileWindow,
TraceComponentValue,
TraceContext,
event_dependencies,
trace_time_ns,
)
-from omnidreams.interactive_drive.simulation.backend import SimulationBackend
-from omnidreams.interactive_drive.types import DriverCommand, PresentedFrame
-from omnidreams.interactive_drive.video_model.chunk_pipeline import (
- ChunkPipeline,
- ChunkRequest,
- QueuedFrame,
-)
_PROFILE_INPUT_TO_PRESENT_ENV = "INTERACTIVE_DRIVE_PROFILE_INPUT_TO_PRESENT"
_PROFILE_INPUT_TO_PRESENT_INTERVAL_S_ENV = (
"INTERACTIVE_DRIVE_PROFILE_INPUT_TO_PRESENT_INTERVAL_S"
)
-_PROFILE_E2E_SUM_RAW_MS: float = 0.0
-_PROFILE_E2E_SUM_ADJ_MS: float = 0.0
-_PROFILE_E2E_COUNT: int = 0
-_PROFILE_E2E_WINDOW_START: float | None = None
+_PROFILE_E2E_WINDOW = InputToPresentProfileWindow()
def _profile_input_to_present_enabled() -> bool:
@@ -56,15 +57,7 @@ def _profile_input_to_present_interval_s() -> float:
def reset_input_to_present_profile_window() -> None:
"""Clear accumulated e2e samples when the main loop starts."""
- global _PROFILE_E2E_SUM_RAW_MS
- global _PROFILE_E2E_SUM_ADJ_MS
- global _PROFILE_E2E_COUNT
- global _PROFILE_E2E_WINDOW_START
-
- _PROFILE_E2E_SUM_RAW_MS = 0.0
- _PROFILE_E2E_SUM_ADJ_MS = 0.0
- _PROFILE_E2E_COUNT = 0
- _PROFILE_E2E_WINDOW_START = None
+ _PROFILE_E2E_WINDOW.reset(interval_s=_profile_input_to_present_interval_s())
def _chunk_frame_interval_s(chunk_times: ChunkTimes) -> float:
@@ -83,42 +76,15 @@ def _record_input_to_present_for_profile(
frame_index: int,
frame_interval_s: float,
) -> None:
- global _PROFILE_E2E_SUM_RAW_MS
- global _PROFILE_E2E_SUM_ADJ_MS
- global _PROFILE_E2E_COUNT
- global _PROFILE_E2E_WINDOW_START
-
- raw_ms = (present_time - input_sample_time) * 1000.0
- scheduled_ms = frame_index * (frame_interval_s * 1000.0)
- adj_ms = raw_ms - scheduled_ms
- _PROFILE_E2E_SUM_RAW_MS += raw_ms
- _PROFILE_E2E_SUM_ADJ_MS += adj_ms
- _PROFILE_E2E_COUNT += 1
- if _PROFILE_E2E_WINDOW_START is None:
- _PROFILE_E2E_WINDOW_START = present_time
-
- interval_s = _profile_input_to_present_interval_s()
- if present_time - _PROFILE_E2E_WINDOW_START < interval_s:
- return
-
- count = _PROFILE_E2E_COUNT
- if count <= 0:
- return
- window_s = present_time - _PROFILE_E2E_WINDOW_START
- wall_present_fps = float(count) / window_s if window_s > 1e-9 else 0.0
- avg_raw_ms = _PROFILE_E2E_SUM_RAW_MS / float(count)
- avg_adj_ms = _PROFILE_E2E_SUM_ADJ_MS / float(count)
- logger.info(
- "[profile] e2e "
- f"wall_present_fps={wall_present_fps:.1f} "
- f"avg_adj_control_to_present_ms={avg_adj_ms:.2f} "
- f"avg_raw_control_to_present_ms={avg_raw_ms:.2f} "
- f"samples={count}",
+ _PROFILE_E2E_WINDOW.interval_s = _profile_input_to_present_interval_s()
+ summary = _PROFILE_E2E_WINDOW.record(
+ present_time=present_time,
+ input_sample_time=input_sample_time,
+ frame_index=frame_index,
+ frame_interval_s=frame_interval_s,
)
- _PROFILE_E2E_SUM_RAW_MS = 0.0
- _PROFILE_E2E_SUM_ADJ_MS = 0.0
- _PROFILE_E2E_COUNT = 0
- _PROFILE_E2E_WINDOW_START = present_time
+ if summary is not None:
+ logger.info(summary.log_message())
class PresenterBackend(Protocol):
@@ -434,22 +400,18 @@ def _prepare_queued_frame(
def _drain_pipeline_frames(
*,
pipeline: ChunkPipeline,
- ready_frames: "deque[QueuedFrame]",
+ ready_frames: PresentationQueue[QueuedFrame],
presenter: PresenterBackend,
view_mode: str,
) -> None:
current_generation = pipeline.current_generation
- while True:
- try:
- queued_frame = pipeline.frame_queue.get_nowait()
- except queue.Empty:
- return
- if queued_frame.generation != current_generation:
- # Stale frame from a superseded rollout/scene (generation bumped);
- # drop it so old content isn't flashed over the new load.
- continue
- _prepare_queued_frame(queued_frame, presenter, view_mode)
- ready_frames.append(queued_frame)
+ ready_frames.drain_nowait(
+ pipeline.frame_queue,
+ include=lambda queued_frame: queued_frame.generation == current_generation,
+ prepare=lambda queued_frame: _prepare_queued_frame(
+ queued_frame, presenter, view_mode
+ ),
+ )
def run_main_loop(
@@ -478,7 +440,7 @@ def run_main_loop(
"""
state = MainLoopState()
last_presented_frame: PresentedFrame = initial_presented_frame
- ready_frames: deque[QueuedFrame] = deque()
+ ready_frames = PresentationQueue[QueuedFrame]()
chunk_history = ChunkHistory(config.history_capacity)
last_input_sample_event: int | None = None
last_present_wait_event: int | None = None
@@ -535,18 +497,16 @@ def run_main_loop(
view_mode=view_mode,
)
- now = time.perf_counter()
- if now < state.next_present_time:
- wait_begin = now
- time.sleep(
- min(config.poll_timeout_s, max(0.0, state.next_present_time - now))
- )
- wait_end = time.perf_counter()
+ present_wait = wait_until_present_time(
+ state.next_present_time,
+ poll_timeout_s=config.poll_timeout_s,
+ )
+ if present_wait is not None:
last_present_wait_event = _trace_main_range(
active_trace,
"present_wait",
- begin_time=wait_begin,
- end_time=wait_end,
+ begin_time=present_wait.begin_time,
+ end_time=present_wait.end_time,
depends_on=[],
)
continue
diff --git a/integrations/omnidreams/omnidreams/interactive_drive/runtime/timing.py b/integrations/omnidreams/omnidreams/interactive_drive/runtime/timing.py
index 2fa0dea3f..cb1b0af4d 100644
--- a/integrations/omnidreams/omnidreams/interactive_drive/runtime/timing.py
+++ b/integrations/omnidreams/omnidreams/interactive_drive/runtime/timing.py
@@ -1,186 +1,50 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
-"""Timing records for latency measurement.
-
-``FrameTimes`` / ``ChunkTimes`` are mutable on purpose: a single instance
-travels through the pipeline accumulating per-stage timestamps, so the record
-created at request time is the same one that records present time (direct
-correlation, no copying or re-association).
-"""
-
-from collections import deque
-from dataclasses import dataclass
-from threading import Lock
-from typing import Protocol
-
-TraceComponentValue = str | int | float | bool | None
-
-
-def trace_time_ns(seconds: float) -> int:
- return int(seconds * 1_000_000_000)
-
-
-def event_dependencies(*events: int | None) -> list[int]:
- return [event for event in events if event is not None]
-
-
-@dataclass
-class FrameTimes:
- frame_index: int
- intended_present_time: float
- image_ready_time: float | None = None
- sample_display_pose_time: float | None = None
- present_time: float | None = None
-
-
-@dataclass(frozen=True)
-class ChunkPrediction:
- """Predicted timestamps for a chunk's pipeline stages.
-
- Stage 1 only predicts ``first_present`` (when the chunk's first frame
- will reach the screen). The Stage 2 design adds intermediate
- EMA-summed milestones (``request -> render_start -> chunk_ready ->
- decode_first -> first_present``) per ``alpasim.frame_timing``; new
- fields land here as named attributes when that work arrives.
-
- Use :meth:`create` rather than constructing directly: the prediction
- formula lives there so it stays beside the data it produces.
- """
-
- first_present: float
-
- @classmethod
- def create(
- cls, *, request_time: float, frame_interval_s: float
- ) -> "ChunkPrediction":
- """Stage 1 prediction: ``first_present = request_time + frame_interval_s``.
-
- Placeholder for the Stage 2 EMA-summed prediction chain; the
- signature stays the same so callers don't change when Stage 2
- lands and the body grows to consume EMA latency stats.
- """
- return cls(first_present=request_time + frame_interval_s)
-
-
-@dataclass
-class ChunkTimes:
- chunk_index: int
- input_sample_time: float
- request_time: float
- request_poses_ready_time: float
- frames: list[FrameTimes]
- prediction: ChunkPrediction | None = None
- chunk_render_start_time: float | None = None
- chunk_ready_time: float | None = None
-
- @classmethod
- def create(
- cls,
- chunk_index: int,
- input_sample_time: float,
- request_time: float,
- request_poses_ready_time: float,
- intended_present_times: list[float],
- prediction: ChunkPrediction | None = None,
- ) -> "ChunkTimes":
- frames = [
- FrameTimes(frame_index=index, intended_present_time=time_value)
- for index, time_value in enumerate(intended_present_times)
- ]
- return cls(
- chunk_index=chunk_index,
- input_sample_time=input_sample_time,
- request_time=request_time,
- request_poses_ready_time=request_poses_ready_time,
- frames=frames,
- prediction=prediction,
- )
-
-
-class ChunkHistory:
- def __init__(self, capacity: int) -> None:
- self._deque: deque[ChunkTimes] = deque(maxlen=capacity)
-
- def append(self, chunk: ChunkTimes) -> None:
- self._deque.append(chunk)
-
-
-class TraceSink(Protocol):
- def add_thread(self, name: str) -> int: ...
-
- def add_instant(
- self,
- name: str,
- *,
- thread: int,
- time_ns: int,
- depends_on: list[int] | None = None,
- **components: TraceComponentValue,
- ) -> int: ...
-
- def add_range(
- self,
- name: str,
- *,
- thread: int,
- begin_ns: int,
- end_ns: int,
- depends_on: list[int] | None = None,
- **components: TraceComponentValue,
- ) -> int: ...
-
-
-@dataclass(frozen=True)
-class TraceContext:
- sink: TraceSink
- main_thread: int
- worker_thread: int
- lock: Lock
-
- @classmethod
- def create(cls, sink: TraceSink) -> "TraceContext":
- return cls(
- sink=sink,
- main_thread=sink.add_thread("main"),
- worker_thread=sink.add_thread("pipeline-worker"),
- lock=Lock(),
- )
-
- def add_instant(
- self,
- name: str,
- *,
- thread: int,
- time_ns: int,
- depends_on: list[int] | None = None,
- **components: TraceComponentValue,
- ) -> int:
- with self.lock:
- return self.sink.add_instant(
- name,
- thread=thread,
- time_ns=time_ns,
- depends_on=depends_on,
- **components,
- )
-
- def add_range(
- self,
- name: str,
- *,
- thread: int,
- begin_ns: int,
- end_ns: int,
- depends_on: list[int] | None = None,
- **components: TraceComponentValue,
- ) -> int:
- with self.lock:
- return self.sink.add_range(
- name,
- thread=thread,
- begin_ns=begin_ns,
- end_ns=end_ns,
- depends_on=depends_on,
- **components,
- )
+"""Compatibility exports for shared realtime timing helpers."""
+
+from flashdreams.serving.realtime.timing import (
+ ChunkHistory,
+ ChunkPrediction,
+ ChunkTimes,
+ FrameTimes,
+ InputToPresentProfileWindow,
+ InputToPresentSummary,
+ RecentTimingSummary,
+ RollingChunkTimingSummary,
+ StageDurationSummary,
+ TraceComponentValue,
+ TraceContext,
+ TraceSink,
+ VideoModelTimings,
+ VideoModelTraceEvents,
+ chunk_stage_durations_ms,
+ emit_video_model_timing_ranges,
+ event_dependencies,
+ summarize_chunk_history,
+ summarize_stage_durations,
+ trace_time_ns,
+)
+
+__all__ = [
+ "ChunkHistory",
+ "ChunkPrediction",
+ "ChunkTimes",
+ "FrameTimes",
+ "InputToPresentProfileWindow",
+ "InputToPresentSummary",
+ "RecentTimingSummary",
+ "RollingChunkTimingSummary",
+ "StageDurationSummary",
+ "TraceComponentValue",
+ "TraceContext",
+ "TraceSink",
+ "VideoModelTimings",
+ "VideoModelTraceEvents",
+ "chunk_stage_durations_ms",
+ "emit_video_model_timing_ranges",
+ "event_dependencies",
+ "summarize_chunk_history",
+ "summarize_stage_durations",
+ "trace_time_ns",
+]
diff --git a/integrations/omnidreams/omnidreams/interactive_drive/streaming_presenter.py b/integrations/omnidreams/omnidreams/interactive_drive/streaming_presenter.py
index 8de0cd46e..aeed422d4 100644
--- a/integrations/omnidreams/omnidreams/interactive_drive/streaming_presenter.py
+++ b/integrations/omnidreams/omnidreams/interactive_drive/streaming_presenter.py
@@ -10,7 +10,6 @@
from __future__ import annotations
-import io
import json
import shutil
import subprocess
@@ -28,7 +27,17 @@
from omnidreams.interactive_drive.input.keyboard import KeyboardState
from omnidreams.interactive_drive.loading_overlay import render_loading_overlay
from omnidreams.interactive_drive.types import DriverCommand, PresentedFrame
-from PIL import Image
+
+from flashdreams.serving.realtime.frame_bus import LatestFrameBus
+from flashdreams.serving.realtime.media import encode_rgb_frame_to_jpeg
+from flashdreams.serving.realtime.mjpeg import (
+ WaitForMjpegFrame,
+ publish_latest_jpeg,
+ send_mjpeg_response_headers,
+ wait_for_latest_jpeg,
+ write_mjpeg_stream,
+)
+from flashdreams.serving.realtime.presenter import materialize_rgb_host_uint8
# Boundary marker embedded in the multipart response. The exact string
# doesn't matter as long as it never appears inside a JPEG payload (they
@@ -540,19 +549,11 @@ def __init__(
self._keyboard = keyboard
self._jpeg_quality = int(jpeg_quality)
self._stop_event = threading.Event()
- # Guarded by ``_frame_cond`` so a sending thread can ``wait()``
- # for the next frame rather than spinning.
- self._latest_jpeg: bytes | None = None
- self._frame_count = 0
- self._frame_cond = threading.Condition()
+ self._frame_bus = LatestFrameBus[bytes]()
# BEV minimap stream lives on its own JPEG buffer so connected
# clients of /bev_stream can paginate at a different rate than
- # /stream (e.g. if the HUD process throttles). We reuse the same
- # condition variable as the main stream because frames are only
- # published when ``present_frame`` runs anyway, so notifications
- # to either waiter are always safe.
- self._latest_bev_jpeg: bytes | None = None
- self._bev_frame_count = 0
+ # /stream (e.g. if the HUD process throttles).
+ self._bev_frame_bus = LatestFrameBus[bytes]()
# Scene options surfaced to the browser dropdown via /scenes.
# Each entry is a dict with ``label``, ``path``, ``variants``;
# the demo wrapper builds these from its scene-discovery layer
@@ -736,10 +737,8 @@ def present_frame(self, frame: PresentedFrame, view_mode: str) -> None:
def close(self) -> None:
self._stop_event.set()
- # Wake any /stream handlers blocked in ``_frame_cond.wait`` so
- # they observe ``should_close`` and exit their per-connection loop.
- with self._frame_cond:
- self._frame_cond.notify_all()
+ self._frame_bus.close()
+ self._bev_frame_bus.close()
self._server.shutdown()
self._server.server_close()
if self._server_thread.is_alive():
@@ -748,15 +747,12 @@ def close(self) -> None:
# -- Internals --------------------------------------------------
def _publish(self, rgb_host_uint8: object) -> None:
- buf = io.BytesIO()
- Image.fromarray(_as_rgb_host_uint8(rgb_host_uint8)).save(
- buf, format="JPEG", quality=self._jpeg_quality
+ jpeg = encode_rgb_frame_to_jpeg(
+ _as_rgb_host_uint8(rgb_host_uint8),
+ quality=self._jpeg_quality,
+ value_range="uint8",
)
- jpeg = buf.getvalue()
- with self._frame_cond:
- self._latest_jpeg = jpeg
- self._frame_count += 1
- self._frame_cond.notify_all()
+ _publish_if_open(self._frame_bus, jpeg, stop_event=self._stop_event)
def _publish_bev(self, bev_rgb_host_uint8: object) -> None:
"""Encode the BEV minimap (quality 95, not 85) and stash it for ``/bev_stream``.
@@ -764,44 +760,34 @@ def _publish_bev(self, bev_rgb_host_uint8: object) -> None:
Higher quality avoids JPEG ringing the HUD's Google-Maps filter would
otherwise surface as grey halos around lane / vehicle edges.
"""
- buf = io.BytesIO()
- Image.fromarray(_as_rgb_host_uint8(bev_rgb_host_uint8)).save(
- buf, format="JPEG", quality=95
+ jpeg = encode_rgb_frame_to_jpeg(
+ _as_rgb_host_uint8(bev_rgb_host_uint8),
+ quality=95,
+ value_range="uint8",
)
- jpeg = buf.getvalue()
- with self._frame_cond:
- self._latest_bev_jpeg = jpeg
- self._bev_frame_count += 1
- self._frame_cond.notify_all()
+ _publish_if_open(self._bev_frame_bus, jpeg, stop_event=self._stop_event)
def _wait_for_new_frame(self, last_seen_count: int) -> tuple[bytes, int] | None:
"""Block until a frame newer than ``last_seen_count`` is ready or
the server is shutting down. Returns ``(jpeg_bytes, frame_count)``
on success, ``None`` when closing.
"""
- with self._frame_cond:
- while self._latest_jpeg is None or self._frame_count <= last_seen_count:
- if self._stop_event.is_set():
- return None
- self._frame_cond.wait(timeout=1.0)
- return self._latest_jpeg, self._frame_count
+ return _wait_for_bus_frame(
+ self._frame_bus,
+ last_seen_count=last_seen_count,
+ stop_event=self._stop_event,
+ )
def _wait_for_new_bev_frame(self, last_seen_count: int) -> tuple[bytes, int] | None:
"""Same as :meth:`_wait_for_new_frame` but for the BEV stream.
- Returns ``None`` when the server is closing. Sharing the condition
- variable means the waiter wakes immediately on every published
- frame; the loop body then re-checks the BEV-specific counter.
+ Returns ``None`` when the server is closing.
"""
- with self._frame_cond:
- while (
- self._latest_bev_jpeg is None
- or self._bev_frame_count <= last_seen_count
- ):
- if self._stop_event.is_set():
- return None
- self._frame_cond.wait(timeout=1.0)
- return self._latest_bev_jpeg, self._bev_frame_count
+ return _wait_for_bus_frame(
+ self._bev_frame_bus,
+ last_seen_count=last_seen_count,
+ stop_event=self._stop_event,
+ )
def _apply_control(self, key: str, down: bool) -> None:
# Direction keys (W/A/S/D + arrows + Space) flow through the
@@ -867,24 +853,11 @@ def _request_scene_change(self, scene_path_str: str, variant: str) -> bool:
resolved_variant = (
variant if variant in entry_variants else entry_variants[0]
)
- # Wake any handlers waiting on the frame condition so
- # they observe ``should_close`` flipping and exit their
- # per-connection loop promptly. Not strictly required
- # for correctness (the existing 1 s timeout would
- # eventually retry) but it makes the scene transition
- # feel snappier.
self._pending_scene_change = (Path(entry_path), str(resolved_variant))
- with self._frame_cond:
- self._frame_cond.notify_all()
return True
return False
-# Type alias for ``_serve_mjpeg``'s blocking getter parameter. ``None``
-# means the server is shutting down; ``(jpeg, count)`` is a fresh frame.
-_WaitForFrame = Callable[[int], tuple[bytes, int] | None]
-
-
def _make_handler(presenter: MJPEGStreamingPresenter) -> type[BaseHTTPRequestHandler]:
"""Build a BaseHTTPRequestHandler subclass closed over ``presenter``.
@@ -1007,48 +980,19 @@ def _serve_stream(self) -> None:
def _serve_bev_stream(self) -> None:
self._serve_mjpeg(presenter._wait_for_new_bev_frame)
- def _serve_mjpeg(self, wait_fn: _WaitForFrame) -> None:
+ def _serve_mjpeg(self, wait_fn: WaitForMjpegFrame) -> None:
"""Generic ``multipart/x-mixed-replace`` writer used by /stream and
/bev_stream. ``wait_fn(last_seen)`` is the per-stream blocking
getter that returns ``(jpeg, frame_count)`` or ``None`` on
shutdown.
"""
- self.send_response(HTTPStatus.OK)
- self.send_header(
- "Cache-Control", "no-store, no-cache, must-revalidate, max-age=0"
- )
- self.send_header("Pragma", "no-cache")
- self.send_header(
- "Content-Type",
- f"multipart/x-mixed-replace; boundary={_MULTIPART_BOUNDARY}",
- )
- self.end_headers()
- last_seen = 0
- try:
- # Loop until shutdown (``wait_fn`` returns None only on
- # ``_stop_event``). NOT gated on ``should_close``: that also
- # flips True on a pending scene/variant change, and closing the
- # connection there would freeze the browser's multipart
- # (it never auto-reconnects) mid-switch.
- while True:
- result = wait_fn(last_seen)
- if result is None:
- break
- jpeg, last_seen = result
- part = (
- (
- f"--{_MULTIPART_BOUNDARY}\r\n"
- f"Content-Type: image/jpeg\r\n"
- f"Content-Length: {len(jpeg)}\r\n\r\n"
- ).encode("ascii")
- + jpeg
- + b"\r\n"
- )
- self.wfile.write(part)
- self.wfile.flush()
- except (BrokenPipeError, ConnectionResetError):
- # Client disconnected; that's normal, not an error.
- return
+ send_mjpeg_response_headers(self, boundary=_MULTIPART_BOUNDARY)
+ # Loop until shutdown (``wait_fn`` returns None only on
+ # ``_stop_event``). NOT gated on ``should_close``: that also
+ # flips True on a pending scene/variant change, and closing the
+ # connection there would freeze the browser's multipart
+ # (it never auto-reconnects) mid-switch.
+ write_mjpeg_stream(self.wfile, wait_fn, boundary=_MULTIPART_BOUNDARY)
def _serve_control(self, query: dict[str, list[str]]) -> None:
key = query.get("key", [""])[0]
@@ -1070,13 +1014,29 @@ def _as_rgb_host_uint8(frame: object) -> np.ndarray:
"""Materialize a frame to ``(H, W, 3)`` uint8.
World-model frames are lazy GPU handles (``_LazyRGBFrame``) with
- ``to_numpy()`` but no ``__array_interface__``, so ``Image.fromarray`` can't
- take them directly. Mirrors the slangpy presenter.
+ ``to_numpy()`` but no ``__array_interface__``, so shared media helpers
+ need an explicit materialization step. Mirrors the slangpy presenter.
"""
- to_numpy = getattr(frame, "to_numpy", None)
- if callable(to_numpy):
- frame = to_numpy()
- return np.ascontiguousarray(np.asarray(frame, dtype=np.uint8)[..., :3])
+ return materialize_rgb_host_uint8(frame)
+
+
+def _publish_if_open(
+ bus: LatestFrameBus[bytes], jpeg: bytes, *, stop_event: threading.Event
+) -> None:
+ publish_latest_jpeg(bus, jpeg, stop_event=stop_event)
+
+
+def _wait_for_bus_frame(
+ bus: LatestFrameBus[bytes],
+ *,
+ last_seen_count: int,
+ stop_event: threading.Event,
+) -> tuple[bytes, int] | None:
+ return wait_for_latest_jpeg(
+ bus,
+ last_seen_count=last_seen_count,
+ stop_event=stop_event,
+ )
def _with_status_overlay(rgb_host_uint8: object, message: str | None) -> np.ndarray:
diff --git a/integrations/omnidreams/omnidreams/interactive_drive/types.py b/integrations/omnidreams/omnidreams/interactive_drive/types.py
index 1ffbabfcb..efac4f5ea 100644
--- a/integrations/omnidreams/omnidreams/interactive_drive/types.py
+++ b/integrations/omnidreams/omnidreams/interactive_drive/types.py
@@ -11,6 +11,8 @@
import numpy as np
import numpy.typing as npt
+from flashdreams.serving.realtime.timing import VideoModelTimings
+
FloatArray = npt.NDArray[np.float32]
UInt8Array = npt.NDArray[np.uint8]
Int32Array = npt.NDArray[np.int32]
@@ -207,16 +209,6 @@ class PresentedFrame:
status_message: str | None = None
-@dataclass(frozen=True)
-class VideoModelTimings:
- condition_start_time: float
- condition_ready_time: float
- model_start_time: float
- model_ready_time: float
- merge_start_time: float
- merge_ready_time: float
-
-
@dataclass(frozen=True)
class FrameChunk:
frames: tuple[PresentedFrame, ...]
diff --git a/integrations/omnidreams/omnidreams/interactive_drive/video_model/chunk_pipeline.py b/integrations/omnidreams/omnidreams/interactive_drive/video_model/chunk_pipeline.py
index ec45e0812..a1799fcfb 100644
--- a/integrations/omnidreams/omnidreams/interactive_drive/video_model/chunk_pipeline.py
+++ b/integrations/omnidreams/omnidreams/interactive_drive/video_model/chunk_pipeline.py
@@ -9,12 +9,6 @@
from typing import Protocol
from loguru import logger
-from omnidreams.interactive_drive.runtime.timing import (
- ChunkTimes,
- TraceContext,
- event_dependencies,
- trace_time_ns,
-)
from omnidreams.interactive_drive.types import (
FrameChunk,
PresentedFrame,
@@ -22,6 +16,14 @@
TrajectoryChunk,
)
+from flashdreams.serving.realtime.timing import (
+ ChunkTimes,
+ TraceContext,
+ emit_video_model_timing_ranges,
+ event_dependencies,
+ trace_time_ns,
+)
+
class VideoModelBackend(Protocol):
"""Video-model interface called from the pipeline worker thread.
@@ -238,30 +240,14 @@ def render_command(backend: VideoModelBackend) -> bool:
worker_ready_event_id = chunk_render_event
timings = frame_chunk.video_model_timings
if timings is not None:
- condition_event = trace_context.add_range(
- "condition_raster",
+ timing_events = emit_video_model_timing_ranges(
+ trace_context,
+ timings=timings,
thread=trace_context.worker_thread,
- begin_ns=trace_time_ns(timings.condition_start_time),
- end_ns=trace_time_ns(timings.condition_ready_time),
depends_on=event_dependencies(queue_wait_event),
chunk_index=chunk_times.chunk_index,
)
- model_event = trace_context.add_range(
- "model_generate",
- thread=trace_context.worker_thread,
- begin_ns=trace_time_ns(timings.model_start_time),
- end_ns=trace_time_ns(timings.model_ready_time),
- depends_on=event_dependencies(condition_event),
- chunk_index=chunk_times.chunk_index,
- )
- worker_ready_event_id = trace_context.add_range(
- "frame_merge",
- thread=trace_context.worker_thread,
- begin_ns=trace_time_ns(timings.merge_start_time),
- end_ns=trace_time_ns(timings.merge_ready_time),
- depends_on=event_dependencies(model_event),
- chunk_index=chunk_times.chunk_index,
- )
+ worker_ready_event_id = timing_events.final_event_id
# Drop the output if a reset / scene switch superseded this chunk
# while it was queued or rendering -- its frames belong to a
# rollout the user has already moved on from.
diff --git a/integrations/omnidreams/tests/interactive_drive/test_chunk_pipeline.py b/integrations/omnidreams/tests/interactive_drive/test_chunk_pipeline.py
index 5f1379f9a..02c327a24 100644
--- a/integrations/omnidreams/tests/interactive_drive/test_chunk_pipeline.py
+++ b/integrations/omnidreams/tests/interactive_drive/test_chunk_pipeline.py
@@ -12,17 +12,18 @@
make_trajectory,
minimal_scene,
)
-from omnidreams.interactive_drive.runtime.timing import (
- ChunkTimes,
- TraceComponentValue,
- TraceContext,
-)
from omnidreams.interactive_drive.types import FrameChunk, PresentedFrame, SceneBundle
from omnidreams.interactive_drive.video_model.chunk_pipeline import (
ChunkPipeline,
ChunkRequest,
)
+from flashdreams.serving.realtime.timing import (
+ ChunkTimes,
+ TraceComponentValue,
+ TraceContext,
+)
+
@dataclass(frozen=True)
class _TraceEvent:
diff --git a/integrations/omnidreams/tests/interactive_drive/test_keyboard_state.py b/integrations/omnidreams/tests/interactive_drive/test_keyboard_state.py
index eb08877a4..4f1203730 100644
--- a/integrations/omnidreams/tests/interactive_drive/test_keyboard_state.py
+++ b/integrations/omnidreams/tests/interactive_drive/test_keyboard_state.py
@@ -46,6 +46,17 @@ def test_view_mode_reflects_set_view_mode() -> None:
assert keyboard.view_mode == "hdmap"
+def test_keyboard_state_uses_shared_key_normalization() -> None:
+ keyboard = KeyboardState()
+ keyboard.set_key("ArrowUp", True)
+ keyboard.set_key("ArrowLeft", True)
+
+ command = keyboard.command()
+
+ assert command.throttle == 1.0
+ assert command.steer == 1.0
+
+
def test_consume_exit_scene_request_returns_false_when_none_pending() -> None:
keyboard = KeyboardState()
assert keyboard.consume_exit_scene_request() is False
diff --git a/integrations/omnidreams/tests/interactive_drive/test_latency_loop.py b/integrations/omnidreams/tests/interactive_drive/test_latency_loop.py
index f7140b20a..190b6fba4 100644
--- a/integrations/omnidreams/tests/interactive_drive/test_latency_loop.py
+++ b/integrations/omnidreams/tests/interactive_drive/test_latency_loop.py
@@ -20,12 +20,6 @@
present_queued_frame,
run_main_loop,
)
-from omnidreams.interactive_drive.runtime.timing import (
- ChunkPrediction,
- ChunkTimes,
- TraceComponentValue,
- TraceContext,
-)
from omnidreams.interactive_drive.types import (
DriverCommand,
PresentedFrame,
@@ -37,6 +31,13 @@
QueuedFrame,
)
+from flashdreams.serving.realtime.timing import (
+ ChunkPrediction,
+ ChunkTimes,
+ TraceComponentValue,
+ TraceContext,
+)
+
def _on_ci() -> bool:
"""True when running under CI (GitHub Actions et al. set ``CI=true``)."""
diff --git a/integrations/omnidreams/tests/interactive_drive/test_latency_timing.py b/integrations/omnidreams/tests/interactive_drive/test_latency_timing.py
index feacd19e3..bdefc5f09 100644
--- a/integrations/omnidreams/tests/interactive_drive/test_latency_timing.py
+++ b/integrations/omnidreams/tests/interactive_drive/test_latency_timing.py
@@ -3,7 +3,7 @@
import time
-from omnidreams.interactive_drive.runtime.timing import ChunkTimes
+from flashdreams.serving.realtime.timing import ChunkTimes
def _make_chunk(chunk_index: int = 0, chunk_size: int = 4) -> ChunkTimes:
diff --git a/integrations/omnidreams/tests/interactive_drive/test_streaming_presenter_realtime.py b/integrations/omnidreams/tests/interactive_drive/test_streaming_presenter_realtime.py
new file mode 100644
index 000000000..4d39be82f
--- /dev/null
+++ b/integrations/omnidreams/tests/interactive_drive/test_streaming_presenter_realtime.py
@@ -0,0 +1,57 @@
+# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+# SPDX-License-Identifier: Apache-2.0
+
+from __future__ import annotations
+
+import threading
+
+import numpy as np
+from omnidreams.interactive_drive.streaming_presenter import (
+ _as_rgb_host_uint8,
+ _publish_if_open,
+ _wait_for_bus_frame,
+)
+
+from flashdreams.serving.realtime.frame_bus import LatestFrameBus
+
+
+def test_streaming_presenter_materializes_lazy_rgba_frames() -> None:
+ class LazyFrame:
+ def to_numpy(self) -> np.ndarray:
+ return np.array(
+ [[[1, 2, 3, 255], [4, 5, 6, 255]]],
+ dtype=np.uint8,
+ )
+
+ frame = _as_rgb_host_uint8(LazyFrame())
+
+ assert frame.flags.c_contiguous
+ np.testing.assert_array_equal(
+ frame,
+ np.array([[[1, 2, 3], [4, 5, 6]]], dtype=np.uint8),
+ )
+
+
+def test_streaming_presenter_publishes_jpeg_on_latest_frame_bus() -> None:
+ bus = LatestFrameBus[bytes]()
+
+ _publish_if_open(bus, b"jpeg", stop_event=threading.Event())
+
+ latest = bus.latest()
+ assert latest is not None
+ assert latest.payload == b"jpeg"
+ assert latest.count == 1
+
+
+def test_streaming_presenter_frame_wait_returns_none_after_bus_close() -> None:
+ bus = LatestFrameBus[bytes]()
+ bus.publish(b"old")
+ bus.close()
+
+ frame = _wait_for_bus_frame(
+ bus,
+ last_seen_count=1,
+ stop_event=threading.Event(),
+ )
+
+ assert frame is None