diff --git a/flashdreams/flashdreams/demo/bridge.py b/flashdreams/flashdreams/demo/bridge.py index 317c37103..fd668d772 100644 --- a/flashdreams/flashdreams/demo/bridge.py +++ b/flashdreams/flashdreams/demo/bridge.py @@ -40,9 +40,9 @@ DeviceConverter, InputCanonicalizer, KeyboardToCameraCommand, - KeyboardToDriverCommand, ) from flashdreams.runtime.config import InferenceConfig +from flashdreams.runtime.gamepad import DrivingInputConverter from flashdreams.runtime.demo.drivers import BatchSessionDriver from flashdreams.runtime.demo.host import ( ModelWarmupPlan, @@ -715,7 +715,7 @@ def application_scenario( requested = frozenset(modality.name for modality in schema.modalities) converters: list[DeviceConverter] = [] if DRIVER_COMMAND.name in requested: - converters.append(KeyboardToDriverCommand()) + converters.append(DrivingInputConverter()) if CAMERA_COMMAND.name in requested: converters.append(KeyboardToCameraCommand()) return PreparedScenario( diff --git a/flashdreams/flashdreams/demo/local_input.py b/flashdreams/flashdreams/demo/local_input.py index 59e211347..80ff4c7cf 100644 --- a/flashdreams/flashdreams/demo/local_input.py +++ b/flashdreams/flashdreams/demo/local_input.py @@ -28,7 +28,13 @@ from flashdreams.runtime.canonical import ( DRIVER_COMMAND, InputCanonicalizer, - KeyboardToDriverCommand, +) +from flashdreams.runtime.gamepad import ( + GAMEPAD_STATE_CAPABILITY, + GAMEPAD_STATE_EVENT, + DrivingInputConverter, + GamepadState, + gamepad_state_payload, ) from flashdreams.runtime.inputs import ( CanonicalInputSchema, @@ -39,7 +45,7 @@ UserInputSchema, ) -_KEYBOARD_SOURCE_SCHEMA = UserInputSchema( +_LOCAL_SOURCE_SCHEMA = UserInputSchema( capabilities=( UserInputCapability( event_type="key_down", @@ -51,14 +57,12 @@ input_modality="keyboard", payload_fields=frozenset({"key"}), ), + GAMEPAD_STATE_CAPABILITY, ), - description="SlangPy local-window keyboard events.", + description="SlangPy local-window keyboard and gamepad events.", ) """Raw event schema emitted by the SlangPy window callback.""" -_GAMEPAD_DEADZONE = 0.05 -"""Minimum SDL gamepad axis magnitude treated as active input.""" - class SlangPyLocalInputHandler(InputHandler): """Convert SlangPy window events into application canonical inputs.""" @@ -91,7 +95,7 @@ def __init__( unsupported.append(modality.name) continue if not converters: - converters.append(KeyboardToDriverCommand()) + converters.append(DrivingInputConverter()) if unsupported: raise ValueError( "Local-window input cannot provide canonical modalities: " @@ -109,8 +113,6 @@ def __init__( self._session_start_s = 0.0 self._window_start_s = 0.0 self._opened = False - self._gamepad_connected = False - self._gamepad_state: dict[str, float] | None = None @property def accepts_window_events(self) -> bool: @@ -123,8 +125,6 @@ def open(self, session_info: SessionInfo) -> None: self._canonicalizer.reset() with self._event_lock: self._events.clear() - self._gamepad_connected = False - self._gamepad_state = None self._session_start_s = self._clock() self._window_start_s = 0.0 self._opened = True @@ -147,7 +147,7 @@ def current_inputs(self) -> CanonicalInputWindow: canonical = self._canonicalizer.canonicalize( UserInputs(events=events), window=window, - source_schema=_KEYBOARD_SOURCE_SCHEMA, + source_schema=_LOCAL_SOURCE_SCHEMA, ) values = { name: value @@ -156,10 +156,6 @@ def current_inputs(self) -> CanonicalInputWindow: } metadata = dict(canonical.metadata) - gamepad_command = self._current_gamepad_command() - if gamepad_command is not None and DRIVER_COMMAND.name in self._requested_names: - values[DRIVER_COMMAND.name] = gamepad_command - metadata["canonical_sources"] = {DRIVER_COMMAND.name: "gamepad"} return CanonicalInputWindow( values=values, metadata=metadata, @@ -171,8 +167,6 @@ def close(self) -> None: self._opened = False with self._event_lock: self._events.clear() - self._gamepad_connected = False - self._gamepad_state = None def on_keyboard_event(self, event: Any) -> None: """Record one SlangPy keyboard edge from the window event pump.""" @@ -200,51 +194,41 @@ def on_gamepad_event(self, event: Any) -> None: """Track SlangPy gamepad connection changes.""" if not self._opened: return - with self._event_lock: - if _event_flag(event, "is_connect"): - self._gamepad_connected = True - elif _event_flag(event, "is_disconnect"): - self._gamepad_connected = False - self._gamepad_state = None + if _event_flag(event, "is_disconnect"): + self._record_gamepad_state(GamepadState(False, 0.0, 0.0, 0.0)) def on_gamepad_state(self, state: Any) -> None: - """Record the latest SDL gamepad axes for driving control.""" + """Record the latest SDL gamepad driving state.""" if not self._opened or DRIVER_COMMAND.name not in self._requested_names: return - with self._event_lock: - self._gamepad_connected = True - self._gamepad_state = { - "left_x": _clamp(float(getattr(state, "left_x", 0.0)), -1.0, 1.0), - "left_trigger": _clamp( - float(getattr(state, "left_trigger", 0.0)), 0.0, 1.0 + self._record_gamepad_state( + GamepadState( + connected=True, + steer=-_clamp(float(getattr(state, "left_x", 0.0)), -1.0, 1.0), + throttle=_clamp( + float(getattr(state, "right_trigger", 0.0)), + 0.0, + 1.0, ), - "right_trigger": _clamp( - float(getattr(state, "right_trigger", 0.0)), 0.0, 1.0 + brake=_clamp( + float(getattr(state, "left_trigger", 0.0)), + 0.0, + 1.0, ), - } - - def _current_gamepad_command(self) -> dict[str, object] | None: - with self._event_lock: - if not self._gamepad_connected or self._gamepad_state is None: - return None - state = dict(self._gamepad_state) - if not any(abs(value) > _GAMEPAD_DEADZONE for value in state.values()): - return None - steer = -state["left_x"] - if abs(steer) <= _GAMEPAD_DEADZONE: - steer = 0.0 - return dict( - DRIVER_COMMAND.value( - { - "throttle": state["right_trigger"], - "brake": state["left_trigger"], - "steer": steer, - "stop": False, - "reverse": False, - } ) ) + def _record_gamepad_state(self, state: GamepadState) -> None: + """Append one normalized gamepad event.""" + event = UserInputEvent( + timestamp_s=max(0.0, self._clock() - self._session_start_s), + event_type=GAMEPAD_STATE_EVENT, + payload=gamepad_state_payload(state), + source="slangpy-gamepad", + ) + with self._event_lock: + self._events.append(event) + def _event_flag(event: Any, method_name: str) -> bool: method = getattr(event, method_name, None) diff --git a/flashdreams/flashdreams/runtime/__init__.py b/flashdreams/flashdreams/runtime/__init__.py index 1465058b1..b6b18afd7 100644 --- a/flashdreams/flashdreams/runtime/__init__.py +++ b/flashdreams/flashdreams/runtime/__init__.py @@ -20,6 +20,14 @@ ScriptedModality, ) from flashdreams.runtime.config import ExecutionBackend, InferenceConfig, Precision +from flashdreams.runtime.gamepad import ( + GAMEPAD_STATE_CAPABILITY, + GAMEPAD_STATE_EVENT, + DrivingInputConverter, + GamepadState, + gamepad_state_payload, + parse_gamepad_state, +) from flashdreams.runtime.inputs import ( INPUT_PHASES, CanonicalInputs, @@ -100,6 +108,11 @@ "DRIVING_SUPPORTED_KEYS", "DRIVER_COMMAND", "ExecutionBackend", + "GAMEPAD_STATE_CAPABILITY", + "GAMEPAD_STATE_EVENT", + "DrivingInputConverter", + "GamepadState", + "gamepad_state_payload", "IdentityInputMapping", "InferenceConfig", "InferenceInput", @@ -128,6 +141,7 @@ "NullOutputTarget", "OutputArtifact", "OutputTarget", + "parse_gamepad_state", "Precision", "PromptRequest", "ResetRequest", diff --git a/flashdreams/flashdreams/runtime/gamepad.py b/flashdreams/flashdreams/runtime/gamepad.py new file mode 100644 index 000000000..cc6e06b86 --- /dev/null +++ b/flashdreams/flashdreams/runtime/gamepad.py @@ -0,0 +1,253 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Gamepad input validation and canonical driving conversion.""" + +from __future__ import annotations + +import math +from collections.abc import Mapping +from dataclasses import dataclass +from typing import Any + +from flashdreams.infra.time import TimeWindow +from flashdreams.runtime.canonical import ( + DRIVER_COMMAND, + DeviceConverterSchema, + KeyboardToDriverCommand, +) +from flashdreams.runtime.inputs import UserInputCapability, UserInputs + +GAMEPAD_STATE_EVENT = "gamepad_state" +"""Event type shared by local and browser gamepad sources.""" + +_GAMEPAD_DEADZONE = 0.05 +"""Deadzone shared by generation activation and canonical conversion.""" + +GAMEPAD_STATE_CAPABILITY = UserInputCapability( + event_type=GAMEPAD_STATE_EVENT, + input_modality="gamepad", + payload_fields=frozenset({"connected", "steer", "throttle", "brake"}), + description="Normalized gamepad driving state.", +) + + +@dataclass(frozen=True, slots=True) +class GamepadState: + """Normalized driving state reported by one gamepad source.""" + + connected: bool + """Whether the source still reports an attached gamepad.""" + + steer: float + """Steering in ``[-1, 1]`` with positive values turning left.""" + + throttle: float + """Throttle engagement in ``[0, 1]``.""" + + brake: float + """Brake engagement in ``[0, 1]``.""" + + def is_active(self) -> bool: + """Return whether input exceeds the gamepad deadzone.""" + return self.connected and ( + self.throttle > _GAMEPAD_DEADZONE + or self.brake > _GAMEPAD_DEADZONE + or abs(self.steer) > _GAMEPAD_DEADZONE + ) + + +def parse_gamepad_state(payload: Mapping[str, Any]) -> GamepadState: + """Validate and normalize one gamepad payload. + + Args: + payload: JSON-like gamepad state. + + Returns: + Validated gamepad state. + + Raises: + TypeError: A required field has the wrong type. + ValueError: A numeric field is non-finite or outside its range. + """ + connected = payload["connected"] + if not isinstance(connected, bool): + raise TypeError("Gamepad field 'connected' must be boolean.") + steer = _number(payload["steer"], name="steer", minimum=-1.0) + throttle = _number(payload["throttle"], name="throttle", minimum=0.0) + brake = _number(payload["brake"], name="brake", minimum=0.0) + if not connected: + return GamepadState( + connected=False, + steer=0.0, + throttle=0.0, + brake=0.0, + ) + return GamepadState( + connected=True, + steer=steer, + throttle=throttle, + brake=brake, + ) + + +def gamepad_state_payload(state: GamepadState) -> dict[str, bool | float]: + """Return a JSON-compatible payload for ``state``. + + Args: + state: Validated gamepad state. + + Returns: + Complete payload accepted by :func:`parse_gamepad_state`. + """ + return { + "connected": state.connected, + "steer": state.steer, + "throttle": state.throttle, + "brake": state.brake, + } + + +class DrivingInputConverter: + """Merge keyboard and gamepad timelines into one driving command.""" + + def __init__(self) -> None: + """Configure live driving input conversion.""" + self._gamepad_state = GamepadState(False, 0.0, 0.0, 0.0) + self._keyboard = KeyboardToDriverCommand() + self._schema = DeviceConverterSchema( + name="live-driving-input", + produces=DRIVER_COMMAND, + consumes=( + UserInputCapability( + event_type="key_down", + payload_fields=frozenset({"key"}), + ), + UserInputCapability( + event_type="key_up", + payload_fields=frozenset({"key"}), + ), + GAMEPAD_STATE_CAPABILITY, + ), + device_kind="driving-input", + accepted_keys=self._keyboard.schema.accepted_keys, + ) + + @property + def schema(self) -> DeviceConverterSchema: + """Return the raw-input and canonical-output contract.""" + return self._schema + + def reset(self) -> None: + """Clear held keyboard and gamepad state.""" + self._gamepad_state = GamepadState(False, 0.0, 0.0, 0.0) + self._keyboard.reset() + + def convert( + self, + user_inputs: UserInputs, + window: TimeWindow, + ) -> Mapping[str, Any]: + """Merge keyboard and gamepad state at every timeline boundary. + + Args: + user_inputs: Raw keyboard edges and gamepad snapshots. + window: Half-open interval represented by the command. + + Returns: + Canonical driving state with a complete merged segment timeline. + """ + keyboard = self._keyboard.convert(UserInputs(), window) + if keyboard is None: + raise RuntimeError("Keyboard conversion did not produce a command.") + + segment_start = window.start_s + level = self._gamepad_level() if self._gamepad_state.is_active() else keyboard + segments: list[tuple[float, float, Mapping[str, Any]]] = [] + for event in user_inputs.events: + if event.event_type not in { + "key_down", + "key_up", + GAMEPAD_STATE_EVENT, + }: + continue + event_time = min( + max(float(event.timestamp_s), window.start_s), + window.end_s, + ) + if event_time > segment_start: + segments.append((segment_start, event_time, level)) + if event.event_type == GAMEPAD_STATE_EVENT: + self._gamepad_state = parse_gamepad_state(event.payload) + else: + keyboard = self._keyboard.convert( + UserInputs(events=(event,)), + window, + ) + if keyboard is None: + raise RuntimeError("Keyboard conversion did not produce a command.") + level = ( + self._gamepad_level() if self._gamepad_state.is_active() else keyboard + ) + segment_start = event_time + if window.end_s > segment_start or not segments: + segments.append((segment_start, window.end_s, level)) + + merged: list[tuple[float, float, Mapping[str, Any]]] = [] + for start, end, segment_level in segments: + if merged and merged[-1][2] == segment_level: + previous_start, _previous_end, previous_level = merged[-1] + merged[-1] = (previous_start, end, previous_level) + else: + merged.append((start, end, segment_level)) + final_level = merged[-1][2] + return DRIVER_COMMAND.value({**final_level, "segments": tuple(merged)}) + + def _gamepad_level(self) -> Mapping[str, Any]: + """Return the current post-deadzone gamepad level.""" + state = self._gamepad_state + return { + "throttle": ( + state.throttle + if state.connected and state.throttle > _GAMEPAD_DEADZONE + else 0.0 + ), + "brake": ( + state.brake + if state.connected and state.brake > _GAMEPAD_DEADZONE + else 0.0 + ), + "steer": ( + state.steer + if state.connected and abs(state.steer) > _GAMEPAD_DEADZONE + else 0.0 + ), + "stop": False, + "reverse": False, + } + + +def _number( + value: object, + *, + name: str, + minimum: float, +) -> float: + if isinstance(value, bool) or not isinstance(value, int | float): + raise TypeError(f"Gamepad field {name!r} must be numeric.") + parsed = float(value) + if not math.isfinite(parsed): + raise ValueError(f"Gamepad field {name!r} must be finite.") + if not minimum <= parsed <= 1.0: + raise ValueError(f"Gamepad field {name!r} must be in [{minimum:g}, 1].") + return parsed + + +__all__ = [ + "GAMEPAD_STATE_CAPABILITY", + "GAMEPAD_STATE_EVENT", + "DrivingInputConverter", + "GamepadState", + "gamepad_state_payload", + "parse_gamepad_state", +] diff --git a/flashdreams/flashdreams/serving/webrtc/manager.py b/flashdreams/flashdreams/serving/webrtc/manager.py index 06a00512a..b9b2a1061 100644 --- a/flashdreams/flashdreams/serving/webrtc/manager.py +++ b/flashdreams/flashdreams/serving/webrtc/manager.py @@ -46,6 +46,7 @@ WebRTCOutputSpec, run_demo_session_async, ) +from flashdreams.runtime.gamepad import GAMEPAD_STATE_EVENT from flashdreams.runtime.inputs import ( CanonicalInputSchema, InferenceInput, @@ -777,8 +778,9 @@ def __init__( self._lifecycle = WebRTCManagerLifecycle( busy_message=busy_message, client_liveness_timeout_s=client_liveness_timeout_s, - health_check=lambda: self._shared_host is None - or self._shared_host.is_healthy, + health_check=lambda: ( + self._shared_host is None or self._shared_host.is_healthy + ), ) @property @@ -806,7 +808,10 @@ def runtime(self) -> _RuntimeT: def browser_ui_config(self) -> dict[str, object]: """Return accepted control keys for the generic browser UI.""" accepted_keys = self._effective_supported_control_keys() or () - return {"accepted_keys": sorted(accepted_keys)} + return { + "accepted_keys": sorted(accepted_keys), + "gamepad_enabled": self._supports_gamepad_input(), + } def set_pending_session_input(self, session_input: Any) -> None: """Store validated model input for the next session.""" @@ -874,6 +879,14 @@ def _converter_supported_control_keys(self) -> frozenset[str] | None: ) return accepted_keys or None + def _supports_gamepad_input(self) -> bool: + """Return whether the prepared converter stack consumes gamepad state.""" + return any( + capability.event_type == GAMEPAD_STATE_EVENT + for schema in self._feedable_converter_schemas() + for capability in schema.consumes + ) + @staticmethod def _positive_int_runtime_value(value: Any, *, label: str) -> int: try: @@ -1915,6 +1928,12 @@ async def _handle_shared_datachannel_payload( if handled: managed_session.first_action_received.set() return + if message_type == GAMEPAD_STATE_EVENT and not self._supports_gamepad_input(): + self._send_json( + channel, + make_error_payload("This application does not accept gamepad input."), + ) + return result = input_source.handle_browser_payload( payload, timestamp_s=asyncio.get_running_loop().time(), diff --git a/flashdreams/flashdreams/serving/webrtc/services.py b/flashdreams/flashdreams/serving/webrtc/services.py index c56f06735..1cba4b45e 100644 --- a/flashdreams/flashdreams/serving/webrtc/services.py +++ b/flashdreams/flashdreams/serving/webrtc/services.py @@ -40,6 +40,12 @@ UserInputs, UserInputSchema, ) +from flashdreams.runtime.gamepad import ( + GAMEPAD_STATE_CAPABILITY, + GAMEPAD_STATE_EVENT, + gamepad_state_payload, + parse_gamepad_state, +) from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.demo import ( AsyncSessionDriver, @@ -83,6 +89,7 @@ "action", "disconnect", "event", + "gamepad", "heartbeat", "error", ] @@ -534,6 +541,8 @@ def handle_browser_payload( return WebRTCMessageResult(kind="disconnect") if message_type == MESSAGE_TYPE_EVENT: return self._record_text_event(payload, timestamp_s=timestamp_s) + if message_type == GAMEPAD_STATE_EVENT: + return self._record_gamepad(payload, timestamp_s=timestamp_s) if message_type == MESSAGE_TYPE_ACTION: action_payload = payload.get("action", payload) if not isinstance(action_payload, Mapping): @@ -549,7 +558,8 @@ def handle_browser_payload( kind="error", error=( "Unsupported message type, expected " - "'action', 'event', 'heartbeat', or 'disconnect'." + "'action', 'event', 'gamepad_state', 'heartbeat', or " + "'disconnect'." ), ) @@ -668,6 +678,42 @@ def _activate(self, timestamp_s: float) -> None: self._activation_timestamp_s = timestamp_s self._activation_signal.set() + def _record_gamepad( + self, + payload: Mapping[str, object], + *, + timestamp_s: float, + ) -> WebRTCMessageResult: + """Validate and record one browser gamepad state.""" + raw_state = payload.get("gamepad") + if not isinstance(raw_state, Mapping): + return WebRTCMessageResult( + kind="error", + error="'gamepad' must be an object.", + ) + try: + state = parse_gamepad_state( + {str(key): value for key, value in raw_state.items()} + ) + if not self._activation_signal.is_set(): + self._events = deque( + event + for event in self._events + if event.event_type != GAMEPAD_STATE_EVENT + ) + self.record_user_event( + timestamp_s=timestamp_s, + event_type=GAMEPAD_STATE_EVENT, + payload=gamepad_state_payload(state), + activate=state.is_active(), + ) + except (KeyError, TypeError, ValueError) as exc: + return WebRTCMessageResult(kind="error", error=str(exc)) + return WebRTCMessageResult( + kind="gamepad", + activated=state.is_active(), + ) + def _record_text_event( self, payload: Mapping[str, object], @@ -1220,6 +1266,7 @@ def _discard_task(self, task: asyncio.Task[RunResult]) -> None: input_modality="text", payload_fields=frozenset({"event_id", "state"}), ), + GAMEPAD_STATE_CAPABILITY, ), description="browser WebRTC data-channel events", ) diff --git a/flashdreams/flashdreams/serving/webrtc/web/request_session.html b/flashdreams/flashdreams/serving/webrtc/web/request_session.html index e370bfce5..02c25347f 100644 --- a/flashdreams/flashdreams/serving/webrtc/web/request_session.html +++ b/flashdreams/flashdreams/serving/webrtc/web/request_session.html @@ -87,6 +87,6 @@