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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion flashdreams/flashdreams/api_v2/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@ Protocols for the FlashDreams v2 API.
`IModelLoop` and `IUILoop` define model and UI work.
- `input_source.py` / `output_sink.py` / `client_window.py`: `IClientWindow`
groups one client's input and output.
- `user_input_event_data.py`: base class for input event data.
- `user_input_event.py`: base class for timestamped input events.

Running an application
----------------------
Expand Down
7 changes: 2 additions & 5 deletions flashdreams/flashdreams/api_v2/loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@

from flashdreams.runtime_v2.event_buffer import EventBuffer
from flashdreams.runtime_v2.step_result import StepResult
from flashdreams.runtime_v2.user_input_event import CloseUserInputEventData
from flashdreams.runtime_v2.user_input_event import CloseUserInputEvent
from flashdreams.runtime_v2.user_input_events import UserInputEvents
from flashdreams.runtime_v2.video_tensor import VideoTensorLayout

Expand Down Expand Up @@ -267,10 +267,7 @@ def presented_model_frames(self) -> tuple[Tensor, ...]:


def _contains_close(events: UserInputEvents) -> bool:
return any(
isinstance(event.get_event_data(), CloseUserInputEventData)
for event in events.get_events()
)
return any(isinstance(event, CloseUserInputEvent) for event in events.get_events())


def _model_results(
Expand Down
74 changes: 74 additions & 0 deletions flashdreams/flashdreams/api_v2/user_input_event.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,74 @@
# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: Apache-2.0
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""User input event protocol."""

from __future__ import annotations

from abc import ABC, abstractmethod
from dataclasses import dataclass
from typing import ClassVar, final

from numpy import uint64


@dataclass(frozen=True, slots=True, eq=False, kw_only=True)
class UserInputEvent(ABC):
"""Base class for timestamped user input events."""
Comment thread
greptile-apps[bot] marked this conversation as resolved.

_type_name_owners: ClassVar[dict[str, str]] = {}
"""ClassVar tracking all registered UserInputEvent's."""

timestamp: uint64
"""Timestamp in microseconds since the start of the session."""

def __init_subclass__(cls, **kwargs: object) -> None:
"""Register and validate the concrete event type name.

Raises:
TypeError: The event type name is not a non-empty string.
ValueError: Another event class uses the same type name.
"""
super(UserInputEvent, cls).__init_subclass__(**kwargs)
type_name = cls.get_type_name()
if not isinstance(type_name, str) or not type_name:
raise TypeError("User input event type names must be non-empty strings.")

owner = f"{cls.__module__}.{cls.__qualname__}"
registered_owner = cls._type_name_owners.get(type_name)
if registered_owner is not None and registered_owner != owner:
raise ValueError(
f"User input event type name {type_name!r} is already registered "
f"by {registered_owner}."
)
cls._type_name_owners[type_name] = owner

@classmethod
@abstractmethod
def get_type_name(cls) -> str:
"""Return the event type name."""
...

@final
def get_timestamp(self) -> uint64:
"""Return the timestamp of the event."""
return self.timestamp

def __hash__(self) -> int:
"""Return the hash of the concrete class name.

The value is not stable across processes.
"""
return hash(type(self).__name__)
29 changes: 0 additions & 29 deletions flashdreams/flashdreams/api_v2/user_input_event_data.py

This file was deleted.

5 changes: 2 additions & 3 deletions flashdreams/flashdreams/runtime_v2/event_buffer.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
import threading

from flashdreams.runtime_v2.user_input_event import (
ResetUserInputEventData,
ResetUserInputEvent,
UserInputEvent,
)
from flashdreams.runtime_v2.user_input_events import UserInputEvents
Expand Down Expand Up @@ -39,8 +39,7 @@ def append(self, events: UserInputEvents) -> None:
with self._lock:
self._events.extend(received)
self._generation += sum(
isinstance(event.get_event_data(), ResetUserInputEventData)
for event in received
isinstance(event, ResetUserInputEvent) for event in received
)

def read(self, reader_id: int) -> tuple[UserInputEvents, int]:
Expand Down
59 changes: 31 additions & 28 deletions flashdreams/flashdreams/runtime_v2/serving/webrtc_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,12 +26,12 @@
from flashdreams.runtime_v2.session_desc import SessionDesc
from flashdreams.runtime_v2.step_result import StepResult
from flashdreams.runtime_v2.user_input_event import (
CloseUserInputEventData,
FocusUserInputEventData,
CloseUserInputEvent,
FocusUserInputEvent,
KeyboardInputState,
KeyboardUserInputEventData,
MouseUserInputEventData,
ResetUserInputEventData,
KeyboardUserInputEvent,
MouseUserInputEvent,
ResetUserInputEvent,
UserInputEvent,
)
from flashdreams.runtime_v2.video_tensor import VideoTensorLayout
Expand Down Expand Up @@ -380,6 +380,9 @@ def _buffer_browser_message(self, raw_message: object) -> None:
if not isinstance(payload, dict):
raise ValueError("Browser event must be a JSON object.")

timestamp_us = self._timestamp_us()
if timestamp_us is None:
return
event_type = payload.get("type")
if event_type == "keyboard":
key = payload.get("key")
Expand All @@ -388,7 +391,8 @@ def _buffer_browser_message(self, raw_message: object) -> None:
raise ValueError("Keyboard event requires a non-empty key.")
if not isinstance(pressed, bool):
raise ValueError("Keyboard event requires a boolean pressed value.")
event_data = KeyboardUserInputEventData(
event = KeyboardUserInputEvent(
timestamp=timestamp_us,
key=key,
state=(
KeyboardInputState.PRESSED
Expand All @@ -412,7 +416,8 @@ def _buffer_browser_message(self, raw_message: object) -> None:
raise ValueError("Mouse button must be a non-negative integer.")
if not isinstance(pressed, bool):
raise ValueError("Mouse pressed must be a boolean.")
event_data = MouseUserInputEventData(
event = MouseUserInputEvent(
timestamp=timestamp_us,
action=action,
x=x,
y=y,
Expand All @@ -425,31 +430,20 @@ def _buffer_browser_message(self, raw_message: object) -> None:
focused = payload.get("focused")
if not isinstance(focused, bool):
raise ValueError("Focus event requires a boolean focused value.")
event_data = FocusUserInputEventData(focused=focused)
event = FocusUserInputEvent(
timestamp=timestamp_us,
focused=focused,
)
elif event_type == "reset":
event_data = ResetUserInputEventData()
event = ResetUserInputEvent(timestamp=timestamp_us)
elif event_type == "close":
event_data = CloseUserInputEventData()
event = CloseUserInputEvent(timestamp=timestamp_us)
else:
raise ValueError("Unsupported browser event type.")
self._append_event(event_data)
self._append_event(event)

def _append_event(
self,
event_data: (
KeyboardUserInputEventData
| MouseUserInputEventData
| FocusUserInputEventData
| ResetUserInputEventData
| CloseUserInputEventData
),
) -> None:
"""Timestamp and buffer one validated browser event."""
session_start_ns = self._session_start_ns
if session_start_ns is None:
return
timestamp_us = np.uint64((time.monotonic_ns() - session_start_ns) // 1_000)
event = UserInputEvent(timestamp=timestamp_us, event_data=event_data)
def _append_event(self, event: UserInputEvent) -> None:
"""Buffer one validated browser event."""
callback = self._input_callback
if callback is None:
raise RuntimeError("WebRTC input callback is not registered.")
Expand All @@ -463,7 +457,16 @@ def _record_client_disconnect(self) -> None:
return
self._client_connected = False
if not self._closed:
self._append_event(CloseUserInputEventData())
timestamp_us = self._timestamp_us()
if timestamp_us is not None:
self._append_event(CloseUserInputEvent(timestamp=timestamp_us))

def _timestamp_us(self) -> np.uint64 | None:
"""Return the current session-relative event timestamp."""
session_start_ns = self._session_start_ns
if session_start_ns is None:
return None
return np.uint64((time.monotonic_ns() - session_start_ns) // 1_000)

async def _enqueue_frames(
self, frames: tuple[np.ndarray[Any, np.dtype[np.uint8]], ...]
Expand Down
14 changes: 6 additions & 8 deletions flashdreams/flashdreams/runtime_v2/session_runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,11 @@
from flashdreams.api_v2.loop import IModelLoop, IUILoop
from flashdreams.api_v2.output_sink import OutputSink
from flashdreams.api_v2.session import ISession
from flashdreams.api_v2.user_input_event_data import UserInputEventData
from flashdreams.api_v2.user_input_event import UserInputEvent
from flashdreams.runtime_v2.event_buffer import EventBuffer
from flashdreams.runtime_v2.session_desc import PresentationMode
from flashdreams.runtime_v2.step_result import StepResult
from flashdreams.runtime_v2.user_input_event import CloseUserInputEventData
from flashdreams.runtime_v2.user_input_event import CloseUserInputEvent
from flashdreams.runtime_v2.user_input_events import UserInputEvents

_LOGGER = logging.getLogger(__name__)
Expand All @@ -23,11 +23,9 @@
_MODEL_READER_ID = 1


def _contains(events: UserInputEvents, event_type: type[UserInputEventData]) -> bool:
"""Return whether any event in ``events`` carries ``event_type`` data."""
return any(
isinstance(event.get_event_data(), event_type) for event in events.get_events()
)
def _contains(events: UserInputEvents, event_type: type[UserInputEvent]) -> bool:
"""Return whether ``events`` contains an instance of ``event_type``."""
return any(isinstance(event, event_type) for event in events.get_events())


def _log_secondary_failure(message: str, error: BaseException) -> None:
Expand Down Expand Up @@ -84,7 +82,7 @@ def run_session(
def collect_input() -> UserInputEvents:
events = window.get_user_input_events()
event_buffer.append(events)
if _contains(events, CloseUserInputEventData):
if _contains(events, CloseUserInputEvent):
stop.set()
return events

Expand Down
2 changes: 1 addition & 1 deletion flashdreams/flashdreams/runtime_v2/slangpy_ui_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
from torch import Tensor

from flashdreams.api_v2.loop import IUILoop
from flashdreams.runtime_v2._slangpy_ui_renderer import (
from flashdreams.runtime_v2.slangpy_ui_renderer import (
_SlangPyUIRenderer,
_UIRenderer,
)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,8 @@

from flashdreams.runtime_v2.user_input_event import (
KeyboardInputState,
KeyboardUserInputEventData,
MouseUserInputEventData,
KeyboardUserInputEvent,
MouseUserInputEvent,
)
from flashdreams.runtime_v2.user_input_events import UserInputEvents

Expand Down Expand Up @@ -235,10 +235,9 @@ def _route_input_events(
) -> None:
"""Route supported runtime input events into SlangPy's UI context."""
for event in events.get_events():
data = event.get_event_data()
if isinstance(data, KeyboardUserInputEventData):
pressed = data.state is KeyboardInputState.PRESSED
key = _resolve_slangpy_key(slangpy, data.key)
if isinstance(event, KeyboardUserInputEvent):
pressed = event.state is KeyboardInputState.PRESSED
key = _resolve_slangpy_key(slangpy, event.key)
if key is not None:
key_event = slangpy.KeyboardEvent()
key_event.type = (
Expand All @@ -249,33 +248,33 @@ def _route_input_events(
key_event.key = key
key_event.mods = slangpy.KeyModifierFlags.none
ui_context.handle_keyboard_event(key_event)
if pressed and len(data.key) == 1:
if pressed and len(event.key) == 1:
text_event = slangpy.KeyboardEvent()
text_event.type = slangpy.KeyboardEventType.input
text_event.codepoint = ord(data.key)
text_event.codepoint = ord(event.key)
text_event.mods = slangpy.KeyModifierFlags.none
ui_context.handle_keyboard_event(text_event)
elif isinstance(data, MouseUserInputEventData):
elif isinstance(event, MouseUserInputEvent):
mouse_event = slangpy.MouseEvent()
mouse_event.pos = (data.x * width, data.y * height)
mouse_event.pos = (event.x * width, event.y * height)
mouse_event.mods = slangpy.KeyModifierFlags.none
if data.action == "button":
if event.action == "button":
buttons = (
slangpy.MouseButton.left,
slangpy.MouseButton.middle,
slangpy.MouseButton.right,
)
if not 0 <= data.button < len(buttons):
if not 0 <= event.button < len(buttons):
continue
mouse_event.type = (
slangpy.MouseEventType.button_down
if data.pressed
if event.pressed
else slangpy.MouseEventType.button_up
)
mouse_event.button = buttons[data.button]
elif data.action == "wheel":
mouse_event.button = buttons[event.button]
elif event.action == "wheel":
mouse_event.type = slangpy.MouseEventType.scroll
mouse_event.scroll = (data.wheel_x, data.wheel_y)
mouse_event.scroll = (event.wheel_x, event.wheel_y)
else:
mouse_event.type = slangpy.MouseEventType.move
ui_context.handle_mouse_event(mouse_event)
Expand Down
Loading
Loading