diff --git a/.github/workflows/pr-tests.yml b/.github/workflows/pr-tests.yml index ba0cf46..ef753b9 100644 --- a/.github/workflows/pr-tests.yml +++ b/.github/workflows/pr-tests.yml @@ -40,9 +40,6 @@ jobs: - name: Run unit tests run: uv run --frozen pytest tests/unit -q - - name: Install native FlatBuffers headers - run: uv run --frozen python integrations/unreal/OpenUSDConnect/setup_flatbuffers.py - - name: Build native client tests run: cmake -S tests/native -B build/native-tests diff --git a/CHANGELOG.md b/CHANGELOG.md index b6e32bb..094c256 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -64,7 +64,8 @@ Return values and defaults: Behavior: - Token, metadata, and playback notifications run during `update()` or - `close()` on the calling thread instead of on network threads. + `close()` on the calling thread instead of on network threads; stage + metadata is delivered when it changes. - Stage edits made in `on_resync` are no longer published, matching `on_applied`. - `UsdPublisher.update()` raises before `start()`. While disconnected it @@ -78,8 +79,17 @@ Behavior: like `UsdPublisher.flush()`, instead of returning `False`. - `UsdReceiver.status` reports `CONNECTING` instead of `READY` while it reconnects, like the other clients. - -The low-level `EventSender`, `ReceiverThread`, and `EventDispatcher` keep their +- `ReceiverThread` is `EventReceiver` and no longer a `threading.Thread`: + call `close(timeout)` instead of `stop()` and `join()`, and read `running` + instead of `is_alive()`. On it and on `EventSender`, settings and state are + read-only properties (`token`, and the receiver's `reconnect`, stay + assignable) and `sock` is gone; read `connected`. Both close on leaving a + `with` block, and a collected one closes too, so keep the handle while it + should run. `EventSender.close()` writes the transactions already queued + and the Quit message, and does not wait for acknowledgements; call + `flush()` first for those. + +The low-level `EventSender`, `EventReceiver`, and `EventDispatcher` keep their callable arguments and properties. ### Added @@ -98,12 +108,12 @@ callable arguments and properties. `no_pending_recovery_stage`). - `claim_playback()` and `send_playback_control()` on `SharedStageClient` and `UsdPublisher`. -- Opt-in `background_send=True` moves transaction writes to a worker thread. - The worker needs the GIL, adding about 5 ms per write while the host's main - thread runs Python. -- `token_provider=` on `EventSender` and `ReceiverThread` supplies the token +- `token_provider=` on `EventSender` and `EventReceiver` supplies the token for each connection attempt. -- `EventDispatcher.drained_message_count`, `ReceiverThread.stopped`, and +- `notifications=` on `EventSender` and `EventReceiver` pushes their + notifications into a queue the owner drains, and `snapshot()` returns the + native status in one call. +- `EventDispatcher.drained_message_count`, `EventReceiver.stopped`, and `NoticeEmitter.has_local_changes`. - Receiver replay identity and optional post-commit transaction checkpoints. @@ -111,12 +121,29 @@ callable arguments and properties. - `EventDispatcher` starts its cursor at `receiver.sync_from - 1`, so integrations no longer seed `last_seq` for continuation. +- `EventSender` and `EventReceiver` run their connections on native threads + in the client core: transaction writes leave the calling thread and need no + GIL, callbacks and the token provider run on the connection thread, and + `close()` writes what was already queued, then stops the thread; a pending + connect is interrupted at once. +- `EventSender.connect()` waits for an attempt already in flight before + making its own; a token provider that raises is logged and fails that + attempt instead of raising. While recovery is required, `rejection_reason` + names the failure. +- Building the native extension fetches the pinned FlatBuffers headers on the + first configure, which needs network access unless + `FETCHCONTENT_SOURCE_DIR_FLATBUFFERS` names a local copy. +- The Unreal plugin runs on the native client engine: its receiver and emitter + threads drive the client core's `ReceiverEndpoint` and `ProducerEndpoint`, + built as the plugin's `OpenUSDConnectClientCore` module. Auth tokens stay in + the user's Unreal config under the same keys. ### Fixed -- A receiver continuing from a live-open snapshot replayed the full history - over it when the integration did not seed the dispatcher cursor. - MCP writes were confirmed before the mirror applied them. +- `EventSender.send_events()` returns `False` for a transaction above the + 16 MiB frame limit instead of queueing one that made the server close the + connection on every replay. - Replay completion markers were lost when a resync reset applied progress. - The emitter dropped property edits absorbed by a prim resync. - Bidirectional clients read the token file on every `update()` while their @@ -125,6 +152,9 @@ callable arguments and properties. - `NoticeEmitter.rebind_stage()` compared stage metadata against the previous stage, and `cleanup()` kept pending stage-metadata changes. - A `ManagedClient` constructed with invalid options left the stage modified. +- The Unreal plugin sent edits made while the server was down with the + server's values after a replay, because it read the values when emitting + instead of when it observed the edit. ## [0.4.0] - 2026-09-18 diff --git a/CMakeLists.txt b/CMakeLists.txt index cc821ad..de711ed 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -26,8 +26,11 @@ add_subdirectory(native/client_core) nanobind_add_module(_native_client STABLE_ABI native/python/client_module.cpp + native/python/driver_bindings.cpp + native/python/producer_bindings.cpp + native/python/receiver_bindings.cpp ) -target_link_libraries(_native_client PRIVATE OpenUSDConnectClientCore) +target_link_libraries(_native_client PRIVATE OpenUSDConnect::ClientDriver) target_compile_features(_native_client PRIVATE cxx_std_17) if(MSVC) diff --git a/docs/README.md b/docs/README.md index a6f4227..7cfda00 100644 --- a/docs/README.md +++ b/docs/README.md @@ -44,6 +44,8 @@ linked from the shorter workflow guides. stage ownership, native adapters, publishers, receivers, and resolver behavior. - [Shared stage architecture](shared-stage-architecture.md): exact file-layer synchronization and its protocol. +- [Native client core](../native/client_core/README.md): C++ targets, the + sans-IO endpoints, and the reference driver. - [MCP integration layout](../integrations/mcp/README.md#layout): extension points and module ownership. - [Unreal plugin developer notes](../integrations/unreal/OpenUSDConnect/PLUGIN_DEV.md): diff --git a/docs/mcp-server-usage.md b/docs/mcp-server-usage.md index 1f7610d..5580665 100644 --- a/docs/mcp-server-usage.md +++ b/docs/mcp-server-usage.md @@ -5,7 +5,7 @@ operations as local stdio tools. A client can author USD transactions and inspect the composed result through an in-memory mirror. The MCP process is a network client built on the core library (`EventSender` + -`ReceiverThread` + `EventDispatcher` + `UsdStageAdapter`), the same shape as the +`EventReceiver` + `EventDispatcher` + `UsdStageAdapter`), the same shape as the `usdview` integration. Every scene event it sends uses the core protocol. Its USD mirror also negotiates the optional layered-replay capability so authored logical-layer opinions retain their server strength ordering during live sync @@ -229,7 +229,7 @@ Verify any network with ## Implementation notes - **Emit + mirror.** The MCP emits via `EventSender` and keeps a read-only - `Usd.Stage` mirror through `ReceiverThread`, replaying from sequence 1. The + `Usd.Stage` mirror through `EventReceiver`, replaying from sequence 1. The server broadcasts committed records to every receiver, including the producer's receiver, so the mirror contains the authoritative result of both local and remote edits. The emitter and receiver use distinct diagnostic diff --git a/docs/usd-native-integration.md b/docs/usd-native-integration.md index 5b2d954..e1cc29d 100644 --- a/docs/usd-native-integration.md +++ b/docs/usd-native-integration.md @@ -1,9 +1,9 @@ # Python client and host-integration API These APIs attach OpenUSDConnect to an application-owned `pxr.Usd.Stage`. -Call `update()` from the stage-owning thread. Socket reads and reconnects run -on background threads; encoding, USD work, and (by default) transaction writes -run on the calling thread. +Call `update()` from the stage-owning thread. Socket reads, transaction +writes, and reconnects run on native threads that do not need the GIL; +encoding and USD work run on the calling thread. ## Choose an API @@ -82,7 +82,7 @@ Pass one `ClientObserver` subclass as `observer=` and override only what the host needs. Methods never run on a network thread: notifications arrive in `update()` or `close()`, and delivery methods run wherever the client applies authoritative state (`update()`, `refresh_asset_dependency()`, recovery). The -client wires only overridden methods, so unused notifications cost nothing: +client calls only overridden methods: ```python class HostObserver(ClientObserver): @@ -132,11 +132,6 @@ Receiving pauses while a `ManagedClient` edit target is foreign, because its accepts any edit target (session-layer edits stay local), so its loop calls `update()` unconditionally. -`background_send=True` moves transaction writes to a worker so a full socket -buffer cannot block the UI thread. The worker needs the GIL: while the host's -main thread runs Python, each write waits for Python's thread switch interval -(about 5 ms), so keep the default for latency-sensitive editing on fast links. - Before closing, stop authoring and call `submit_and_wait()`. Success means the edits are durable, not that their echo has been applied locally. @@ -168,8 +163,7 @@ instead of silently degrading to flat replay. Open the original base scene. A generated live-open snapshot already contains composed server state and is rejected because replaying the complete managed -history over it would duplicate opinions. Snapshot continuation is a separate -flat integration path used by the live-open host plugins. +history over it would duplicate opinions. Use `rebind_stage(new_stage)` when a host replaces its stage. Passing `None` parks stage application (phase `PARKED`) while the network queue continues to @@ -285,9 +279,8 @@ with ManagedClient( Construction creates `client.authoring_layer`, inserts it below the authoritative managed block, and makes it the edit target. Keep that target while the client is active. `update()` freezes local edits, applies the queued -authoritative prefix, then submits the frozen local batch. The dispatcher -suppresses and invalidates the emitter while applying server records, so -authoritative echoes do not become new local submissions. +authoritative prefix, then submits the frozen batch; applied server records are +never republished as local edits. `publish_current_edit_target()` queues a snapshot of the authoring layer for the next `update()` that can publish; a zero return can mean it is still queued. @@ -328,8 +321,8 @@ happens next depends on `DCCAdapter.targets_stage()`: Custom stage-backed adapters must override `targets_stage()` explicitly. -Shader mapping interfaces live in `openusdconnect.shader_mapping`; existing -imports from `openusdconnect.adapters` remain supported. Integrations that +Shader mapping interfaces live in `openusdconnect.shader_mapping` and are also +importable from `openusdconnect.adapters`. Integrations that author shader inputs directly can use `set_connectable_input_value` and `resolve_shader_port_type` from `openusdconnect.usd_authoring`. They operate under the stage's current edit target and do not send network events. @@ -424,14 +417,10 @@ managed receiver, call `refresh_asset_dependency(path)` after an asset becomes available or its resolver mapping changes; omit the path to retry all pending dependencies. -A context-only resolver remap is a special case for adapters targeting a -non-USD native scene. It can recompose both the live and previous-state stages -before projection observes the old topology. The dispatcher then sets -`native_scene_rebuild_required` and stops incremental delivery. The high-level -receiver reports it as `RECOVERY_REQUIRED` in `client.status`. Rebuild the -native destination and call -`client.acknowledge_native_scene_rebuilt()` before resuming. An ordinary -reconnect does not clear this guard. +For an adapter targeting a non-USD native scene, a context-only resolver remap +can recompose both the live and previous-state stages before projection +observes the old topology. That is the `RECOVERY_REQUIRED` case in +[Observing the client](#observing-the-client); a reconnect does not clear it. ## Identity and authentication @@ -446,9 +435,9 @@ store. ## Low-level APIs -`NoticeEmitter`, `EventSender`, `ReceiverThread`, and `EventDispatcher` remain +`NoticeEmitter`, `EventSender`, `EventReceiver`, and `EventDispatcher` are public for integrations whose scheduling or continuation requirements cannot -use the high-level clients. `ReceiverThread` requests layered replay by default; +use the high-level clients. `EventReceiver` requests layered replay by default; passing `layered_replay=False` selects the single-layer flat contract. Ordinary native-scene integrations should use `UsdReceiver(adapter=...)` instead of assembling these components. @@ -464,6 +453,19 @@ rejection/recovery status on subsequent ticks. `cancel_connect()` invalidates pending attempts and reports whether they have finished; `disconnect()` also closes an established connection. Neither discards the transaction outbox. +A sender's connection thread starts on the first connection request, a +receiver's on `start()`; callbacks run on that thread. `close(timeout=None)`, +or leaving a `with` block, stops the thread, which cannot be restarted, and +returns whether it exited in time (`False` at once from a callback). A sender +first writes the transactions already queued and the Quit message, only while +the socket is open and within one second or a shorter `timeout`, so `close(0)` +does not wait for them; closing does not wait for acknowledgements, which +`flush()` does. Keep a reference while the object should run: a collected one +closes. +Either object also takes `notifications=`, a `NotificationQueue` its owner +drains instead of every callback but `on_token_issued` (combining them raises +`ValueError`), and offers `snapshot()`, its native status read in one call. + ## Embed a server Use `ServerRuntime` when the application owns startup and shutdown: diff --git a/examples/fourier_waves/author.py b/examples/fourier_waves/author.py index f4c7cb3..aaa68b4 100644 --- a/examples/fourier_waves/author.py +++ b/examples/fourier_waves/author.py @@ -143,7 +143,8 @@ def run_author(args) -> int: print("\nstopping.") return 0 finally: - sender.disconnect() + sender.flush(timeout=5.0) + sender.close() def main() -> int: diff --git a/examples/fourier_waves/wave_client.py b/examples/fourier_waves/wave_client.py index a957716..335d399 100644 --- a/examples/fourier_waves/wave_client.py +++ b/examples/fourier_waves/wave_client.py @@ -28,7 +28,7 @@ from openusdconnect.adapters import UsdStageAdapter # noqa: E402 from openusdconnect.cli_common import add_sync_endpoint_args # noqa: E402 from openusdconnect.dispatcher import EventDispatcher # noqa: E402 -from openusdconnect.receiver import ReceiverThread # noqa: E402 +from openusdconnect.receiver import EventReceiver # noqa: E402 from openusdconnect.sender import EventSender # noqa: E402 DEFAULTS = { @@ -91,7 +91,7 @@ def __init__(self, host, port, proc_path): from pxr import Usd self.mirror = Usd.Stage.CreateInMemory() - self.receiver = ReceiverThread( + self.receiver = EventReceiver( host=host, port=port, sync_from=1, client_id="fourier-wave-client", origin=f"{origin}-recv", ) @@ -171,8 +171,8 @@ def run(self): time.sleep(1.0 / 30.0) def stop(self): - self.receiver.stop() - self.sender.disconnect() + self.receiver.close() + self.sender.close() def build_parser(add_help: bool = True) -> argparse.ArgumentParser: diff --git a/examples/instancing_dance/dance.py b/examples/instancing_dance/dance.py index 217b3aa..788c2a7 100644 --- a/examples/instancing_dance/dance.py +++ b/examples/instancing_dance/dance.py @@ -143,7 +143,7 @@ def run_dance(args: argparse.Namespace) -> int: print(f"sending setup events ({args.instances} instances)...") if not sender.send_events(setup_events(asset, args.instances)): print("setup send failed") - sender.disconnect() + sender.close() return 1 print("setup complete.") @@ -173,7 +173,8 @@ def run_dance(args: argparse.Namespace) -> int: except KeyboardInterrupt: print("\nstopping.") finally: - sender.disconnect() + sender.flush(timeout=5.0) + sender.close() return 0 diff --git a/integrations/blender/capture.py b/integrations/blender/capture.py index c386cf4..e6eef61 100644 --- a/integrations/blender/capture.py +++ b/integrations/blender/capture.py @@ -1136,7 +1136,7 @@ def _try_send_dirty_events(): """Build and send dirty events if emitter and sender are both connected.""" if _state.author is not None and _state.author._applying_remote: return - if _state.notice_emitter is None or _state.sender is None or _state.sender.sock is None: + if _state.notice_emitter is None or _state.sender is None or not _state.sender.connected: return events = _state.notice_emitter.prepare_events_for_send() if events: @@ -1580,7 +1580,7 @@ def execute(self, context): if not _cancel_emitter_reconnect(): self.report({"WARNING"}, "Emitter reconnect cancellation is still finishing") return {"CANCELLED"} - if _state.sender is not None and _state.sender.sock is not None: + if _state.sender is not None and _state.sender.connected: self.report({"INFO"}, "Already connected") return {"CANCELLED"} sender = _state.sender diff --git a/integrations/blender/receiver_addon.py b/integrations/blender/receiver_addon.py index 42a27f8..45f08fc 100644 --- a/integrations/blender/receiver_addon.py +++ b/integrations/blender/receiver_addon.py @@ -1,6 +1,6 @@ """Blender receiver applies incoming network events to Blender objects. -Uses ReceiverThread from openusdconnect.receiver and drains the queue +Uses EventReceiver from openusdconnect.receiver and drains the queue on the main thread via bpy.app.timers. """ @@ -28,14 +28,14 @@ K_SET_GPRIM_ATTRS, K_SET_XFORM_TRS, ) -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver from . import SESSION_ORIGIN as _ORIGIN from .blender_adapter import BlenderAdapter, apply_stage_metadata_to_scene LOG = logging.getLogger(__name__) -_RECEIVER: ReceiverThread | None = None +_RECEIVER: EventReceiver | None = None _DISPATCHER: EventDispatcher | None = None _ADAPTER: BlenderAdapter | None = None _MIRROR_STAGE = None @@ -231,15 +231,6 @@ def _set_remote_apply_guard(value: bool): capture_module.set_emitter_feedback_guard(value) -def _stop_receiver_thread(receiver: ReceiverThread) -> None: - """Stop a receiver and tolerate joining a thread that never started.""" - receiver.stop() - try: - receiver.join(timeout=2.0) - except RuntimeError: - LOG.debug("Receiver thread could not be joined", exc_info=True) - - def _store_last_sequence(scene, value: int) -> None: """Persist replay progress when the Blender scene is still writable.""" try: @@ -310,7 +301,7 @@ def _refresh_shader_reverse_sync_state(author, adapter, prim_path: str): ) -# Latest-wins PlaybackState handoff: written by the receiver thread, +# Latest-wins PlaybackState handoff: written by the receiver's connection thread, # read-and-cleared by Blender's main-thread timer. The lock keeps the # read-then-clear pair atomic so an update arriving in between can't be # silently wiped under the GIL this is rare in practice, under @@ -325,7 +316,7 @@ def _on_playback_state(state: dict) -> None: """Receive an authoritative PlaybackState broadcast. Stashes the payload for the next timer tick to apply on the Blender - main thread (callbacks run on the receiver thread). + main thread (callbacks run on the receiver's connection thread). """ global _LATEST_PLAYBACK_STATE snapshot = dict(state) @@ -492,7 +483,7 @@ def _on_resync() -> None: _DISPATCHER.adapter.mirror_stage = mirror_stage -def _build_dispatcher(receiver: ReceiverThread) -> EventDispatcher: +def _build_dispatcher(receiver: EventReceiver) -> EventDispatcher: """Construct a dispatcher over the receiver-owned layered USD mirror.""" mirror_stage = _ensure_mirror_stage() adapter = _ensure_adapter() @@ -540,7 +531,7 @@ def _apply_received_events_timer(): LOG.error("OpenUSDConnect receiver rejected: %s", reason) receiver = _RECEIVER _RECEIVER = None - _stop_receiver_thread(receiver) + receiver.close(timeout=2.0) _RECEIVER_TIMER_REGISTERED = False scene = bpy.context.scene if scene is not None: @@ -675,7 +666,7 @@ def execute(self, context): try: from . import STABLE_CLIENT_ID - _RECEIVER = ReceiverThread( + _RECEIVER = EventReceiver( host=host, port=port, sync_from=plan.sync_from, @@ -711,7 +702,7 @@ def execute(self, context): if _RECEIVER is not None and getattr(_RECEIVER, "auth_rejected", False): token_client.delete_token(host, port) if _RECEIVER is not None: - _stop_receiver_thread(_RECEIVER) + _RECEIVER.close(timeout=2.0) _RECEIVER = None _set_remote_apply_guard(False) _unregister_receiver_timer() @@ -730,7 +721,7 @@ def execute(self, context): if _RECEIVER is not None: receiver = _RECEIVER _RECEIVER = None - _stop_receiver_thread(receiver) + receiver.close(timeout=2.0) _store_last_sequence(context.scene, _LAST_SEQ) _unregister_receiver_timer() _set_remote_apply_guard(False) @@ -899,7 +890,7 @@ def unregister(): if _RECEIVER is not None: receiver = _RECEIVER _RECEIVER = None - _stop_receiver_thread(receiver) + receiver.close(timeout=2.0) _discard_replay_state() _unregister_receiver_timer() if BPY_AVAILABLE: diff --git a/integrations/mcp/__init__.py b/integrations/mcp/__init__.py index 329b0ec..2160e4d 100644 --- a/integrations/mcp/__init__.py +++ b/integrations/mcp/__init__.py @@ -6,7 +6,7 @@ and stream them to the sync server, which fans them out to every connected DCC. The server is a network client built on the core library (``EventSender`` + -``ReceiverThread`` + ``EventDispatcher`` + ``UsdStageAdapter``), the same shape +``EventReceiver`` + ``EventDispatcher`` + ``UsdStageAdapter``), the same shape as the ``usdview`` integration. It introduces no protocol changes. Launch with ``uv run python -m integrations.mcp`` (stdio transport). diff --git a/integrations/mcp/session.py b/integrations/mcp/session.py index df0ece9..b24937e 100644 --- a/integrations/mcp/session.py +++ b/integrations/mcp/session.py @@ -148,7 +148,7 @@ def _seed_metadata(self, metadata: dict | None) -> None: apply_events(self.mirror_stage, [{"k": K_SET_STAGE_METADATA, **payload}]) def _teardown(self) -> None: - """Stop the receiver thread and sender, clearing all connection state.""" + """Close the mirror receiver and disconnect the sender, clearing all connection state.""" if self.receiver is not None: self.receiver.close() if self.sender is not None: diff --git a/integrations/unreal/OpenUSDConnect/.gitignore b/integrations/unreal/OpenUSDConnect/.gitignore index f20d400..2c55b72 100644 --- a/integrations/unreal/OpenUSDConnect/.gitignore +++ b/integrations/unreal/OpenUSDConnect/.gitignore @@ -8,6 +8,10 @@ Saved/ # to the committed flatc-generated bindings under Public/Schema) Source/*/ThirdParty/flatbuffers/ +# native/client_core sources the Unreal harness stages before BuildPlugin +Source/OpenUSDConnectClientCore/include/ +Source/OpenUSDConnectClientCore/src/ + # Visual Studio user files (in case someone runs GenerateProjectFiles from this dir) *.suo *.user diff --git a/integrations/unreal/OpenUSDConnect/OpenUSDConnect.uplugin b/integrations/unreal/OpenUSDConnect/OpenUSDConnect.uplugin index 30e4776..f7a2ace 100644 --- a/integrations/unreal/OpenUSDConnect/OpenUSDConnect.uplugin +++ b/integrations/unreal/OpenUSDConnect/OpenUSDConnect.uplugin @@ -14,6 +14,11 @@ "CanContainContent": false, "Installed": false, "Modules": [ + { + "Name": "OpenUSDConnectClientCore", + "Type": "Runtime", + "LoadingPhase": "None" + }, { "Name": "OpenUSDConnectPXR", "Type": "Runtime", diff --git a/integrations/unreal/OpenUSDConnect/PLUGIN_DEV.md b/integrations/unreal/OpenUSDConnect/PLUGIN_DEV.md index becd4b5..6278c58 100644 --- a/integrations/unreal/OpenUSDConnect/PLUGIN_DEV.md +++ b/integrations/unreal/OpenUSDConnect/PLUGIN_DEV.md @@ -16,28 +16,26 @@ the plugin's modules, threading model, protocol implementation, and known gaps. │ ─► fires deferred Connect() │ │ ─► finds AUsdStageActor → AttachToStageActor() subscribes to │ │ FUsdListener::OnObjectsChanged │ -│ ─► pops validated frames with bSuppressEmit=true while applying │ +│ ─► drains notifications and receiver frames, applying with │ +│ bSuppressEmit=true │ │ │ │ OnObjectsChanged() ─► queues exact changed Sdf paths │ -│ Tick() ─► drains those paths unless bSuppressEmit is true │ -│ ─► reads TRS / visibility / shader inputs │ -│ ─► encodes Txn frames and pushes to FEmitClient │ +│ Tick() ─► reads their TRS / visibility / shader inputs before │ +│ applying received frames (latest value wins) │ +│ ─► encodes Txn frames and appends them to the producer │ │ │ └──────────────────┬─────────────────────────────────┬─────────────────────────┘ │ │ ┌───────────▼──────────┐ ┌────────────▼────────────┐ - │ FSyncClient │ │ FEmitClient │ + │ FEndpointRunner │ │ FEndpointRunner │ + │ │ │ │ │ (FRunnable) │ │ (FRunnable) │ - │ role = "receiver" │ │ role = "emitter" │ │ │ │ │ - │ TCP recv loop: │ │ TCP loop: │ - │ – read framed FB │ │ – claim shared frames │ - │ – verify once │ │ from native outbox │ - │ – enqueue bytes + │ │ – wake on enqueue │ - │ trusted metadata │ │ – peek with │ - │ – direct TryPop │ │ HasPendingData │ - │ on game thread │ │ for inbound results │ - │ │ │ and rate limits │ + │ applies the │ │ applies the │ + │ endpoint's actions │ │ endpoint's actions │ + │ with an FSocket and │ │ with an FSocket; an │ + │ reports bytes, │ │ Append wakes it │ + │ timeouts, and time │ │ │ └──────────────────────┘ └─────────────────────────┘ │ ▲ ▼ │ @@ -58,17 +56,16 @@ applies received frames. | File | Class / Symbol | Role | |------|----------------|------| | `Public/USDConnectSettings.h` | `UUSDConnectSettings` | UDeveloperSettings exposed at *Edit → Project Settings → Plugins → OpenUSD Connect*. | -| `Public/USDConnectSubsystem.h` | `UUSDConnectSubsystem` | UTickableWorldSubsystem that owns both clients and the stage-actor attachment; it drains the receiver session on the game thread. | +| `Public/USDConnectSubsystem.h` | `UUSDConnectSubsystem` | UTickableWorldSubsystem that owns both endpoints, their runners, and the stage-actor attachment; it drains notifications and receiver frames on the game thread. | | `OpenUSDConnectPXR/Public/USDConnectProtocol.h` | `namespace OUC` | Wraps the generated FlatBuffers bindings with framing limits and small Unreal helpers. | -| `Private/SyncClient.h/.cpp` | `FSyncClient` | Receiver TCP thread. Handles HELLO, verifies each frame once, and queues bytes with trusted sequence/event metadata. | -| `Private/EmitClient.h/.cpp` | `FEmitClient` | Emitter TCP thread. Claims shared immutable frames, wakes immediately on enqueue, and peeks for inbound results via `HasPendingData`. | -| `native/client_core` (repository root) | `OrderedProducerSession`, `OrderedReceiverSession`, `FrameDecoder` | Canonical C++ ordering, reconnect generation, recovery, replay, queue, and framing state shared by nanobind and the staged Unreal build. | +| `Private/EndpointRunner.h/.cpp` | `FEndpointRunner` | One `FRunnable` per role. Applies the endpoint's actions with an `FSocket` and reports bytes, read timeouts, and time; it holds no protocol state. | +| `native/client_core` (repository root) | `ReceiverEndpoint`, `ProducerEndpoint` | Sans-IO connection protocol (handshake, replay, outbox, recovery, reconnect) shared with the Python module, built here as the `OpenUSDConnectClientCore` module. | | `Private/TxnBuilder.h/.cpp` | `BuildXformTxnFrame`, `BuildVisibilityTxnFrame`, `BuildConnectableInputTxnFrame` | FlatBuffers Txn frame builders for the supported emitter event kinds. | | `OpenUSDConnectPXR/Private/OpenUSDConnectPXR.cpp` | `IMPLEMENT_MODULE` | Registers the PXR dynamic module with Unreal's module manager. A successful link does not replace this runtime entry point. | | `OpenUSDConnectPXR/Public/USDEventApplier.h`, `Private/USDEventApplier.cpp` | `FUSDEventApplier::ApplyValidatedFrame` | Applies a boundary-verified BroadcastEvent without repeating FlatBuffers verification; the subsystem manages `pxr::SdfChangeBlock` runs from queued metadata. | | `OpenUSDConnectPXR/Public/USDStageBridge.h`, `Private/USDStageBridge.cpp` | `FUSDStageBridge` | Keeps direct pxr stage reads and writes out of the no-RTTI UObject module. | | `OpenUSDConnectPXR/Public/USDMaterialXMaterializer.h`, `Private/USDMaterialXMaterializer.cpp` | `FUSDMaterialXMaterializer` | Maintains Unreal-local MaterialX documents for inline networks. | -| `OpenUSDConnect.uplugin`, `Source/*/*.Build.cs` | - | Registers the runtime and PXR modules and their engine dependencies. | +| `OpenUSDConnect.uplugin`, `Source/*/*.Build.cs` | - | Registers the client core, runtime, and PXR modules and their engine dependencies. | ## Wire protocol @@ -103,8 +100,9 @@ The shared native core includes the flatc-generated C++ bindings under `include/openusdconnect/client/schema/`. `protocol_codec.h` provides transport-neutral, borrowed receive views and caller-owned builders. `USDConnectProtocol.h` adds only Unreal-friendly string and array helpers. `TxnBuilder.cpp` converts Unreal-native -values into the shared stateless event builders; receive code classifies verified -handshake and control messages through the shared views. +values into the shared stateless event builders. The endpoints handle handshake and +control messages; the subsystem only tells `BroadcastEvent` from `Resync` in the +frames it drains. Run `scripts/generate_flatbuffers.sh` after changing either schema and commit the regenerated Python and C++ bindings together. The generated C++ header pins @@ -116,12 +114,12 @@ headers. - The frame-length prefix is big-endian (`struct.pack(">I", ...)` on the server), but the FlatBuffers payload itself is little-endian as always. -- Emitter builders call `FinishSizePrefixed`, rewrite only that four-byte prefix - to big-endian, and detach FlatBuffers' allocation into `FWireFrame`. The outbox - shares that immutable allocation through reconnect and acknowledgement; do not - materialize a second `TArray`. -- The frame size limit is 16 MiB (`OUC::kMaxFrameSize`) and must match the - server. +- Emitter builders finish frames with the core's `FinishTransactionFrame`, which + prepends the big-endian length, and copy each once into the vector + `ProducerEndpoint::Append` takes. The outbox shares that allocation through + reconnect and acknowledgement. +- The frame size limit is the core's `kDefaultMaxFrameSize` (16 MiB) and must + match the server. - Emitter and receiver each open their own TCP socket. `client_id` is the stable authentication and producer identity. `origin` is diagnostic metadata, while the emitter's `producer_session_id` provides exactly-once transaction identity. @@ -148,19 +146,18 @@ deadlock loading at ~90 %. ## Echo / feedback-loop guards -Two independent guards keep changes from bouncing forever: +`UUSDConnectSubsystem::bSuppressEmit` (a `std::atomic`) is set while +`DrainAndApply()` and plugin-owned USD authoring are running. The attached +`FUsdListener::OnObjectsChanged` callback ignores notices during that window, +preventing received changes and local MaterialX support opinions from being +emitted back to the server. -1. `FSyncClient` compares the incoming `BroadcastEvent.origin` with its own - `SessionOrigin` and drops matching frames. -2. `UUSDConnectSubsystem::bSuppressEmit` (a `std::atomic`) is set while - `DrainAndApply()` and plugin-owned USD authoring are running. The attached - `FUsdListener::OnObjectsChanged` callback ignores notices during that window, - preventing received changes and local MaterialX support opinions from being - emitted back to the server. - -The listener reports exact Sdf paths. They are coalesced in `PendingEmitPaths` -and drained once per tick, avoiding the ancestor roll-up behavior of -`AUsdStageActor::OnPrimChanged`. +The listener reports exact Sdf paths. They are coalesced in `PendingEmitPaths`, +and once per tick, before any received frame applies, the subsystem reads their +values into captured events (the latest per prim and input). Those values, not a +later read of the stage, are sent once the emitter can publish, so a replay after +a reconnect cannot replace an edit made offline. Exact paths also avoid the +ancestor roll-up behavior of `AUsdStageActor::OnPrimChanged`. ## Build configuration @@ -180,23 +177,17 @@ reports boundary failures through status values and uses assertions only for internal invariants. The canonical implementation lives at `native/client_core` in the repository. -`producer_session.h` and `receiver_session.h` are payload-generic templates: -Python instantiates them with owned references to immutable Python `bytes`, while -Unreal instantiates them with `TSharedPtr` and -`FValidatedReceiverFrame`. This removes Python boundary copies, preserves -Unreal's zero-copy producer frame and move-only receiver queue, and avoids -maintaining a second connection-state implementation. -`protocol_codec.h` also owns the shared handshake/control classification and -transaction envelope construction. It borrows receive buffers and operates on a -caller-supplied FlatBuffers builder, leaving transport, allocation, threading, -and offset storage to the integration. -The Unreal packaging harness copies it into the temporary plugin source tree -before `BuildPlugin`; the staged copy is an artifact and is never maintained as -a second source. Repository CMake compiles the canonical files directly into -the nanobind extension. - -The PXR module's `PublicSystemIncludePaths` exposes the pinned, plugin-local -FlatBuffers headers installed by `setup_flatbuffers.py` to both modules. +The Unreal packaging harness stages its `include/`, `src/frame_codec.cpp`, and +`src/engine/` into `Source/OpenUSDConnectClientCore` before `BuildPlugin`; the +staged copy is an artifact and is never maintained as a second source. That +module builds the endpoints into their own DLL and exports them through +`OPENUSDCONNECT_CLIENT_API`. It links Core so the containers that cross into the +other modules share Unreal's allocator. Repository CMake compiles the canonical +files directly into the nanobind extension. + +The client core module's `PublicSystemIncludePaths` exposes the pinned FlatBuffers +headers that `setup_flatbuffers.py` installs under `OpenUSDConnectPXR/ThirdParty` +to every module that depends on it. ## Known gaps @@ -218,7 +209,6 @@ the Output Log: ``` Log LogUSDConnect Verbose Log LogUSDConnectSubsystem Verbose -Log LogUSDEmit Verbose Log LogUSDEventApplier Verbose ``` diff --git a/integrations/unreal/OpenUSDConnect/README.md b/integrations/unreal/OpenUSDConnect/README.md index 9beed7e..5245b8a 100644 --- a/integrations/unreal/OpenUSDConnect/README.md +++ b/integrations/unreal/OpenUSDConnect/README.md @@ -87,7 +87,7 @@ same but has not been validated. | Auto-start Receiver from Metadata | `true` | Start the receiver automatically when live metadata is detected. | | Auto-start Emitter from Metadata | `true` | Start the emitter automatically when live metadata is detected. | | Persist Auth Tokens | `true` | Save server-issued TOFU tokens in the user's Unreal config and reuse them on reconnect. | -| Reconnect Delay (s) | `3.0` | Wait time between reconnect attempts | +| Reconnect Delay (s) | `3.0` | Wait before the first reconnect attempt; later attempts back off exponentially | The plugin receives through flat replay. Use it with one unmuted collaboration layer and no server department policy. A department @@ -176,10 +176,8 @@ The **Output Log** should contain: ``` LogUSDConnectSubsystem: Detected OpenUSDConnect live metadata on stage: 127.0.0.1:7200 snapshot_seq=... LogUSDConnectSubsystem: Using USD live metadata; receiver will sync from seq=... -LogUSDConnect: Connected to OpenUSDConnect server at 127.0.0.1:7200 (receiver, sync_from=...) -LogUSDConnect: HELLO_OK received entering receive loop -LogUSDEmit: Emitter connected to 127.0.0.1:7200 -LogUSDEmit: Emitter HELLO_OK ready to send +LogUSDConnect: Receiver: connected (sync_from=...) +LogUSDConnect: Emitter: connected to 127.0.0.1:7200 (session=..., pending=0) LogUSDConnectSubsystem: Attached to AUsdStageActor (UsdStageActor_0) ``` Per-event messages use the `Verbose` level. Enable them with @@ -258,14 +256,13 @@ Live-open-specific checks: | Inline MaterialX materials render gray or black | UE 5.8 translates referenced `.mtlx` documents but not values from inline `ND_*` networks. The plugin materializes supported networks automatically; see [MaterialX rendering](#materialx-rendering-auto-materializer). Enable **Substrate Adaptive GBuffer** for fuller `standard_surface` support. UsdPreviewSurface uses the universal context and is unaffected. | | Generated meshes are named after a **container** prim instead of the individual objects (e.g. `SM_World1`, `SM_World2`, … for a root prim called `World`), every synced edit rebuilds them, and materials jump between objects on visibility changes | That container prim is being **collapsed**: it has no `kind`, and UE collapses kind-less subtrees by default (`USD.CollapsePrimsWithoutKind` is true), folding the whole subtree into **one** static mesh whose sections and material slots re-index on every rebuild. Author `kind = "group"` on scene-root Xforms (correct USD model hierarchy), or set `USD.CollapsePrimsWithoutKind 0`, or uncheck **Use Prim Kinds For Collapsing** on the stage actor. | | Generated assets churn constantly (new transient packages per edit); appearance drifts until a full stage reload | No persistent asset cache: each stage actor defaults to a throwaway transient cache. Create a **USD Asset Cache** asset and assign it on the stage actor (or Project Settings → USDCore → Default Asset Cache). Consider also disabling **Share Assets for Identical Prims**, so prims with identical geometry but different materials don't share one mesh asset. | -| Edits in Unreal don't reach Blender | Confirm the **Emitter HELLO_OK** line appears in the log; if not, the emitter socket failed. Check the dashboard's *Clients* tab. | +| Edits in Unreal don't reach Blender | Confirm the **Emitter: connected** line appears in the log; if not, the emitter socket failed. Check the dashboard's *Clients* tab. | | Plugin engine-version warning | The checked-in descriptor targets Unreal 5.8. Rebuild and deliberately port the plugin for another engine version rather than editing only the descriptor. | For deeper diagnostics, enable verbose logging in the editor console: ``` Log LogUSDConnect Verbose Log LogUSDConnectSubsystem Verbose -Log LogUSDEmit Verbose Log LogUSDEventApplier Verbose ``` diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/OpenUSDConnect.Build.cs b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/OpenUSDConnect.Build.cs index ae31e89..1f622d2 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/OpenUSDConnect.Build.cs +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/OpenUSDConnect.Build.cs @@ -1,6 +1,5 @@ // Copyright OpenUSDConnect Contributors. All Rights Reserved. -using System.IO; using UnrealBuildTool; public class OpenUSDConnect : ModuleRules @@ -8,9 +7,6 @@ public class OpenUSDConnect : ModuleRules public OpenUSDConnect(ReadOnlyTargetRules Target) : base(Target) { PCHUsage = PCHUsageMode.UseExplicitOrSharedPCHs; - PublicIncludePaths.Add(Path.GetFullPath(Path.Combine( - ModuleDirectory, - "../ThirdParty/OpenUSDConnectClientCore/include"))); PublicDependencyModuleNames.AddRange(new string[] { @@ -20,6 +16,7 @@ public OpenUSDConnect(ReadOnlyTargetRules Target) : base(Target) "Sockets", // FSocket, ISocketSubsystem "Networking", // FInternetAddr "DeveloperSettings", + "OpenUSDConnectClientCore", }); PrivateDependencyModuleNames.AddRange(new string[] diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EmitClient.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EmitClient.cpp deleted file mode 100644 index 7812600..0000000 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EmitClient.cpp +++ /dev/null @@ -1,681 +0,0 @@ -// Copyright OpenUSDConnect Contributors. All Rights Reserved. - -#include "EmitClient.h" - -#include "USDConnectProtocol.h" -#include "USDConnectSubsystem.h" -#include "USDWireFraming.h" - -#include "Logging/LogMacros.h" -#include "Sockets.h" -#include "SocketSubsystem.h" -#include "HAL/PlatformProcess.h" -#include "HAL/PlatformTime.h" -#include "Misc/ScopeLock.h" - -#include "flatbuffers/flatbuffer_builder.h" - -DEFINE_LOG_CATEGORY_STATIC(LogUSDEmit, Log, All); - -using namespace OUC; - -namespace -{ -constexpr double FrameReadTimeoutSeconds = 10.0; -constexpr size_t MaxPendingTransactions = 10'000; - -bool IsPeerClosed(FSocket* Socket) -{ - if (!Socket) - { - return true; - } - - uint32 PendingBytes = 0; - if (Socket->HasPendingData(PendingBytes)) - { - return false; - } - - // An orderly TCP close becomes readable with no payload. HasPendingData() - // alone cannot distinguish that state from an idle connection on Windows. - if (!Socket->Wait(ESocketWaitConditions::WaitForRead, FTimespan::Zero())) - { - return false; - } - - uint8 Probe = 0; - int32 BytesRead = 0; - return !Socket->Recv(&Probe, 1, BytesRead, ESocketReceiveFlags::Peek) || BytesRead == 0; -} -} // namespace - -// --------------------------------------------------------------------------- -FProducerEndpointState::FProducerEndpointState(const FString& InHost, int32 InPort, - const FString& InDepartment, - const FString& InSessionId) - : Host(InHost) - , Port(InPort) - , Department(InDepartment) - , SessionId(InSessionId) - , Session(MaxPendingTransactions) -{ -} - -bool FProducerEndpointState::MatchesEndpoint(const FString& InHost, int32 InPort, - const FString& InDepartment) const -{ - return Host == InHost && Port == InPort && Department == InDepartment; -} - -uint64 FProducerEndpointState::GetNextTransactionId() const -{ - return Session.NextTransactionId(); -} - -bool FProducerEndpointState::BeginConnection(uint64& OutGeneration) -{ - const std::optional Start = - Session.BeginConnection(); - if (!Start.has_value()) - { - return false; - } - OutGeneration = Start->Generation; - return true; -} - -void FProducerEndpointState::Disconnect(uint64 Generation) -{ - Session.Disconnect(Generation); -} - -bool FProducerEndpointState::EnqueueFrame(uint64 Generation, uint64 TxnId, FWireFrame&& Frame) -{ - if (Frame.IsEmpty()) - { - return false; - } - FProducerFrame SharedFrame = MakeShared(MoveTemp(Frame)); - return Session.Append(Generation, TxnId, MoveTemp(SharedFrame), 1, {}) == - openusdconnect::client::ProducerResult::Accepted; -} - -bool FProducerEndpointState::ClaimNextUnsent(uint64 Generation, FQueuedProducerTxn& OutTxn) -{ - FProducerSession::Entry Entry; - if (Session.ClaimNextUnsent(Generation, Entry) != - openusdconnect::client::ProducerResult::Accepted) - { - return false; - } - OutTxn.TxnId = Entry.TransactionId; - OutTxn.Frame = std::move(Entry.Payload); - return true; -} - -bool FProducerEndpointState::AcceptServerHighwater(uint64 Generation, uint64 CommittedThrough, - FString& OutError) -{ - const openusdconnect::client::ProducerResult Result = - Session.AcceptHello(Generation, CommittedThrough); - if (Result == openusdconnect::client::ProducerResult::Accepted) - { - return true; - } - - if (Result == openusdconnect::client::ProducerResult::HighwaterAhead) - { - OutError = FString::Printf(TEXT("Server reports transaction %llu for producer session " - "whose local highwater is %llu"), - CommittedThrough, Session.NextTransactionId() - 1); - } - else if (Result == openusdconnect::client::ProducerResult::HighwaterRegressed) - { - OutError = FString::Printf(TEXT("Server producer highwater regressed from %llu to %llu"), - Session.LastAcknowledgedTransactionId(), CommittedThrough); - } - else - { - OutError = TEXT("Producer session rejected HELLO_OK in its current state"); - } - FScopeLock Lock(&RecoveryCS); - RecoveryReason = OutError; - return false; -} - -void FProducerEndpointState::RetireThrough(uint64 Generation, uint64 AckId) -{ - const openusdconnect::client::ProducerResult Result = - Session.AcknowledgeThrough(Generation, AckId); - if (Result != openusdconnect::client::ProducerResult::Accepted && - Result != openusdconnect::client::ProducerResult::StaleGeneration) - { - FScopeLock Lock(&RecoveryCS); - RecoveryReason = FString::Printf( - TEXT("Invalid cumulative acknowledgement through transaction %llu"), AckId); - } -} - -void FProducerEndpointState::MarkRejected(uint64 Generation, uint64 TxnId, uint8 RejectionCode, - const FString& Reason) -{ - openusdconnect::client::ProducerRecoveryDisposition Disposition; - if (RejectionCode == - static_cast(OpenUSDConnect::TransactionRejectionCode::StaleLayerGraph)) - { - Disposition = openusdconnect::client::ProducerRecoveryDisposition::RecoverableConflict; - } - else if (RejectionCode == - static_cast(OpenUSDConnect::TransactionRejectionCode::InvalidTransaction)) - { - Disposition = openusdconnect::client::ProducerRecoveryDisposition::InvalidOperation; - } - else - { - Disposition = openusdconnect::client::ProducerRecoveryDisposition::SessionFatal; - } - Session.Reject(Generation, TxnId, Disposition); - FScopeLock Lock(&RecoveryCS); - RecoveryReason = Reason.IsEmpty() - ? FString::Printf(TEXT("Transaction %llu rejected"), TxnId) - : FString::Printf(TEXT("Transaction %llu rejected: %s"), TxnId, *Reason); -} - -uint64 FProducerEndpointState::GetSubmittedTransactionCount() const -{ - return Session.SubmittedTransactionCount(); -} - -uint64 FProducerEndpointState::GetAcknowledgedTransactionCount() const -{ - return Session.AcknowledgedTransactionCount(); -} - -uint64 FProducerEndpointState::GetPendingTransactionCount() const -{ - return static_cast(Session.PendingTransactionCount()); -} - -bool FProducerEndpointState::IsRecoveryRequired() const -{ - return Session.RecoveryRequired(); -} - -FString FProducerEndpointState::GetRecoveryReason() const -{ - FScopeLock Lock(&RecoveryCS); - return RecoveryReason; -} - -EUSDConnectRecoveryDisposition FProducerEndpointState::GetRecoveryDisposition() const -{ - switch (Session.RecoveryDisposition()) - { - case openusdconnect::client::ProducerRecoveryDisposition::RecoverableConflict: - return EUSDConnectRecoveryDisposition::RecoverableConflict; - case openusdconnect::client::ProducerRecoveryDisposition::InvalidOperation: - return EUSDConnectRecoveryDisposition::InvalidOperation; - case openusdconnect::client::ProducerRecoveryDisposition::SessionFatal: - return EUSDConnectRecoveryDisposition::SessionFatal; - default: - return EUSDConnectRecoveryDisposition::None; - } -} - -// --------------------------------------------------------------------------- -FEmitClient::FEmitClient(UUSDConnectSubsystem* InOwner, const FString& InClientId, - const TSharedRef& InProducerState, - float InReconnectDelaySecs, const FString& InAuthToken) - : Owner(InOwner) - , ProducerState(InProducerState) - , ClientId(InClientId) - , AuthToken(InAuthToken) - , ReconnectDelaySecs(InReconnectDelaySecs) - , Socket(nullptr) - , Thread(nullptr) - , WorkEvent(FPlatformProcess::GetSynchEventFromPool(false)) - , bShouldStop(false) - , bConnected(false) - , ConnectionGeneration(0) - , SessionGeneration(0) -{ -} - -FEmitClient::~FEmitClient() -{ - StopAndWait(); - if (WorkEvent) - { - FPlatformProcess::ReturnSynchEventToPool(WorkEvent); - WorkEvent = nullptr; - } -} - -bool FEmitClient::Start() -{ - Thread = FRunnableThread::Create(this, TEXT("OpenUSDConnect_EmitClient"), 0, TPri_BelowNormal); - return Thread != nullptr; -} - -void FEmitClient::StopAndWait() -{ - Stop(); - if (Thread) - { - Thread->WaitForCompletion(); - delete Thread; - Thread = nullptr; - } - CloseSocket(); -} - -bool FEmitClient::FlushPending(double TimeoutSeconds) const -{ - const double Deadline = FPlatformTime::Seconds() + FMath::Max(0.0, TimeoutSeconds); - while (ProducerState->GetPendingTransactionCount() > 0 && - !ProducerState->IsRecoveryRequired() && FPlatformTime::Seconds() < Deadline) - { - FPlatformProcess::Sleep(0.005f); - } - return ProducerState->GetPendingTransactionCount() == 0 && !ProducerState->IsRecoveryRequired(); -} - -bool FEmitClient::Init() -{ - return true; -} - -uint32 FEmitClient::Run() -{ - const FString& Host = ProducerState->GetHost(); - const int32 Port = ProducerState->GetPort(); - const FString& Department = ProducerState->GetDepartment(); - const FString& SessionOrigin = ProducerState->GetSessionId(); - while (!bShouldStop.load(std::memory_order_relaxed)) - { - uint64 ActiveGeneration = 0; - if (!ProducerState->BeginConnection(ActiveGeneration)) - { - return 0; - } - - // --- Connect --- - ISocketSubsystem* SS = ISocketSubsystem::Get(PLATFORM_SOCKETSUBSYSTEM); - if (!SS) - { - ProducerState->Disconnect(ActiveGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - - FSocket* NewSocket = SS->CreateSocket(NAME_Stream, TEXT("USDConnectEmit"), false); - if (!NewSocket) - { - ProducerState->Disconnect(ActiveGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - { - FScopeLock Lock(&SocketCS); - Socket = NewSocket; - } - if (bShouldStop.load(std::memory_order_relaxed)) - { - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - break; - } - - TSharedRef Addr = SS->CreateInternetAddr(); - bool bValid = false; - Addr->SetIp(*Host, bValid); - Addr->SetPort(Port); - if (!bValid || !Socket->Connect(*Addr)) - { - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - UE_LOG(LogUSDEmit, Log, TEXT("Emitter connected to %s:%d"), *Host, Port); - - // --- Send HELLO as emitter --- - FWireFrame HelloFrame; - const openusdconnect::client::FrameResult HelloResult = - OUC::BuildHelloFrame(TEXT("emitter"), 0, ClientId, SessionOrigin, Department, - HelloFrame, AuthToken, SessionOrigin); - if (HelloResult != openusdconnect::client::FrameResult::Success || - !SendAll(HelloFrame.GetData(), HelloFrame.Num())) - { - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - - // --- Read HELLO_OK (blocking) --- - { - TArray Frame; - if (!RecvFrame(Frame)) - { - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - - openusdconnect::client::EnvelopeView Envelope; - if (openusdconnect::client::DecodeEnvelope( - Frame.GetData(), static_cast(Frame.Num()), Envelope) != - openusdconnect::client::ProtocolResult::Success) - { - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - const openusdconnect::client::HandshakeResponseView Response(Envelope); - if (Response.Kind() == - openusdconnect::client::HandshakeResponseKind::AuthenticationRejected) - { - UE_LOG(LogUSDEmit, Error, TEXT("Emitter auth rejected")); - if (Owner) - { - Owner->OnClientAuthRejected(TEXT("emitter")); - } - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - return 0; - } - if (Response.Kind() == - openusdconnect::client::HandshakeResponseKind::ConfigurationRejected) - { - const OpenUSDConnect::HelloRejected* Rejection = Response.ConfigurationRejection(); - const OpenUSDConnect::HelloRejectionCode Code = Rejection->code(); - const FString CodeName = - Code == OpenUSDConnect::HelloRejectionCode::Unspecified - ? FString() - : FString(UTF8_TO_TCHAR(OpenUSDConnect::EnumNameHelloRejectionCode(Code))); - const FString Reason = ToFString(Rejection->reason()); - if (Owner) - { - Owner->OnClientHelloRejected(TEXT("emitter"), CodeName, Reason); - } - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - return 0; - } - if (Response.Kind() != openusdconnect::client::HandshakeResponseKind::Accepted) - { - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - const OpenUSDConnect::HelloOk* HelloOk = Response.Accepted(); - FString HighwaterError; - if (!ProducerState->AcceptServerHighwater(ActiveGeneration, - HelloOk->committed_through(), HighwaterError)) - { - UE_LOG(LogUSDEmit, Error, TEXT("Producer recovery required: %s"), *HighwaterError); - if (Owner) - { - Owner->OnEmitterTransactionRejected(HelloOk->committed_through(), - HighwaterError); - } - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - return 0; - } - const FString IssuedToken = ToFString(HelloOk->token()); - if (Owner) - { - Owner->OnClientHelloOk(TEXT("emitter")); - } - if (!IssuedToken.IsEmpty()) - { - AuthToken = IssuedToken; - if (Owner) - { - Owner->OnClientTokenIssued(IssuedToken); - } - } - UE_LOG(LogUSDEmit, Log, TEXT("Emitter HELLO_OK ready to send")); - } - - ConnectionGeneration.fetch_add(1, std::memory_order_relaxed); - SessionGeneration.store(ActiveGeneration, std::memory_order_relaxed); - bConnected.store(true, std::memory_order_release); - - // --- Main loop: drain send queue + poll for incoming messages --- - bool bShouldDisconnect = false; - float RateLimitDelay = 0.0f; - while (!bShouldStop.load(std::memory_order_relaxed) && !bShouldDisconnect) - { - if (IsPeerClosed(Socket)) - { - bShouldDisconnect = true; - break; - } - - // 1. Claim and send every endpoint-state transaction not yet written - // on this connection. Items remain owned by ProducerState until a - // committed/duplicate TransactionResult arrives. - bool bDidWork = false; - FQueuedProducerTxn Pending; - while (ProducerState->ClaimNextUnsent(ActiveGeneration, Pending)) - { - if (!SendAll(Pending.Frame->GetData(), Pending.Frame->Num())) - { - UE_LOG(LogUSDEmit, Warning, TEXT("Emitter send failed for txn %llu"), - Pending.TxnId); - bShouldDisconnect = true; - break; - } - bDidWork = true; - } - if (bShouldDisconnect) - break; - - // 2. Drain a bounded batch of complete result/control frames. The - // server may coalesce many individually framed acknowledgements into - // one TCP write; consuming one frame per 5 ms loop would cap quiescent - // acknowledgement progress at roughly 200 transactions/second. - constexpr int32 MaxResultsPerIteration = 256; - for (int32 ResultIndex = 0; ResultIndex < MaxResultsPerIteration && !bShouldDisconnect; - ++ResultIndex) - { - uint32 PendingBytes = 0; - if (!Socket || !Socket->HasPendingData(PendingBytes) || PendingBytes == 0) - { - break; - } - // RecvFrame handles fragmented headers and bodies. Enter it as - // soon as any byte is readable so a peer close after a partial - // header is observed instead of leaving this connection stuck. - TArray InFrame; - if (!RecvFrame(InFrame)) - { - bShouldDisconnect = true; - break; - } - bDidWork = true; - openusdconnect::client::EnvelopeView Envelope; - if (openusdconnect::client::DecodeEnvelope( - InFrame.GetData(), static_cast(InFrame.Num()), Envelope) != - openusdconnect::client::ProtocolResult::Success) - { - bShouldDisconnect = true; - break; - } - const openusdconnect::client::ControlMessageView Message(Envelope); - if (Message.Kind() == openusdconnect::client::ControlMessageKind::TransactionResult) - { - const OpenUSDConnect::TransactionResult* Result = Message.TransactionResult(); - const uint64 AckId = Result->txn_id(); - const OpenUSDConnect::TransactionStatus Status = Result->status(); - if (Status == OpenUSDConnect::TransactionStatus::Rejected) - { - const FString Reason = ToFString(Result->reason()); - UE_LOG(LogUSDEmit, Error, TEXT("Transaction %llu rejected: %s"), AckId, - *Reason); - ProducerState->MarkRejected(ActiveGeneration, AckId, - static_cast(Result->rejection_code()), - Reason); - if (Owner) - { - Owner->OnEmitterTransactionRejected(AckId, Reason); - } - bShouldStop.store(true, std::memory_order_relaxed); - bShouldDisconnect = true; - } - else - { - ProducerState->RetireThrough(ActiveGeneration, AckId); - } - } - else if (Message.Kind() == openusdconnect::client::ControlMessageKind::RateLimited) - { - const float Retry = Message.RateLimit()->retry_after(); - UE_LOG(LogUSDEmit, Warning, TEXT("Emitter rate limited sleeping %.1fs"), Retry); - RateLimitDelay = Retry; - bShouldDisconnect = true; - } - } - - // 3. Sleep until a producer enqueue wakes us. While acknowledgements - // are outstanding, retain a short bounded wait so socket results are - // consumed promptly. With no in-flight work, only perform a low-rate - // connection health check. - if (!bDidWork) - { - const uint32 WaitMilliseconds = - ProducerState->GetPendingTransactionCount() > 0 ? 5U : 1000U; - WorkEvent->Wait(WaitMilliseconds); - } - } - - bConnected.store(false, std::memory_order_relaxed); - CloseSocket(); - ProducerState->Disconnect(ActiveGeneration); - if (!bShouldStop.load(std::memory_order_relaxed)) - { - UE_LOG(LogUSDEmit, Log, TEXT("Emitter disconnected reconnecting in %.1fs"), - ReconnectDelaySecs); - FPlatformProcess::Sleep(RateLimitDelay > 0.0f ? RateLimitDelay : ReconnectDelaySecs); - } - } - return 0; -} - -void FEmitClient::Stop() -{ - bShouldStop.store(true, std::memory_order_relaxed); - if (WorkEvent) - { - WorkEvent->Trigger(); - } - InterruptSocket(); -} - -void FEmitClient::Exit() -{ - bConnected.store(false, std::memory_order_relaxed); - CloseSocket(); -} - -bool FEmitClient::EnqueueFrame(uint64 TxnId, FWireFrame&& Frame) -{ - const bool bAccepted = ProducerState->EnqueueFrame( - SessionGeneration.load(std::memory_order_relaxed), TxnId, MoveTemp(Frame)); - if (bAccepted && WorkEvent) - { - WorkEvent->Trigger(); - } - return bAccepted; -} - -// --------------------------------------------------------------------------- -// Private helpers -// --------------------------------------------------------------------------- -void FEmitClient::InterruptSocket() -{ - FScopeLock Lock(&SocketCS); - if (Socket) - { - Socket->Shutdown(ESocketShutdownMode::ReadWrite); - } -} - -void FEmitClient::CloseSocket() -{ - FSocket* SocketToDestroy = nullptr; - { - FScopeLock Lock(&SocketCS); - SocketToDestroy = Socket; - Socket = nullptr; - } - if (SocketToDestroy) - { - SocketToDestroy->Close(); - if (ISocketSubsystem* SS = ISocketSubsystem::Get(PLATFORM_SOCKETSUBSYSTEM)) - { - SS->DestroySocket(SocketToDestroy); - } - } -} - -bool FEmitClient::RecvExact(uint8* Buf, int32 Needed) -{ - int32 Got = 0; - while (Got < Needed) - { - if (bShouldStop.load(std::memory_order_relaxed) || !Socket) - return false; - if (!Socket->Wait(ESocketWaitConditions::WaitForRead, - FTimespan::FromSeconds(FrameReadTimeoutSeconds))) - { - UE_LOG(LogUSDEmit, Warning, TEXT("Timed out while receiving emitter frame")); - return false; - } - int32 Read = 0; - if (!Socket->Recv(Buf + Got, Needed - Got, Read) || Read <= 0) - return false; - Got += Read; - } - return true; -} - -bool FEmitClient::RecvFrame(TArray& OutFrame) -{ - uint8 LenBuf[4]; - if (!RecvExact(LenBuf, 4)) - { - return false; - } - std::size_t PayloadLen = 0; - if (!openusdconnect::client::TryReadFrameHeader(LenBuf, kMaxFrameSize, PayloadLen)) - { - return false; - } - OutFrame.SetNumUninitialized(static_cast(PayloadLen)); - return RecvExact(OutFrame.GetData(), static_cast(PayloadLen)); -} - -bool FEmitClient::SendAll(const uint8* Data, int32 Len) -{ - int32 Sent = 0; - while (Sent < Len) - { - if (!Socket) - return false; - int32 ThisSent = 0; - if (!Socket->Send(Data + Sent, Len - Sent, ThisSent) || ThisSent <= 0) - { - return false; - } - Sent += ThisSent; - } - return true; -} diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EmitClient.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EmitClient.h deleted file mode 100644 index 936a848..0000000 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EmitClient.h +++ /dev/null @@ -1,162 +0,0 @@ -// Copyright OpenUSDConnect Contributors. All Rights Reserved. -#pragma once - -#include "HAL/Runnable.h" -#include "HAL/RunnableThread.h" -#include "HAL/Event.h" -#include "HAL/CriticalSection.h" -#include "Sockets.h" -#include "USDConnectRecovery.h" -#include "USDWireFraming.h" -#include "openusdconnect/client/producer_session.h" -#include - -class UUSDConnectSubsystem; - -using FProducerFrame = TSharedPtr; -using FProducerSession = openusdconnect::client::OrderedProducerSession; - -struct FQueuedProducerTxn -{ - uint64 TxnId = 0; - FProducerFrame Frame; -}; - -/** - * Endpoint-scoped producer identity and durable in-memory outbox. - * - * The subsystem owns this object, not an individual socket thread. Replacing - * FEmitClient for the same endpoint therefore preserves transaction identity, - * encoded bytes, acknowledgements, and recovery state. A different endpoint - * must use a different state object whose transaction sequence starts at one. - */ -class FProducerEndpointState -{ -public: - FProducerEndpointState(const FString& InHost, int32 InPort, const FString& InDepartment, - const FString& InSessionId); - - bool MatchesEndpoint(const FString& InHost, int32 InPort, const FString& InDepartment) const; - const FString& GetHost() const - { - return Host; - } - int32 GetPort() const - { - return Port; - } - const FString& GetDepartment() const - { - return Department; - } - const FString& GetSessionId() const - { - return SessionId; - } - - uint64 GetNextTransactionId() const; - bool BeginConnection(uint64& OutGeneration); - void Disconnect(uint64 Generation); - bool EnqueueFrame(uint64 Generation, uint64 TxnId, OUC::FWireFrame&& Frame); - bool ClaimNextUnsent(uint64 Generation, FQueuedProducerTxn& OutTxn); - bool AcceptServerHighwater(uint64 Generation, uint64 CommittedThrough, FString& OutError); - void RetireThrough(uint64 Generation, uint64 AckId); - void MarkRejected(uint64 Generation, uint64 TxnId, uint8 RejectionCode, const FString& Reason); - - uint64 GetSubmittedTransactionCount() const; - uint64 GetAcknowledgedTransactionCount() const; - uint64 GetPendingTransactionCount() const; - bool IsRecoveryRequired() const; - FString GetRecoveryReason() const; - EUSDConnectRecoveryDisposition GetRecoveryDisposition() const; - -private: - FString Host; - int32 Port; - FString Department; - FString SessionId; - - FProducerSession Session; - mutable FCriticalSection RecoveryCS; - FString RecoveryReason; -}; - -/** - * Background TCP thread that connects to the OpenUSDConnect server as an **emitter**. - * - * Protocol: - * 1. Connect TCP and send Envelope{Hello, role="emitter"} - * 2. Await Envelope{HelloOk} - * 3. Loop: - * - Claim unsent frames from the endpoint-scoped producer state → send - * - HasPendingData() → if data available, read one frame and dispatch - * (RateLimited is handled; other control messages are ignored) - * - Wait on an enqueue event while idle with a bounded socket health check - * 4. On error: - * retain the in-flight frame, close socket, wait, and retry it before later queued frames on the - * next connection. - * - * Frames pushed via EnqueueFrame() include the 4-byte big-endian length prefix. - * FProducerEndpointState serializes game-thread submission with socket-thread - * acknowledgement and survives replacement of this client object. - */ -class FEmitClient : public FRunnable -{ -public: - FEmitClient(UUSDConnectSubsystem* InOwner, const FString& InClientId, - const TSharedRef& InProducerState, - float InReconnectDelaySecs, const FString& InAuthToken = FString()); - - virtual ~FEmitClient(); - - // FRunnable - virtual bool Init() override; - virtual uint32 Run() override; - virtual void Stop() override; - virtual void Exit() override; - - bool Start(); - void StopAndWait(); - /** Wait for the server's cumulative acknowledgement without blocking normal emits. */ - bool FlushPending(double TimeoutSeconds) const; - - bool IsConnected() const - { - return bConnected.load(std::memory_order_acquire); - } - - /** - * Monotonically increasing identifier for each successful HELLO handshake. - * The game thread uses this to replay per-connection structural prerequisites - * before sending value-only transform events to a fresh server session. - */ - uint64 GetConnectionGeneration() const - { - return ConnectionGeneration.load(std::memory_order_relaxed); - } - - /** Push a complete pre-framed Envelope{Txn} into the endpoint outbox. */ - bool EnqueueFrame(uint64 TxnId, OUC::FWireFrame&& Frame); - -private: - bool RecvExact(uint8* Buf, int32 Needed); - bool RecvFrame(TArray& OutFrame); - bool SendAll(const uint8* Data, int32 Len); - void InterruptSocket(); - void CloseSocket(); - - UUSDConnectSubsystem* Owner; - TSharedRef ProducerState; - FString ClientId; - FString AuthToken; - float ReconnectDelaySecs; - - FSocket* Socket; - FCriticalSection SocketCS; - FRunnableThread* Thread; - FEvent* WorkEvent; - std::atomic bShouldStop; - std::atomic bConnected; - std::atomic ConnectionGeneration; - std::atomic SessionGeneration; -}; diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EndpointRunner.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EndpointRunner.cpp new file mode 100644 index 0000000..ecc723c --- /dev/null +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EndpointRunner.cpp @@ -0,0 +1,144 @@ +// Copyright OpenUSDConnect Contributors. All Rights Reserved. + +#include "EndpointRunner.h" + +#include "Logging/LogMacros.h" +#include "SocketSubsystem.h" +#include "Sockets.h" + +DEFINE_LOG_CATEGORY_STATIC(LogUSDConnect, Log, All); + +namespace OUC::Runner +{ +namespace +{ +// A receiver blocks on its socket so data is read at once and a Wake waits up to a slice; a +// producer blocks on its event so a submit is sent at once, and polls its socket for results. +constexpr uint32 ReceiveSliceMilliseconds = 100; +constexpr uint32 PollMilliseconds = 20; + +ISocketSubsystem& Sockets() +{ + return *ISocketSubsystem::Get(PLATFORM_SOCKETSUBSYSTEM); +} + +uint32 MillisecondsUntil(TimePoint Until, uint32 Limit) +{ + const int64 Remaining = + std::chrono::ceil(Until - std::chrono::steady_clock::now()) + .count(); + return static_cast(FMath::Clamp(Remaining, 0, Limit)); +} + +FTimespan PollSlice(TimePoint Until) +{ + return FTimespan::FromMilliseconds(MillisecondsUntil(Until, PollMilliseconds)); +} +} // namespace + +FSocket* Connect(const openusdconnect::client::ConnectAction& Action, + const std::atomic& bStopping) +{ + const FAddressInfoResult Resolved = + Sockets().GetAddressInfo(*ToFString(Action.Host), nullptr, EAddressInfoFlags::Default, + NAME_None, SOCKTYPE_Streaming); + if (Resolved.ReturnCode != SE_NO_ERROR || Resolved.Results.IsEmpty()) + { + return nullptr; + } + const TSharedRef Address = Resolved.Results[0].Address; + Address->SetPort(Action.Port); + FSocket* Socket = + Sockets().CreateSocket(NAME_Stream, TEXT("OpenUSDConnect"), Address->GetProtocolType()); + if (Socket && Socket->SetNonBlocking(true) && Socket->Connect(*Address)) + { + for (;;) + { + const ESocketConnectionState State = Socket->GetConnectionState(); + if (State == SCS_Connected) + { + return Socket; + } + if (State == SCS_ConnectionError || bStopping || + std::chrono::steady_clock::now() >= Action.Deadline) + { + break; + } + Socket->Wait(ESocketWaitConditions::WaitForWrite, PollSlice(Action.Deadline)); + } + } + Destroy(Socket); + return nullptr; +} + +bool SendAll(FSocket& Socket, const std::vector& Bytes, TimePoint Deadline, + const std::atomic& bStopping) +{ + size_t Sent = 0; + while (Sent < Bytes.size()) + { + int32 Count = 0; + const int32 Size = static_cast(FMath::Min(Bytes.size() - Sent, MAX_int32)); + if (Socket.Send(Bytes.data() + Sent, Size, Count)) + { + Sent += static_cast(Count); + continue; + } + if (Sockets().GetLastErrorCode() != SE_EWOULDBLOCK || bStopping || + std::chrono::steady_clock::now() >= Deadline) + { + return false; + } + Socket.Wait(ESocketWaitConditions::WaitForWrite, PollSlice(Deadline)); + } + return true; +} + +bool Receive(FSocket& Socket, TArray& Buffer, FTimespan WaitTime, int32& OutReceived) +{ + OutReceived = 0; + return !Socket.Wait(ESocketWaitConditions::WaitForRead, WaitTime) || + Socket.Recv(Buffer.GetData(), Buffer.Num(), OutReceived); +} + +FTimespan ReceiveSlice(std::optional Until) +{ + return FTimespan::FromMilliseconds(Until ? MillisecondsUntil(*Until, ReceiveSliceMilliseconds) + : ReceiveSliceMilliseconds); +} + +void Destroy(FSocket* Socket) +{ + if (Socket) + { + Socket->Close(); + Sockets().DestroySocket(Socket); + } +} + +void Wait(FEvent& Event, std::optional Until, bool bSocketOpen) +{ + const uint32 Limit = bSocketOpen ? PollMilliseconds : MAX_uint32; + Event.Wait(Until ? MillisecondsUntil(*Until, Limit) : Limit); +} + +void Log(const TCHAR* Role, const openusdconnect::client::LogAction& Action) +{ + const FString Message = ToFString(Action.Message); + switch (Action.Level) + { + case openusdconnect::client::LogLevel::Debug: + UE_LOG(LogUSDConnect, Verbose, TEXT("%s: %s"), Role, *Message); + break; + case openusdconnect::client::LogLevel::Info: + UE_LOG(LogUSDConnect, Log, TEXT("%s: %s"), Role, *Message); + break; + case openusdconnect::client::LogLevel::Warning: + UE_LOG(LogUSDConnect, Warning, TEXT("%s: %s"), Role, *Message); + break; + case openusdconnect::client::LogLevel::Error: + UE_LOG(LogUSDConnect, Error, TEXT("%s: %s"), Role, *Message); + break; + } +} +} // namespace OUC::Runner diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EndpointRunner.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EndpointRunner.h new file mode 100644 index 0000000..40dbef7 --- /dev/null +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/EndpointRunner.h @@ -0,0 +1,255 @@ +// Copyright OpenUSDConnect Contributors. All Rights Reserved. +#pragma once + +#include "CoreMinimal.h" +#include "HAL/Event.h" +#include "HAL/PlatformProcess.h" +#include "HAL/Runnable.h" +#include "HAL/RunnableThread.h" +#include "USDConnectProtocol.h" + +THIRD_PARTY_INCLUDES_START +#include "openusdconnect/client/engine/producer_endpoint.h" +#include "openusdconnect/client/engine/receiver_endpoint.h" +THIRD_PARTY_INCLUDES_END + +#include +#include +#include +#include +#include +#include +#include + +class FSocket; + +// The socket plumbing both roles share, implemented in EndpointRunner.cpp. +namespace OUC::Runner +{ +using openusdconnect::client::TimePoint; + +// A nonblocking socket connected before Action.Deadline, or nullptr. +FSocket* Connect(const openusdconnect::client::ConnectAction& Action, + const std::atomic& bStopping); +bool SendAll(FSocket& Socket, const std::vector& Bytes, TimePoint Deadline, + const std::atomic& bStopping); +// Waits up to WaitTime for data. False once the peer closed or the socket failed; +// OutReceived is zero when nothing arrived. +bool Receive(FSocket& Socket, TArray& Buffer, FTimespan WaitTime, int32& OutReceived); +// How long a receiver blocks on its socket before it applies a Wake. +FTimespan ReceiveSlice(std::optional Until); +void Destroy(FSocket* Socket); +// Waits for the event until Until, and while a socket is open at most until its next poll. +void Wait(FEvent& Event, std::optional Until, bool bSocketOpen); +void Log(const TCHAR* Role, const openusdconnect::client::LogAction& Action); +} // namespace OUC::Runner + +/** + * Drives one sans-IO client endpoint on its own thread: applies the endpoint's + * actions with an FSocket, then reports bytes, read timeouts, and time. Only + * ReadToken calls back into the owner; everything else reaches the game thread + * through the endpoint's notification queue and Status(). + */ +template +class FEndpointRunner final : public FRunnable +{ +public: + FEndpointRunner(Endpoint& InTarget, const TCHAR* InRole, TFunction InReadToken) + : Target(InTarget) + , Role(InRole) + , ReadToken(MoveTemp(InReadToken)) + , SocketTimeout(SocketTimeoutOf(InTarget)) + , WakeEvent(FPlatformProcess::GetSynchEventFromPool(false)) + { + Buffer.SetNumUninitialized(64 * 1024); + } + + virtual ~FEndpointRunner() override + { + StopAndWait(); + FPlatformProcess::ReturnSynchEventToPool(WakeEvent); + } + + bool Start() + { + Thread = FRunnableThread::Create(this, *FString::Printf(TEXT("OpenUSDConnect%s"), Role), 0, + TPri_BelowNormal); + return Thread != nullptr; + } + + // After another thread queues actions through the endpoint. + void Wake() + { + WakeEvent->Trigger(); + } + + virtual void Stop() override + { + Target.Stop(); + bStopping = true; + WakeEvent->Trigger(); + } + + void StopAndWait() + { + Stop(); + if (Thread) + { + Thread->WaitForCompletion(); + delete Thread; + Thread = nullptr; + } + } + + virtual uint32 Run() override + { + for (;;) + { + ApplyActions(); + if (Socket) + { + Read(); + } + else if (Target.Status().Stopped) + { + return 0; + } + else + { + const std::optional Due = Target.NextWake(); + OUC::Runner::Wait(*WakeEvent, Due, false); + TickIfDue(Due); + } + } + } + +private: + using TimePoint = openusdconnect::client::TimePoint; + using DisconnectReason = openusdconnect::client::DisconnectReason; + static constexpr bool bReadsTimeOut = + std::is_same_v; + + static std::chrono::milliseconds SocketTimeoutOf(const Endpoint& InTarget) + { + if constexpr (bReadsTimeOut) + { + return InTarget.Configuration().SocketTimeout; + } + else + { + return InTarget.Configuration().HandshakeTimeout; + } + } + + static TimePoint Now() + { + return std::chrono::steady_clock::now(); + } + + void ApplyActions() + { + for (std::vector Actions = Target.TakeActions(); + !Actions.empty(); Actions = Target.TakeActions()) + { + for (const openusdconnect::client::Action& Action : Actions) + { + std::visit( + [this](const auto& Value) + { + Apply(Value); + }, + Action); + } + } + } + + void Apply(const openusdconnect::client::ConnectAction& Action) + { + Socket = OUC::Runner::Connect(Action, bStopping); + if (!Socket) + { + Target.OnDisconnected(DisconnectReason::ConnectFailed, Now()); + return; + } + LastReceived = Now(); + Target.OnConnected(OUC::ToUtf8(ReadToken())); + } + + void Apply(const openusdconnect::client::SendAction& Action) + { + if (Socket && + !OUC::Runner::SendAll(*Socket, *Action.Bytes, Now() + SocketTimeout, bStopping)) + { + Close(DisconnectReason::TransportError); + } + } + + void Apply(const openusdconnect::client::CloseAction& Action) + { + Close(Action.Reason); + } + + void Apply(const openusdconnect::client::LogAction& Action) + { + OUC::Runner::Log(Role, Action); + } + + void Close(DisconnectReason Reason) + { + OUC::Runner::Destroy(std::exchange(Socket, nullptr)); + Target.OnDisconnected(Reason, Now()); + } + + void Read() + { + const std::optional Due = Target.NextWake(); + int32 Received = 0; + const FTimespan WaitTime = + bReadsTimeOut ? OUC::Runner::ReceiveSlice(Due) : FTimespan::Zero(); + if (!OUC::Runner::Receive(*Socket, Buffer, WaitTime, Received)) + { + Close(DisconnectReason::PeerClosed); + return; + } + if (Received > 0) + { + LastReceived = Now(); + Target.OnBytes(Buffer.GetData(), static_cast(Received)); + return; + } + if constexpr (bReadsTimeOut) + { + if (Now() - LastReceived >= SocketTimeout) + { + LastReceived = Now(); + Target.OnReadTimeout(); + return; + } + } + else + { + OUC::Runner::Wait(*WakeEvent, Due, true); + } + TickIfDue(Due); + } + + void TickIfDue(std::optional Due) + { + if (const TimePoint Current = Now(); Due && Current >= *Due) + { + Target.OnTick(Current); + } + } + + Endpoint& Target; + const TCHAR* const Role; + const TFunction ReadToken; + const std::chrono::milliseconds SocketTimeout; + FEvent* const WakeEvent; + FRunnableThread* Thread = nullptr; + std::atomic bStopping = false; + // Touched only by the runner thread. + FSocket* Socket = nullptr; + TimePoint LastReceived; + TArray Buffer; +}; diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/OpenUSDConnectClientCore.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/OpenUSDConnectClientCore.cpp deleted file mode 100644 index fe2e8a7..0000000 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/OpenUSDConnectClientCore.cpp +++ /dev/null @@ -1,5 +0,0 @@ -// Copyright OpenUSDConnect Contributors. All Rights Reserved. - -// UnrealBuildTool compiles module-local translation units. Include the portable -// implementation here so Unreal and the nanobind module execute the same code. -#include "../../ThirdParty/OpenUSDConnectClientCore/src/frame_codec.cpp" diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/SyncClient.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/SyncClient.cpp deleted file mode 100644 index be38c49..0000000 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/SyncClient.cpp +++ /dev/null @@ -1,546 +0,0 @@ -// Copyright OpenUSDConnect Contributors. All Rights Reserved. - -#include "SyncClient.h" - -#include "USDConnectProtocol.h" -#include "USDConnectSubsystem.h" -#include "USDWireFraming.h" -#include "USDEventApplier.h" - -#include "Logging/LogMacros.h" -#include "Sockets.h" -#include "SocketSubsystem.h" -#include "HAL/PlatformProcess.h" -#include "Misc/ScopeLock.h" - -DEFINE_LOG_CATEGORY_STATIC(LogUSDConnect, Log, All); - -using namespace OUC; - -namespace -{ -constexpr size_t MaxReceiverFrames = 50'000; -} - -// --------------------------------------------------------------------------- -// FSyncClient -// --------------------------------------------------------------------------- - -FSyncClient::FSyncClient(UUSDConnectSubsystem* InOwner, const FString& InHost, int32 InPort, - const FString& InDepartment, const FString& InClientId, - const FString& InSessionOrigin, float InReconnectDelaySecs, - int32 InInitialLastSeq, const FString& InAuthToken) - : Owner(InOwner) - , Host(InHost) - , Port(InPort) - , Department(InDepartment) - , ClientId(InClientId) - , SessionOrigin(InSessionOrigin) - , AuthToken(InAuthToken) - , ReconnectDelaySecs(InReconnectDelaySecs) - , ReceiverSession(InInitialLastSeq + 1, MaxReceiverFrames, true) - , ActiveGeneration(0) - , Socket(nullptr) - , Thread(nullptr) - , bShouldStop(false) - , bConnected(false) -{ -} - -FSyncClient::~FSyncClient() -{ - StopAndWait(); -} - -bool FSyncClient::Start() -{ - Thread = FRunnableThread::Create(this, TEXT("OpenUSDConnect_SyncClient"), 0, TPri_BelowNormal); - return Thread != nullptr; -} - -void FSyncClient::StopAndWait() -{ - Stop(); - if (Thread) - { - Thread->WaitForCompletion(); - delete Thread; - Thread = nullptr; - } - CloseSocket(); -} - -bool FSyncClient::Init() -{ - return true; -} - -uint32 FSyncClient::Run() -{ - // Only log connection failures on state transitions to avoid spamming the - // log every ReconnectDelaySecs when the server is down. The first attempt - // after construction logs at Log; subsequent retries log at Verbose until - // we succeed once. - bool bHasAnnouncedFailure = false; - - while (!bShouldStop.load(std::memory_order_relaxed)) - { - const openusdconnect::client::ConnectionStart Connection = - ReceiverSession.BeginConnection(); - const uint64 ConnectionGeneration = Connection.Generation; - const int32 SyncFrom = Connection.SyncFrom; - - // --- Create socket and connect --- - ISocketSubsystem* SS = ISocketSubsystem::Get(PLATFORM_SOCKETSUBSYSTEM); - if (!SS) - { - ReceiverSession.Disconnect(ConnectionGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - - FSocket* NewSocket = SS->CreateSocket(NAME_Stream, TEXT("USDConnectRecv"), false); - if (!NewSocket) - { - UE_LOG(LogUSDConnect, Warning, TEXT("Failed to create socket")); - ReceiverSession.Disconnect(ConnectionGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - { - FScopeLock Lock(&SocketCS); - Socket = NewSocket; - } - if (bShouldStop.load(std::memory_order_relaxed)) - { - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - break; - } - - int32 ActualRcvBuf = 0; - Socket->SetReceiveBufferSize(256 * 1024, ActualRcvBuf); // best-effort, ignore actual - - TSharedRef Addr = SS->CreateInternetAddr(); - bool bValid = false; - Addr->SetIp(*Host, bValid); - if (!bValid) - { - UE_LOG(LogUSDConnect, Warning, TEXT("Invalid server address: %s"), *Host); - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - Addr->SetPort(Port); - - if (!Socket->Connect(*Addr)) - { - if (bHasAnnouncedFailure) - { - UE_LOG(LogUSDConnect, Verbose, TEXT("Could not connect to %s:%d retrying in %.1fs"), - *Host, Port, ReconnectDelaySecs); - } - else - { - UE_LOG(LogUSDConnect, Log, TEXT("Could not connect to %s:%d retrying in %.1fs"), - *Host, Port, ReconnectDelaySecs); - bHasAnnouncedFailure = true; - } - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - bHasAnnouncedFailure = false; - - // --- Send HELLO --- - UE_LOG(LogUSDConnect, Log, - TEXT("Connected to OpenUSDConnect server at %s:%d (receiver, sync_from=%d)"), *Host, - Port, SyncFrom); - OUC::FWireFrame HelloFrame; - openusdconnect::client::ReplayPrefixClaim ReplayPrefix; - { - FScopeLock Lock(&ReplayIdentityCS); - ReplayPrefix = ReplayIdentityState.BeginConnection(); - } - const openusdconnect::client::FrameResult HelloResult = OUC::BuildHelloFrame( - TEXT("receiver"), SyncFrom, ClientId, SessionOrigin, Department, HelloFrame, AuthToken, - FString(), std::move(ReplayPrefix)); - if (HelloResult != openusdconnect::client::FrameResult::Success || - !SendAll(HelloFrame.GetData(), HelloFrame.Num())) - { - UE_LOG(LogUSDConnect, Warning, TEXT("Failed to send HELLO")); - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - - // --- Read HELLO_OK --- - { - TArray Frame; - if (!RecvFrame(Frame)) - { - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - - openusdconnect::client::EnvelopeView Envelope; - if (openusdconnect::client::DecodeEnvelope( - Frame.GetData(), static_cast(Frame.Num()), Envelope) != - openusdconnect::client::ProtocolResult::Success) - { - UE_LOG(LogUSDConnect, Warning, TEXT("Invalid response to receiver HELLO")); - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - const openusdconnect::client::HandshakeResponseView Response(Envelope); - if (Response.Kind() == - openusdconnect::client::HandshakeResponseKind::AuthenticationRejected) - { - UE_LOG(LogUSDConnect, Error, TEXT("OpenUSDConnect: auth rejected by server")); - if (Owner) - { - Owner->OnClientAuthRejected(TEXT("receiver")); - } - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - return 0; - } - if (Response.Kind() == - openusdconnect::client::HandshakeResponseKind::ConfigurationRejected) - { - const OpenUSDConnect::HelloRejected* Rejection = Response.ConfigurationRejection(); - const OpenUSDConnect::HelloRejectionCode Code = Rejection->code(); - const FString CodeName = - Code == OpenUSDConnect::HelloRejectionCode::Unspecified - ? FString() - : FString(UTF8_TO_TCHAR(OpenUSDConnect::EnumNameHelloRejectionCode(Code))); - const FString Reason = ToFString(Rejection->reason()); - UE_LOG(LogUSDConnect, Error, - TEXT("OpenUSDConnect: receiver rejected by server (%s): %s"), *CodeName, - *Reason); - if (Owner) - { - Owner->OnClientHelloRejected(TEXT("receiver"), CodeName, Reason); - } - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - return 0; - } - if (Response.Kind() != openusdconnect::client::HandshakeResponseKind::Accepted) - { - UE_LOG(LogUSDConnect, Warning, TEXT("Unexpected response to HELLO (type=%u)"), - static_cast(Envelope.PayloadType())); - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - FPlatformProcess::Sleep(ReconnectDelaySecs); - continue; - } - const OpenUSDConnect::HelloOk* HelloOk = Response.Accepted(); - const FString IssuedToken = ToFString(HelloOk->token()); - const flatbuffers::String* HelloServerInstance = HelloOk->server_instance(); - const flatbuffers::Optional HelloReplayEpoch = HelloOk->replay_epoch(); - const std::string_view ServerInstance = HelloServerInstance - ? std::string_view(HelloServerInstance->c_str(), HelloServerInstance->size()) - : std::string_view(); - const std::optional ReplayEpoch = HelloReplayEpoch.has_value() - ? std::optional(*HelloReplayEpoch) - : std::nullopt; - { - FScopeLock Lock(&ReplayIdentityCS); - ReplayIdentityState.AcceptHello( - SyncFrom, HelloOk->replay_identity(), ServerInstance, ReplayEpoch); - } - ActiveGeneration.store(ConnectionGeneration, std::memory_order_release); - if (Owner) - { - Owner->OnReceiverReplayGenerationChanged(ConnectionGeneration); - Owner->OnClientHelloOk(TEXT("receiver")); - } - if (!IssuedToken.IsEmpty()) - { - AuthToken = IssuedToken; - if (Owner) - { - Owner->OnClientTokenIssued(IssuedToken); - } - } - UE_LOG(LogUSDConnect, Log, TEXT("HELLO_OK received entering receive loop")); - } - - bConnected.store(true, std::memory_order_relaxed); - - // --- Receive loop --- - while (!bShouldStop.load(std::memory_order_relaxed)) - { - TArray Frame; - if (!RecvFrame(Frame)) - break; - if (!HandleFrame(ConnectionGeneration, MoveTemp(Frame))) - break; - } - - bConnected.store(false, std::memory_order_relaxed); - CloseSocket(); - ReceiverSession.Disconnect(ConnectionGeneration); - - if (!bShouldStop.load(std::memory_order_relaxed)) - { - if (ReceiverSession.Generation() == ConnectionGeneration) - { - const bool bRequested = - ReceiverSession.RequestReplayFrom(ReceiverSession.LastAppliedSequence() + 1); - check(bRequested); - } - if (Owner) - { - Owner->OnReceiverReplayGenerationChanged(ReceiverSession.Generation()); - } - UE_LOG(LogUSDConnect, Log, TEXT("Receiver disconnected reconnecting in %.1fs"), - ReconnectDelaySecs); - FPlatformProcess::Sleep(ReconnectDelaySecs); - } - } - return 0; -} - -void FSyncClient::Stop() -{ - bShouldStop.store(true, std::memory_order_relaxed); - InterruptSocket(); -} - -bool FSyncClient::MarkAppliedThrough(int32 Seq) -{ - return ReceiverSession.MarkAppliedThrough(ActiveGeneration.load(std::memory_order_acquire), - Seq); -} - -bool FSyncClient::TryPopFrame(FValidatedReceiverFrame& OutFrame) -{ - return ReceiverSession.TryPop(OutFrame); -} - -bool FSyncClient::MarkReplayApplied() -{ - FScopeLock Lock(&ReplayIdentityCS); - if (!ReceiverSession.TryMarkReplayApplied()) - { - return false; - } - ReplayIdentityState.MarkReplayApplied(); - return true; -} - -void FSyncClient::ResetAppliedProgress() -{ - ReceiverSession.ResetAppliedProgress(); -} - -void FSyncClient::RequestReplayFromApplied() -{ - const bool bRequested = - ReceiverSession.RequestReplayFrom(ReceiverSession.LastAppliedSequence() + 1); - check(bRequested); - InterruptSocket(); -} - -void FSyncClient::Exit() -{ - bConnected.store(false, std::memory_order_relaxed); - CloseSocket(); -} - -// --------------------------------------------------------------------------- -// Private helpers -// --------------------------------------------------------------------------- - -void FSyncClient::InterruptSocket() -{ - FScopeLock Lock(&SocketCS); - if (Socket) - { - Socket->Shutdown(ESocketShutdownMode::ReadWrite); - } -} - -void FSyncClient::CloseSocket() -{ - FSocket* SocketToDestroy = nullptr; - { - FScopeLock Lock(&SocketCS); - SocketToDestroy = Socket; - Socket = nullptr; - } - if (SocketToDestroy) - { - SocketToDestroy->Close(); - if (ISocketSubsystem* SS = ISocketSubsystem::Get(PLATFORM_SOCKETSUBSYSTEM)) - { - SS->DestroySocket(SocketToDestroy); - } - } -} - -bool FSyncClient::RecvExact(uint8* Buf, int32 Needed) -{ - int32 Got = 0; - while (Got < Needed) - { - if (bShouldStop.load(std::memory_order_relaxed) || !Socket) - return false; - - int32 Read = 0; - const bool bOK = Socket->Recv(Buf + Got, Needed - Got, Read); - if (!bOK || Read <= 0) - return false; - Got += Read; - } - return true; -} - -bool FSyncClient::RecvFrame(TArray& OutFrame) -{ - uint8 LenBuf[4]; - if (!RecvExact(LenBuf, 4)) - { - return false; - } - - std::size_t PayloadLen = 0; - if (!openusdconnect::client::TryReadFrameHeader(LenBuf, kMaxFrameSize, PayloadLen)) - { - UE_LOG(LogUSDConnect, Warning, TEXT("Bad frame length disconnecting")); - return false; - } - - OutFrame.SetNumUninitialized(static_cast(PayloadLen)); - return RecvExact(OutFrame.GetData(), static_cast(PayloadLen)); -} - -bool FSyncClient::SendAll(const uint8* Data, int32 Len) -{ - int32 Sent = 0; - while (Sent < Len) - { - if (!Socket) - return false; - int32 ThisSent = 0; - if (!Socket->Send(Data + Sent, Len - Sent, ThisSent) || ThisSent <= 0) - { - return false; - } - Sent += ThisSent; - } - return true; -} - -bool FSyncClient::HandleFrame(uint64 Generation, TArray&& Frame) -{ - openusdconnect::client::EnvelopeView Envelope; - if (openusdconnect::client::DecodeEnvelope(Frame.GetData(), static_cast(Frame.Num()), - Envelope) != - openusdconnect::client::ProtocolResult::Success) - { - UE_LOG(LogUSDConnect, Error, TEXT("Invalid receiver frame requesting replay")); - return false; - } - const openusdconnect::client::ControlMessageView Message(Envelope); - - if (Message.Kind() == openusdconnect::client::ControlMessageKind::BroadcastEvent) - { - const OpenUSDConnect::BroadcastEvent* BcEvent = Message.BroadcastEvent(); - const OpenUSDConnect::EventWrapper* Event = BcEvent->event(); - const int32 Seq = BcEvent->seq(); - const FString Origin = ToFString(BcEvent->origin()); - FValidatedReceiverFrame ValidatedFrame; - ValidatedFrame.Bytes = MoveTemp(Frame); - ValidatedFrame.Sequence = Seq; - ValidatedFrame.EventKind = Event->event_type(); - ValidatedFrame.bUsesChangeBlock = - FUSDEventApplier::EventUsesChangeBlock(ValidatedFrame.EventKind); - const openusdconnect::client::AcceptResult Result = - ReceiverSession.Accept(Generation, openusdconnect::client::ReceiverMessageKind::Event, - Seq, MoveTemp(ValidatedFrame)); - if (Result == openusdconnect::client::AcceptResult::Accepted) - { - UE_LOG(LogUSDConnect, Verbose, TEXT("Received BroadcastEvent seq=%d origin='%s'"), Seq, - *Origin); - return true; - } - if (Result == openusdconnect::client::AcceptResult::Duplicate) - { - UE_LOG(LogUSDConnect, Verbose, TEXT("Ignoring duplicate receiver seq=%d"), Seq); - return true; - } - if (Result == openusdconnect::client::AcceptResult::SequenceGap) - { - UE_LOG(LogUSDConnect, Error, - TEXT("Receiver sequence gap before seq=%d requesting replay"), Seq); - } - else if (Result == openusdconnect::client::AcceptResult::QueueFull) - { - UE_LOG(LogUSDConnect, Error, TEXT("Receiver queue is full requesting replay")); - } - else if (Result == openusdconnect::client::AcceptResult::InvalidSequence) - { - UE_LOG(LogUSDConnect, Error, TEXT("Invalid receiver sequence %d requesting replay"), - Seq); - } - return false; - } - else if (Message.Kind() == openusdconnect::client::ControlMessageKind::Ping) - { - // Heartbeat receivers ignore (server is just checking the connection). - } - else if (Message.Kind() == openusdconnect::client::ControlMessageKind::RateLimited) - { - const float Retry = Message.RateLimit()->retry_after(); - UE_LOG(LogUSDConnect, Warning, TEXT("Rate limited sleeping %.1fs"), Retry); - FPlatformProcess::Sleep(Retry); - } - else if (Message.Kind() == openusdconnect::client::ControlMessageKind::Resync) - { - UE_LOG(LogUSDConnect, Log, TEXT("Resync received resetting seq counter")); - FValidatedReceiverFrame ValidatedFrame; - ValidatedFrame.Bytes = MoveTemp(Frame); - ValidatedFrame.bResync = true; - const openusdconnect::client::AcceptResult Result = - ReceiverSession.Accept(Generation, openusdconnect::client::ReceiverMessageKind::Resync, - 0, MoveTemp(ValidatedFrame)); - if (Result == openusdconnect::client::AcceptResult::Accepted) - { - FScopeLock Lock(&ReplayIdentityCS); - ReplayIdentityState.AcceptResync(); - } - if (Owner) - { - Owner->OnReceiverReplayGenerationChanged(Generation); - } - return Result == openusdconnect::client::AcceptResult::Accepted; - } - else if (Message.Kind() == openusdconnect::client::ControlMessageKind::ReplayComplete) - { - const OpenUSDConnect::ReplayComplete* Complete = Message.ReplayComplete(); - FScopeLock Lock(&ReplayIdentityCS); - const openusdconnect::client::AcceptResult Result = - ReceiverSession.AcceptReplayComplete(Generation, Complete->head_seq(), Complete->epoch()); - if (Result == openusdconnect::client::AcceptResult::Accepted) - { - ReplayIdentityState.AcceptReplayComplete(Complete->epoch()); - } - if (Result == openusdconnect::client::AcceptResult::InvalidSequence) - { - UE_LOG(LogUSDConnect, Error, TEXT("Invalid replay head %d"), Complete->head_seq()); - } - return Result == openusdconnect::client::AcceptResult::Accepted; - } - // Handshake responses do not appear in the receive loop. - return true; -} diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/SyncClient.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/SyncClient.h deleted file mode 100644 index dc8e8a9..0000000 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/SyncClient.h +++ /dev/null @@ -1,118 +0,0 @@ -// Copyright OpenUSDConnect Contributors. All Rights Reserved. -#pragma once - -#include "HAL/Runnable.h" -#include "HAL/RunnableThread.h" -#include "HAL/CriticalSection.h" -#include "Sockets.h" -#include "USDConnectProtocol.h" -#include "openusdconnect/client/receiver_session.h" -#include "openusdconnect/client/replay_identity.h" -#include - -class UUSDConnectSubsystem; - -struct FValidatedReceiverFrame -{ - TArray Bytes; - int32 Sequence = 0; - OpenUSDConnect::EventPayload EventKind = OpenUSDConnect::EventPayload::NONE; - bool bResync = false; - bool bUsesChangeBlock = false; -}; - -using FReceiverSession = openusdconnect::client::OrderedReceiverSession; - -/** - * Background TCP thread that connects to the OpenUSDConnect server as a **receiver**. - * - * Protocol: - * 1. Connect TCP to host:port - * 2. Send Envelope{Hello, role="receiver"} with shared ClientId + SessionOrigin - * 3. Read Envelope{HelloOk} auth handshake - * 4. Loop: read 4-byte big-endian length + FlatBuffers payload - * - BroadcastEvent → validate contiguous receipt and enqueue raw bytes - * - Ping → ignore - * - RateLimited → sleep retry_after seconds - * - Resync → begin a new replay generation from sequence zero - * 5. On error: close socket, wait ReconnectDelaySecs, goto 1. - * - * Stop signalling: setting bShouldStop + closing the socket unblocks any pending Recv. - */ -class FSyncClient : public FRunnable -{ -public: - FSyncClient(UUSDConnectSubsystem* InOwner, const FString& InHost, int32 InPort, - const FString& InDepartment, const FString& InClientId, - const FString& InSessionOrigin, float InReconnectDelaySecs, - int32 InInitialLastSeq = 0, const FString& InAuthToken = FString()); - - virtual ~FSyncClient(); - - // FRunnable - virtual bool Init() override; - virtual uint32 Run() override; - virtual void Stop() override; - virtual void Exit() override; - - bool Start(); - void StopAndWait(); - - bool IsConnected() const - { - return bConnected.load(std::memory_order_relaxed); - } - int32 GetLastAppliedSeq() const - { - return ReceiverSession.LastAppliedSequence(); - } - int32 GetReplayHeadSeq() const - { - return ReceiverSession.ReplayHeadSequence(); - } - uint64 GetReplayEpoch() const - { - return ReceiverSession.ReplayEpoch(); - } - uint64 GetGeneration() const - { - return ReceiverSession.Generation(); - } - int32 GetPendingFrameCount() const - { - return static_cast(ReceiverSession.Size()); - } - bool TryPopFrame(FValidatedReceiverFrame& OutFrame); - bool MarkReplayApplied(); - bool MarkAppliedThrough(int32 Seq); - void ResetAppliedProgress(); - void RequestReplayFromApplied(); - -private: - bool RecvExact(uint8* Buf, int32 Needed); - bool RecvFrame(TArray& OutFrame); - bool SendAll(const uint8* Data, int32 Len); - bool HandleFrame(uint64 Generation, TArray&& Frame); - void InterruptSocket(); - void CloseSocket(); - - UUSDConnectSubsystem* Owner; - - FString Host; - int32 Port; - FString Department; - FString ClientId; // shared with FEmitClient - FString SessionOrigin; // shared with FEmitClient for attribution/reconciliation - FString AuthToken; - float ReconnectDelaySecs; - FReceiverSession ReceiverSession; - openusdconnect::client::ReceiverReplayIdentity ReplayIdentityState; - std::atomic ActiveGeneration; - - FSocket* Socket; - FCriticalSection SocketCS; - FCriticalSection ReplayIdentityCS; - FRunnableThread* Thread; - std::atomic bShouldStop; - std::atomic bConnected; -}; diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/Tests/ProducerEndpointStateTests.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/Tests/ProducerEndpointStateTests.cpp deleted file mode 100644 index fe9c4ca..0000000 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/Tests/ProducerEndpointStateTests.cpp +++ /dev/null @@ -1,174 +0,0 @@ -// Copyright OpenUSDConnect Contributors. All Rights Reserved. - -#if WITH_DEV_AUTOMATION_TESTS - -#include "EmitClient.h" -#include "Misc/AutomationTest.h" -#include "USDWireFraming.h" - -IMPLEMENT_SIMPLE_AUTOMATION_TEST(FOpenUSDConnectProducerOutboxReconnectTest, - "OpenUSDConnect.Producer.OutboxSurvivesClientReplacement", - EAutomationTestFlags::EditorContext | - EAutomationTestFlags::EngineFilter) - -bool FOpenUSDConnectProducerOutboxReconnectTest::RunTest(const FString& Parameters) -{ - TSharedRef State = - MakeShared(TEXT("127.0.0.1"), 7200, TEXT(""), TEXT("session-a")); - TestTrue(TEXT("same endpoint matches"), - State->MatchesEndpoint(TEXT("127.0.0.1"), 7200, TEXT(""))); - TestFalse(TEXT("different port is a different producer endpoint"), - State->MatchesEndpoint(TEXT("127.0.0.1"), 7201, TEXT(""))); - uint64 FirstGeneration = 0; - TestTrue(TEXT("first connection starts"), State->BeginConnection(FirstGeneration)); - FString Error; - TestTrue(TEXT("first handshake is accepted"), - State->AcceptServerHighwater(FirstGeneration, 0, Error)); - - FQueuedProducerTxn FirstClaim; - { - OUC::FWireFrame Frame; - TestEqual(TEXT("test frame builds"), - OUC::BuildHelloFrame(TEXT("emitter"), 0, TEXT("client-a"), TEXT("session-a"), - TEXT(""), Frame), - openusdconnect::client::FrameResult::Success); - TestTrue(TEXT("transaction one is accepted"), - State->EnqueueFrame(FirstGeneration, 1, MoveTemp(Frame))); - TestEqual(TEXT("next transaction advances"), State->GetNextTransactionId(), uint64(2)); - TestEqual(TEXT("one transaction remains pending"), State->GetPendingTransactionCount(), - uint64(1)); - TestTrue(TEXT("first client claims the frame"), - State->ClaimNextUnsent(FirstGeneration, FirstClaim)); - TestEqual(TEXT("claimed transaction identity"), FirstClaim.TxnId, uint64(1)); - TestTrue(TEXT("claimed encoded frame remains owned by endpoint state"), - FirstClaim.Frame.IsValid()); - } - - // The subsystem marks the endpoint outbox unsent when one socket client - // stops. A replacement object then claims the same identity and exact bytes. - State->Disconnect(FirstGeneration); - uint64 ReplacementGeneration = 0; - TestTrue(TEXT("replacement connection starts"), State->BeginConnection(ReplacementGeneration)); - TestTrue(TEXT("replacement handshake is accepted"), - State->AcceptServerHighwater(ReplacementGeneration, 0, Error)); - FQueuedProducerTxn ReplacementClaim; - { - FEmitClient ReplacementClient(nullptr, TEXT("client-a"), State, 0.01f); - TestTrue(TEXT("replacement client reclaims the pending frame"), - State->ClaimNextUnsent(ReplacementGeneration, ReplacementClaim)); - TestFalse(TEXT("zero-timeout flush reports outstanding durability"), - ReplacementClient.FlushPending(0.0)); - } - TestEqual(TEXT("replacement keeps transaction identity"), ReplacementClaim.TxnId, uint64(1)); - TestTrue(TEXT("replacement keeps exact encoded bytes"), - ReplacementClaim.Frame == FirstClaim.Frame); - - State->RetireThrough(ReplacementGeneration, 1); - TestEqual(TEXT("acknowledgement empties the outbox"), State->GetPendingTransactionCount(), - uint64(0)); - TestEqual(TEXT("acknowledgement counter advances"), State->GetAcknowledgedTransactionCount(), - uint64(1)); - FEmitClient AcknowledgedClient(nullptr, TEXT("client-a"), State, 0.01f); - TestTrue(TEXT("flush succeeds once the endpoint outbox is acknowledged"), - AcknowledgedClient.FlushPending(0.0)); - return true; -} - -IMPLEMENT_SIMPLE_AUTOMATION_TEST(FOpenUSDConnectProducerEndpointIsolationTest, - "OpenUSDConnect.Producer.EndpointIsolationAndHighwater", - EAutomationTestFlags::EditorContext | - EAutomationTestFlags::EngineFilter) - -bool FOpenUSDConnectProducerEndpointIsolationTest::RunTest(const FString& Parameters) -{ - FProducerEndpointState First(TEXT("server-a"), 7200, TEXT(""), TEXT("session-a")); - uint64 FirstGeneration = 0; - TestTrue(TEXT("first endpoint starts a connection"), First.BeginConnection(FirstGeneration)); - FString Error; - TestTrue(TEXT("first endpoint accepts initial highwater"), - First.AcceptServerHighwater(FirstGeneration, 0, Error)); - OUC::FWireFrame Frame; - TestEqual(TEXT("test frame builds"), - OUC::BuildHelloFrame(TEXT("emitter"), 0, TEXT("client-a"), TEXT("session-a"), - TEXT(""), Frame), - openusdconnect::client::FrameResult::Success); - TestTrue(TEXT("first endpoint accepts transaction one"), - First.EnqueueFrame(FirstGeneration, 1, MoveTemp(Frame))); - First.Disconnect(FirstGeneration); - uint64 ReconnectGeneration = 0; - TestTrue(TEXT("reconnect starts"), First.BeginConnection(ReconnectGeneration)); - TestTrue(TEXT("matching durable highwater is accepted"), - First.AcceptServerHighwater(ReconnectGeneration, 1, Error)); - First.Disconnect(ReconnectGeneration); - uint64 RegressionGeneration = 0; - TestTrue(TEXT("regression connection starts"), First.BeginConnection(RegressionGeneration)); - TestFalse(TEXT("durable highwater regression is rejected"), - First.AcceptServerHighwater(RegressionGeneration, 0, Error)); - TestTrue(TEXT("regression requires explicit recovery"), First.IsRecoveryRequired()); - - FProducerEndpointState Second(TEXT("server-b"), 7200, TEXT(""), TEXT("session-b")); - TestEqual(TEXT("new endpoint starts a new ordered session"), Second.GetNextTransactionId(), - uint64(1)); - TestFalse(TEXT("new endpoint does not inherit recovery state"), Second.IsRecoveryRequired()); - return true; -} - -IMPLEMENT_SIMPLE_AUTOMATION_TEST(FOpenUSDConnectProducerRejectionDispositionTest, - "OpenUSDConnect.Producer.RejectionDisposition", - EAutomationTestFlags::EditorContext | - EAutomationTestFlags::EngineFilter) - -bool FOpenUSDConnectProducerRejectionDispositionTest::RunTest(const FString& Parameters) -{ - FProducerEndpointState Conflict(TEXT("server"), 7200, TEXT(""), TEXT("conflict")); - uint64 ConflictGeneration = 0; - TestTrue(TEXT("conflict session starts"), Conflict.BeginConnection(ConflictGeneration)); - FString Error; - TestTrue(TEXT("conflict session handshake"), - Conflict.AcceptServerHighwater(ConflictGeneration, 0, Error)); - OUC::FWireFrame ConflictFrame; - TestEqual(TEXT("conflict frame builds"), - OUC::BuildHelloFrame(TEXT("emitter"), 0, TEXT("client"), TEXT("conflict"), TEXT(""), - ConflictFrame), - openusdconnect::client::FrameResult::Success); - TestTrue(TEXT("conflict transaction queued"), - Conflict.EnqueueFrame(ConflictGeneration, 1, MoveTemp(ConflictFrame))); - Conflict.MarkRejected(ConflictGeneration, 1, 3, TEXT("obsolete layer graph")); - TestEqual(TEXT("stale graph is recoverable"), Conflict.GetRecoveryDisposition(), - EUSDConnectRecoveryDisposition::RecoverableConflict); - - FProducerEndpointState Invalid(TEXT("server"), 7200, TEXT(""), TEXT("invalid")); - uint64 InvalidGeneration = 0; - TestTrue(TEXT("invalid session starts"), Invalid.BeginConnection(InvalidGeneration)); - TestTrue(TEXT("invalid session handshake"), - Invalid.AcceptServerHighwater(InvalidGeneration, 0, Error)); - OUC::FWireFrame InvalidFrame; - TestEqual(TEXT("invalid frame builds"), - OUC::BuildHelloFrame(TEXT("emitter"), 0, TEXT("client"), TEXT("invalid"), TEXT(""), - InvalidFrame), - openusdconnect::client::FrameResult::Success); - TestTrue(TEXT("invalid transaction queued"), - Invalid.EnqueueFrame(InvalidGeneration, 1, MoveTemp(InvalidFrame))); - Invalid.MarkRejected(InvalidGeneration, 1, 4, TEXT("invalid event")); - TestEqual(TEXT("malformed operation is an integration fault"), Invalid.GetRecoveryDisposition(), - EUSDConnectRecoveryDisposition::InvalidOperation); - - FProducerEndpointState Sequence(TEXT("server"), 7200, TEXT(""), TEXT("sequence")); - uint64 SequenceGeneration = 0; - TestTrue(TEXT("sequence session starts"), Sequence.BeginConnection(SequenceGeneration)); - TestTrue(TEXT("sequence session handshake"), - Sequence.AcceptServerHighwater(SequenceGeneration, 0, Error)); - OUC::FWireFrame SequenceFrame; - TestEqual(TEXT("sequence frame builds"), - OUC::BuildHelloFrame(TEXT("emitter"), 0, TEXT("client"), TEXT("sequence"), TEXT(""), - SequenceFrame), - openusdconnect::client::FrameResult::Success); - TestTrue(TEXT("sequence transaction queued"), - Sequence.EnqueueFrame(SequenceGeneration, 1, MoveTemp(SequenceFrame))); - Sequence.MarkRejected(SequenceGeneration, 1, 2, TEXT("unexpected transaction id")); - TestEqual(TEXT("sequence contradiction is session-fatal"), Sequence.GetRecoveryDisposition(), - EUSDConnectRecoveryDisposition::SessionFatal); - return true; -} - -#endif // WITH_DEV_AUTOMATION_TESTS diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/Tests/ReceiverCursorTests.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/Tests/ReceiverCursorTests.cpp deleted file mode 100644 index 1d53a29..0000000 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/Tests/ReceiverCursorTests.cpp +++ /dev/null @@ -1,49 +0,0 @@ -// Copyright OpenUSDConnect Contributors. All Rights Reserved. - -#if WITH_DEV_AUTOMATION_TESTS - -#include "SyncClient.h" -#include "Misc/AutomationTest.h" - -IMPLEMENT_SIMPLE_AUTOMATION_TEST(FOpenUSDConnectReceiverAppliedCursorTest, - "OpenUSDConnect.Receiver.AppliedCursorIsMonotonicAndResettable", - EAutomationTestFlags::EditorContext | - EAutomationTestFlags::EngineFilter) - -bool FOpenUSDConnectReceiverAppliedCursorTest::RunTest(const FString& Parameters) -{ - FReceiverSession Session(42, 4); - - TestEqual(TEXT("reconnect starts from the last successfully applied sequence"), - Session.LastAppliedSequence(), 41); - const openusdconnect::client::ConnectionStart Connection = Session.BeginConnection(); - TestEqual(TEXT("invalid replay metadata is rejected without throwing"), - Session.AcceptReplayComplete(Connection.Generation, -1, 0), - openusdconnect::client::AcceptResult::InvalidSequence); - FValidatedReceiverFrame Frame; - Frame.Bytes = {1}; - Frame.Sequence = 42; - TestEqual(TEXT("the receive thread accepts the next ordered frame"), - Session.Accept(Connection.Generation, - openusdconnect::client::ReceiverMessageKind::Event, 42, - MoveTemp(Frame)), - openusdconnect::client::AcceptResult::Accepted); - FValidatedReceiverFrame Drained; - TestTrue(TEXT("the game thread pops one frame without allocating a batch"), - Session.TryPop(Drained)); - TestEqual(TEXT("the popped frame retains its sequence metadata"), Drained.Sequence, 42); - TestTrue(TEXT("successful game-thread application advances the cursor"), - Session.MarkAppliedThrough(Connection.Generation, 42)); - TestEqual(TEXT("successful game-thread application advances the cursor"), - Session.LastAppliedSequence(), 42); - TestFalse(TEXT("an older observation is rejected"), - Session.MarkAppliedThrough(Connection.Generation, 40)); - TestEqual(TEXT("an older observation cannot move the cursor backward"), - Session.LastAppliedSequence(), 42); - Session.ResetAppliedProgress(); - TestEqual(TEXT("an explicit server resync resets the applied cursor"), - Session.LastAppliedSequence(), 0); - return true; -} - -#endif // WITH_DEV_AUTOMATION_TESTS diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/TxnBuilder.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/TxnBuilder.cpp index fdc4dbf..e2b708a 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/TxnBuilder.cpp +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/TxnBuilder.cpp @@ -1,43 +1,41 @@ // Copyright OpenUSDConnect Contributors. All Rights Reserved. #include "TxnBuilder.h" -#include "USDConnectProtocol.h" -#include "USDWireFraming.h" -using namespace OUC; +using openusdconnect::client::ProtocolResult; + +static std::string_view ToStringView(const FTCHARToUTF8& Value) +{ + return {Value.Get(), static_cast(Value.Length())}; +} // --------------------------------------------------------------------------- // Shared: Envelope{Txn{events}} wrapping // --------------------------------------------------------------------------- -static openusdconnect::client::FrameResult +static ProtocolResult FinishTxnFrame(flatbuffers::FlatBufferBuilder& Builder, uint64 TxnId, const TArray>& Events, - FWireFrame& OutFrame) + std::vector& OutFrame) { - const openusdconnect::client::ProtocolResult Result = - openusdconnect::client::FinishTransactionFrame(Builder, TxnId, Events.GetData(), - static_cast(Events.Num())); - if (Result != openusdconnect::client::ProtocolResult::Success) + const ProtocolResult Result = openusdconnect::client::FinishTransactionFrame( + Builder, TxnId, Events.GetData(), static_cast(Events.Num())); + if (Result == ProtocolResult::Success) { - OutFrame = FWireFrame(); - return ToFrameResult(Result); + const uint8* Bytes = Builder.GetBufferPointer(); + OutFrame.assign(Bytes, Bytes + Builder.GetSize()); } - OutFrame = FWireFrame(Builder.Release()); - return openusdconnect::client::FrameResult::Success; + return Result; } // --------------------------------------------------------------------------- // Build Envelope { Txn { events: [EventWrapper{SetXformTrs}, ...] } } // --------------------------------------------------------------------------- -openusdconnect::client::FrameResult BuildXformTxnFrame(uint64 TxnId, - const TArray& Xforms, - FWireFrame& OutFrame, - bool bIncludeEnsureXformOps) +ProtocolResult BuildXformTxnFrame(uint64 TxnId, const TArray& Xforms, + std::vector& OutFrame, bool bIncludeEnsureXformOps) { - OutFrame = FWireFrame(); if (Xforms.IsEmpty()) - return openusdconnect::client::FrameResult::EmptyPayload; + return ProtocolResult::EmptyTransaction; flatbuffers::FlatBufferBuilder Builder(512 + Xforms.Num() * (bIncludeEnsureXformOps ? 192 : 128)); @@ -52,11 +50,11 @@ openusdconnect::client::FrameResult BuildXformTxnFrame(uint64 TxnId, if (bIncludeEnsureXformOps) { flatbuffers::Offset Ensure; - const openusdconnect::client::ProtocolResult Result = + const ProtocolResult Result = openusdconnect::client::BuildEnsureXformOpsEvent(Builder, Prim, Ensure); - if (Result != openusdconnect::client::ProtocolResult::Success) + if (Result != ProtocolResult::Success) { - return ToFrameResult(Result); + return Result; } Events.Add(Ensure); } @@ -64,11 +62,11 @@ openusdconnect::client::FrameResult BuildXformTxnFrame(uint64 TxnId, const openusdconnect::client::XformTrsEventView View{ToStringView(PrimUtf8), X.T, X.R, X.S, X.Fields}; flatbuffers::Offset Event; - const openusdconnect::client::ProtocolResult Result = + const ProtocolResult Result = openusdconnect::client::BuildXformTrsEvent(Builder, View, Prim, Event); - if (Result != openusdconnect::client::ProtocolResult::Success) + if (Result != ProtocolResult::Success) { - return ToFrameResult(Result); + return Result; } Events.Add(Event); } @@ -79,13 +77,11 @@ openusdconnect::client::FrameResult BuildXformTxnFrame(uint64 TxnId, // --------------------------------------------------------------------------- // Build Envelope { Txn { events: [EventWrapper{SetVisibility}, ...] } } // --------------------------------------------------------------------------- -openusdconnect::client::FrameResult -BuildVisibilityTxnFrame(uint64 TxnId, const TArray& Visibilities, - FWireFrame& OutFrame) +ProtocolResult BuildVisibilityTxnFrame(uint64 TxnId, const TArray& Visibilities, + std::vector& OutFrame) { - OutFrame = FWireFrame(); if (Visibilities.IsEmpty()) - return openusdconnect::client::FrameResult::EmptyPayload; + return ProtocolResult::EmptyTransaction; flatbuffers::FlatBufferBuilder Builder(256 + Visibilities.Num() * 64); @@ -97,11 +93,11 @@ BuildVisibilityTxnFrame(uint64 TxnId, const TArray& Visibilitie const FTCHARToUTF8 PrimUtf8(*V.PrimPath); const openusdconnect::client::VisibilityEventView View{ToStringView(PrimUtf8), V.bVisible}; flatbuffers::Offset Event; - const openusdconnect::client::ProtocolResult Result = + const ProtocolResult Result = openusdconnect::client::BuildVisibilityEvent(Builder, View, Event); - if (Result != openusdconnect::client::ProtocolResult::Success) + if (Result != ProtocolResult::Success) { - return ToFrameResult(Result); + return Result; } Events.Add(Event); } @@ -112,13 +108,12 @@ BuildVisibilityTxnFrame(uint64 TxnId, const TArray& Visibilitie // --------------------------------------------------------------------------- // Build Envelope { Txn { events: [EventWrapper{SetConnectableInput}, ...] } } // --------------------------------------------------------------------------- -openusdconnect::client::FrameResult -BuildConnectableInputTxnFrame(uint64 TxnId, const TArray& InEvents, - FWireFrame& OutFrame) +ProtocolResult BuildConnectableInputTxnFrame(uint64 TxnId, + const TArray& InEvents, + std::vector& OutFrame) { - OutFrame = FWireFrame(); if (InEvents.IsEmpty()) - return openusdconnect::client::FrameResult::EmptyPayload; + return ProtocolResult::EmptyTransaction; flatbuffers::FlatBufferBuilder Builder(512 + InEvents.Num() * 256); @@ -147,11 +142,11 @@ BuildConnectableInputTxnFrame(uint64 TxnId, const TArray& static_cast(In.Floats.Num()), }; flatbuffers::Offset Input; - const openusdconnect::client::ProtocolResult Result = + const ProtocolResult Result = openusdconnect::client::BuildConnectableInputValue(Builder, View, Input); - if (Result != openusdconnect::client::ProtocolResult::Success) + if (Result != ProtocolResult::Success) { - return ToFrameResult(Result); + return Result; } Inputs.Add(Input); } @@ -159,13 +154,12 @@ BuildConnectableInputTxnFrame(uint64 TxnId, const TArray& const FTCHARToUTF8 PrimUtf8(*Ev.PrimPath); const FTCHARToUTF8 InfoIdUtf8(*Ev.InfoId); flatbuffers::Offset Event; - const openusdconnect::client::ProtocolResult Result = - openusdconnect::client::BuildConnectableInputEvent( - Builder, ToStringView(PrimUtf8), ToStringView(InfoIdUtf8), Inputs.GetData(), - static_cast(Inputs.Num()), Event); - if (Result != openusdconnect::client::ProtocolResult::Success) + const ProtocolResult Result = openusdconnect::client::BuildConnectableInputEvent( + Builder, ToStringView(PrimUtf8), ToStringView(InfoIdUtf8), Inputs.GetData(), + static_cast(Inputs.Num()), Event); + if (Result != ProtocolResult::Success) { - return ToFrameResult(Result); + return Result; } Events.Add(Event); } diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/TxnBuilder.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/TxnBuilder.h index 213595d..84bd1a1 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/TxnBuilder.h +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/TxnBuilder.h @@ -2,29 +2,31 @@ #pragma once #include "CoreMinimal.h" +#include "USDConnectProtocol.h" #include "USDStageValues.h" -#include "USDWireFraming.h" + +#include /** * Encode a batch of SetXformTrs events into a complete Envelope{Txn} FlatBuffers frame, * including the 4-byte big-endian length prefix. When bIncludeEnsureXformOps is true, * each value event is preceded by its structural xform-op prerequisite. */ -openusdconnect::client::FrameResult BuildXformTxnFrame(uint64 TxnId, - const TArray& Xforms, - OUC::FWireFrame& OutFrame, - bool bIncludeEnsureXformOps = false); +openusdconnect::client::ProtocolResult BuildXformTxnFrame(uint64 TxnId, + const TArray& Xforms, + std::vector& OutFrame, + bool bIncludeEnsureXformOps = false); /** * Encode a batch of SetVisibility events into a complete Envelope{Txn} frame. */ -openusdconnect::client::FrameResult +openusdconnect::client::ProtocolResult BuildVisibilityTxnFrame(uint64 TxnId, const TArray& Visibilities, - OUC::FWireFrame& OutFrame); + std::vector& OutFrame); /** * Encode a batch of SetConnectableInput events into a complete Envelope{Txn} frame. */ -openusdconnect::client::FrameResult +openusdconnect::client::ProtocolResult BuildConnectableInputTxnFrame(uint64 TxnId, const TArray& Events, - OUC::FWireFrame& OutFrame); + std::vector& OutFrame); diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/USDConnectSubsystem.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/USDConnectSubsystem.cpp index e3366d6..7f31c67 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/USDConnectSubsystem.cpp +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/USDConnectSubsystem.cpp @@ -2,8 +2,7 @@ #include "USDConnectSubsystem.h" -#include "SyncClient.h" -#include "EmitClient.h" +#include "EndpointRunner.h" #include "USDConnectSettings.h" #include "USDEventApplier.h" #include "USDMaterialXMaterializer.h" @@ -16,12 +15,16 @@ #include "EngineUtils.h" #include "Logging/LogMacros.h" #include "Stats/Stats.h" -#include "Async/Async.h" #include "Misc/Guid.h" #include "Misc/App.h" #include "Misc/ConfigCacheIni.h" #include "Misc/Crc.h" +#include "Misc/ScopeLock.h" #include "HAL/PlatformProcess.h" +#include "HAL/PlatformTime.h" + +#include +#include #if WITH_EDITOR #include "ScopedTransaction.h" @@ -32,6 +35,48 @@ DEFINE_LOG_CATEGORY_STATIC(LogUSDConnectSubsystem, Log, All); DECLARE_STATS_GROUP(TEXT("OpenUSDConnect"), STATGROUP_OpenUSDConnect, STATCAT_Advanced); DECLARE_CYCLE_STAT(TEXT("USDConnect Tick"), STAT_USDConnectTick, STATGROUP_OpenUSDConnect); +namespace ClientCore = openusdconnect::client; + +// The latest captured value per prim, and per input of a prim. +struct FUSDConnectCapturedEdits +{ + TMap Xforms; + TMap Visibilities; + TMap Inputs; +}; + +namespace +{ +ClientCore::TimePoint Now() +{ + return std::chrono::steady_clock::now(); +} + +EUSDConnectRecoveryDisposition +ToRecoveryDisposition(ClientCore::ProducerRecoveryDisposition Disposition) +{ + switch (Disposition) + { + case ClientCore::ProducerRecoveryDisposition::RecoverableConflict: + return EUSDConnectRecoveryDisposition::RecoverableConflict; + case ClientCore::ProducerRecoveryDisposition::InvalidOperation: + return EUSDConnectRecoveryDisposition::InvalidOperation; + case ClientCore::ProducerRecoveryDisposition::SessionFatal: + return EUSDConnectRecoveryDisposition::SessionFatal; + case ClientCore::ProducerRecoveryDisposition::None: + break; + } + return EUSDConnectRecoveryDisposition::None; +} + +bool TargetsEndpoint(const ClientCore::ProducerConfig& Config, const FString& Host, int32 Port, + const FString& Department) +{ + return Config.Host == OUC::ToUtf8(Host) && Config.Port == Port && + Config.Department == OUC::ToUtf8(Department); +} +} // namespace + // --------------------------------------------------------------------------- // Authentication helpers // --------------------------------------------------------------------------- @@ -54,8 +99,9 @@ void UUSDConnectSubsystem::Initialize(FSubsystemCollectionBase& Collection) { Super::Initialize(Collection); bSuppressEmit.store(false); - bReplaySynchronized.store(false); - ActiveReplayGeneration.store(0); + ReceiverNotifications = MakeShared(); + ProducerNotifications = MakeShared(); + CapturedEdits = MakeShared(); // Generate the stable client ID. Producer session identity is endpoint-scoped // and is created lazily when ConnectResolved selects an endpoint. @@ -78,6 +124,7 @@ void UUSDConnectSubsystem::Initialize(FSubsystemCollectionBase& Collection) void UUSDConnectSubsystem::Deinitialize() { Disconnect(); + ReleaseProducer(); Super::Deinitialize(); } @@ -90,30 +137,48 @@ void UUSDConnectSubsystem::Connect() ConnectResolved(false); } -void UUSDConnectSubsystem::StopClients() +void UUSDConnectSubsystem::ReleaseReceiver() { - if (SyncClient) + if (ReceiverRunner) { - SyncClient->StopAndWait(); - SyncClient.Reset(); + ReceiverRunner->StopAndWait(); + ReceiverRunner.Reset(); } - if (EmitClient) + Receiver.Reset(); +} + +void UUSDConnectSubsystem::ReleaseProducer() +{ + if (ProducerRunner) { - const uint64 Pending = ProducerState ? ProducerState->GetPendingTransactionCount() : 0; - if (Pending > 0 && !EmitClient->FlushPending(2.0)) + ProducerRunner->StopAndWait(); + ProducerRunner.Reset(); + } + Producer.Reset(); +} + +void UUSDConnectSubsystem::StopClients() +{ + ReleaseReceiver(); + if (Producer) + { + if (bActiveEmitterStarted && !Flush(2.0f)) { UE_LOG(LogUSDConnectSubsystem, Warning, TEXT("Disconnecting with %llu unacknowledged producer transactions"), - ProducerState->GetPendingTransactionCount()); + static_cast(Producer->Status().PendingTransactions)); + } + Producer->Disconnect(); + ProducerRunner->Wake(); + // Lets the runner write the Quit before ReleaseProducer's Stop aborts its sends. + const double CloseDeadline = FPlatformTime::Seconds() + 0.3; + while (Producer->Status().Closing && FPlatformTime::Seconds() < CloseDeadline) + { + FPlatformProcess::Sleep(0.005f); } - EmitClient->StopAndWait(); - EmitClient.Reset(); } + DeliverNotifications(); EmittedXformPrims.Reset(); - LastEmitConnectionGeneration = 0; - bReplaySynchronized.store(false); - ReplayHeadSeq = 0; - ReplayEpoch = 0; ActiveServerHost.Empty(); ActiveServerPort = 0; @@ -122,7 +187,7 @@ void UUSDConnectSubsystem::StopClients() bActiveUsingLiveMetadata = false; bDeferredEmitterForToken = false; ActiveSnapshotSeq = 0; - ActiveAuthToken.Empty(); + SetAuthToken(FString()); } void UUSDConnectSubsystem::ConnectResolved(bool bRespectLiveMetadataAutoStart) @@ -174,21 +239,29 @@ void UUSDConnectSubsystem::ConnectResolved(bool bRespectLiveMetadataAutoStart) } const bool bSameProducerEndpoint = - ProducerState && - ProducerState->MatchesEndpoint(TargetHost, TargetPort, Settings->Department); - if (ProducerState && !bSameProducerEndpoint && ProducerState->GetPendingTransactionCount() > 0) + Producer && + TargetsEndpoint(Producer->Configuration(), TargetHost, TargetPort, Settings->Department); + if (Producer && !bSameProducerEndpoint) { - SetStatusMessage(TEXT("pending_transactions"), - FString::Printf(TEXT("Cannot switch endpoints with %llu unacknowledged " - "transaction(s); reconnect to %s:%d and flush first"), - ProducerState->GetPendingTransactionCount(), - *ProducerState->GetHost(), ProducerState->GetPort())); - return; + const uint64 Pending = Producer->Status().PendingTransactions; + if (Pending > 0) + { + SetStatusMessage( + TEXT("pending_transactions"), + FString::Printf(TEXT("Cannot switch endpoints with %llu unacknowledged " + "transaction(s); reconnect to %s:%d and flush first"), + Pending, *OUC::ToFString(Producer->Configuration().Host), + static_cast(Producer->Configuration().Port))); + return; + } } - if (bSameProducerEndpoint && ProducerState->IsRecoveryRequired()) + if (bSameProducerEndpoint) { - SetStatusMessage(TEXT("recovery_required"), ProducerState->GetRecoveryReason()); - return; + if (const std::optional Failure = Producer->Failure()) + { + SetStatusMessage(TEXT("recovery_required"), OUC::ToFString(Failure->Describe())); + return; + } } const FString TargetToken = Settings->bPersistAuthTokens @@ -204,7 +277,7 @@ void UUSDConnectSubsystem::ConnectResolved(bool bRespectLiveMetadataAutoStart) ActiveServerPort = TargetPort; bActiveUsingLiveMetadata = bUsingLiveMetadata; ActiveSnapshotSeq = ReceiverInitialLastSeq; - ActiveAuthToken = TargetToken; + SetAuthToken(TargetToken); SetStatusMessage(bTargetRequiresToken ? TEXT("token_required") : TEXT("not_connected"), TEXT("Live metadata configured; auto-start disabled")); UE_LOG(LogUSDConnectSubsystem, Log, @@ -213,7 +286,7 @@ void UUSDConnectSubsystem::ConnectResolved(bool bRespectLiveMetadataAutoStart) return; } - if ((SyncClient || EmitClient) && ActiveServerHost == TargetHost && + if ((Receiver || bActiveEmitterStarted) && ActiveServerHost == TargetHost && ActiveServerPort == TargetPort && bActiveReceiverStarted == bStartReceiver && bActiveEmitterStarted == bStartEmitter) { @@ -222,11 +295,42 @@ void UUSDConnectSubsystem::ConnectResolved(bool bRespectLiveMetadataAutoStart) } StopClients(); + const std::chrono::milliseconds ReconnectDelay( + FMath::Max(1, FMath::RoundToInt64(Settings->ReconnectDelaySecs * 1000.0))); + const TFunction ReadToken = [this] + { + return ReadAuthToken(); + }; if (!bSameProducerEndpoint) { - ProducerState = - MakeShared(TargetHost, TargetPort, Settings->Department, - FGuid::NewGuid().ToString(EGuidFormats::Digits)); + ReleaseProducer(); + ClientCore::ProducerConfig ProducerSettings; + ProducerSettings.Host = OUC::ToUtf8(TargetHost); + ProducerSettings.Port = static_cast(TargetPort); + ProducerSettings.ClientId = OUC::ToUtf8(ClientId); + ProducerSettings.SessionId = OUC::ToUtf8(FGuid::NewGuid().ToString(EGuidFormats::Digits)); + ProducerSettings.Origin = ProducerSettings.SessionId; + ProducerSettings.Department = OUC::ToUtf8(Settings->Department); + ProducerSettings.ReconnectBaseDelay = ReconnectDelay; + ProducerSettings.ReconnectMaxDelay = + std::max(ProducerSettings.ReconnectMaxDelay, ReconnectDelay); + if (TargetPort < 1 || TargetPort > MAX_uint16 || + !ClientCore::ProducerEndpoint::IsValidConfiguration(ProducerSettings)) + { + SetStatusMessage(TEXT("error"), FString::Printf(TEXT("Invalid server endpoint %s:%d"), + *TargetHost, TargetPort)); + return; + } + Producer = MakeShared(MoveTemp(ProducerSettings), + *ProducerNotifications); + ProducerRunner = MakeShared>( + *Producer, TEXT("Emitter"), ReadToken); + if (!ProducerRunner->Start()) + { + ReleaseProducer(); + SetStatusMessage(TEXT("error"), TEXT("Failed to start the emitter thread")); + return; + } } UE_LOG(LogUSDConnectSubsystem, Log, TEXT("Connecting to %s:%d (client_id=%s)"), *TargetHost, TargetPort, *ClientId); @@ -244,7 +348,7 @@ void UUSDConnectSubsystem::ConnectResolved(bool bRespectLiveMetadataAutoStart) bActiveUsingLiveMetadata = bUsingLiveMetadata; bDeferredEmitterForToken = bDelayEmitterForToken; ActiveSnapshotSeq = ReceiverInitialLastSeq; - ActiveAuthToken = TargetToken; + SetAuthToken(TargetToken); SetStatusMessage(bTargetRequiresToken && TargetToken.IsEmpty() ? TEXT("token_required") : TEXT("connecting"), bDelayEmitterForToken ? TEXT("Starting receiver first to obtain auth token") @@ -252,71 +356,45 @@ void UUSDConnectSubsystem::ConnectResolved(bool bRespectLiveMetadataAutoStart) if (bStartReceiver) { - SyncClient = - MakeShared(this, TargetHost, TargetPort, Settings->Department, ClientId, - ProducerState->GetSessionId(), Settings->ReconnectDelaySecs, - ReceiverInitialLastSeq, TargetToken); - if (SyncClient->Start()) + ClientCore::ReceiverConfig ReceiverSettings; + ReceiverSettings.Host = OUC::ToUtf8(TargetHost); + ReceiverSettings.Port = static_cast(TargetPort); + ReceiverSettings.ClientId = OUC::ToUtf8(ClientId); + ReceiverSettings.Origin = Producer->Configuration().SessionId; + ReceiverSettings.Department = OUC::ToUtf8(Settings->Department); + ReceiverSettings.LayeredReplay = false; + ReceiverSettings.SyncFrom = ReceiverInitialLastSeq + 1; + ReceiverSettings.ReconnectBaseDelay = ReconnectDelay; + ReceiverSettings.ReconnectMaxDelay = + std::max(ReceiverSettings.ReconnectMaxDelay, ReconnectDelay); + Receiver = MakeShared(MoveTemp(ReceiverSettings), + *ReceiverNotifications); + ReceiverRunner = MakeShared>( + *Receiver, TEXT("Receiver"), ReadToken); + static_cast(Receiver->Start(Now())); + if (ReceiverRunner->Start()) { bActiveReceiverStarted = true; } else { - SyncClient.Reset(); + ReleaseReceiver(); } } if (bStartEmitter && !bDelayEmitterForToken) { - EmitClient = MakeShared(this, ClientId, ProducerState.ToSharedRef(), - Settings->ReconnectDelaySecs, TargetToken); - if (EmitClient->Start()) - { - bActiveEmitterStarted = true; - } - else + bActiveEmitterStarted = true; + // Unlike Tick's requests, an explicit connect also clears a handshake rejection. + const ClientCore::TimePoint Current = Now(); + if (Producer->Connect(Current, Current + Producer->Configuration().HandshakeTimeout) == + ClientCore::ConnectResult::Started) { - EmitClient.Reset(); + ProducerRunner->Wake(); } } } -void UUSDConnectSubsystem::TryStartDeferredEmitter() -{ - if (!bDeferredEmitterForToken || EmitClient || !ProducerState || ActiveServerHost.IsEmpty() || - ActiveServerPort <= 0) - { - return; - } - - const UUSDConnectSettings* Settings = GetDefault(); - if (!Settings) - return; - - FString Token = ActiveAuthToken; - if (Token.IsEmpty() && Settings->bPersistAuthTokens) - { - Token = LoadAuthToken(ActiveServerHost, ActiveServerPort, Settings->Department); - ActiveAuthToken = Token; - } - if (Token.IsEmpty()) - return; - - bDeferredEmitterForToken = false; - EmitClient = MakeShared(this, ClientId, ProducerState.ToSharedRef(), - Settings->ReconnectDelaySecs, Token); - if (EmitClient->Start()) - { - bActiveEmitterStarted = true; - SetStatusMessage(TEXT("connected"), TEXT("Auth token available; emitter started")); - } - else - { - EmitClient.Reset(); - SetStatusMessage(TEXT("error"), TEXT("Failed to start deferred emitter")); - } -} - void UUSDConnectSubsystem::Disconnect() { DetachFromStageActor(); @@ -328,10 +406,29 @@ void UUSDConnectSubsystem::Disconnect() bool UUSDConnectSubsystem::Flush(float TimeoutSeconds) const { - if (EmitClient) - return EmitClient->FlushPending(TimeoutSeconds); - return !ProducerState || (ProducerState->GetPendingTransactionCount() == 0 && - !ProducerState->IsRecoveryRequired()); + if (!Producer) + { + return true; + } + const double Deadline = FPlatformTime::Seconds() + FMath::Max(0.0f, TimeoutSeconds); + for (;;) + { + const ClientCore::ProducerStatus ProducerState = Producer->Status(); + if (ProducerState.Failure || ProducerState.PendingTransactions == 0) + { + return !ProducerState.Failure; + } + if (FPlatformTime::Seconds() >= Deadline) + { + return false; + } + // The game thread waits here, so Tick cannot reconnect the producer. + if (bActiveEmitterStarted) + { + RequestProducerConnect(); + } + FPlatformProcess::Sleep(0.005f); + } } void UUSDConnectSubsystem::RefreshLiveMetadataFromStage(AUsdStageActor* Actor) @@ -362,7 +459,7 @@ void UUSDConnectSubsystem::RefreshLiveMetadataFromStage(AUsdStageActor* Actor) bool UUSDConnectSubsystem::IsConnected() const { - return SyncClient && SyncClient->IsConnected(); + return Receiver && Receiver->Status().Connected; } FUSDConnectStatus UUSDConnectSubsystem::GetStatus() const @@ -373,26 +470,36 @@ FUSDConnectStatus UUSDConnectSubsystem::GetStatus() const Status.bUsingLiveMetadata = bActiveUsingLiveMetadata; Status.SnapshotSeq = ActiveSnapshotSeq; Status.bReceiverStarted = bActiveReceiverStarted; - Status.bReceiverConnected = SyncClient && SyncClient->IsConnected(); - Status.bReceiverSynchronized = bReplaySynchronized.load(); Status.bEmitterStarted = bActiveEmitterStarted; - Status.bEmitterConnected = EmitClient && EmitClient->IsConnected(); - if (ProducerState) + if (Receiver) { - Status.SubmittedTransactions = - static_cast(ProducerState->GetSubmittedTransactionCount()); - Status.AcknowledgedTransactions = - static_cast(ProducerState->GetAcknowledgedTransactionCount()); - Status.PendingTransactions = static_cast(FMath::Min( - ProducerState->GetPendingTransactionCount(), static_cast(MAX_int32))); - Status.bRecoveryRequired = ProducerState->IsRecoveryRequired(); - Status.RecoveryDisposition = ProducerState->GetRecoveryDisposition(); + const ClientCore::ReceiverStatus ReceiverState = Receiver->Status(); + Status.bReceiverConnected = ReceiverState.Connected; + Status.bReceiverSynchronized = ReceiverState.Synchronized; } { FScopeLock Lock(&StatusCS); Status.AuthState = LastAuthState; Status.LastMessage = LastStatusMessage; } + if (Producer) + { + const ClientCore::ProducerStatus ProducerState = Producer->Status(); + Status.bEmitterConnected = ProducerState.Connected; + Status.SubmittedTransactions = static_cast(ProducerState.NextTransactionId - 1); + Status.AcknowledgedTransactions = + static_cast(ProducerState.AcknowledgedTransactions); + Status.PendingTransactions = static_cast( + FMath::Min(ProducerState.PendingTransactions, static_cast(MAX_int32))); + if (ProducerState.Failure) + { + Status.bRecoveryRequired = true; + Status.RecoveryDisposition = + ToRecoveryDisposition(ProducerState.Failure->Disposition()); + Status.AuthState = TEXT("recovery_required"); + Status.LastMessage = OUC::ToFString(ProducerState.Failure->Describe()); + } + } return Status; } @@ -425,144 +532,123 @@ void UUSDConnectSubsystem::SetStatusMessage(const FString& AuthState, const FStr LastStatusMessage = Message; } -void UUSDConnectSubsystem::OnClientTokenIssued(const FString& Token) +FString UUSDConnectSubsystem::ReadAuthToken() const { - if (!IsInGameThread()) - { - TWeakObjectPtr WeakThis(this); - AsyncTask(ENamedThreads::GameThread, - [WeakThis, Token]() - { - if (WeakThis.IsValid()) - { - WeakThis->OnClientTokenIssued(Token); - } - }); - return; - } - - const UUSDConnectSettings* Settings = GetDefault(); - ActiveAuthToken = Token; - bool bPersisted = false; - if (Settings && Settings->bPersistAuthTokens) - { - SaveAuthToken(ActiveServerHost, ActiveServerPort, Settings->Department, Token); - bPersisted = true; - } - SetStatusMessage(bPersisted ? TEXT("token_saved") : TEXT("token_issued"), - bPersisted ? TEXT("Auth token issued and saved") - : TEXT("Auth token issued for this session")); - if (bDeferredEmitterForToken) - { - UE_LOG(LogUSDConnectSubsystem, Log, - TEXT("Auth token issued; emitter will start on the next tick")); - } + FScopeLock Lock(&AuthTokenCS); + return ActiveAuthToken; } -void UUSDConnectSubsystem::OnClientHelloOk(const FString& Role) +void UUSDConnectSubsystem::SetAuthToken(const FString& Token) { - if (!IsInGameThread()) - { - TWeakObjectPtr WeakThis(this); - AsyncTask(ENamedThreads::GameThread, - [WeakThis, Role]() - { - if (WeakThis.IsValid()) - { - WeakThis->OnClientHelloOk(Role); - } - }); - return; - } - - SetStatusMessage(TEXT("connected"), FString::Printf(TEXT("%s connected"), *Role)); -} - -void UUSDConnectSubsystem::OnReceiverReplayGenerationChanged(uint64 ReplayGeneration) -{ - uint64 Current = ActiveReplayGeneration.load(std::memory_order_acquire); - while (ReplayGeneration > Current && - !ActiveReplayGeneration.compare_exchange_weak( - Current, ReplayGeneration, std::memory_order_acq_rel, std::memory_order_acquire)) - { - } - if (ReplayGeneration < Current) - return; - bReplaySynchronized.store(false); -} - -void UUSDConnectSubsystem::OnEmitterTransactionRejected(uint64 TxnId, const FString& Reason) -{ - if (!IsInGameThread()) - { - TWeakObjectPtr WeakThis(this); - AsyncTask(ENamedThreads::GameThread, - [WeakThis, TxnId, Reason]() - { - if (WeakThis.IsValid()) - { - WeakThis->OnEmitterTransactionRejected(TxnId, Reason); - } - }); - return; - } - - SetStatusMessage(TEXT("recovery_required"), - Reason.IsEmpty() - ? FString::Printf(TEXT("Transaction %llu rejected"), TxnId) - : FString::Printf(TEXT("Transaction %llu rejected: %s"), TxnId, *Reason)); + FScopeLock Lock(&AuthTokenCS); + ActiveAuthToken = Token; } -void UUSDConnectSubsystem::OnClientAuthRejected(const FString& Role) +void UUSDConnectSubsystem::DeliverNotifications() { - if (!IsInGameThread()) + for (const bool bProducer : {false, true}) { - TWeakObjectPtr WeakThis(this); - AsyncTask(ENamedThreads::GameThread, - [WeakThis, Role]() - { - if (WeakThis.IsValid()) - { - WeakThis->OnClientAuthRejected(Role); - } - }); - return; + ClientCore::NotificationQueue& Queue = + bProducer ? *ProducerNotifications : *ReceiverNotifications; + const TCHAR* Role = bProducer ? TEXT("Emitter") : TEXT("Receiver"); + for (const ClientCore::Notification& Notification : Queue.Drain()) + { + std::visit( + [this, bProducer, Role](const auto& Value) + { + using FKind = std::decay_t; + if constexpr (std::is_same_v) + { + const FString Token = OUC::ToFString(Value.Token); + SetAuthToken(Token); + const UUSDConnectSettings* Settings = GetDefault(); + const bool bPersisted = Settings && Settings->bPersistAuthTokens; + if (bPersisted) + { + SaveAuthToken(ActiveServerHost, ActiveServerPort, Settings->Department, + Token); + } + SetStatusMessage(bPersisted ? TEXT("token_saved") : TEXT("token_issued"), + bPersisted ? TEXT("Auth token issued and saved") + : TEXT("Auth token issued for this session")); + if (bDeferredEmitterForToken) + { + bDeferredEmitterForToken = false; + bActiveEmitterStarted = true; + SetStatusMessage(TEXT("connected"), + TEXT("Auth token available; emitter started")); + } + } + else if constexpr (std::is_same_v) + { + if (bProducer) + { + // A replacement server has no guarantee that it saw prerequisites from + // the previous connection, so each prim's next transform carries them + // again. A captured edit is newer than the current values and wins. + AUsdStageActor* StageActor = CachedStageActor.Get(); + for (const FString& PrimPath : EmittedXformPrims) + { + if (IsValid(StageActor) && + !CapturedEdits->Xforms.Contains(PrimPath) && + !CapturedEdits->Visibilities.Contains(PrimPath)) + { + CapturePrim(StageActor, PrimPath); + } + } + EmittedXformPrims.Reset(); + } + SetStatusMessage(TEXT("connected"), + FString::Printf(TEXT("%s connected to %s:%d"), Role, + *ActiveServerHost, ActiveServerPort)); + } + else if constexpr (std::is_same_v) + { + SetStatusMessage(TEXT("not_connected"), + FString::Printf(TEXT("%s disconnected from %s:%d"), Role, + *ActiveServerHost, ActiveServerPort)); + } + else if constexpr (std::is_same_v) + { + const FString Reason = OUC::ToFString(Value.Reason); + if (Value.Authentication) + { + SetStatusMessage( + TEXT("auth_rejected"), + FString::Printf(TEXT("%s auth rejected: %s"), Role, *Reason)); + } + else + { + const OpenUSDConnect::HelloRejectionCode Code = Value.Code; + SetStatusMessage( + Code == OpenUSDConnect::HelloRejectionCode::Unspecified + ? FString(TEXT("connection_rejected")) + : FString(UTF8_TO_TCHAR( + OpenUSDConnect::EnumNameHelloRejectionCode(Code))), + FString::Printf(TEXT("%s connection rejected: %s"), Role, *Reason)); + } + } + // Stage metadata and playback notifications have no consumer in Unreal. + }, + Notification); + } } - - SetStatusMessage(TEXT("auth_rejected"), FString::Printf(TEXT("%s auth rejected"), *Role)); } -void UUSDConnectSubsystem::OnClientHelloRejected(const FString& Role, const FString& Code, - const FString& Reason) +void UUSDConnectSubsystem::RequestProducerConnect() const { - if (!IsInGameThread()) + const ClientCore::TimePoint Current = Now(); + if (Producer->RequestConnect(Current, Current + Producer->Configuration().HandshakeTimeout)) { - TWeakObjectPtr WeakThis(this); - AsyncTask(ENamedThreads::GameThread, - [WeakThis, Role, Code, Reason]() - { - if (WeakThis.IsValid()) - { - WeakThis->OnClientHelloRejected(Role, Code, Reason); - } - }); - return; + ProducerRunner->Wake(); } - - const FString Message = - Reason.IsEmpty() ? FString::Printf(TEXT("%s connection rejected"), *Role) - : FString::Printf(TEXT("%s connection rejected: %s"), *Role, *Reason); - SetStatusMessage(Code.IsEmpty() ? TEXT("connection_rejected") : Code, Message); } void UUSDConnectSubsystem::RequestReceiverReplay(const FString& Reason) { SetStatusMessage(TEXT("receiver_recovering"), Reason); - if (SyncClient) - { - SyncClient->RequestReplayFromApplied(); - OnReceiverReplayGenerationChanged(SyncClient->GetGeneration()); - } + verify(Receiver->RequestReplayFrom(Receiver->Status().LastAppliedSequence + 1)); + ReceiverRunner->Wake(); } // --------------------------------------------------------------------------- @@ -580,7 +666,7 @@ void UUSDConnectSubsystem::Tick(float DeltaTime) // Guard: don't touch anything until the world is fully initialized and not // being torn down. UTickableWorldSubsystem can tick during the tail end of - // world load running heavy work here can deadlock startup. + // world load, and heavy work there can deadlock startup. UWorld* World = GetWorld(); if (!World || !World->bIsWorldInitialized || World->bIsTearingDown) { @@ -611,9 +697,21 @@ void UUSDConnectSubsystem::Tick(float DeltaTime) ConnectResolved(true); } + // Local edits are read before received frames apply, so a replay cannot + // replace them with older server values. + if (StageActor) + { + CaptureEdits(StageActor); + } + DrainAndApply(); - TryStartDeferredEmitter(); + // Token, connection, and rejection reports from both runner threads. + DeliverNotifications(); + if (Producer && bActiveEmitterStarted) + { + RequestProducerConnect(); + } // Emit any user edits captured by the USD notice listener since last tick. DrainAndEmit(); @@ -669,7 +767,7 @@ void UUSDConnectSubsystem::AttachToStageActor(AUsdStageActor* Actor) if (Path.IsEmpty() || Path == TEXT("/")) continue; - // Strip property suffix ".attrName" we want the prim path. + // Strip the property suffix ".attrName" to get the prim path. // Changed "inputs:*" properties keep their name so the drain // can emit just the edited shader inputs. int32 DotIdx = INDEX_NONE; @@ -715,6 +813,7 @@ void UUSDConnectSubsystem::DetachFromStageActor() PendingEmitInputs.Reset(); } + *CapturedEdits = FUSDConnectCapturedEdits(); EmittedXformPrims.Reset(); CachedStageActor = nullptr; LastMaterializedRootLayerIdentifier.Empty(); @@ -726,7 +825,7 @@ void UUSDConnectSubsystem::DetachFromStageActor() void UUSDConnectSubsystem::DrainAndApply() { - if (!SyncClient) + if (!Receiver) { return; } @@ -734,8 +833,8 @@ void UUSDConnectSubsystem::DrainAndApply() AUsdStageActor* StageActor = CachedStageActor.Get(); if (!StageActor || !IsValid(StageActor)) { - constexpr int32 MaxBufferedFrames = 5000; - if (SyncClient->GetPendingFrameCount() > MaxBufferedFrames) + constexpr size_t MaxBufferedFrames = 5000; + if (Receiver->Status().QueuedFrames > MaxBufferedFrames) { RequestReceiverReplay( TEXT("Receiver queue overflowed before a USD stage was available; replay requested " @@ -744,34 +843,12 @@ void UUSDConnectSubsystem::DrainAndApply() return; } - auto PublishReplayIfApplied = [this]() - { - if (!SyncClient || !SyncClient->MarkReplayApplied()) - { - return; - } - ReplayHeadSeq = SyncClient->GetReplayHeadSeq(); - ReplayEpoch = SyncClient->GetReplayEpoch(); - bReplaySynchronized.store(true); - UE_LOG(LogUSDConnectSubsystem, Log, - TEXT("Receiver replay applied through seq=%d epoch=%llu publishing enabled"), - ReplayHeadSeq, static_cast(ReplayEpoch)); - }; - - PublishReplayIfApplied(); - if (SyncClient->GetPendingFrameCount() == 0) - { - return; - } - constexpr int32 MaxApplyPerTick = 512; constexpr double MaxApplySecondsPerTick = 0.016; // 16 ms preserves ~60 fps const double Start = FPlatformTime::Seconds(); int32 Applied = 0; - int32 LastApplied = SyncClient->GetLastAppliedSeq(); - bool bNeedsReplay = false; - FString RecoveryReason; + FString FailureReason; bSuppressEmit.store(true); { @@ -779,35 +856,31 @@ void UUSDConnectSubsystem::DrainAndApply() while (Applied < MaxApplyPerTick && FPlatformTime::Seconds() - Start <= MaxApplySecondsPerTick) { - FValidatedReceiverFrame Frame; - if (!SyncClient->TryPopFrame(Frame)) + // One frame per drain, so the time budget never strands drained frames. + const uint64 Generation = Receiver->Generation(); + const std::vector> Drained = Receiver->DrainFrames(1); + if (Drained.empty()) { break; } - if (Frame.bResync) + ++Applied; + const std::vector& Frame = Drained.front(); + // The endpoint verified every queued envelope. + const OpenUSDConnect::Envelope* Envelope = OpenUSDConnect::GetEnvelope(Frame.data()); + if (Envelope->payload_type() == OpenUSDConnect::Payload::Resync) { RunBlock.Reset(); - SyncClient->ResetAppliedProgress(); - LastApplied = 0; - bReplaySynchronized.store(false); - ++Applied; + Receiver->ResetAppliedProgress(); continue; } - const int32 Seq = Frame.Sequence; - if (Seq <= LastApplied) + const OpenUSDConnect::BroadcastEvent* Broadcast = Envelope->payload_as_BroadcastEvent(); + if (!Broadcast) { - ++Applied; continue; } - if (Seq != LastApplied + 1) - { - bNeedsReplay = true; - RecoveryReason = FString::Printf( - TEXT("Receiver apply gap: expected=%d dequeued=%d"), LastApplied + 1, Seq); - break; - } - - if (Frame.bUsesChangeBlock) + const int32 Seq = Broadcast->seq(); + const OpenUSDConnect::EventPayload EventKind = Broadcast->event()->event_type(); + if (FUSDEventApplier::EventUsesChangeBlock(EventKind)) { if (!RunBlock) { @@ -819,82 +892,58 @@ void UUSDConnectSubsystem::DrainAndApply() RunBlock.Reset(); } FString TouchedPrim; - if (!FUSDEventApplier::ApplyValidatedFrame(Frame.Bytes, StageActor, &TouchedPrim)) + if (!FUSDEventApplier::ApplyValidatedFrame( + MakeArrayView(Frame.data(), static_cast(Frame.size())), StageActor, + &TouchedPrim)) { RunBlock.Reset(); - bNeedsReplay = true; - RecoveryReason = FString::Printf(TEXT("Failed to apply receiver sequence %d"), Seq); - break; - } - LastApplied = Seq; - if (!SyncClient->MarkAppliedThrough(Seq)) - { - RunBlock.Reset(); - bNeedsReplay = true; - RecoveryReason = FString::Printf( - TEXT("Receiver could not advance the applied cursor through seq=%d"), Seq); + FailureReason = FString::Printf(TEXT("Failed to apply receiver sequence %d"), Seq); break; } + // False only when a reconnect replaced the stream, which then resumes + // from its own cursor. + static_cast(Receiver->MarkAppliedThrough(Generation, Seq)); // Received network edits dirty their owning material for the // materializer. EnsurePrim is included because shader-node // creation changes the network without a connectable event. if (!TouchedPrim.IsEmpty() && - (Frame.EventKind == OpenUSDConnect::EventPayload::SetConnectableInput || - Frame.EventKind == OpenUSDConnect::EventPayload::SetConnectableConnection || - Frame.EventKind == OpenUSDConnect::EventPayload::EnsurePrim)) + (EventKind == OpenUSDConnect::EventPayload::SetConnectableInput || + EventKind == OpenUSDConnect::EventPayload::SetConnectableConnection || + EventKind == OpenUSDConnect::EventPayload::EnsurePrim)) { PendingMaterializePrims.Add(MoveTemp(TouchedPrim)); } - ++Applied; } } bSuppressEmit.store(false); - if (bNeedsReplay) + if (!FailureReason.IsEmpty()) { - RequestReceiverReplay(RecoveryReason); + RequestReceiverReplay(FailureReason); return; } - PublishReplayIfApplied(); - const int32 QueueRemaining = SyncClient->GetPendingFrameCount(); - UE_LOG(LogUSDConnectSubsystem, Verbose, - TEXT("Applied %d event(s) this tick (queue remaining: %d)"), Applied, QueueRemaining); + if (Receiver->MarkReplayApplied()) + { + const ClientCore::ReceiverStatus ReceiverState = Receiver->Status(); + UE_LOG(LogUSDConnectSubsystem, Log, + TEXT("Receiver replay applied through seq=%d epoch=%llu publishing enabled"), + ReceiverState.ReplayHeadSequence, static_cast(ReceiverState.ReplayEpoch)); + ReceiverRunner->Wake(); + } + if (Applied > 0) + { + UE_LOG(LogUSDConnectSubsystem, Verbose, + TEXT("Applied %d frame(s) this tick (queue remaining: %llu)"), Applied, + static_cast(Receiver->Status().QueuedFrames)); + } } // --------------------------------------------------------------------------- -// DrainAndEmit (emitter ← USD stage, via TfNotice listener) +// CaptureEdits and DrainAndEmit (emitter ← USD stage, via TfNotice listener) // --------------------------------------------------------------------------- -void UUSDConnectSubsystem::DrainAndEmit() +void UUSDConnectSubsystem::CaptureEdits(AUsdStageActor* StageActor) { - if (!EmitClient || !EmitClient->IsConnected()) - return; - if (!bReplaySynchronized.load()) - return; - if (bSuppressEmit.load()) - return; - - const uint64 ConnectionGeneration = EmitClient->GetConnectionGeneration(); - if (ConnectionGeneration != LastEmitConnectionGeneration) - { - // A replacement server has no guarantee that it saw prerequisites from - // the previous TCP connection. Requeue the current values and make their - // first transaction on this connection self-contained. - { - FScopeLock Lock(&PendingEmitPathsCS); - for (const FString& PrimPath : EmittedXformPrims) - { - PendingEmitPaths.Add(PrimPath); - } - } - EmittedXformPrims.Reset(); - LastEmitConnectionGeneration = ConnectionGeneration; - } - - AUsdStageActor* StageActor = CachedStageActor.Get(); - if (!StageActor || !IsValid(StageActor)) - return; - TSet Changed; TMap> ChangedInputs; { @@ -908,123 +957,152 @@ void UUSDConnectSubsystem::DrainAndEmit() } UE_LOG(LogUSDConnectSubsystem, Verbose, - TEXT("Draining %d changed prim path(s) from FUsdListener"), Changed.Num()); + TEXT("Capturing %d changed prim path(s) from FUsdListener"), Changed.Num()); for (const FString& Path : Changed) { - EmitPrimChange(StageActor, Path); + CapturePrim(StageActor, Path); } for (const auto& Pair : ChangedInputs) { // Edits on a Material's document-projected interface inputs are // local artifacts; reroute them onto the inline shader instead of // emitting an orphan material-level event. The shader authoring - // re-enters this path next tick and emits/rematerializes normally. + // re-enters this path next tick and is captured normally. if (FUSDMaterialXMaterializer::RerouteMaterialInterfaceEdit(StageActor, Pair.Key, Pair.Value)) { continue; } - EmitConnectableInputs(StageActor, Pair.Key, Pair.Value); + FEmitConnectableInput Event; + if (FUSDStageBridge::ReadConnectableInputs(StageActor, Pair.Key, Pair.Value, Event)) + { + FEmitConnectableInput& Captured = CapturedEdits->Inputs.FindOrAdd(Pair.Key); + Captured.PrimPath = MoveTemp(Event.PrimPath); + Captured.InfoId = MoveTemp(Event.InfoId); + for (FEmitConnectableValue& Value : Event.Inputs) + { + FEmitConnectableValue* Existing = Captured.Inputs.FindByPredicate( + [&Value](const FEmitConnectableValue& Candidate) + { + return Candidate.Name == Value.Name; + }); + if (Existing) + { + *Existing = MoveTemp(Value); + } + else + { + Captured.Inputs.Add(MoveTemp(Value)); + } + } + } // Local shader edits also dirty their owning material so the // materializer refreshes the local .mtlx document. PendingMaterializePrims.Add(Pair.Key); } } -void UUSDConnectSubsystem::EmitPrimChange(AUsdStageActor* StageActor, const FString& PrimPath) +void UUSDConnectSubsystem::CapturePrim(AUsdStageActor* StageActor, const FString& PrimPath) { - // Try to read and emit TRS + FEmitXformTrs Xform; + bool bFromMatrixOp = false; + if (FUSDStageBridge::ReadXformTrs(StageActor, PrimPath, Xform, &bFromMatrixOp)) { - FEmitXformTrs Xform; - bool bFromMatrixOp = false; - if (FUSDStageBridge::ReadXformTrs(StageActor, PrimPath, Xform, &bFromMatrixOp)) + if (bFromMatrixOp) { - if (bFromMatrixOp) - { - // Suppress our own listener: the restore fires notices, but - // it re-authors the exact values being emitted below. - bSuppressEmit.store(true); - FUSDStageBridge::RestoreCanonicalXformOps(StageActor, PrimPath, Xform); - bSuppressEmit.store(false); - } - - TArray Batch = {Xform}; - const bool bIncludeEnsureXformOps = !EmittedXformPrims.Contains(PrimPath); - const uint64 TxnId = ProducerState->GetNextTransactionId(); - OUC::FWireFrame Frame; - const openusdconnect::client::FrameResult FrameResult = - BuildXformTxnFrame(TxnId, Batch, Frame, bIncludeEnsureXformOps); - if (FrameResult != openusdconnect::client::FrameResult::Success) - { - UE_LOG(LogUSDConnectSubsystem, Error, - TEXT("EmitPrimChange(%s): failed to build TRS frame"), *PrimPath); - return; - } - UE_LOG(LogUSDConnectSubsystem, Verbose, - TEXT("EmitPrimChange(%s): TRS frame built (%d bytes, fields=0x%02x%s%s); " - "enqueueing"), - *PrimPath, Frame.Num(), Xform.Fields, - bFromMatrixOp ? TEXT(", decomposed from matrix op") : TEXT(""), - bIncludeEnsureXformOps ? TEXT(", includes ensure_xform_ops") : TEXT("")); - if (EmitClient->EnqueueFrame(TxnId, MoveTemp(Frame))) - { - EmittedXformPrims.Add(PrimPath); - } + // Suppress our own listener: the restore fires notices, but + // it re-authors the exact values being captured. + bSuppressEmit.store(true); + FUSDStageBridge::RestoreCanonicalXformOps(StageActor, PrimPath, Xform); + bSuppressEmit.store(false); } + CapturedEdits->Xforms.Add(PrimPath, Xform); } - // Try to read and emit visibility + FEmitVisibility Visibility; + if (FUSDStageBridge::ReadVisibility(StageActor, PrimPath, Visibility)) { - FEmitVisibility Vis; - if (FUSDStageBridge::ReadVisibility(StageActor, PrimPath, Vis)) + CapturedEdits->Visibilities.Add(PrimPath, Visibility); + } +} + +void UUSDConnectSubsystem::DrainAndEmit() +{ + if (!Producer || !Receiver || bSuppressEmit.load()) + return; + // New transactions wait until the receiver's replay is applied. + if (!Producer->Status().Connected || !Receiver->Status().Synchronized) + return; + + for (auto It = CapturedEdits->Xforms.CreateIterator(); It; ++It) + { + const TArray Batch = {It.Value()}; + const bool bIncludeEnsureXformOps = !EmittedXformPrims.Contains(It.Key()); + if (SubmitTransaction(It.Key(), TEXT("TRS"), bIncludeEnsureXformOps ? 2 : 1, + [&](uint64 TxnId, std::vector& Frame) + { + return BuildXformTxnFrame(TxnId, Batch, Frame, + bIncludeEnsureXformOps) == + ClientCore::ProtocolResult::Success; + })) { - TArray Batch = {Vis}; - const uint64 TxnId = ProducerState->GetNextTransactionId(); - OUC::FWireFrame Frame; - const openusdconnect::client::FrameResult FrameResult = - BuildVisibilityTxnFrame(TxnId, Batch, Frame); - if (FrameResult != openusdconnect::client::FrameResult::Success) - { - UE_LOG(LogUSDConnectSubsystem, Error, - TEXT("EmitPrimChange(%s): failed to build visibility frame"), *PrimPath); - return; - } - UE_LOG( - LogUSDConnectSubsystem, Verbose, - TEXT( - "EmitPrimChange(%s): Visibility frame built (%d bytes, visible=%d) enqueueing"), - *PrimPath, Frame.Num(), Vis.bVisible ? 1 : 0); - EmitClient->EnqueueFrame(TxnId, MoveTemp(Frame)); + EmittedXformPrims.Add(It.Key()); + It.RemoveCurrent(); + } + } + for (auto It = CapturedEdits->Visibilities.CreateIterator(); It; ++It) + { + const TArray Batch = {It.Value()}; + if (SubmitTransaction(It.Key(), TEXT("visibility"), 1, + [&](uint64 TxnId, std::vector& Frame) + { + return BuildVisibilityTxnFrame(TxnId, Batch, Frame) == + ClientCore::ProtocolResult::Success; + })) + { + It.RemoveCurrent(); + } + } + for (auto It = CapturedEdits->Inputs.CreateIterator(); It; ++It) + { + const TArray Batch = {It.Value()}; + if (SubmitTransaction(It.Key(), TEXT("connectable input"), 1, + [&](uint64 TxnId, std::vector& Frame) + { + return BuildConnectableInputTxnFrame(TxnId, Batch, Frame) == + ClientCore::ProtocolResult::Success; + })) + { + It.RemoveCurrent(); } } } -void UUSDConnectSubsystem::EmitConnectableInputs(AUsdStageActor* StageActor, - const FString& PrimPath, - const TSet& InputAttrNames) +bool UUSDConnectSubsystem::SubmitTransaction( + const FString& PrimPath, const TCHAR* Kind, int32 EventCount, + TFunctionRef&)> BuildFrame) { - FEmitConnectableInput Event; - if (!FUSDStageBridge::ReadConnectableInputs(StageActor, PrimPath, InputAttrNames, Event)) + FScopeLock Lock(&SubmitCS); + const uint64 TxnId = Producer->NextTransactionId(); + std::vector Frame; + if (!BuildFrame(TxnId, Frame)) { - return; + UE_LOG(LogUSDConnectSubsystem, Error, TEXT("Failed to build the %s transaction for %s"), + Kind, *PrimPath); + return false; } - - TArray Batch = {MoveTemp(Event)}; - const uint64 TxnId = ProducerState->GetNextTransactionId(); - OUC::FWireFrame Frame; - const openusdconnect::client::FrameResult FrameResult = - BuildConnectableInputTxnFrame(TxnId, Batch, Frame); - if (FrameResult != openusdconnect::client::FrameResult::Success) + UE_LOG(LogUSDConnectSubsystem, Verbose, TEXT("Appending %s transaction %llu for %s (%d bytes)"), + Kind, TxnId, *PrimPath, static_cast(Frame.size())); + if (Producer->Append(TxnId, MoveTemp(Frame), static_cast(EventCount), std::string()) != + ClientCore::ProducerResult::Accepted) { - UE_LOG(LogUSDConnectSubsystem, Error, - TEXT("EmitConnectableInputs(%s): failed to build frame"), *PrimPath); - return; + UE_LOG(LogUSDConnectSubsystem, Warning, + TEXT("The emitter refused the %s transaction for %s"), Kind, *PrimPath); + return false; } - UE_LOG(LogUSDConnectSubsystem, Verbose, - TEXT("EmitConnectableInputs(%s): %d input(s), frame %d bytes enqueueing"), *PrimPath, - Batch[0].Inputs.Num(), Frame.Num()); - EmitClient->EnqueueFrame(TxnId, MoveTemp(Frame)); + ProducerRunner->Wake(); + return true; } // --------------------------------------------------------------------------- @@ -1124,7 +1202,7 @@ void UUSDConnectSubsystem::ProcessPendingMaterializations() // The session-layer authoring below fires stage notices; suppress our own // listener so they don't loop back into the emit path. The stage actor's - // listener still sees them that's what triggers the re-import. + // listener still sees them, which triggers the re-import. bSuppressEmit.store(true); for (const FString& Material : Materials) { diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/USDWireFraming.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/USDWireFraming.h deleted file mode 100644 index 8fed083..0000000 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Private/USDWireFraming.h +++ /dev/null @@ -1,130 +0,0 @@ -// Copyright OpenUSDConnect Contributors. All Rights Reserved. -#pragma once - -#include "CoreMinimal.h" -#include "USDConnectProtocol.h" -#include "openusdconnect/client/frame_codec.h" -#include "openusdconnect/client/protocol_codec.h" -#include "flatbuffers/detached_buffer.h" - -#include - -namespace OUC -{ -inline openusdconnect::client::FrameResult -ToFrameResult(openusdconnect::client::ProtocolResult Result) noexcept -{ - switch (Result) - { - case openusdconnect::client::ProtocolResult::Success: - return openusdconnect::client::FrameResult::Success; - case openusdconnect::client::ProtocolResult::InvalidMaxFrameSize: - return openusdconnect::client::FrameResult::InvalidMaxFrameSize; - case openusdconnect::client::ProtocolResult::PayloadTooLarge: - return openusdconnect::client::FrameResult::PayloadTooLarge; - default: - return openusdconnect::client::FrameResult::EmptyPayload; - } -} - -inline std::string_view ToStringView(const FTCHARToUTF8& Value) noexcept -{ - return {Value.Get(), static_cast(Value.Length())}; -} - -/** - * Immutable contiguous wire frame backed directly by FlatBuffers' allocation. - * The first four bytes are the network-order payload length; the remaining - * bytes are the FlatBuffer. Moving this object transfers allocation ownership. - */ -class FWireFrame final -{ -public: - FWireFrame() = default; - explicit FWireFrame(flatbuffers::DetachedBuffer&& InBuffer) - : Buffer(std::move(InBuffer)) - { - } - - FWireFrame(FWireFrame&&) = default; - FWireFrame& operator=(FWireFrame&&) = default; - FWireFrame(const FWireFrame&) = delete; - FWireFrame& operator=(const FWireFrame&) = delete; - - const uint8* GetData() const - { - return Buffer.data(); - } - int32 Num() const - { - return static_cast(Buffer.size()); - } - bool IsEmpty() const - { - return Buffer.size() == 0; - } - -private: - flatbuffers::DetachedBuffer Buffer; -}; - -/** - * Finish with FlatBuffers' in-allocation size prefix, rewrite that prefix to - * the protocol's big-endian form, and detach the original allocation. - * No serialized payload bytes are copied. - */ -inline openusdconnect::client::FrameResult -FinishWireFrame(flatbuffers::FlatBufferBuilder& Builder, - flatbuffers::Offset RootOffset, FWireFrame& OutFrame) -{ - OutFrame = FWireFrame(); - const openusdconnect::client::ProtocolResult Result = - openusdconnect::client::FinishEnvelopeFrame(Builder, RootOffset); - if (Result != openusdconnect::client::ProtocolResult::Success) - { - return ToFrameResult(Result); - } - OutFrame = FWireFrame(Builder.Release()); - return openusdconnect::client::FrameResult::Success; -} - -// Build a complete framed Envelope{Hello} message. -// Role = "receiver" or "emitter" -// SyncFrom = receiver-side resume point (0 = full replay; emitters pass 0) -// Token = saved TOFU auth token, or empty for first-connect issuance -inline openusdconnect::client::FrameResult -BuildHelloFrame(const FString& Role, int32 SyncFrom, const FString& ClientId, - const FString& SessionOrigin, const FString& Department, FWireFrame& OutFrame, - const FString& Token = FString(), const FString& ProducerSessionId = FString(), - openusdconnect::client::ReplayPrefixClaim ReplayPrefix = std::nullopt) -{ - flatbuffers::FlatBufferBuilder Builder(512); - const FTCHARToUTF8 RoleUtf8(*Role); - const FTCHARToUTF8 ClientIdUtf8(*ClientId); - const FTCHARToUTF8 SessionOriginUtf8(*SessionOrigin); - const FTCHARToUTF8 DepartmentUtf8(*Department); - const FTCHARToUTF8 TokenUtf8(*Token); - const FTCHARToUTF8 ProducerSessionIdUtf8(*ProducerSessionId); - const openusdconnect::client::HelloParameters Parameters{ - ToStringView(RoleUtf8), - SyncFrom, - ToStringView(ClientIdUtf8), - ToStringView(SessionOriginUtf8), - ToStringView(DepartmentUtf8), - ToStringView(TokenUtf8), - false, - OpenUSDConnect::LayerMode::Managed, - ToStringView(ProducerSessionIdUtf8), - std::move(ReplayPrefix), - }; - const openusdconnect::client::ProtocolResult Result = - openusdconnect::client::BuildHelloFrame(Builder, Parameters); - if (Result != openusdconnect::client::ProtocolResult::Success) - { - OutFrame = FWireFrame(); - return ToFrameResult(Result); - } - OutFrame = FWireFrame(Builder.Release()); - return openusdconnect::client::FrameResult::Success; -} -} // namespace OUC diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Public/USDConnectSettings.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Public/USDConnectSettings.h index 87692ff..655fb07 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Public/USDConnectSettings.h +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Public/USDConnectSettings.h @@ -55,7 +55,7 @@ class OPENUSDCONNECT_API UUSDConnectSettings : public UDeveloperSettings UPROPERTY(config, EditAnywhere, Category="Authentication", meta=(DisplayName="Persist Auth Tokens")) bool bPersistAuthTokens = true; - /** Seconds between reconnection attempts after a disconnect */ + /** Seconds before the first reconnection attempt; later attempts back off exponentially */ UPROPERTY(config, EditAnywhere, Category="Connection", meta=(DisplayName="Reconnect Delay (s)", ClampMin=1, ClampMax=60)) float ReconnectDelaySecs = 3.0f; diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Public/USDConnectSubsystem.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Public/USDConnectSubsystem.h index e9979e3..ab657c6 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Public/USDConnectSubsystem.h +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnect/Public/USDConnectSubsystem.h @@ -5,14 +5,23 @@ #include "HAL/CriticalSection.h" #include "Containers/Set.h" #include "Delegates/IDelegateInstance.h" +#include "Templates/Function.h" #include "USDConnectRecovery.h" #include +#include #include "USDConnectSubsystem.generated.h" -class FSyncClient; -class FEmitClient; -class FProducerEndpointState; class AUsdStageActor; +struct FUSDConnectCapturedEdits; +template +class FEndpointRunner; + +namespace openusdconnect::client +{ +class NotificationQueue; +class ProducerEndpoint; +class ReceiverEndpoint; +} // namespace openusdconnect::client USTRUCT(BlueprintType) struct OPENUSDCONNECT_API FUSDConnectStatus @@ -69,26 +78,9 @@ struct OPENUSDCONNECT_API FUSDConnectStatus }; /** - * World subsystem that manages the OpenUSD Connect two-way sync. - * - * Receiver side (server → Unreal): - * FSyncClient background thread receives BroadcastEvent frames and - * pushes raw bytes to the event queue. Tick() drains the queue and - * applies each event to the open AUsdStageActor stage via FUSDEventApplier. - * - * Emitter side (Unreal → server): - * The stage listener fires whenever the USD stage changes locally (viewport - * transforms, USD Stage panel property edits). The subsystem reads the - * current TRS, visibility, and changed shader inputs from the pxr stage and - * sends SetXformTrs/SetVisibility/SetConnectableInput events via - * FEmitClient. A feedback loop guard (bSuppressEmit) prevents echoing - * events received from the server back out. - * - * Usage: - * 1. Place an AUsdStageActor in the level; set its RootLayer to the USD file - * the server is managing. - * 2. Configure host/port in Edit > Project Settings > Plugins > OpenUSD Connect. - * 3. Press Play the subsystem auto-connects both receiver and emitter. + * Two-way sync for the level's AUsdStageActor: Tick applies the events a receiver runner thread + * queues and appends local stage edits, never the applied ones, to the producer a second runner + * thread sends. */ UCLASS() class OPENUSDCONNECT_API UUSDConnectSubsystem : public UTickableWorldSubsystem @@ -108,8 +100,8 @@ class OPENUSDCONNECT_API UUSDConnectSubsystem : public UTickableWorldSubsystem return true; } - // Live sync must run in the editor (without PIE), not just during Play. - // UTickableWorldSubsystem defaults this to false we override to true. + // Live sync must run in the editor (without PIE), not just during Play; + // UTickableWorldSubsystem defaults this to false. virtual bool IsTickableInEditor() const override { return true; @@ -141,67 +133,73 @@ class OPENUSDCONNECT_API UUSDConnectSubsystem : public UTickableWorldSubsystem UFUNCTION(BlueprintPure, Category = "OpenUSD Connect") FUSDConnectStatus GetStatus() const; - /** Called from client background threads when the server issues a TOFU token. */ - void OnClientTokenIssued(const FString& Token); - - /** Called from client background threads after HELLO_OK. */ - void OnClientHelloOk(const FString& Role); - - /** Select a receiver replay generation and discard queued frames from older streams. */ - void OnReceiverReplayGenerationChanged(uint64 ReplayGeneration); - - /** Called when the server deterministically rejects a producer transaction. */ - void OnEmitterTransactionRejected(uint64 TxnId, const FString& Reason); - - /** Called from client background threads when auth is rejected. */ - void OnClientAuthRejected(const FString& Role); - - /** Called from client background threads when the requested mode is rejected. */ - void OnClientHelloRejected(const FString& Role, const FString& Code, const FString& Reason); - private: AUsdStageActor* FindStageActor() const; void AttachToStageActor(AUsdStageActor* Actor); void DetachFromStageActor(); void StopClients(); + void ReleaseReceiver(); + void ReleaseProducer(); void ConnectResolved(bool bRespectLiveMetadataAutoStart); void RefreshLiveMetadataFromStage(AUsdStageActor* Actor); - void TryStartDeferredEmitter(); void QueueInitialMaterializations(AUsdStageActor* Actor); FString LoadAuthToken(const FString& Host, int32 Port, const FString& Department) const; void SaveAuthToken(const FString& Host, int32 Port, const FString& Department, const FString& Token) const; + /** The token the runner threads present at their next handshake. */ + FString ReadAuthToken() const; + void SetAuthToken(const FString& Token); void SetStatusMessage(const FString& AuthState, const FString& Message); + /** Reacts on the game thread to what both endpoints reported since the last call. */ + void DeliverNotifications(); + /** Starts a producer connection attempt unless one is in flight or backing off. */ + void RequestProducerConnect() const; void RequestReceiverReplay(const FString& Reason); void DrainAndApply(); - /** Drain accumulated SdfPaths from the stage listener and emit one frame per unique path */ - void DrainAndEmit(); + /** Read the values of the paths the stage listener reported into CapturedEdits */ + void CaptureEdits(AUsdStageActor* StageActor); - /** Build and send a Txn event for a changed prim (emitter side) */ - void EmitPrimChange(AUsdStageActor* StageActor, const FString& PrimPath); + /** Capture a prim's current transform and visibility */ + void CapturePrim(AUsdStageActor* StageActor, const FString& PrimPath); - /** Build and send a SetConnectableInput Txn for changed shader inputs on one prim */ - void EmitConnectableInputs(AUsdStageActor* StageActor, const FString& PrimPath, - const TSet& InputAttrNames); + /** Send the captured edits, keeping each one the producer refuses */ + void DrainAndEmit(); + + /** + * Pairs the next transaction ID with the frame BuildFrame encodes for it and + * appends the frame to the producer outbox; true when the producer accepted it. + */ + bool SubmitTransaction(const FString& PrimPath, const TCHAR* Kind, int32 EventCount, + TFunctionRef&)> BuildFrame); /** Refresh local .mtlx documents for materials dirtied this tick */ void ProcessPendingMaterializations(); + /** One queue per endpoint, so a notification names its role; Tick drains both. */ + TSharedPtr ReceiverNotifications; + TSharedPtr ProducerNotifications; + // --- Receiver --- - TSharedPtr SyncClient; + TSharedPtr Receiver; + TSharedPtr> ReceiverRunner; // --- Emitter --- - TSharedPtr EmitClient; - TSharedPtr ProducerState; + /** + * Endpoint-scoped producer identity and outbox. It outlives Disconnect() + * and is replaced only when the endpoint changes, so a reconnect to the + * same endpoint keeps transaction identity and unacknowledged frames. + */ + TSharedPtr Producer; + TSharedPtr> ProducerRunner; + FCriticalSection SubmitCS; /** * Transform prims whose structural xform-op prerequisite has been sent on * the current emitter connection. Game-thread only. */ TSet EmittedXformPrims; - uint64 LastEmitConnectionGeneration = 0; /** Stable client ID shared by both receiver and emitter connections */ FString ClientId; @@ -216,20 +214,9 @@ class OPENUSDCONNECT_API UUSDConnectSubsystem : public UTickableWorldSubsystem /** * Set to true while DrainAndApply() is applying received events. * Prevents OnPrimChanged from echoing those changes back to the server. - * - * Uses default seq_cst ordering. The other socket-thread atomics in this - * module use relaxed because they only flag a state for polling; this one - * fences a code region around stage mutation, so the stronger barrier is - * the safer default and not on a hot path. */ std::atomic bSuppressEmit; - /** New native transactions remain gated until replay is applied on the game thread. */ - std::atomic bReplaySynchronized; - std::atomic ActiveReplayGeneration; - int32 ReplayHeadSeq = 0; - uint64 ReplayEpoch = 0; - /** Cached weak reference to the currently attached stage actor */ TWeakObjectPtr CachedStageActor; @@ -251,11 +238,17 @@ class OPENUSDCONNECT_API UUSDConnectSubsystem : public UTickableWorldSubsystem /** * Changed "inputs:*" property names per prim (same lock as PendingEmitPaths). - * Keeping the property names lets the drain read and emit only the edited - * shader inputs instead of the whole network. + * Keeping the property names lets the capture read only the edited shader + * inputs instead of the whole network. */ TMap> PendingEmitInputs; + /** + * Values of local edits, read before received frames apply so a replay + * cannot overwrite them before they are sent. Game thread only. + */ + TSharedPtr CapturedEdits; + /** Active TCP endpoint for the currently running clients. */ FString ActiveServerHost; int32 ActiveServerPort = 0; @@ -264,6 +257,7 @@ class OPENUSDCONNECT_API UUSDConnectSubsystem : public UTickableWorldSubsystem bool bActiveUsingLiveMetadata = false; bool bDeferredEmitterForToken = false; int32 ActiveSnapshotSeq = 0; + mutable FCriticalSection AuthTokenCS; FString ActiveAuthToken; /** Last live metadata key seen on the attached stage root layer. */ @@ -279,8 +273,8 @@ class OPENUSDCONNECT_API UUSDConnectSubsystem : public UTickableWorldSubsystem /** * Prims whose material networks changed this tick (received events and * local edits), resolved to owning materials and materialized to local - * .mtlx documents at the end of Tick. Game thread only both producers - * (DrainAndApply, DrainAndEmit) and the consumer run there. + * .mtlx documents at the end of Tick. Game thread only, where both producers + * (DrainAndApply, DrainAndEmit) and the consumer run. */ TSet PendingMaterializePrims; }; diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectClientCore/OpenUSDConnectClientCore.Build.cs b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectClientCore/OpenUSDConnectClientCore.Build.cs new file mode 100644 index 0000000..8e20cea --- /dev/null +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectClientCore/OpenUSDConnectClientCore.Build.cs @@ -0,0 +1,35 @@ +// Copyright OpenUSDConnect Contributors. All Rights Reserved. + +using System.IO; +using UnrealBuildTool; + +// The repository's native/client_core, staged into include/ and src/ before +// BuildPlugin. +public class OpenUSDConnectClientCore : ModuleRules +{ + public OpenUSDConnectClientCore(ReadOnlyTargetRules Target) : base(Target) + { + bRequiresImplementModule = false; + PCHUsage = PCHUsageMode.NoPCHs; + bUseUnity = false; + + PublicIncludePaths.Add(Path.Combine(ModuleDirectory, "include")); + + // Other modules call into this DLL, so the core's API carries the export attribute that + // Platform.h defines. Linking Core gives the DLL Unreal's operator new and delete, which + // the containers crossing the API to other modules require. + PublicDefinitions.Add("OPENUSDCONNECT_CLIENT_API=OPENUSDCONNECTCLIENTCORE_API"); + PrivateDependencyModuleNames.Add("Core"); + ForceIncludeFiles.Add("HAL/Platform.h"); + + string FlatBuffersInclude = Path.GetFullPath(Path.Combine( + ModuleDirectory, "..", "OpenUSDConnectPXR", "ThirdParty", "flatbuffers", "include")); + if (!File.Exists(Path.Combine(FlatBuffersInclude, "flatbuffers", "flatbuffer_builder.h"))) + { + throw new BuildException( + "OpenUSDConnect: FlatBuffers headers not found. Run " + + "python /setup_flatbuffers.py once, then rebuild."); + } + PublicSystemIncludePaths.Add(FlatBuffersInclude); + } +} diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/OpenUSDConnectPXR.Build.cs b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/OpenUSDConnectPXR.Build.cs index 6b46bf3..96bb348 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/OpenUSDConnectPXR.Build.cs +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/OpenUSDConnectPXR.Build.cs @@ -1,16 +1,12 @@ // Copyright OpenUSDConnect Contributors. All Rights Reserved. using UnrealBuildTool; -using System.IO; public class OpenUSDConnectPXR : ModuleRules { public OpenUSDConnectPXR(ReadOnlyTargetRules Target) : base(Target) { PCHUsage = PCHUsageMode.UseExplicitOrSharedPCHs; - PublicIncludePaths.Add(Path.GetFullPath(Path.Combine( - ModuleDirectory, - "../ThirdParty/OpenUSDConnectClientCore/include"))); // pxr headers use typeid. Keep RTTI confined to this pure C++ module; // Unreal's UObject modules and base classes are built without it. @@ -19,6 +15,7 @@ public OpenUSDConnectPXR(ReadOnlyTargetRules Target) : base(Target) PublicDependencyModuleNames.AddRange(new string[] { "Core", + "OpenUSDConnectClientCore", }); PrivateDependencyModuleNames.AddRange(new string[] @@ -33,14 +30,5 @@ public OpenUSDConnectPXR(ReadOnlyTargetRules Target) : base(Target) }); UnrealBuildTool.Rules.UnrealUSDWrapper.CheckAndSetupUsdSdk(Target, this); - - string LocalInclude = Path.Combine(ModuleDirectory, "ThirdParty", "flatbuffers", "include"); - if (!File.Exists(Path.Combine(LocalInclude, "flatbuffers", "flatbuffer_builder.h"))) - { - throw new BuildException( - "OpenUSDConnect: FlatBuffers headers not found. Run " + - "python /setup_flatbuffers.py once, then rebuild."); - } - PublicSystemIncludePaths.Add(LocalInclude); } } diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Private/USDEventApplier.cpp b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Private/USDEventApplier.cpp index b6385a2..8be51f6 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Private/USDEventApplier.cpp +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Private/USDEventApplier.cpp @@ -1789,24 +1789,14 @@ bool UsesChangeBlock(OpenUSDConnect::EventPayload EventKind) } } -// BroadcastEvent's EventWrapper from a raw frame; nullptr on malformed input. -const OpenUSDConnect::EventWrapper* GetFrameEventWrapper(const TArray& RawFrame) -{ - const OpenUSDConnect::Envelope* Env = OUC::GetEnvelopeFromFrame(RawFrame); - const OpenUSDConnect::BroadcastEvent* BcEvent = - Env ? Env->payload_as_BroadcastEvent() : nullptr; - return BcEvent ? BcEvent->event() : nullptr; -} - -const OpenUSDConnect::EventWrapper* GetValidatedFrameEventWrapper(const TArray& RawFrame) +const OpenUSDConnect::EventWrapper* GetValidatedFrameEventWrapper(TConstArrayView RawFrame) { const OpenUSDConnect::Envelope* Env = OpenUSDConnect::GetEnvelope(RawFrame.GetData()); return Env->payload_as_BroadcastEvent()->event(); } bool ApplyEventWrapper(const OpenUSDConnect::EventWrapper* Wrapper, AUsdStageActor* StageActor, - FString* OutTouchedPrim, OpenUSDConnect::EventPayload* OutEventKind, - bool bManageChangeBlock) + FString* OutTouchedPrim, OpenUSDConnect::EventPayload* OutEventKind) { if (!StageActor || !Wrapper) { @@ -1832,21 +1822,7 @@ bool ApplyEventWrapper(const OpenUSDConnect::EventWrapper* Wrapper, AUsdStageAct *OutTouchedPrim = GetEventPrim(Wrapper); } - // Structural events must apply outside an SdfChangeBlock: recomposition is - // deferred until the block closes, so UsdStage::DefinePrim cannot return - // the newly defined prim and arc edits act on a stale composed view. Value - // writes on existing prims are ChangeBlock-safe and batch into one - // consolidated ObjectsChanged notice, which the stage actor's FUsdListener - // turns into a single scene refresh. - if (bManageChangeBlock && UsesChangeBlock(Wrapper->event_type())) - { - pxr::SdfChangeBlock ChangeBlock; - DispatchEvent(PxrStage, Wrapper); - } - else - { - DispatchEvent(PxrStage, Wrapper); - } + DispatchEvent(PxrStage, Wrapper); return true; #else UE_LOG(LogUSDEventApplier, Warning, TEXT("USD SDK not available cannot apply USD events")); @@ -1856,38 +1832,12 @@ bool ApplyEventWrapper(const OpenUSDConnect::EventWrapper* Wrapper, AUsdStageAct } // namespace -// --------------------------------------------------------------------------- -// FUSDEventApplier::FrameUsesChangeBlock -// --------------------------------------------------------------------------- -bool FUSDEventApplier::FrameUsesChangeBlock(const TArray& RawFrame) -{ - if (RawFrame.Num() < 8) - return false; - const OpenUSDConnect::EventWrapper* Wrapper = GetFrameEventWrapper(RawFrame); - return Wrapper && UsesChangeBlock(Wrapper->event_type()); -} - bool FUSDEventApplier::EventUsesChangeBlock(OpenUSDConnect::EventPayload EventKind) { return UsesChangeBlock(EventKind); } -// --------------------------------------------------------------------------- -// FUSDEventApplier::ApplyFrame -// --------------------------------------------------------------------------- -bool FUSDEventApplier::ApplyFrame(const TArray& RawFrame, AUsdStageActor* StageActor, - FString* OutTouchedPrim, - OpenUSDConnect::EventPayload* OutEventKind) -{ - if (!StageActor || RawFrame.Num() < 8) - { - return false; - } - const OpenUSDConnect::EventWrapper* Wrapper = GetFrameEventWrapper(RawFrame); - return ApplyEventWrapper(Wrapper, StageActor, OutTouchedPrim, OutEventKind, true); -} - -bool FUSDEventApplier::ApplyValidatedFrame(const TArray& RawFrame, +bool FUSDEventApplier::ApplyValidatedFrame(TConstArrayView RawFrame, AUsdStageActor* StageActor, FString* OutTouchedPrim, OpenUSDConnect::EventPayload* OutEventKind) { @@ -1896,5 +1846,5 @@ bool FUSDEventApplier::ApplyValidatedFrame(const TArray& RawFrame, return false; } return ApplyEventWrapper(GetValidatedFrameEventWrapper(RawFrame), StageActor, OutTouchedPrim, - OutEventKind, false); + OutEventKind); } diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Public/USDConnectProtocol.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Public/USDConnectProtocol.h index 36b4138..d2f5636 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Public/USDConnectProtocol.h +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Public/USDConnectProtocol.h @@ -3,7 +3,7 @@ // Wire protocol access for the plugin. The message/event types are the // flatc-generated bindings in the shared native core (regenerate with // scripts/generate_flatbuffers.sh after schema changes); this header only -// adds the framing limit and small UE-flavored helpers on top. +// adds small UE-flavored helpers on top. // // The generated header pins the FlatBuffers runtime version it was produced // with (a static_assert fires on mismatch). setup_flatbuffers.py fetches the @@ -17,33 +17,25 @@ THIRD_PARTY_INCLUDES_START #include "openusdconnect/client/protocol_codec.h" THIRD_PARTY_INCLUDES_END +#include +#include + namespace OUC { -inline constexpr uint32 kMaxFrameSize = - static_cast(openusdconnect::client::kDefaultMaxFrameSize); -inline constexpr uint16 kSchemaVersion = openusdconnect::client::kSchemaVersion; -inline constexpr int32 kProtocolVersion = openusdconnect::client::kProtocolVersion; - inline FString ToFString(const ::flatbuffers::String* S) { return S ? FString(UTF8_TO_TCHAR(S->c_str())) : FString(); } -// Root Envelope of a raw (already de-framed) buffer; nullptr when the -// buffer is too small to hold one. -inline const OpenUSDConnect::Envelope* GetEnvelopeFromFrame(const TArray& Frame) +inline FString ToFString(std::string_view S) { - openusdconnect::client::EnvelopeView View; - return openusdconnect::client::DecodeEnvelope(Frame.GetData(), static_cast(Frame.Num()), - View) == - openusdconnect::client::ProtocolResult::Success - ? View.Get() - : nullptr; + return FString::ConstructFromPtrSize(reinterpret_cast(S.data()), + static_cast(S.size())); } -inline OpenUSDConnect::Payload GetEnvelopePayloadType(const TArray& Frame) +inline std::string ToUtf8(const FString& S) { - const OpenUSDConnect::Envelope* Env = GetEnvelopeFromFrame(Frame); - return Env ? Env->payload_type() : OpenUSDConnect::Payload::NONE; + const FTCHARToUTF8 Converted(*S); + return std::string(Converted.Get(), static_cast(Converted.Length())); } } // namespace OUC diff --git a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Public/USDEventApplier.h b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Public/USDEventApplier.h index da56ffb..95ccfe9 100644 --- a/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Public/USDEventApplier.h +++ b/integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/Public/USDEventApplier.h @@ -24,38 +24,17 @@ class AUsdStageActor; class OPENUSDCONNECTPXR_API FUSDEventApplier { public: - /** - * Decode and apply a single BroadcastEvent frame to the given stage actor. - * @param RawFrame Complete FlatBuffers Envelope bytes (no framing prefix). - * @param StageActor The AUsdStageActor whose pxr stage to modify. - * @param OutTouchedPrim Optional: receives the event's target prim path - * (empty for stage-scoped events). - * @param OutEventKind Optional: receives the wire event kind. - */ - static bool ApplyFrame(const TArray& RawFrame, AUsdStageActor* StageActor, - FString* OutTouchedPrim = nullptr, - OpenUSDConnect::EventPayload* OutEventKind = nullptr); - /** * Apply a frame already schema-verified at the network boundary. The caller * owns any * surrounding SdfChangeBlock according to EventUsesChangeBlock(). */ - static bool ApplyValidatedFrame(const TArray& RawFrame, AUsdStageActor* StageActor, + static bool ApplyValidatedFrame(TConstArrayView RawFrame, AUsdStageActor* StageActor, FString* OutTouchedPrim = nullptr, OpenUSDConnect::EventPayload* OutEventKind = nullptr); /** Whether a validated event kind can join an existing SdfChangeBlock. */ static bool EventUsesChangeBlock(OpenUSDConnect::EventPayload EventKind); - - /** - * Whether the frame's event kind is safe to apply inside an SdfChangeBlock. - * Value writes on existing prims are; structural events (prim definition, - * composition arcs, variants, schema application) are not they need the - * stage to recompose as they apply. Callers batching multiple frames must - * close any open block before applying a frame this returns false for. - */ - static bool FrameUsesChangeBlock(const TArray& RawFrame); }; /** diff --git a/integrations/unreal/test_harness.py b/integrations/unreal/test_harness.py index 6c9ef58..2983ad6 100644 --- a/integrations/unreal/test_harness.py +++ b/integrations/unreal/test_harness.py @@ -20,6 +20,10 @@ REPO_ROOT = Path(__file__).resolve().parents[2] PLUGIN_SOURCE = Path(__file__).resolve().parent / "OpenUSDConnect" CLIENT_CORE_SOURCE = REPO_ROOT / "native" / "client_core" +CLIENT_CORE_MODULE = Path("Source/OpenUSDConnectClientCore") +# UBT compiles every source file under a module, so the driver, platform, and +# testing sources stay out of it. +CLIENT_CORE_STAGED = ("include", "src/frame_codec.cpp", "src/engine") EDITOR_DRIVER = REPO_ROOT / "tests" / "integration" / "scripts" / "unreal_e2e_driver.py" ENGINE_CONFIG = REPO_ROOT / "unreal.test.cfg" FLATBUFFERS_HEADER = Path( @@ -383,11 +387,7 @@ def _plugin_fingerprint( if root.is_dir(): paths.extend((path, plugin_source) for path in root.rglob("*") if path.is_file()) if client_core_source is not None: - paths.extend( - (path, client_core_source) - for path in client_core_source.rglob("*") - if path.is_file() - ) + paths.extend((path, client_core_source) for path in _client_core_files(client_core_source)) for path, relative_root in sorted(paths): if any(part in {"Binaries", "Intermediate"} for part in path.parts): continue @@ -398,6 +398,15 @@ def _plugin_fingerprint( return digest.hexdigest() +def _client_core_files(client_core_source: Path) -> list[Path]: + """The client core files the Unreal client core module compiles.""" + files = [] + for entry in CLIENT_CORE_STAGED: + path = client_core_source / entry + files.extend(sorted(p for p in path.rglob("*") if p.is_file()) if path.is_dir() else [path]) + return files + + def _stage_plugin_source( plugin_source: Path, destination: Path, @@ -412,12 +421,14 @@ def _stage_plugin_source( destination, ignore=shutil.ignore_patterns("Binaries", "Intermediate"), ) - core_destination = ( - destination / "Source" / "ThirdParty" / "OpenUSDConnectClientCore" - ) - if core_destination.exists(): - shutil.rmtree(core_destination) - shutil.copytree(client_core_source, core_destination) + core_destination = destination / CLIENT_CORE_MODULE + for staged in ("include", "src"): + if (core_destination / staged).exists(): + shutil.rmtree(core_destination / staged) + for path in _client_core_files(client_core_source): + target = core_destination / path.relative_to(client_core_source) + target.parent.mkdir(parents=True, exist_ok=True) + shutil.copy2(path, target) return destination @@ -481,7 +492,9 @@ def package_plugin( return package cache_root.mkdir(parents=True, exist_ok=True) - build_dir = cache_root / f".{package.name}.building-{os.getpid()}" + # BuildPlugin compiles inside this directory, and UBT refuses action paths + # over 260 characters on Windows, so its name stays short. + build_dir = cache_root / f".build-{os.getpid()}" staged_source = cache_root / f".{package.name}.source-{os.getpid()}" if build_dir.exists(): shutil.rmtree(build_dir) diff --git a/integrations/unreal/usd_connect.py b/integrations/unreal/usd_connect.py index 1e31813..46f3af2 100644 --- a/integrations/unreal/usd_connect.py +++ b/integrations/unreal/usd_connect.py @@ -36,7 +36,7 @@ # -- Module state -------------------------------------------------------- -_receiver = None # ReceiverThread +_receiver = None # EventReceiver _emitter = None # NoticeEmitter _sender = None # EventSender _dispatcher = None # EventDispatcher @@ -339,9 +339,9 @@ def _remember_token(token: str) -> None: # -- Start receiver -------------------------------------------------- if receive: - from openusdconnect.receiver import ReceiverThread + from openusdconnect.receiver import EventReceiver - _receiver = ReceiverThread( + _receiver = EventReceiver( host=host, port=port, sync_from=sync_from, @@ -421,9 +421,9 @@ def stop(): unreal.unregister_slate_post_tick_callback(_tick_handle) _tick_handle = None - # Stop receiver + # Close receiver if _receiver is not None: - _receiver.stop() + _receiver.close(timeout=2.0) _receiver = None # Drop dispatcher (no separate state to flush last_seq dies with it) diff --git a/integrations/usdview/connection.py b/integrations/usdview/connection.py index a7b6813..19415c8 100644 --- a/integrations/usdview/connection.py +++ b/integrations/usdview/connection.py @@ -3,7 +3,7 @@ Receive-only: no emitter, no feedback-loop guard. Stage mutation triggers ``Usd.Notice.ObjectsChanged`` which refreshes the viewport on its own. Built on ``UsdReceiver`` token persistence, identity, and the underlying -``ReceiverThread`` + ``EventDispatcher`` composition are handled by the client. +``EventReceiver`` + ``EventDispatcher`` composition are handled by the client. """ from __future__ import annotations diff --git a/native/client_core/CMakeLists.txt b/native/client_core/CMakeLists.txt index 8f73072..e56aab9 100644 --- a/native/client_core/CMakeLists.txt +++ b/native/client_core/CMakeLists.txt @@ -1,26 +1,28 @@ +function(openusdconnect_configure_client_library target) + set_target_properties(${target} PROPERTIES + POSITION_INDEPENDENT_CODE ON + CXX_VISIBILITY_PRESET hidden + VISIBILITY_INLINES_HIDDEN YES + ) + if(MSVC) + target_compile_options(${target} PRIVATE /W4 /permissive-) + else() + target_compile_options(${target} PRIVATE -Wall -Wextra -Wpedantic) + endif() +endfunction() + add_library(OpenUSDConnectClientCore STATIC src/frame_codec.cpp ) add_library(OpenUSDConnect::ClientCore ALIAS OpenUSDConnectClientCore) - target_compile_features(OpenUSDConnectClientCore PUBLIC cxx_std_17) target_include_directories(OpenUSDConnectClientCore PUBLIC ${CMAKE_CURRENT_SOURCE_DIR}/include ) -set_target_properties(OpenUSDConnectClientCore PROPERTIES - POSITION_INDEPENDENT_CODE ON - CXX_VISIBILITY_PRESET hidden - VISIBILITY_INLINES_HIDDEN YES -) - -if(MSVC) - target_compile_options(OpenUSDConnectClientCore PRIVATE /W4 /permissive-) -else() - target_compile_options(OpenUSDConnectClientCore PRIVATE -Wall -Wextra -Wpedantic) -endif() +openusdconnect_configure_client_library(OpenUSDConnectClientCore) # Header-only FlatBuffers protocol views/builders. Consumers retain control of -# their FlatBuffers allocator and supply the FlatBuffers include target/path. +# their FlatBuffers allocator. add_library(OpenUSDConnectClientProtocol INTERFACE) add_library(OpenUSDConnect::ClientProtocol ALIAS OpenUSDConnectClientProtocol) target_compile_features(OpenUSDConnectClientProtocol INTERFACE cxx_std_17) @@ -31,4 +33,48 @@ target_link_libraries(OpenUSDConnectClientProtocol INTERFACE OpenUSDConnectClien if(TARGET flatbuffers::flatbuffers) target_link_libraries(OpenUSDConnectClientProtocol INTERFACE flatbuffers::flatbuffers) +else() + include(FetchContent) + # include/ has no CMakeLists.txt, so the FlatBuffers project is never configured. + FetchContent_Declare(flatbuffers + URL https://github.com/google/flatbuffers/archive/refs/tags/v25.12.19.tar.gz + URL_HASH SHA256=f81c3162b1046fe8b84b9a0dbdd383e24fdbcf88583b9cb6028f90d04d90696a + SOURCE_SUBDIR include + ) + FetchContent_MakeAvailable(flatbuffers) + target_include_directories(OpenUSDConnectClientProtocol SYSTEM + INTERFACE "${flatbuffers_SOURCE_DIR}/include" + ) endif() + +# Sans-IO endpoints: socket bytes and host time in, actions and notifications out. +add_library(OpenUSDConnectClientEngine STATIC + src/engine/producer_endpoint.cpp + src/engine/receiver_endpoint.cpp +) +add_library(OpenUSDConnect::ClientEngine ALIAS OpenUSDConnectClientEngine) +target_link_libraries(OpenUSDConnectClientEngine PUBLIC OpenUSDConnectClientProtocol) +openusdconnect_configure_client_library(OpenUSDConnectClientEngine) + +# Optional reference driver: one thread and blocking sockets drive the engine. +# Hosts with their own I/O and scheduler drive the engine directly instead. +add_library(OpenUSDConnectClientDriver STATIC + src/driver/threaded_producer_driver.cpp + src/platform/bsd_socket.cpp +) +add_library(OpenUSDConnect::ClientDriver ALIAS OpenUSDConnectClientDriver) +find_package(Threads REQUIRED) +target_link_libraries(OpenUSDConnectClientDriver + PUBLIC OpenUSDConnectClientEngine Threads::Threads + PRIVATE $<$:ws2_32> +) +openusdconnect_configure_client_library(OpenUSDConnectClientDriver) + +# Scripted sockets that let the C++ tests drive the reference driver without a +# network. Built only for a target that links it; no shipped module does. +add_library(OpenUSDConnectClientDriverTesting STATIC EXCLUDE_FROM_ALL + src/driver/testing/scripted_socket.cpp +) +add_library(OpenUSDConnect::ClientDriverTesting ALIAS OpenUSDConnectClientDriverTesting) +target_link_libraries(OpenUSDConnectClientDriverTesting PUBLIC OpenUSDConnectClientDriver) +openusdconnect_configure_client_library(OpenUSDConnectClientDriverTesting) diff --git a/native/client_core/README.md b/native/client_core/README.md index d8c6b83..1c41a10 100644 --- a/native/client_core/README.md +++ b/native/client_core/README.md @@ -1,31 +1,72 @@ # Native client core -The native core is split into two composable C++17 targets: - -- `OpenUSDConnect::ClientCore` provides framing, receiver ordering/replay, and producer outbox - state. It has no FlatBuffers or OpenUSD dependency. -- `OpenUSDConnect::ClientProtocol` adds the generated FlatBuffers schema plus transport-neutral - handshake, control-message, and transaction construction helpers. - -The protocol layer deliberately does not own transport, threads, queues, event-offset storage, or -serialized buffers. Decoded views borrow the caller's receive buffer. Builders operate on a -caller-owned `flatbuffers::FlatBufferBuilder`, so an integrator may supply a custom allocator, -construct schema events directly, keep offsets in its native container, and send from the builder -or detach its allocation without copying serialized bytes. - -Typical construction is: - -1. Create a `flatbuffers::FlatBufferBuilder` with the desired allocator and initial capacity. -2. Build event offsets with the stateless helpers or the generated schema API. -3. Call `FinishTransactionFrame` with the caller-owned contiguous offset range. -4. send `builder.GetBufferPointer()` / `builder.GetSize()`, or call `builder.Release()` to transfer - the exact allocation. - -On receive, call `DecodeEnvelope` once at the untrusted-buffer boundary. `HandshakeResponseView` -and `ControlMessageView` then classify the verified envelope without further validation or copies. -All borrowed pointers remain valid only while the original receive buffer remains alive and -unchanged. - -When included with `add_subdirectory`, link `OpenUSDConnect::ClientProtocol`. FlatBuffers remains a -consumer-provided header dependency; if `flatbuffers::flatbuffers` already exists, the target links -it automatically. +Four C++17 targets, each building on the previous one: + +- `OpenUSDConnect::ClientCore`: framing, receiver ordering and replay, the producer outbox and + its rejection policy, and the client phase. No FlatBuffers or OpenUSD dependency. +- `OpenUSDConnect::ClientProtocol`: the generated FlatBuffers schema plus transport-neutral + handshake, control-message, and transaction construction helpers (`protocol_codec.h`). +- `OpenUSDConnect::ClientEngine`: sans-IO endpoints (`engine/receiver_endpoint.h`, + `engine/producer_endpoint.h`) that own the connection protocol from handshake through + reconnect backoff. They start no threads and open no sockets. +- `OpenUSDConnect::ClientDriver`: an optional reference host loop (`driver/`) that drives one + endpoint on one thread with blocking Winsock or BSD sockets. The Python module links it; a + host with its own I/O loop and scheduler drives the endpoints directly instead. + +Either way the host builds transactions, decodes and applies drained frames on its +stage-owning thread, stores tokens, and handles notifications. + +Add the directory with `add_subdirectory` and link the highest target you use. If a +`flatbuffers::flatbuffers` target exists, the protocol target links it; otherwise CMake fetches +the pinned FlatBuffers headers with `FetchContent`, which needs network access on the first +configure unless `FETCHCONTENT_SOURCE_DIR_FLATBUFFERS` names a copy containing +`include/flatbuffers`. + +## Protocol layer + +The protocol layer owns no transport, threads, queues, event-offset storage, or serialized +buffers. Builders operate on a caller-owned `flatbuffers::FlatBufferBuilder`: build event +offsets with the stateless helpers or the generated schema API, call `FinishTransactionFrame` +with the contiguous offset range, then send from the builder or `Release()` its allocation +without copying. On receive, call `DecodeEnvelope` once at the untrusted-buffer boundary; +`HandshakeResponseView` and `ControlMessageView` then classify the verified envelope without +further validation or copies. Views borrow the receive buffer and are valid only while it is +alive and unchanged. + +## Endpoints + +An endpoint is thread-safe, never blocks, and never calls into the host. The host feeds it +socket events and time and applies what it returns: + +1. Call `ReceiverEndpoint::Start(now)`; a `ProducerEndpoint` connects only when asked + (`RequestConnect` or `Connect`). +2. Apply `TakeActions()` in order, on the thread that reports socket events: `ConnectAction` + opens a socket, `SendAction` writes, `CloseAction` closes, `LogAction` logs. +3. Report `OnConnected(token)`, each read with `OnBytes`, the end of a socket or attempt with + `OnDisconnected`, the time with `OnTick` once `NextWake()` passes, and on the receiver a + read that waited `SocketTimeout` with `OnReadTimeout`. Apply the actions again after each + report. +4. Drain the `NotificationQueue` on any thread. The stage-owning thread drains receiver frames + with `DrainFrames` and reports progress as the header describes; a producer host appends + complete length-prefixed `Txn` frames with `Append`. + +The headers state each method's contract. + +## Reference driver + +`ThreadedReceiverDriver` and `ThreadedProducerDriver` run that loop for one endpoint on one +thread from `Start()`. The host supplies a `SocketFactory` (`TcpSocketFactory` or its own) and +optional `DriverCallbacks`: the token for each handshake, a hook for an issued token, a sink +that receives notifications after every endpoint call (without one, the host drains the +queue), and a log sink. Callbacks run on the driver thread with no lock held and must not +destroy the driver. + +- After an endpoint call from another thread that queues actions (`Append`, `QueueControl`, + `RequestConnect`, `CancelConnect`, `Disconnect`, `RequestReplayFrom`, `MarkReplayApplied`), + call `Wake()` so the loop applies them. +- `Close(timeout)` writes what is queued while the socket is open, within a one-second grace, + then stops the endpoint and joins the thread; destroying the driver does the same. `Stop()` + never blocks: it stops the endpoint and interrupts a pending connect or read. + `Join(timeout)` waits for the thread. +- The receiver's `WaitConnected` and `WaitSynchronized` and the producer's `Connect` and + `Flush` block, so they run on a thread other than the loop. diff --git a/native/client_core/include/openusdconnect/client/driver/socket.h b/native/client_core/include/openusdconnect/client/driver/socket.h new file mode 100644 index 0000000..5c569aa --- /dev/null +++ b/native/client_core/include/openusdconnect/client/driver/socket.h @@ -0,0 +1,113 @@ +#pragma once + +#include "openusdconnect/client/engine/actions.h" +#include "openusdconnect/client/engine/notification.h" + +#include +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +enum class SocketResult : std::uint8_t +{ + Success, + // The call passed its deadline. + Timeout, + Interrupted, + // The peer closed the connection. + Closed, + // SystemError describes the failure. + Failed, +}; + +// A blocking TCP connection owned by one thread. Destroying it closes it. +class Socket +{ +public: + virtual ~Socket() = default; + + [[nodiscard]] virtual SocketResult Connect(const std::string& host, std::uint16_t port, + TimePoint deadline) = 0; + // Timeout once it has to wait for the peer past deadline. + [[nodiscard]] virtual SocketResult SendAll(const std::uint8_t* data, std::size_t size, + TimePoint deadline) = 0; + // Waits for at least one byte until deadline, or without one until it arrives. + [[nodiscard]] virtual SocketResult Receive(std::uint8_t* buffer, std::size_t capacity, + std::optional deadline, + std::size_t& received) = 0; + // The operating system error of the last Timeout or Failed result. + [[nodiscard]] virtual int SystemError() const noexcept = 0; + + // Thread-safe. Ends the blocked call and every later one with Interrupted. + virtual void Interrupt() noexcept = 0; + // Thread-safe. Ends the blocked or the next Receive with Interrupted. + virtual void Wake() noexcept = 0; + +protected: + Socket() = default; + Socket(const Socket&) = delete; + Socket& operator=(const Socket&) = delete; +}; + +class SocketFactory +{ +public: + virtual ~SocketFactory() = default; + + // An unconnected socket, so another thread can interrupt its Connect. + [[nodiscard]] virtual std::unique_ptr Create() = 0; + +protected: + SocketFactory() = default; + SocketFactory(const SocketFactory&) = delete; + SocketFactory& operator=(const SocketFactory&) = delete; +}; + +// Winsock or BSD sockets with TCP_NODELAY, so small control frames go out at once. +class TcpSocketFactory final : public SocketFactory +{ +public: + TcpSocketFactory(); + + [[nodiscard]] std::unique_ptr Create() override; +}; + +enum class SocketOperation : std::uint8_t +{ + Connect, + Send, + Receive, +}; + +struct TransportFailure final +{ + SocketOperation Operation = SocketOperation::Connect; + // Timeout or Failed. + SocketResult Result = SocketResult::Failed; + int SystemError = 0; +}; + +[[nodiscard]] std::string DescribeSystemError(int system_error); +[[nodiscard]] std::string Describe(const TransportFailure& failure); + +// What a reference driver calls on its host. Every callback is optional, must +// return normally, and runs on the driver thread, which holds no driver or +// endpoint lock while calling it. +struct DriverCallbacks final +{ + // Read just before each handshake; nullopt abandons that connection attempt. + std::function()> Token; + // The token an accepted Hello issued, before the next Token read. + std::function TokenIssued; + // When set, the driver drains the notification queue into it after every + // endpoint call, so a notification is handled before the next attempt. + std::function Notifications; + std::function Log; +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/driver/testing/scripted_socket.h b/native/client_core/include/openusdconnect/client/driver/testing/scripted_socket.h new file mode 100644 index 0000000..cb7272b --- /dev/null +++ b/native/client_core/include/openusdconnect/client/driver/testing/scripted_socket.h @@ -0,0 +1,65 @@ +#pragma once + +#include "openusdconnect/client/driver/socket.h" + +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +namespace detail +{ +struct ScriptState; +struct ScriptedChannel; +} // namespace detail + +// The server's end of one accepted scripted connection. +class ScriptedConnection final +{ +public: + ScriptedConnection(std::shared_ptr state, + std::shared_ptr channel); + + // One Receive returns each delivery. False once the client closed its socket. + [[nodiscard]] bool Deliver(std::vector bytes); + // Receive reports Closed once the earlier deliveries are read. + void Close(); + // Every later SendAll waits, as for a peer that stopped reading, until its + // deadline passes and then reports Timeout. + void StallSends(); + + [[nodiscard]] std::vector Sent() const; + // Waits until the client has sent at least size bytes. + [[nodiscard]] bool WaitSent(std::size_t size, std::chrono::milliseconds timeout) const; + // The client waits in Receive with every delivery read, so it handled them all. + [[nodiscard]] bool WaitIdle(std::chrono::milliseconds timeout) const; + [[nodiscard]] bool WaitClosed(std::chrono::milliseconds timeout) const; + [[nodiscard]] bool ClosedByClient() const; + +private: + std::shared_ptr State; + std::shared_ptr Channel; +}; + +// A test seam: each Connect waits until the test accepts or refuses it. +class ScriptedSocketFactory final : public SocketFactory +{ +public: + ScriptedSocketFactory(); + + [[nodiscard]] std::unique_ptr Create() override; + + // Completes the oldest waiting Connect; nullptr when none begins within timeout. + [[nodiscard]] std::shared_ptr Accept(std::chrono::milliseconds timeout); + [[nodiscard]] bool Refuse(std::chrono::milliseconds timeout, int system_error); + [[nodiscard]] std::size_t Attempts() const; + +private: + std::shared_ptr State; +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/driver/threaded_driver.h b/native/client_core/include/openusdconnect/client/driver/threaded_driver.h new file mode 100644 index 0000000..ed72bb4 --- /dev/null +++ b/native/client_core/include/openusdconnect/client/driver/threaded_driver.h @@ -0,0 +1,503 @@ +#pragma once + +#include "openusdconnect/client/driver/socket.h" +#include "openusdconnect/client/engine/actions.h" +#include "openusdconnect/client/engine/notification.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +// Reference host loop: one thread with blocking sockets applies a started endpoint's actions, +// ticks at NextWake, and exits once the endpoint stops. Wake it after queuing actions from +// another thread. Never destroy it on its own thread, such as from one of its callbacks. +template +class ThreadedDriver +{ +public: + using EndpointType = Endpoint; + + ThreadedDriver(const ThreadedDriver&) = delete; + ThreadedDriver& operator=(const ThreadedDriver&) = delete; + + // Starts the loop; false when already started. + [[nodiscard]] bool Start() + { + std::lock_guard lock(Mutex); + if (Started) + { + return false; + } + Thread = std::thread( + [this] + { + Run(); + }); + Started = true; + ThreadIdentity = Thread.get_id(); + return true; + } + + // Stops the endpoint and wakes the loop, which then exits. Never blocks. + void Stop() + { + // First, so a loop that StopRequested wakes finds the endpoint stopped. + Target.Stop(); + std::shared_ptr socket; + { + std::lock_guard lock(Mutex); + StopRequested = true; + socket = Connection; + } + if (socket) + { + socket->Interrupt(); + } + Changed.notify_all(); + } + + // Like Stop, but the loop applies the queued actions, such as sends, before it exits. + void StopAfterQueued() + { + Target.Stop(); + Wake(); + } + + void Wake() + { + std::shared_ptr socket; + { + std::lock_guard lock(Mutex); + WakeRequested = true; + socket = Connection; + } + if (socket) + { + socket->Wake(); + } + Changed.notify_all(); + } + + // Waits for the loop to exit; false on timeout or on the driver thread. + [[nodiscard]] bool Join(std::optional timeout = std::nullopt) + { + std::unique_lock lock(Mutex); + if (!Started) + { + return true; + } + if (ThreadIdentity == std::this_thread::get_id()) + { + return false; + } + const auto exited = [this] + { + return Exited; + }; + if (!timeout) + { + Changed.wait(lock, exited); + } + else if (!Changed.wait_for(lock, *timeout, exited)) + { + return false; + } + // The loop takes Mutex no more once it exited, so joining under it cannot block it. + if (Thread.joinable()) + { + Thread.join(); + } + return true; + } + + static constexpr std::chrono::milliseconds kCloseGrace{1'000}; + + // With a connected socket, StopAfterQueued and Join for kCloseGrace or a shorter + // timeout; then Stop and Join the rest of timeout. Returns whether the loop exited; + // on the loop thread, false at once, and the loop exits after the callback returns. + [[nodiscard]] bool Close(std::optional timeout) + { + const std::optional deadline = + timeout ? std::optional(Now() + *timeout) : std::nullopt; + if (IsConnectionOpen()) + { + StopAfterQueued(); + if (Join(std::min(timeout.value_or(kCloseGrace), kCloseGrace))) + { + return true; + } + } + if (ThreadId() == std::this_thread::get_id()) + { + StopAfterQueued(); + return false; + } + Stop(); + return Join(Until(deadline)); + } + + [[nodiscard]] bool Running() const + { + std::lock_guard lock(Mutex); + return Started && !Exited; + } + + // The loop ran and exited. + [[nodiscard]] bool Stopped() const + { + std::lock_guard lock(Mutex); + return Started && Exited; + } + + [[nodiscard]] std::optional ThreadId() const + { + std::lock_guard lock(Mutex); + if (ThreadIdentity == std::thread::id()) + { + return std::nullopt; + } + return ThreadIdentity; + } + + // Why the latest connection attempt failed, if it did. + [[nodiscard]] std::optional LastFailure() const + { + std::lock_guard lock(Mutex); + return Failure; + } + +protected: + // notifications must be the queue the endpoint pushes to. A write still + // waiting for the peer socket_timeout after it began is a TransportError; + // with on_read_timeout set, a read that waits as long without a byte calls it. + ThreadedDriver(Endpoint& endpoint, NotificationQueue& notifications, + std::shared_ptr sockets, DriverCallbacks callbacks, + std::chrono::milliseconds socket_timeout, + void (Endpoint::*on_read_timeout)() = nullptr) + : Target(endpoint) + , Notifications(notifications) + , Sockets(std::move(sockets)) + , Callbacks(std::move(callbacks)) + , SocketTimeout(socket_timeout) + , OnReadTimeout(on_read_timeout) + , Buffer(kReadBufferSize) + { + } + + ~ThreadedDriver() + { + static_cast(Close(std::nullopt)); + } + + // Waits until ready holds, the loop stops or exits, or timeout passes. + template + [[nodiscard]] bool Wait(Predicate ready, std::optional timeout) + { + std::unique_lock lock(Mutex); + const auto done = [&] + { + return StopRequested || Exited || ready(); + }; + if (!timeout) + { + Changed.wait(lock, done); + } + else + { + Changed.wait_for(lock, *timeout, done); + } + lock.unlock(); + return ready(); + } + + [[nodiscard]] static TimePoint Now() noexcept + { + return std::chrono::steady_clock::now(); + } + + [[nodiscard]] static std::chrono::milliseconds Until(TimePoint deadline) noexcept + { + return std::max(std::chrono::ceil(deadline - Now()), + std::chrono::milliseconds::zero()); + } + + [[nodiscard]] static std::optional + Until(std::optional deadline) + { + return deadline ? std::optional(Until(*deadline)) : std::nullopt; + } + + Endpoint& Target; + +private: + static constexpr std::size_t kReadBufferSize = 64 * 1024; + + void Run() + { + for (;;) + { + Dispatch(); + if (Connection) + { + Read(); + continue; + } + if (Target.Status().Stopped) + { + break; + } + Sleep(Target.NextWake()); + Target.OnTick(Now()); + } + { + std::lock_guard lock(Mutex); + Exited = true; + } + Changed.notify_all(); + } + + // Reports an issued token and delivers notifications before applying each + // batch of actions, so a token issued by one handshake is stored before the + // next connects. + void Dispatch() + { + { + std::lock_guard lock(Mutex); + WakeRequested = false; + } + for (;;) + { + if (std::optional token = Target.TakeIssuedToken(); + token && Callbacks.TokenIssued) + { + Callbacks.TokenIssued(*token); + } + Deliver(); + std::vector actions = Target.TakeActions(); + if (actions.empty()) + { + break; + } + for (Action& action : actions) + { + std::visit( + [this](auto& value) + { + Apply(value); + }, + action); + } + } + // Taking the lock orders this notification after a waiter's check. + { + std::lock_guard lock(Mutex); + } + Changed.notify_all(); + } + + void Deliver() + { + if (!Callbacks.Notifications) + { + return; + } + for (Notification& notification : Notifications.Drain()) + { + Callbacks.Notifications(std::move(notification)); + } + } + + void Apply(const ConnectAction& action) + { + std::shared_ptr socket = Sockets->Create(); + { + std::lock_guard lock(Mutex); + Failure.reset(); + Connection = socket; + if (StopRequested) + { + socket->Interrupt(); + } + } + const SocketResult result = socket->Connect(action.Host, action.Port, action.Deadline); + if (result != SocketResult::Success) + { + Fail(SocketOperation::Connect, result, *socket, + "could not connect to " + action.Host + ":" + std::to_string(action.Port)); + CloseConnection(DisconnectReason::ConnectFailed); + return; + } + { + std::lock_guard lock(Mutex); + ConnectionOpen = true; + } + const std::optional token = + Callbacks.Token ? Callbacks.Token() : std::optional(std::in_place); + if (!token) + { + CloseConnection(DisconnectReason::ConnectFailed); + return; + } + Target.OnConnected(*token); + } + + void Apply(const SendAction& action) + { + if (!Connection) + { + return; + } + const SocketResult result = + Connection->SendAll(action.Bytes->data(), action.Bytes->size(), Now() + SocketTimeout); + if (result != SocketResult::Success) + { + Fail(SocketOperation::Send, result, *Connection, "send failed"); + CloseConnection(DisconnectReason::TransportError); + } + } + + void Apply(const CloseAction& action) + { + CloseConnection(action.Reason); + } + + void Apply(const LogAction& action) + { + Log(action.Level, action.Message); + } + + void Read() + { + const std::optional wake = Target.NextWake(); + std::optional deadline = wake; + if (OnReadTimeout) + { + const TimePoint timeout = Now() + SocketTimeout; + deadline = wake ? std::min(*wake, timeout) : timeout; + } + std::size_t received = 0; + const SocketResult result = + Connection->Receive(Buffer.data(), Buffer.size(), deadline, received); + switch (result) + { + case SocketResult::Success: + Target.OnBytes(Buffer.data(), received); + return; + case SocketResult::Timeout: + if (const TimePoint now = Now(); wake && now >= *wake) + { + Target.OnTick(now); + } + else if (OnReadTimeout) + { + (Target.*OnReadTimeout)(); + } + return; + case SocketResult::Interrupted: + return; + case SocketResult::Closed: + CloseConnection(DisconnectReason::PeerClosed); + return; + case SocketResult::Failed: + Fail(SocketOperation::Receive, result, *Connection, "receive failed"); + CloseConnection(DisconnectReason::TransportError); + return; + } + } + + void Sleep(std::optional until) + { + const auto woken = [this] + { + return StopRequested || WakeRequested; + }; + std::unique_lock lock(Mutex); + if (until) + { + Changed.wait_until(lock, *until, woken); + } + else + { + Changed.wait(lock, woken); + } + } + + void CloseConnection(DisconnectReason reason) + { + std::shared_ptr closed; + { + std::lock_guard lock(Mutex); + closed = std::exchange(Connection, nullptr); + ConnectionOpen = false; + } + closed.reset(); + Target.OnDisconnected(reason, Now()); + } + + // Records a Timeout or Failed result; an interrupted call is not a failure. + void Fail(SocketOperation operation, SocketResult result, const Socket& socket, + const std::string& context) + { + if (result == SocketResult::Interrupted) + { + return; + } + const TransportFailure failure{operation, result, + result == SocketResult::Failed ? socket.SystemError() : 0}; + { + std::lock_guard lock(Mutex); + Failure = failure; + } + Log(LogLevel::Warning, context + ": " + Describe(failure)); + } + + [[nodiscard]] bool IsConnectionOpen() const + { + std::lock_guard lock(Mutex); + return ConnectionOpen; + } + + void Log(LogLevel level, const std::string& message) + { + if (Callbacks.Log) + { + Callbacks.Log(level, message); + } + } + + NotificationQueue& Notifications; + const std::shared_ptr Sockets; + const DriverCallbacks Callbacks; + const std::chrono::milliseconds SocketTimeout; + void (Endpoint::* const OnReadTimeout)(); + std::vector Buffer; + + mutable std::mutex Mutex; + std::condition_variable Changed; + // Written only by the loop thread under Mutex, so the loop reads it unlocked. + std::shared_ptr Connection; + // Connection finished connecting, so the writes queued for it can still reach the peer. + bool ConnectionOpen = false; + std::optional Failure; + std::thread::id ThreadIdentity; + bool Started = false; + bool Exited = false; + bool StopRequested = false; + bool WakeRequested = false; + std::thread Thread; +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/driver/threaded_producer_driver.h b/native/client_core/include/openusdconnect/client/driver/threaded_producer_driver.h new file mode 100644 index 0000000..c93070b --- /dev/null +++ b/native/client_core/include/openusdconnect/client/driver/threaded_producer_driver.h @@ -0,0 +1,55 @@ +#pragma once + +#include "openusdconnect/client/driver/threaded_driver.h" +#include "openusdconnect/client/engine/producer_endpoint.h" + +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +enum class FlushResult : std::uint8_t +{ + // Every appended transaction is acknowledged. + Flushed, + // The endpoint's Failure must be resolved first. + RecoveryRequired, + // The timeout passed, or the loop stopped, first. + Unfinished, +}; + +// Drives a ProducerEndpoint on one thread, which idles between the connections it is asked for. +// Call Wake after an endpoint call that queues actions (Append, QueueControl, RequestConnect, +// CancelConnect, Disconnect) so the loop applies them. +class ThreadedProducerDriver final : public ThreadedDriver +{ +public: + ThreadedProducerDriver(ProducerEndpoint& endpoint, NotificationQueue& notifications, + std::shared_ptr sockets, DriverCallbacks callbacks = {}) + : ThreadedDriver(endpoint, notifications, std::move(sockets), std::move(callbacks), + endpoint.Configuration().HandshakeTimeout) + { + } + + // The blocking calls need the loop running on another thread; otherwise they + // only report the endpoint's state. + + // Makes one attempt once any attempt or close in flight ends, all within + // timeout and HandshakeTimeout, and returns whether the endpoint is connected. + [[nodiscard]] bool Connect(std::optional timeout); + // Waits until every appended transaction is acknowledged, connecting while + // disconnected and outside the rate-limit window. + [[nodiscard]] FlushResult Flush(std::optional timeout); + +private: + [[nodiscard]] bool CanWait() const; + // Waits until no attempt or close is in flight; false once deadline passes first. + [[nodiscard]] bool WaitSettled(TimePoint deadline); + void Pause(TimePoint until); +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/driver/threaded_receiver_driver.h b/native/client_core/include/openusdconnect/client/driver/threaded_receiver_driver.h new file mode 100644 index 0000000..f2a991f --- /dev/null +++ b/native/client_core/include/openusdconnect/client/driver/threaded_receiver_driver.h @@ -0,0 +1,58 @@ +#pragma once + +#include "openusdconnect/client/driver/threaded_driver.h" +#include "openusdconnect/client/engine/receiver_endpoint.h" + +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +// Drives a ReceiverEndpoint on one thread. Call Wake after a consumer-thread call that can close +// the connection or complete the replay, such as RequestReplayFrom or MarkReplayApplied, so the +// loop applies its actions and the waits re-check. +class ThreadedReceiverDriver final : public ThreadedDriver +{ +public: + ThreadedReceiverDriver(ReceiverEndpoint& endpoint, NotificationQueue& notifications, + std::shared_ptr sockets, DriverCallbacks callbacks = {}) + : ThreadedDriver(endpoint, notifications, std::move(sockets), std::move(callbacks), + endpoint.Configuration().SocketTimeout, &ReceiverEndpoint::OnReadTimeout) + { + } + + // Starts the endpoint, then the loop; false when already started. + [[nodiscard]] bool Start() + { + static_cast(Target.Start(Now())); + return ThreadedDriver::Start(); + } + + // Each returns the state once it holds, the loop stops, or timeout passes. + [[nodiscard]] bool WaitConnected(std::optional timeout) + { + return WaitFor(&ReceiverStatus::Connected, timeout); + } + + [[nodiscard]] bool WaitSynchronized(std::optional timeout) + { + return WaitFor(&ReceiverStatus::Synchronized, timeout); + } + +private: + [[nodiscard]] bool WaitFor(bool ReceiverStatus::* state, + std::optional timeout) + { + return Wait( + [this, state] + { + return Target.Status().*state; + }, + timeout); + } +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/engine/actions.h b/native/client_core/include/openusdconnect/client/engine/actions.h new file mode 100644 index 0000000..49675d8 --- /dev/null +++ b/native/client_core/include/openusdconnect/client/engine/actions.h @@ -0,0 +1,138 @@ +#pragma once + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +// Hosts pass the current time in; the engine never reads a clock. +using TimePoint = std::chrono::steady_clock::time_point; + +// Exponential backoff between connection attempts, expressed as due times. +class ReconnectPolicy final +{ +public: + ReconnectPolicy(bool enabled, std::chrono::milliseconds base_delay, + std::chrono::milliseconds max_delay) noexcept + : EnabledValue(enabled) + , BaseDelay(base_delay) + , MaxDelay(max_delay) + , Delay(base_delay) + { + assert(IsValidConfiguration(BaseDelay, MaxDelay)); + } + + [[nodiscard]] static bool IsValidConfiguration(std::chrono::milliseconds base_delay, + std::chrono::milliseconds max_delay) noexcept + { + return base_delay.count() > 0 && max_delay >= base_delay; + } + + [[nodiscard]] bool Enabled() const noexcept + { + return EnabledValue; + } + + void SetEnabled(bool enabled) noexcept + { + EnabledValue = enabled; + } + + // A session reached its connected state; the next wait starts from the base. + void Reset() noexcept + { + Delay = BaseDelay; + } + + [[nodiscard]] TimePoint NextAttempt(TimePoint now) noexcept + { + const TimePoint due = now + Delay; + Delay = std::min(Delay * 2, MaxDelay); + return due; + } + + // After an overflow the next attempt waits for the consumer to drain the + // queue, but no longer than this deadline. + [[nodiscard]] TimePoint DrainDeadline(TimePoint now) noexcept + { + Reset(); + return now + MaxDelay; + } + +private: + bool EnabledValue; + const std::chrono::milliseconds BaseDelay; + const std::chrono::milliseconds MaxDelay; + std::chrono::milliseconds Delay; +}; + +enum class DisconnectReason : std::uint8_t +{ + // Reported by the host. + ConnectFailed, + PeerClosed, + TransportError, + // Requested by the endpoint with a CloseAction. + Stopped, + HandshakeRejected, + ReplayRequested, + SequenceGap, + QueueFull, + ReadTimeout, + ProtocolError, + Cancelled, + HandshakeTimeout, + RecoveryRequired, + RateLimited, +}; + +enum class LogLevel : std::uint8_t +{ + Debug, + Info, + Warning, + Error, +}; + +// Open a socket, then report OnConnected, or OnDisconnected once the attempt +// fails or Deadline passes. +struct ConnectAction final +{ + std::string Host; + std::uint16_t Port = 0; + TimePoint Deadline; +}; + +// Write complete length-prefixed frames, in order. +struct SendAction final +{ + // Shared with the producer outbox, so replaying a transaction copies nothing. + std::shared_ptr> Bytes; +}; + +// Both endpoints Close with ProtocolError on an undecodable frame or a handshake answer other than +// HelloOk, HelloRejected, or AuthRejected. After the handshake the receiver closes on a payload +// type it does not know; the producer ignores everything but transaction results and RateLimited. + +// Close the socket, or abandon the connect attempt, then report OnDisconnected. +struct CloseAction final +{ + DisconnectReason Reason = DisconnectReason::Stopped; +}; + +struct LogAction final +{ + LogLevel Level = LogLevel::Info; + std::string Message; +}; + +using Action = std::variant; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/engine/notification.h b/native/client_core/include/openusdconnect/client/engine/notification.h new file mode 100644 index 0000000..263e4e3 --- /dev/null +++ b/native/client_core/include/openusdconnect/client/engine/notification.h @@ -0,0 +1,93 @@ +#pragma once + +#include "openusdconnect/client/engine/actions.h" +#include "openusdconnect/client/schema/messages_generated.h" + +#include +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +struct Connected final +{ +}; + +// Follows Connected when that session ends. +struct Disconnected final +{ + DisconnectReason Reason = DisconnectReason::PeerClosed; +}; + +struct HandshakeRejected final +{ + // AuthRejected rather than HelloRejected; Code is then Unspecified. + bool Authentication = false; + OpenUSDConnect::HelloRejectionCode Code = OpenUSDConnect::HelloRejectionCode::Unspecified; + std::string Reason; +}; + +struct TokenIssued final +{ + std::string Token; +}; + +// Only the fields the server authored are set. +struct StageMetadata final +{ + std::optional TimeCodesPerSecond; + std::optional FramesPerSecond; + std::optional StartTimeCode; + std::optional EndTimeCode; + std::optional MetersPerUnit; + std::optional UpAxis; +}; + +struct PlaybackState final +{ + double Time = 0.0; + bool Playing = false; + double Rate = 0.0; + std::string LeaderClientId; +}; + +struct PlaybackClaimed final +{ + std::string LeaderClientId; +}; + +struct PlaybackRejected final +{ + std::string Reason; + std::string CurrentLeaderClientId; +}; + +using Notification = std::variant; + +// Endpoints push while holding their own lock; the host drains on any thread. +class NotificationQueue final +{ +public: + void Push(Notification notification) + { + std::lock_guard lock(Mutex); + Pending.push_back(std::move(notification)); + } + + [[nodiscard]] std::vector Drain() + { + std::lock_guard lock(Mutex); + return std::exchange(Pending, {}); + } + +private: + std::mutex Mutex; + std::vector Pending; +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/engine/producer_endpoint.h b/native/client_core/include/openusdconnect/client/engine/producer_endpoint.h new file mode 100644 index 0000000..db7e6f6 --- /dev/null +++ b/native/client_core/include/openusdconnect/client/engine/producer_endpoint.h @@ -0,0 +1,219 @@ +#pragma once + +#include "openusdconnect/client/engine/actions.h" +#include "openusdconnect/client/engine/notification.h" +#include "openusdconnect/client/producer_recovery.h" +#include "openusdconnect/client/producer_session.h" +#include "openusdconnect/client/protocol_codec.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +struct ProducerConfig final +{ + std::string Host; + std::uint16_t Port = 0; + std::string ClientId; + std::string Origin; + std::string Department; + OpenUSDConnect::LayerMode LayerMode = OpenUSDConnect::LayerMode::Managed; + // The first producer session; AbandonRejectedSession starts each later one. + std::string SessionId; + // Bounds an attempt from its ConnectAction to the handshake result. The host + // reports a write that makes no progress for this long as TransportError. + std::chrono::milliseconds HandshakeTimeout{10'000}; + std::size_t MaxPendingTransactions = 10'000; + std::chrono::milliseconds ReconnectBaseDelay{1'000}; + std::chrono::milliseconds ReconnectMaxDelay{8'000}; +}; + +// The replay position at which the server made the latest acknowledgement durable. +struct MirrorCheckpoint final +{ + std::string ServerInstance; + std::uint64_t Epoch = 0; + std::int32_t HeadSequence = 0; +}; + +// The unacknowledged outbox a failure quarantined, with the frames as appended. +struct RecoveryArtifact final +{ + std::string SessionId; + TransactionFailure Failure; + std::vector Transactions; +}; + +struct ProducerStatus final +{ + // The handshake completed; appended transactions are sent. + bool Connected = false; + // An attempt is connecting or awaiting the handshake response. + bool Handshaking = false; + // The endpoint closed a connection or attempt that the host has not yet + // reported gone; Connect is Busy until it does. + bool Closing = false; + bool Stopped = false; + std::optional Rejection; + OpenUSDConnect::LayerMode LayerModeActive = OpenUSDConnect::LayerMode::Managed; + StageMetadata Metadata; + std::string SessionId; + std::size_t PendingTransactions = 0; + std::size_t PendingEvents = 0; + std::uint64_t AcknowledgedTransactions = 0; + std::uint64_t AcknowledgedEvents = 0; + std::uint64_t NextTransactionId = 1; + // Set exactly while recovery is required. + std::optional Failure; + // Attempts are refused before this time, which the server's RateLimited set. + TimePoint RetryAfter; +}; + +enum class ConnectResult : std::uint8_t +{ + Started, + Connected, + // An attempt or a close is in flight; try again once it ends. + Busy, + // Stopped, recovery required, rate limited, or no time left. + Refused, +}; + +// Sans-IO producer: the host applies the returned actions in order, on the +// thread that reports its socket events, and reports time. It connects only +// when asked. Thread-safe; never blocks or calls into the host. +class OPENUSDCONNECT_CLIENT_API ProducerEndpoint final +{ +public: + ProducerEndpoint(ProducerConfig config, NotificationQueue& notifications); + ProducerEndpoint(const ProducerEndpoint&) = delete; + ProducerEndpoint& operator=(const ProducerEndpoint&) = delete; + + [[nodiscard]] static bool IsValidConfiguration(const ProducerConfig& config) noexcept; + + [[nodiscard]] const ProducerConfig& Configuration() const noexcept; + + // One attempt for a retry loop. Refused while a connection, attempt, or close + // exists, after a handshake rejection, while recovery is required, inside the + // rate-limit window, and inside the backoff that a failed request starts. + [[nodiscard]] bool RequestConnect(TimePoint now, TimePoint deadline); + // One attempt regardless of backoff; it clears a handshake rejection. + [[nodiscard]] ConnectResult Connect(TimePoint now, TimePoint deadline); + // Abandons an attempt in flight and resets the backoff, leaving a connection + // intact. Returns whether nothing remains for the host to close. + bool CancelConnect(); + // Cancels any attempt, then says Quit and closes the connection. The outbox + // is kept for the next connection. + void Disconnect(); + + // Host I/O. Reports that do not match the current connection are ignored. + // Read the token just before calling, so a newly issued one is presented. + void OnConnected(std::string_view token); + void OnBytes(const std::uint8_t* data, std::size_t size); + void OnDisconnected(DisconnectReason reason, TimePoint now); + void OnTick(TimePoint now); + // Like Disconnect, and refuses every later attempt. + void Stop(); + [[nodiscard]] std::vector TakeActions(); + // The token the latest accepted Hello issued, once. + [[nodiscard]] std::optional TakeIssuedToken(); + [[nodiscard]] std::optional NextWake() const; + + // Frames are complete and length-prefixed. Append's frame must encode + // NextTransactionId(), so a host submitting from several threads holds one + // lock from reading the id through Append. + [[nodiscard]] std::uint64_t NextTransactionId() const noexcept; + [[nodiscard]] ProducerResult Append(std::uint64_t transaction_id, + std::vector frame, std::size_t event_count, + std::string layer_key); + // Sent in order with transactions while connected, and never replayed. + [[nodiscard]] bool QueueControl(std::vector frame); + [[nodiscard]] bool OutboxEmpty() const noexcept; + [[nodiscard]] std::uint64_t DrainAcknowledgedEventCount() noexcept; + // Set while every appended transaction is acknowledged, if the latest + // acknowledgement carried a checkpoint. + [[nodiscard]] std::optional AcknowledgedCheckpoint() const; + + [[nodiscard]] std::optional Failure() const; + [[nodiscard]] std::optional Artifact() const; + // Replaces a recoverable rejected transaction with a frame that encodes + // Failure()->TransactionId; later transactions follow it unchanged. + [[nodiscard]] ProducerResult RepairRejected(std::vector frame, + std::size_t event_count, std::string layer_key); + // Discards a failed session and continues as session_id from transaction 1. + // Nullopt without a failure, or for an invalid or unchanged session_id. + [[nodiscard]] std::optional AbandonRejectedSession(std::string session_id); + + [[nodiscard]] ProducerStatus Status() const; + +private: + enum class ConnectionState : std::uint8_t + { + Idle, + Connecting, + Handshaking, + Connected, + Closing, + Stopped, + }; + + [[nodiscard]] bool IsAttempting() const noexcept; + [[nodiscard]] bool IsOpen() const noexcept; + void BeginAttempt(TimePoint now, TimePoint deadline, bool backoff_on_failure); + void ResetBackoff() noexcept; + void Close(DisconnectReason reason); + void EndConnection(DisconnectReason reason); + + void HandleFrame(const std::vector& frame); + void HandleHandshake(EnvelopeView envelope); + void AcceptHello(const OpenUSDConnect::HelloOk& hello); + void Reject(HandshakeRejected rejection); + void Publish(); + std::size_t SendUnsent(); + void HandleMessage(EnvelopeView envelope); + void AcceptResult(const OpenUSDConnect::TransactionResult& result); + void AcceptRateLimit(const OpenUSDConnect::RateLimited& limited); + void Fail(TransactionFailure failure); + [[nodiscard]] std::string HighwaterFailureReason(ProducerResult result, + std::uint64_t transaction_id) const; + [[nodiscard]] RecoveryArtifact CaptureArtifact() const; + + void Notify(Notification notification); + void Log(LogLevel level, std::string message); + + mutable std::mutex Mutex; + const ProducerConfig Config; + // A leaf lock, so pushing while Mutex is held keeps notifications ordered. + NotificationQueue& Notifications; + ProducerSession Session; + FrameDecoder Decoder; + ReconnectPolicy Reconnect; + std::vector Actions; + ConnectionState State = ConnectionState::Idle; + // The session generation of the attempt or connection in flight. + std::uint64_t ConnectionGeneration = 0; + TimePoint AttemptDeadline; + bool BackoffOnFailure = false; + TimePoint BackoffUntil; + TimePoint RetryAfterUntil; + // Applied once the host reports the close that RateLimited requested. + std::optional PendingRetryAfter; + std::optional Rejection; + std::optional IssuedToken; + std::optional SessionFailure; + std::string SessionId; + std::string ServerInstance; + std::optional Checkpoint; + OpenUSDConnect::LayerMode LayerModeActive = OpenUSDConnect::LayerMode::Managed; + StageMetadata Metadata; +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/engine/receiver_endpoint.h b/native/client_core/include/openusdconnect/client/engine/receiver_endpoint.h new file mode 100644 index 0000000..6574ce0 --- /dev/null +++ b/native/client_core/include/openusdconnect/client/engine/receiver_endpoint.h @@ -0,0 +1,163 @@ +#pragma once + +#include "openusdconnect/client/engine/actions.h" +#include "openusdconnect/client/engine/notification.h" +#include "openusdconnect/client/protocol_codec.h" +#include "openusdconnect/client/receiver_session.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +struct ReceiverConfig final +{ + std::string Host; + std::uint16_t Port = 0; + // Optional: a server that requires tokens rejects a receiver without one. + std::string ClientId; + std::string Origin; + std::string Department; + bool LayeredReplay = true; + OpenUSDConnect::LayerMode LayerMode = OpenUSDConnect::LayerMode::Managed; + // Above one, the consumer already holds the prefix, for example from a snapshot. + std::int32_t SyncFrom = 1; + std::size_t MaxQueue = 50'000; + // Bounds a connect and a write, and how long a host read waits before OnReadTimeout. + std::chrono::milliseconds SocketTimeout{30'000}; + std::uint32_t MaxConsecutiveTimeouts = 10; + bool Reconnect = true; + std::chrono::milliseconds ReconnectBaseDelay{1'000}; + std::chrono::milliseconds ReconnectMaxDelay{30'000}; +}; + +struct ReceiverStatus final +{ + bool Connected = false; + // The consumer applied the replay through the head the server advertised. + bool Synchronized = false; + // No further connection attempt will be made. + bool Stopped = false; + std::int32_t ReplayHeadSequence = 0; + std::uint64_t ReplayEpoch = 0; + // The server whose replay the consumer applied; empty when unproven. + std::string ServerInstance; + bool LayeredReplayActive = false; + OpenUSDConnect::LayerMode LayerModeActive = OpenUSDConnect::LayerMode::Managed; + std::optional Rejection; + StageMetadata Metadata; + std::size_t QueuedFrames = 0; + std::int32_t LastSequence = 0; + std::int32_t LastAppliedSequence = 0; +}; + +// Sans-IO receiver: the host applies the returned actions and reports socket +// events and time. Thread-safe; never blocks or calls into the host. +class OPENUSDCONNECT_CLIENT_API ReceiverEndpoint final +{ +public: + ReceiverEndpoint(ReceiverConfig config, NotificationQueue& notifications); + ReceiverEndpoint(const ReceiverEndpoint&) = delete; + ReceiverEndpoint& operator=(const ReceiverEndpoint&) = delete; + + [[nodiscard]] static bool IsValidConfiguration(const ReceiverConfig& config) noexcept; + + [[nodiscard]] const ReceiverConfig& Configuration() const noexcept; + // Applies when the current session ends; a stopped endpoint stays stopped. + void SetReconnect(bool enabled); + + // Host I/O. Reports that do not match the current connection are ignored. + [[nodiscard]] bool Start(TimePoint now); + // Read the token just before calling, so a newly issued one is presented. + void OnConnected(std::string_view token); + void OnBytes(const std::uint8_t* data, std::size_t size); + // A read waited SocketTimeout without receiving a byte. + void OnReadTimeout(); + void OnDisconnected(DisconnectReason reason, TimePoint now); + void OnTick(TimePoint now); + void Stop(); + [[nodiscard]] std::vector TakeActions(); + // The token the latest accepted Hello issued, once. + [[nodiscard]] std::optional TakeIssuedToken(); + [[nodiscard]] std::optional NextWake() const; + + // Stage-owning consumer. Read Generation before draining; report progress + // only once every drained frame applied, and after a failure call + // RequestReplayFrom with the applied cursor instead. + [[nodiscard]] std::vector> + DrainFrames(std::optional max_frames = std::nullopt); + [[nodiscard]] std::uint64_t Generation() const noexcept; + [[nodiscard]] bool MarkAppliedThrough(std::uint64_t generation, std::int32_t sequence); + void ResetAppliedProgress() noexcept; + [[nodiscard]] bool MarkReplayApplied(); + // Discards the queued frames and reconnects to replay from sequence. + [[nodiscard]] bool RequestReplayFrom(std::int32_t sequence); + [[nodiscard]] std::uint64_t FreezeMarker() const noexcept; + [[nodiscard]] bool DrainedThrough(std::uint64_t marker) const noexcept; + + [[nodiscard]] ReceiverStatus Status() const; + +private: + enum class ConnectionState : std::uint8_t + { + Idle, + Connecting, + Handshaking, + Connected, + Closing, + Backoff, + DrainWait, + Stopped, + }; + + [[nodiscard]] bool IsOpen() const noexcept; + void BeginAttempt(TimePoint now); + void ScheduleNextAttempt(TimePoint now); + void PollDrain(TimePoint now); + void Close(DisconnectReason reason); + void EndConnection(DisconnectReason reason); + void RequestReplay(std::int32_t sequence, DisconnectReason reason); + + void HandleFrame(std::vector frame); + void HandleHandshake(EnvelopeView envelope); + void AcceptHello(const OpenUSDConnect::HelloOk& hello); + void Reject(HandshakeRejected rejection); + void HandleMessage(EnvelopeView envelope, std::vector& frame); + void AcceptReplayComplete(const OpenUSDConnect::ReplayComplete& complete); + void AcceptFrame(ReceiverMessageKind kind, std::int32_t sequence, + std::vector frame); + void ReplayAfterGap(const std::string& description); + + void Notify(Notification notification); + void Log(LogLevel level, std::string message); + + mutable std::mutex Mutex; + const ReceiverConfig Config; + // A leaf lock, so pushing while Mutex is held keeps notifications ordered. + NotificationQueue& Notifications; + ReceiverInbox Inbox; + ReceiverReplayIdentity Identity; + FrameDecoder Decoder; + ReconnectPolicy Reconnect; + std::vector Actions; + ConnectionState State = ConnectionState::Idle; + std::uint64_t ConnectionGeneration = 0; + std::int32_t ConnectionSyncFrom = 0; + std::uint32_t ConsecutiveTimeouts = 0; + TimePoint WakeTime; + TimePoint DrainDeadline; + std::optional Rejection; + std::optional IssuedToken; + bool LayeredReplayActive = false; + OpenUSDConnect::LayerMode LayerModeActive = OpenUSDConnect::LayerMode::Managed; + StageMetadata Metadata; +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/engine/status.h b/native/client_core/include/openusdconnect/client/engine/status.h new file mode 100644 index 0000000..3f23e65 --- /dev/null +++ b/native/client_core/include/openusdconnect/client/engine/status.h @@ -0,0 +1,53 @@ +#pragma once + +#include +#include + +namespace openusdconnect::client +{ + +enum class ClientPhase : std::uint8_t +{ + Offline, + Connecting, + Replaying, + Ready, + RecoveryRequired, + Rejected, + Closed, + Parked, +}; + +struct PhaseInputs final +{ + bool Closed = false; + bool RecoveryRequired = false; + bool Rejected = false; + bool Parked = false; + bool Replaying = false; + bool Ready = false; + bool Connecting = false; +}; + +[[nodiscard]] inline ClientPhase ComputePhase(const PhaseInputs& inputs) noexcept +{ + const std::pair precedence[] = { + {inputs.Closed, ClientPhase::Closed}, + {inputs.RecoveryRequired, ClientPhase::RecoveryRequired}, + {inputs.Rejected, ClientPhase::Rejected}, + {inputs.Parked, ClientPhase::Parked}, + {inputs.Replaying, ClientPhase::Replaying}, + {inputs.Ready, ClientPhase::Ready}, + {inputs.Connecting, ClientPhase::Connecting}, + }; + for (const auto& [active, phase] : precedence) + { + if (active) + { + return phase; + } + } + return ClientPhase::Offline; +} + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/frame_codec.h b/native/client_core/include/openusdconnect/client/frame_codec.h index 681a41b..3b5abbf 100644 --- a/native/client_core/include/openusdconnect/client/frame_codec.h +++ b/native/client_core/include/openusdconnect/client/frame_codec.h @@ -4,6 +4,12 @@ #include #include +// A host that builds the core into a shared library defines this as its export +// or import attribute. +#ifndef OPENUSDCONNECT_CLIENT_API +#define OPENUSDCONNECT_CLIENT_API +#endif + namespace openusdconnect::client { @@ -19,7 +25,8 @@ enum class FrameResult : std::uint8_t InvalidHeader, }; -[[nodiscard]] bool IsValidMaxFrameSize(std::size_t max_frame_size) noexcept; +[[nodiscard]] OPENUSDCONNECT_CLIENT_API bool +IsValidMaxFrameSize(std::size_t max_frame_size) noexcept; class FrameDecoder final { @@ -49,7 +56,7 @@ class FrameDecoder final std::vector& frame, std::size_t max_frame_size = kDefaultMaxFrameSize); -[[nodiscard]] FrameResult +[[nodiscard]] OPENUSDCONNECT_CLIENT_API FrameResult WriteFrameHeader(std::size_t payload_size, std::uint8_t* destination, std::size_t max_frame_size = kDefaultMaxFrameSize) noexcept; diff --git a/native/client_core/include/openusdconnect/client/producer_recovery.h b/native/client_core/include/openusdconnect/client/producer_recovery.h new file mode 100644 index 0000000..5e4cd2d --- /dev/null +++ b/native/client_core/include/openusdconnect/client/producer_recovery.h @@ -0,0 +1,102 @@ +#pragma once + +#include +#include +#include +#include + +namespace openusdconnect::client +{ + +enum class ProducerRecoveryDisposition : std::uint8_t +{ + None, + RecoverableConflict, + InvalidOperation, + SessionFatal, +}; + +// Wire values of OpenUSDConnect::TransactionRejectionCode. +enum class RejectionCode : std::uint8_t +{ + None = 0, + InvalidIdentity = 1, + UnexpectedId = 2, + StaleLayerGraph = 3, + InvalidTransaction = 4, +}; + +namespace detail +{ + +struct RejectionPolicy final +{ + RejectionCode Code; + std::string_view Name; + ProducerRecoveryDisposition Disposition; +}; + +inline constexpr RejectionPolicy kRejectionPolicies[] = { + {RejectionCode::InvalidIdentity, "invalid_identity", ProducerRecoveryDisposition::SessionFatal}, + {RejectionCode::UnexpectedId, "unexpected_id", ProducerRecoveryDisposition::SessionFatal}, + {RejectionCode::StaleLayerGraph, "stale_layer_graph", + ProducerRecoveryDisposition::RecoverableConflict}, + {RejectionCode::InvalidTransaction, "invalid_transaction", + ProducerRecoveryDisposition::InvalidOperation}, +}; + +[[nodiscard]] inline const RejectionPolicy* FindRejectionPolicy(std::uint8_t code) noexcept +{ + for (const RejectionPolicy& policy : kRejectionPolicies) + { + if (static_cast(policy.Code) == code) + { + return &policy; + } + } + return nullptr; +} + +} // namespace detail + +[[nodiscard]] inline std::optional RejectionCodeName(std::uint8_t code) noexcept +{ + const detail::RejectionPolicy* policy = detail::FindRejectionPolicy(code); + return policy ? std::optional(policy->Name) : std::nullopt; +} + +// Unknown codes from newer servers fail closed. +[[nodiscard]] inline ProducerRecoveryDisposition RejectionDisposition(std::uint8_t code) noexcept +{ + const detail::RejectionPolicy* policy = detail::FindRejectionPolicy(code); + return policy ? policy->Disposition : ProducerRecoveryDisposition::SessionFatal; +} + +// A rejected transaction, or a server acknowledgement the outbox cannot accept. +struct TransactionFailure final +{ + std::uint64_t TransactionId = 0; + // The wire value, which a newer server may extend. + std::uint8_t Code = 0; + std::string Reason; + std::uint64_t ExpectedTransactionId = 0; + + [[nodiscard]] ProducerRecoveryDisposition Disposition() const noexcept + { + return RejectionDisposition(Code); + } + + [[nodiscard]] std::string Describe() const + { + const std::optional name = RejectionCodeName(Code); + std::string text = "transaction " + std::to_string(TransactionId) + " rejected ("; + text += name ? std::string(*name) : "unknown_" + std::to_string(Code); + if (ExpectedTransactionId != 0) + { + text += ", expected transaction " + std::to_string(ExpectedTransactionId); + } + return text + "): " + (Reason.empty() ? "no reason supplied" : Reason); + } +}; + +} // namespace openusdconnect::client diff --git a/native/client_core/include/openusdconnect/client/producer_session.h b/native/client_core/include/openusdconnect/client/producer_session.h index 1ba14d9..852d4d3 100644 --- a/native/client_core/include/openusdconnect/client/producer_session.h +++ b/native/client_core/include/openusdconnect/client/producer_session.h @@ -1,6 +1,7 @@ #pragma once #include "openusdconnect/client/detail/ordered_outbox_storage.h" +#include "openusdconnect/client/producer_recovery.h" #include #include @@ -38,14 +39,6 @@ enum class ProducerResult : std::uint8_t InvalidArgument, }; -enum class ProducerRecoveryDisposition : std::uint8_t -{ - None, - RecoverableConflict, - InvalidOperation, - SessionFatal, -}; - struct ProducerConnectionStart final { std::uint64_t Generation; diff --git a/native/client_core/include/openusdconnect/client/protocol_codec.h b/native/client_core/include/openusdconnect/client/protocol_codec.h index bedb2c5..22ebf3c 100644 --- a/native/client_core/include/openusdconnect/client/protocol_codec.h +++ b/native/client_core/include/openusdconnect/client/protocol_codec.h @@ -4,6 +4,7 @@ #include "openusdconnect/client/replay_identity.h" #include "openusdconnect/client/schema/messages_generated.h" +#include #include #include #include @@ -210,10 +211,27 @@ struct HelloParameters final ReplayPrefixClaim ReplayPrefix; }; +inline constexpr std::size_t kMaxProducerSessionIdLength = 128; + +// Counts code points, as the server does. +[[nodiscard]] inline bool IsValidProducerSessionId(std::string_view session_id) noexcept +{ + const auto code_points = + std::count_if(session_id.begin(), session_id.end(), + [](char byte) + { + return (static_cast(byte) & 0xC0U) != 0x80U; + }); + return code_points != 0 && static_cast(code_points) <= kMaxProducerSessionIdLength; +} + +// The server requires an emitter's client and producer session ids. Every other +// identity field is optional, and the server decodes empty strings as absent. [[nodiscard]] inline bool IsValidHelloParameters(const HelloParameters& parameters) noexcept { - return (parameters.Role == "receiver" || parameters.Role == "emitter") && - parameters.SyncFrom >= 0 && !parameters.ClientId.empty() && !parameters.Origin.empty(); + const bool identified_emitter = parameters.Role == "emitter" && !parameters.ClientId.empty() && + IsValidProducerSessionId(parameters.ProducerSessionId); + return (parameters.Role == "receiver" || identified_emitter) && parameters.SyncFrom >= 0; } [[nodiscard]] inline flatbuffers::Offset @@ -255,22 +273,20 @@ BuildHelloFrame(flatbuffers::FlatBufferBuilder& builder, const HelloParameters& { return ProtocolResult::InvalidArgument; } - const auto replay_server_instance = parameters.ReplayPrefix - ? CreateString(builder, parameters.ReplayPrefix->ServerInstance()) - : flatbuffers::Offset(); + const auto replay_server_instance = + parameters.ReplayPrefix ? CreateString(builder, parameters.ReplayPrefix->ServerInstance()) + : flatbuffers::Offset(); const std::optional claimed_epoch = parameters.ReplayPrefix ? parameters.ReplayPrefix->Epoch() : std::nullopt; - const auto replay_epoch = claimed_epoch - ? flatbuffers::Optional(*claimed_epoch) - : flatbuffers::nullopt; + const auto replay_epoch = + claimed_epoch ? flatbuffers::Optional(*claimed_epoch) : flatbuffers::nullopt; const auto hello = OpenUSDConnect::CreateHello( builder, CreateString(builder, parameters.Role), kProtocolVersion, parameters.SyncFrom, CreateString(builder, parameters.ClientId), CreateString(builder, parameters.Origin), CreateString(builder, parameters.Department), CreateString(builder, parameters.Token), parameters.LayeredReplay, parameters.LayerMode, - CreateString(builder, parameters.ProducerSessionId), replay_server_instance, - replay_epoch); + CreateString(builder, parameters.ProducerSessionId), replay_server_instance, replay_epoch); const auto envelope = OpenUSDConnect::CreateEnvelope(builder, OpenUSDConnect::Payload::Hello, hello.Union(), kSchemaVersion); return FinishEnvelopeFrame(builder, envelope, max_frame_size); diff --git a/native/client_core/include/openusdconnect/client/receiver_session.h b/native/client_core/include/openusdconnect/client/receiver_session.h index 2683414..71b7a29 100644 --- a/native/client_core/include/openusdconnect/client/receiver_session.h +++ b/native/client_core/include/openusdconnect/client/receiver_session.h @@ -122,6 +122,7 @@ class OrderedReceiverSession final { LastReceivedSequence = 0; ResetSynchronization(); + NewestResetSerial = IncomingSerial + 1; } else if (kind == ReceiverMessageKind::Event || kind == ReceiverMessageKind::LayerGraphState) { @@ -144,6 +145,11 @@ class OrderedReceiverSession final { return AcceptResult::StaleGeneration; } + // The server sends every replay record through the head before the marker. + if (RequireContiguous && head_seq > LastReceivedSequence) + { + return AcceptResult::SequenceGap; + } PendingReplay = ReplayState{generation, head_seq, epoch, IncomingSerial}; return AcceptResult::Accepted; } @@ -242,15 +248,18 @@ class OrderedReceiverSession final LastAppliedSequenceValue = sequence - 1; Frames.clear(); DrainedSerial = IncomingSerial; + ResetsSettledSerial = DrainedSerial; OverflowedValue = false; ResetSynchronization(); return true; } + // The consumer applied every Resync it has drained. void ResetAppliedProgress() noexcept { std::lock_guard lock(Mutex); LastAppliedSequenceValue = 0; + ResetsSettledSerial = DrainedSerial; // The consumer is applying a queued Resync. Its ReplayComplete may // already have arrived, so retain that marker until its frames apply. SynchronizedValue = false; @@ -292,6 +301,13 @@ class OrderedReceiverSession final std::lock_guard lock(Mutex); return OverflowedValue; } + // A Resync is queued, or drained but not yet reported applied, since the + // last replay request. + [[nodiscard]] bool ResetPending() const noexcept + { + std::lock_guard lock(Mutex); + return NewestResetSerial > ResetsSettledSerial; + } [[nodiscard]] std::int32_t ReplayHeadSequence() const noexcept { std::lock_guard lock(Mutex); @@ -302,6 +318,16 @@ class OrderedReceiverSession final std::lock_guard lock(Mutex); return ReplayEpochValue; } + [[nodiscard]] std::uint64_t FreezeMarker() const noexcept + { + std::lock_guard lock(Mutex); + return IncomingSerial; + } + [[nodiscard]] bool DrainedThrough(std::uint64_t marker) const noexcept + { + std::lock_guard lock(Mutex); + return DrainedSerial >= marker; + } private: void ResetSynchronization() noexcept @@ -320,6 +346,8 @@ class OrderedReceiverSession final std::uint64_t GenerationValue = 0; std::uint64_t IncomingSerial = 0; std::uint64_t DrainedSerial = 0; + std::uint64_t NewestResetSerial = 0; + std::uint64_t ResetsSettledSerial = 0; std::int32_t LastReceivedSequence = 0; std::int32_t LastAppliedSequenceValue = 0; std::int32_t ReplayHeadSequenceValue = 0; diff --git a/native/client_core/include/openusdconnect/client/replay_identity.h b/native/client_core/include/openusdconnect/client/replay_identity.h index fc140cd..9d1668a 100644 --- a/native/client_core/include/openusdconnect/client/replay_identity.h +++ b/native/client_core/include/openusdconnect/client/replay_identity.h @@ -63,27 +63,29 @@ class ReplayPrefixIdentity final using ReplayPrefixClaim = std::optional; -// Tracks which replay sequence domain has actually been applied by a receiver. +// Tracks the replay sequence domain of the prefix a receiver holds. The received +// identity covers the applied frames plus the retained queue and is the next +// Hello's claim; the applied identity is published once a replay has applied. // Callers provide synchronization when connection and consumer threads overlap. class ReceiverReplayIdentity final { public: + // The first Hello never claims, so an externally supplied cursor keeps its + // contract without being taken as proof of its prefix. [[nodiscard]] ReplayPrefixClaim BeginConnection() { PendingIdentity.reset(); - ClaimedIdentity = AppliedIdentity; - const bool IncludeClaim = HelloSent; + ClaimedIdentity = ReceivedIdentity; + ClaimIncluded = HelloSent; HelloSent = true; - ClaimIncluded = IncludeClaim; - if (!IncludeClaim) + if (!ClaimIncluded) { return std::nullopt; } - - if (AppliedIdentity) + if (ClaimedIdentity) { - return ReplayPrefixIdentity::Known(AppliedIdentity->ServerInstance, - AppliedIdentity->Epoch); + return ReplayPrefixIdentity::Known(ClaimedIdentity->ServerInstance, + ClaimedIdentity->Epoch); } return ReplayPrefixIdentity::Unknown(); } @@ -91,43 +93,73 @@ class ReceiverReplayIdentity final void AcceptHello(std::int32_t sync_from, bool replay_identity_supported, std::string_view server_instance, std::optional epoch) { - ConnectionIdentity.reset(); - if (replay_identity_supported && !server_instance.empty() && epoch) + ConnectionInstance = + replay_identity_supported ? std::string(server_instance) : std::string(); + HandshakeIdentity.reset(); + if (!ConnectionInstance.empty() && epoch) + { + HandshakeIdentity = ReplayIdentity{ConnectionInstance, *epoch}; + } + ConnectionPrefixProven = sync_from == 1 || (ClaimIncluded && HandshakeIdentity && + ClaimedIdentity == HandshakeIdentity); + // A changed or unknown prefix keeps its identity until the server's Resync. + if (!HandshakeIdentity) { - ConnectionIdentity = ReplayIdentity{std::string(server_instance), *epoch}; + ReceivedIdentity.reset(); + } + else if (ConnectionPrefixProven) + { + ReceivedIdentity = HandshakeIdentity; } - ConnectionPrefixProven = - sync_from == 1 || (ClaimIncluded && ClaimedIdentity && ConnectionIdentity && - *ClaimedIdentity == *ConnectionIdentity); } - void AcceptResync() noexcept + void AcceptResync() { + ReceivedIdentity = std::exchange(HandshakeIdentity, std::nullopt); ConnectionPrefixProven = true; PendingIdentity.reset(); + ResetRequiredValue = false; } void AcceptReplayComplete(std::uint64_t epoch) { - if (ConnectionPrefixProven && ConnectionIdentity) + HandshakeIdentity.reset(); + ReceivedIdentity.reset(); + if (ConnectionPrefixProven && !ConnectionInstance.empty()) { - PendingIdentity = ReplayIdentity{ConnectionIdentity->ServerInstance, epoch}; - } - else - { - PendingIdentity.reset(); + ReceivedIdentity = ReplayIdentity{ConnectionInstance, epoch}; } + PendingIdentity = ReceivedIdentity; } void MarkReplayApplied() { - AppliedIdentity = PendingIdentity; + AppliedIdentity = std::exchange(PendingIdentity, std::nullopt); + } + + // A reset that is discarded or not yet applied separates the consumer's + // prefix from the received frames, so only the applied replay names it. + void RequestReplayFrom(std::int32_t sequence, bool reset_pending) + { + if (reset_pending) + { + ReceivedIdentity = AppliedIdentity; + } + HandshakeIdentity.reset(); PendingIdentity.reset(); + ResetRequiredValue = sequence == 1; } - [[nodiscard]] const std::optional& Applied() const noexcept + // The server never resets a replay from sequence one, so the receiver must + // queue the reset itself once the next Hello is accepted. + [[nodiscard]] bool ResetRequired() const noexcept { - return AppliedIdentity; + return ResetRequiredValue; + } + + [[nodiscard]] const std::optional& Received() const noexcept + { + return ReceivedIdentity; } [[nodiscard]] const std::optional& Pending() const noexcept @@ -135,6 +167,11 @@ class ReceiverReplayIdentity final return PendingIdentity; } + [[nodiscard]] const std::optional& Applied() const noexcept + { + return AppliedIdentity; + } + [[nodiscard]] bool IsConnectionPrefixProven() const noexcept { return ConnectionPrefixProven; @@ -144,8 +181,11 @@ class ReceiverReplayIdentity final bool HelloSent = false; bool ClaimIncluded = false; bool ConnectionPrefixProven = false; + bool ResetRequiredValue = false; + std::string ConnectionInstance; std::optional ClaimedIdentity; - std::optional ConnectionIdentity; + std::optional HandshakeIdentity; + std::optional ReceivedIdentity; std::optional PendingIdentity; std::optional AppliedIdentity; }; diff --git a/native/client_core/src/driver/testing/scripted_socket.cpp b/native/client_core/src/driver/testing/scripted_socket.cpp new file mode 100644 index 0000000..cfaf1d2 --- /dev/null +++ b/native/client_core/src/driver/testing/scripted_socket.cpp @@ -0,0 +1,369 @@ +#include "openusdconnect/client/driver/testing/scripted_socket.h" + +#include +#include +#include +#include +#include +#include +#include + +namespace openusdconnect::client +{ +namespace detail +{ + +struct ScriptedChannel final +{ + enum class Kind : std::uint8_t + { + Bytes, + Closed, + }; + + struct Delivery final + { + Kind Type = Kind::Bytes; + std::vector Bytes; + }; + + std::deque Inbound; + std::vector Outbound; + bool Receiving = false; + bool ClientClosed = false; + bool SendsStalled = false; +}; + +struct ScriptedAttempt final +{ + enum class Outcome : std::uint8_t + { + Pending, + Accepted, + Refused, + }; + + Outcome Result = Outcome::Pending; + std::shared_ptr Channel; + int SystemError = 0; +}; + +// One lock for every socket, channel, and attempt of a factory. +struct ScriptState final +{ + mutable std::mutex Mutex; + mutable std::condition_variable Changed; + std::deque> Pending; + std::size_t Attempts = 0; +}; + +} // namespace detail + +namespace +{ + +using detail::ScriptedAttempt; +using detail::ScriptedChannel; +using detail::ScriptState; + +class ScriptedSocket final : public Socket +{ +public: + explicit ScriptedSocket(std::shared_ptr state) + : State(std::move(state)) + { + } + + ~ScriptedSocket() override + { + std::lock_guard lock(State->Mutex); + if (Channel) + { + Channel->ClientClosed = true; + State->Changed.notify_all(); + } + } + + SocketResult Connect(const std::string&, std::uint16_t, TimePoint deadline) override + { + std::unique_lock lock(State->Mutex); + const auto attempt = std::make_shared(); + State->Pending.push_back(attempt); + ++State->Attempts; + State->Changed.notify_all(); + State->Changed.wait_until(lock, deadline, + [&] + { + return Interrupted || + attempt->Result != ScriptedAttempt::Outcome::Pending; + }); + if (attempt->Result == ScriptedAttempt::Outcome::Pending) + { + State->Pending.erase(std::find(State->Pending.begin(), State->Pending.end(), attempt)); + return Interrupted ? SocketResult::Interrupted : SocketResult::Timeout; + } + if (attempt->Result == ScriptedAttempt::Outcome::Refused) + { + Error = attempt->SystemError; + return SocketResult::Failed; + } + Channel = attempt->Channel; + if (Interrupted) + { + Channel->ClientClosed = true; + State->Changed.notify_all(); + return SocketResult::Interrupted; + } + return SocketResult::Success; + } + + SocketResult SendAll(const std::uint8_t* data, std::size_t size, TimePoint deadline) override + { + std::unique_lock lock(State->Mutex); + if (Interrupted) + { + return SocketResult::Interrupted; + } + if (!Channel) + { + return SocketResult::Failed; + } + if (Channel->SendsStalled) + { + const bool interrupted = State->Changed.wait_until(lock, deadline, + [this] + { + return Interrupted; + }); + return interrupted ? SocketResult::Interrupted : SocketResult::Timeout; + } + Channel->Outbound.insert(Channel->Outbound.end(), data, data + size); + State->Changed.notify_all(); + return SocketResult::Success; + } + + SocketResult Receive(std::uint8_t* buffer, std::size_t capacity, + std::optional deadline, std::size_t& received) override + { + received = 0; + std::unique_lock lock(State->Mutex); + for (;;) + { + if (Interrupted || std::exchange(WakePending, false)) + { + return SocketResult::Interrupted; + } + if (!Channel) + { + return SocketResult::Failed; + } + if (!Channel->Inbound.empty()) + { + return Read(buffer, capacity, received); + } + Channel->Receiving = true; + State->Changed.notify_all(); + bool expired = false; + if (deadline) + { + expired = State->Changed.wait_until(lock, *deadline) == std::cv_status::timeout; + } + else + { + State->Changed.wait(lock); + } + Channel->Receiving = false; + if (expired && Channel->Inbound.empty() && !Interrupted && !WakePending) + { + return SocketResult::Timeout; + } + } + } + + int SystemError() const noexcept override + { + return Error; + } + + void Interrupt() noexcept override + { + std::lock_guard lock(State->Mutex); + Interrupted = true; + State->Changed.notify_all(); + } + + void Wake() noexcept override + { + std::lock_guard lock(State->Mutex); + WakePending = true; + State->Changed.notify_all(); + } + +private: + SocketResult Read(std::uint8_t* buffer, std::size_t capacity, std::size_t& received) + { + ScriptedChannel::Delivery& delivery = Channel->Inbound.front(); + switch (delivery.Type) + { + case ScriptedChannel::Kind::Bytes: + { + received = std::min(capacity, delivery.Bytes.size()); + std::memcpy(buffer, delivery.Bytes.data(), received); + if (received == delivery.Bytes.size()) + { + Channel->Inbound.pop_front(); + } + else + { + delivery.Bytes.erase(delivery.Bytes.begin(), + delivery.Bytes.begin() + + static_cast(received)); + } + return SocketResult::Success; + } + case ScriptedChannel::Kind::Closed: + return SocketResult::Closed; + } + return SocketResult::Failed; + } + + const std::shared_ptr State; + std::shared_ptr Channel; + int Error = 0; + bool Interrupted = false; + bool WakePending = false; +}; + +} // namespace + +ScriptedConnection::ScriptedConnection(std::shared_ptr state, + std::shared_ptr channel) + : State(std::move(state)) + , Channel(std::move(channel)) +{ +} + +bool ScriptedConnection::Deliver(std::vector bytes) +{ + std::lock_guard lock(State->Mutex); + if (Channel->ClientClosed) + { + return false; + } + Channel->Inbound.push_back({ScriptedChannel::Kind::Bytes, std::move(bytes)}); + State->Changed.notify_all(); + return true; +} + +void ScriptedConnection::Close() +{ + std::lock_guard lock(State->Mutex); + Channel->Inbound.push_back({ScriptedChannel::Kind::Closed, {}}); + State->Changed.notify_all(); +} + +void ScriptedConnection::StallSends() +{ + std::lock_guard lock(State->Mutex); + Channel->SendsStalled = true; +} + +std::vector ScriptedConnection::Sent() const +{ + std::lock_guard lock(State->Mutex); + return Channel->Outbound; +} + +bool ScriptedConnection::WaitSent(std::size_t size, std::chrono::milliseconds timeout) const +{ + std::unique_lock lock(State->Mutex); + return State->Changed.wait_for(lock, timeout, + [&] + { + return Channel->Outbound.size() >= size; + }); +} + +bool ScriptedConnection::WaitIdle(std::chrono::milliseconds timeout) const +{ + std::unique_lock lock(State->Mutex); + State->Changed.wait_for(lock, timeout, + [&] + { + return Channel->ClientClosed || + (Channel->Receiving && Channel->Inbound.empty()); + }); + return !Channel->ClientClosed && Channel->Receiving && Channel->Inbound.empty(); +} + +bool ScriptedConnection::WaitClosed(std::chrono::milliseconds timeout) const +{ + std::unique_lock lock(State->Mutex); + return State->Changed.wait_for(lock, timeout, + [&] + { + return Channel->ClientClosed; + }); +} + +bool ScriptedConnection::ClosedByClient() const +{ + std::lock_guard lock(State->Mutex); + return Channel->ClientClosed; +} + +ScriptedSocketFactory::ScriptedSocketFactory() + : State(std::make_shared()) +{ +} + +std::unique_ptr ScriptedSocketFactory::Create() +{ + return std::make_unique(State); +} + +std::shared_ptr ScriptedSocketFactory::Accept(std::chrono::milliseconds timeout) +{ + std::unique_lock lock(State->Mutex); + if (!State->Changed.wait_for(lock, timeout, + [&] + { + return !State->Pending.empty(); + })) + { + return nullptr; + } + const std::shared_ptr attempt = std::move(State->Pending.front()); + State->Pending.pop_front(); + attempt->Channel = std::make_shared(); + attempt->Result = ScriptedAttempt::Outcome::Accepted; + State->Changed.notify_all(); + return std::make_shared(State, attempt->Channel); +} + +bool ScriptedSocketFactory::Refuse(std::chrono::milliseconds timeout, int system_error) +{ + std::unique_lock lock(State->Mutex); + if (!State->Changed.wait_for(lock, timeout, + [&] + { + return !State->Pending.empty(); + })) + { + return false; + } + const std::shared_ptr attempt = std::move(State->Pending.front()); + State->Pending.pop_front(); + attempt->SystemError = system_error; + attempt->Result = ScriptedAttempt::Outcome::Refused; + State->Changed.notify_all(); + return true; +} + +std::size_t ScriptedSocketFactory::Attempts() const +{ + std::lock_guard lock(State->Mutex); + return State->Attempts; +} + +} // namespace openusdconnect::client diff --git a/native/client_core/src/driver/threaded_producer_driver.cpp b/native/client_core/src/driver/threaded_producer_driver.cpp new file mode 100644 index 0000000..bc444ab --- /dev/null +++ b/native/client_core/src/driver/threaded_producer_driver.cpp @@ -0,0 +1,123 @@ +#include "openusdconnect/client/driver/threaded_producer_driver.h" + +#include +#include +#include +#include + +namespace openusdconnect::client +{ +namespace +{ + +// How long Flush waits after a failed attempt before it connects again. +constexpr std::chrono::milliseconds kFlushRetryPause{100}; + +} // namespace + +bool ThreadedProducerDriver::Connect(std::optional timeout) +{ + if (!CanWait()) + { + return Target.Status().Connected; + } + const std::chrono::milliseconds handshake = Target.Configuration().HandshakeTimeout; + const TimePoint deadline = Now() + (timeout ? std::min(*timeout, handshake) : handshake); + for (;;) + { + switch (Target.Connect(Now(), deadline)) + { + case ConnectResult::Connected: + return true; + case ConnectResult::Refused: + return false; + case ConnectResult::Started: + Wake(); + static_cast(WaitSettled(deadline)); + return Target.Status().Connected; + case ConnectResult::Busy: + // Retry once the attempt or close in flight ends. + if (!WaitSettled(deadline)) + { + return Target.Status().Connected; + } + break; + } + } +} + +FlushResult ThreadedProducerDriver::Flush(std::optional timeout) +{ + const std::optional deadline = + timeout ? std::optional(Now() + *timeout) : std::nullopt; + for (;;) + { + const ProducerStatus status = Target.Status(); + if (status.Failure) + { + return FlushResult::RecoveryRequired; + } + if (status.PendingTransactions == 0) + { + return FlushResult::Flushed; + } + const TimePoint now = Now(); + if (status.Stopped || !CanWait() || (deadline && now >= *deadline)) + { + return FlushResult::Unfinished; + } + if (status.Connected) + { + static_cast(Wait( + [this] + { + const ProducerStatus current = Target.Status(); + return !current.Connected || current.PendingTransactions == 0 || + current.Failure; + }, + Until(deadline))); + } + else if (now < status.RetryAfter) + { + if (deadline && status.RetryAfter >= *deadline) + { + return FlushResult::Unfinished; + } + Pause(status.RetryAfter); + } + else if (!Connect(Until(deadline))) + { + const TimePoint retry = Now() + kFlushRetryPause; + Pause(deadline ? std::min(retry, *deadline) : retry); + } + } +} + +bool ThreadedProducerDriver::CanWait() const +{ + // On the loop thread a wait would wait for itself. + return Running() && ThreadId() != std::this_thread::get_id(); +} + +bool ThreadedProducerDriver::WaitSettled(TimePoint deadline) +{ + return Wait( + [this] + { + const ProducerStatus status = Target.Status(); + return !status.Handshaking && !status.Closing; + }, + Until(deadline)); +} + +void ThreadedProducerDriver::Pause(TimePoint until) +{ + static_cast(Wait( + [] + { + return false; + }, + Until(until))); +} + +} // namespace openusdconnect::client diff --git a/native/client_core/src/engine/endpoint_common.h b/native/client_core/src/engine/endpoint_common.h new file mode 100644 index 0000000..87dbad2 --- /dev/null +++ b/native/client_core/src/engine/endpoint_common.h @@ -0,0 +1,203 @@ +#pragma once + +#include "openusdconnect/client/engine/actions.h" +#include "openusdconnect/client/engine/notification.h" +#include "openusdconnect/client/frame_codec.h" +#include "openusdconnect/client/protocol_codec.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +// The connection steps both endpoints share, as free functions over their state. +namespace openusdconnect::client::detail +{ + +[[nodiscard]] inline std::string Text(const flatbuffers::String* value) +{ + return value ? value->str() : std::string(); +} + +template +[[nodiscard]] std::optional Value(flatbuffers::Optional value) noexcept +{ + return value.has_value() ? std::optional(*value) : std::nullopt; +} + +template +[[nodiscard]] std::string Address(const Config& config) +{ + return config.Host + ":" + std::to_string(config.Port); +} + +[[nodiscard]] inline std::shared_ptr> +Share(std::vector bytes) +{ + return std::make_shared>(std::move(bytes)); +} + +[[nodiscard]] inline std::shared_ptr> +Share(const flatbuffers::FlatBufferBuilder& builder) +{ + const std::uint8_t* bytes = builder.GetBufferPointer(); + return Share(std::vector(bytes, bytes + builder.GetSize())); +} + +// The Hello fields every role sends; the caller adds its own before QueueHello. +template +[[nodiscard]] HelloParameters CommonHello(std::string_view role, const Config& config, + std::string_view token) +{ + HelloParameters hello; + hello.Role = role; + hello.ClientId = config.ClientId; + hello.Origin = config.Origin; + hello.Department = config.Department; + hello.Token = token; + hello.LayerMode = config.LayerMode; + return hello; +} + +// False, with an error logged, when the parameters cannot be encoded. +[[nodiscard]] inline bool QueueHello(const HelloParameters& hello, std::vector& actions) +{ + flatbuffers::FlatBufferBuilder builder(256); + if (BuildHelloFrame(builder, hello) != ProtocolResult::Success) + { + actions.push_back(LogAction{LogLevel::Error, "could not build the Hello frame"}); + return false; + } + actions.push_back(SendAction{Share(builder)}); + return true; +} + +// Erases the queued actions of the given kinds; returns whether there were any. +template +bool Discard(std::vector& actions) +{ + const auto kept = std::remove_if(actions.begin(), actions.end(), + [](const Action& action) + { + return (std::holds_alternative(action) || ...); + }); + const bool discarded = kept != actions.end(); + actions.erase(kept, actions.end()); + return discarded; +} + +// Feeds bytes to decoder, then hands each complete frame to handle until it +// returns false. False, with no frame handled, when the bytes do not frame. +template +[[nodiscard]] bool HandleFrames(FrameDecoder& decoder, const std::uint8_t* data, std::size_t size, + Handle handle) +{ + std::vector> frames; + if (decoder.Feed(data, size, frames) != FrameResult::Success) + { + return false; + } + for (std::vector& frame : frames) + { + if (!handle(frame)) + { + break; + } + } + return true; +} + +[[nodiscard]] inline std::string DescribeDecodeFailure(ProtocolResult result) +{ + return result == ProtocolResult::SchemaVersionMismatch + ? "frame uses an unsupported schema version" + : "malformed frame"; +} + +// Accepted or Rejection is set, or neither for a payload that answers no Hello. +struct HandshakeOutcome final +{ + const OpenUSDConnect::HelloOk* Accepted = nullptr; + std::optional Rejection; +}; + +[[nodiscard]] inline HandshakeOutcome ClassifyHandshake(EnvelopeView envelope) +{ + const HandshakeResponseView response(envelope); + switch (response.Kind()) + { + case HandshakeResponseKind::Accepted: + return {response.Accepted(), std::nullopt}; + case HandshakeResponseKind::AuthenticationRejected: + return {nullptr, HandshakeRejected{true, OpenUSDConnect::HelloRejectionCode::Unspecified, + Text(response.AuthenticationRejection()->reason())}}; + case HandshakeResponseKind::ConfigurationRejected: + { + const OpenUSDConnect::HelloRejected& rejected = *response.ConfigurationRejection(); + return {nullptr, HandshakeRejected{false, rejected.code(), Text(rejected.reason())}}; + } + case HandshakeResponseKind::Unexpected: + break; + } + return {}; +} + +[[nodiscard]] inline std::string DescribeRejection(const HandshakeRejected& rejection) +{ + const std::string kind = + rejection.Authentication + ? "authentication rejected" + : "connection rejected (code " + std::to_string(static_cast(rejection.Code)) + ")"; + return kind + ": " + rejection.Reason; +} + +// Only the fields the server authored; nullopt when it authored none. +[[nodiscard]] inline std::optional +DecodeStageMetadata(const OpenUSDConnect::SetStageMetadata* table) +{ + if (!table) + { + return std::nullopt; + } + StageMetadata metadata; + metadata.TimeCodesPerSecond = Value(table->timeCodesPerSecond()); + metadata.FramesPerSecond = Value(table->framesPerSecond()); + metadata.StartTimeCode = Value(table->startTimeCode()); + metadata.EndTimeCode = Value(table->endTimeCode()); + metadata.MetersPerUnit = Value(table->metersPerUnit()); + if (table->upAxis() && table->upAxis()->size() != 0) + { + metadata.UpAxis = table->upAxis()->str(); + } + const bool authored = metadata.TimeCodesPerSecond || metadata.FramesPerSecond || + metadata.StartTimeCode || metadata.EndTimeCode || + metadata.MetersPerUnit || metadata.UpAxis; + return authored ? std::optional(std::move(metadata)) : std::nullopt; +} + +// Notifies the token and the authored stage metadata a HelloOk carries, and +// keeps both. +inline void NotifyHelloFields(const OpenUSDConnect::HelloOk& hello, + std::optional& issued_token, StageMetadata& metadata, + NotificationQueue& notifications, std::vector& actions) +{ + if (std::string token = Text(hello.token()); !token.empty()) + { + actions.push_back(LogAction{LogLevel::Info, "token issued by server"}); + issued_token = token; + notifications.Push(TokenIssued{std::move(token)}); + } + if (std::optional authored = DecodeStageMetadata(hello.stage_metadata())) + { + metadata = *authored; + notifications.Push(std::move(*authored)); + } +} + +} // namespace openusdconnect::client::detail diff --git a/native/client_core/src/engine/producer_endpoint.cpp b/native/client_core/src/engine/producer_endpoint.cpp new file mode 100644 index 0000000..0befc6b --- /dev/null +++ b/native/client_core/src/engine/producer_endpoint.cpp @@ -0,0 +1,664 @@ +#include "openusdconnect/client/engine/producer_endpoint.h" + +#include "endpoint_common.h" + +#include +#include +#include + +namespace openusdconnect::client +{ +namespace +{ + +using detail::Address; +using detail::Discard; +using detail::Share; +using detail::Text; + +// Bounds a hostile retry-after so the deadline arithmetic cannot overflow. +constexpr std::chrono::hours kMaxRetryAfter{1}; + +constexpr auto kUnexpectedId = static_cast(RejectionCode::UnexpectedId); + +[[nodiscard]] bool IsCompleteFrame(const std::vector& frame) noexcept +{ + std::size_t payload_size = 0; + return frame.size() > kFrameHeaderSize && + TryReadFrameHeader(frame.data(), kDefaultMaxFrameSize, payload_size) && + payload_size == frame.size() - kFrameHeaderSize; +} + +[[nodiscard]] SharedByteBuffer BuildQuitFrame() +{ + flatbuffers::FlatBufferBuilder builder(32); + const auto quit = OpenUSDConnect::CreateQuit(builder); + [[maybe_unused]] const ProtocolResult finished = FinishEnvelopeFrame( + builder, OpenUSDConnect::CreateEnvelope(builder, OpenUSDConnect::Payload::Quit, + quit.Union(), kSchemaVersion)); + assert(finished == ProtocolResult::Success); + return Share(builder); +} + +[[nodiscard]] std::string_view LayerModeName(OpenUSDConnect::LayerMode mode) noexcept +{ + switch (mode) + { + case OpenUSDConnect::LayerMode::Managed: + return "managed"; + case OpenUSDConnect::LayerMode::SharedStage: + return "shared_stage"; + } + return "unknown"; +} + +[[nodiscard]] std::chrono::steady_clock::duration RetryAfter(float seconds) +{ + const std::chrono::duration requested(seconds); + if (!(requested > requested.zero())) + { + return std::chrono::steady_clock::duration::zero(); + } + return std::chrono::duration_cast( + std::min>(requested, kMaxRetryAfter)); +} + +} // namespace + +ProducerEndpoint::ProducerEndpoint(ProducerConfig config, NotificationQueue& notifications) + : Config(std::move(config)) + , Notifications(notifications) + , Session(Config.MaxPendingTransactions) + , Reconnect(true, Config.ReconnectBaseDelay, Config.ReconnectMaxDelay) + , SessionId(Config.SessionId) +{ + assert(IsValidConfiguration(Config)); +} + +bool ProducerEndpoint::IsValidConfiguration(const ProducerConfig& config) noexcept +{ + const bool valid_mode = + config.LayerMode == OpenUSDConnect::LayerMode::Managed || + (config.LayerMode == OpenUSDConnect::LayerMode::SharedStage && config.Department.empty()); + return valid_mode && !config.Host.empty() && config.Port != 0 && !config.ClientId.empty() && + IsValidProducerSessionId(config.SessionId) && config.HandshakeTimeout.count() > 0 && + ProducerSession::IsValidConfiguration(config.MaxPendingTransactions) && + ReconnectPolicy::IsValidConfiguration(config.ReconnectBaseDelay, + config.ReconnectMaxDelay); +} + +const ProducerConfig& ProducerEndpoint::Configuration() const noexcept +{ + return Config; +} + +bool ProducerEndpoint::RequestConnect(TimePoint now, TimePoint deadline) +{ + std::lock_guard lock(Mutex); + if (State != ConnectionState::Idle || Rejection || SessionFailure || + now < std::max(BackoffUntil, RetryAfterUntil) || deadline <= now) + { + return false; + } + BeginAttempt(now, deadline, true); + return true; +} + +ConnectResult ProducerEndpoint::Connect(TimePoint now, TimePoint deadline) +{ + std::lock_guard lock(Mutex); + switch (State) + { + case ConnectionState::Connected: + return ConnectResult::Connected; + case ConnectionState::Connecting: + case ConnectionState::Handshaking: + case ConnectionState::Closing: + return ConnectResult::Busy; + case ConnectionState::Stopped: + return ConnectResult::Refused; + case ConnectionState::Idle: + break; + } + if (SessionFailure || now < RetryAfterUntil || deadline <= now) + { + return ConnectResult::Refused; + } + Rejection.reset(); + BeginAttempt(now, deadline, false); + return ConnectResult::Started; +} + +bool ProducerEndpoint::CancelConnect() +{ + std::lock_guard lock(Mutex); + ResetBackoff(); + if (IsAttempting()) + { + Log(LogLevel::Info, "connection attempt cancelled"); + Close(DisconnectReason::Cancelled); + } + return State != ConnectionState::Closing; +} + +void ProducerEndpoint::Disconnect() +{ + std::lock_guard lock(Mutex); + ResetBackoff(); + Close(DisconnectReason::Cancelled); +} + +void ProducerEndpoint::OnConnected(std::string_view token) +{ + std::lock_guard lock(Mutex); + if (State == ConnectionState::Stopped) + { + Actions.push_back(CloseAction{DisconnectReason::Stopped}); + return; + } + if (State != ConnectionState::Connecting) + { + return; + } + // Attempts start only without a failure, and none is recorded before this. + const std::optional connection = Session.BeginConnection(); + assert(connection); + ConnectionGeneration = connection->Generation; + Decoder.Reset(); + State = ConnectionState::Handshaking; + + HelloParameters hello = detail::CommonHello("emitter", Config, token); + hello.ProducerSessionId = SessionId; + if (!detail::QueueHello(hello, Actions)) + { + Close(DisconnectReason::ProtocolError); + } +} + +void ProducerEndpoint::OnBytes(const std::uint8_t* data, std::size_t size) +{ + std::lock_guard lock(Mutex); + if (!IsOpen()) + { + return; + } + const bool framed = detail::HandleFrames(Decoder, data, size, + [this](const std::vector& frame) + { + HandleFrame(frame); + return IsOpen(); + }); + if (!framed) + { + Log(LogLevel::Warning, "invalid frame header"); + Close(DisconnectReason::ProtocolError); + } +} + +void ProducerEndpoint::OnDisconnected(DisconnectReason reason, TimePoint now) +{ + std::lock_guard lock(Mutex); + if (IsOpen()) + { + Log(LogLevel::Warning, "connection to " + Address(Config) + " lost"); + EndConnection(reason); + } + else if (State != ConnectionState::Connecting && State != ConnectionState::Closing) + { + return; + } + // The socket is gone, and a queued close would end the next attempt instead. + Discard(Actions); + if (PendingRetryAfter) + { + RetryAfterUntil = std::max(RetryAfterUntil, now + *PendingRetryAfter); + PendingRetryAfter.reset(); + } + if (BackoffOnFailure) + { + BackoffUntil = Reconnect.NextAttempt(now); + BackoffOnFailure = false; + } + State = ConnectionState::Idle; +} + +void ProducerEndpoint::OnTick(TimePoint now) +{ + std::lock_guard lock(Mutex); + if (IsAttempting() && now >= AttemptDeadline) + { + Log(LogLevel::Warning, "handshake with " + Address(Config) + " timed out"); + Close(DisconnectReason::HandshakeTimeout); + } +} + +void ProducerEndpoint::Stop() +{ + std::lock_guard lock(Mutex); + if (State == ConnectionState::Stopped) + { + return; + } + Close(DisconnectReason::Stopped); + State = ConnectionState::Stopped; + Log(LogLevel::Info, "stopped"); +} + +std::vector ProducerEndpoint::TakeActions() +{ + std::lock_guard lock(Mutex); + return std::exchange(Actions, {}); +} + +std::optional ProducerEndpoint::TakeIssuedToken() +{ + std::lock_guard lock(Mutex); + return std::exchange(IssuedToken, std::nullopt); +} + +std::optional ProducerEndpoint::NextWake() const +{ + std::lock_guard lock(Mutex); + return IsAttempting() ? std::optional(AttemptDeadline) : std::nullopt; +} + +std::uint64_t ProducerEndpoint::NextTransactionId() const noexcept +{ + return Session.NextTransactionId(); +} + +ProducerResult ProducerEndpoint::Append(std::uint64_t transaction_id, + std::vector frame, std::size_t event_count, + std::string layer_key) +{ + if (!IsCompleteFrame(frame)) + { + return ProducerResult::InvalidArgument; + } + SharedByteBuffer payload = Share(std::move(frame)); + std::lock_guard lock(Mutex); + // The session is Ready exactly while connected, so it refuses every other state. + const ProducerResult result = + Session.Append(ConnectionGeneration, transaction_id, std::move(payload), event_count, + std::move(layer_key)); + if (result == ProducerResult::Accepted) + { + SendUnsent(); + } + return result; +} + +bool ProducerEndpoint::QueueControl(std::vector frame) +{ + if (!IsCompleteFrame(frame)) + { + return false; + } + SharedByteBuffer payload = Share(std::move(frame)); + std::lock_guard lock(Mutex); + if (State != ConnectionState::Connected) + { + return false; + } + Actions.push_back(SendAction{std::move(payload)}); + return true; +} + +bool ProducerEndpoint::OutboxEmpty() const noexcept +{ + return Session.Empty(); +} + +std::uint64_t ProducerEndpoint::DrainAcknowledgedEventCount() noexcept +{ + return Session.DrainAcknowledgedEventCount(); +} + +std::optional ProducerEndpoint::AcknowledgedCheckpoint() const +{ + std::lock_guard lock(Mutex); + if (SessionFailure || !Session.Empty()) + { + return std::nullopt; + } + return Checkpoint; +} + +std::optional ProducerEndpoint::Failure() const +{ + std::lock_guard lock(Mutex); + return SessionFailure; +} + +std::optional ProducerEndpoint::Artifact() const +{ + std::lock_guard lock(Mutex); + if (!SessionFailure) + { + return std::nullopt; + } + return CaptureArtifact(); +} + +ProducerResult ProducerEndpoint::RepairRejected(std::vector frame, + std::size_t event_count, std::string layer_key) +{ + if (!IsCompleteFrame(frame)) + { + return ProducerResult::InvalidArgument; + } + SharedByteBuffer payload = Share(std::move(frame)); + std::lock_guard lock(Mutex); + const ProducerResult result = + Session.RepairRejected(std::move(payload), event_count, std::move(layer_key)); + if (result == ProducerResult::Accepted) + { + Log(LogLevel::Info, + "repaired transaction " + std::to_string(SessionFailure->TransactionId)); + SessionFailure.reset(); + RetryAfterUntil = {}; + ResetBackoff(); + } + return result; +} + +std::optional ProducerEndpoint::AbandonRejectedSession(std::string session_id) +{ + if (!IsValidProducerSessionId(session_id)) + { + return std::nullopt; + } + std::lock_guard lock(Mutex); + if (!SessionFailure || session_id == SessionId) + { + return std::nullopt; + } + RecoveryArtifact artifact = CaptureArtifact(); + Log(LogLevel::Info, "abandoned producer session " + SessionId + " for " + session_id); + Session.ResetSession(); + SessionId = std::move(session_id); + SessionFailure.reset(); + RetryAfterUntil = {}; + ResetBackoff(); + return artifact; +} + +ProducerStatus ProducerEndpoint::Status() const +{ + std::lock_guard lock(Mutex); + ProducerStatus status; + status.Connected = State == ConnectionState::Connected; + status.Handshaking = IsAttempting(); + status.Closing = State == ConnectionState::Closing; + status.Stopped = State == ConnectionState::Stopped; + status.Rejection = Rejection; + status.LayerModeActive = LayerModeActive; + status.Metadata = Metadata; + status.SessionId = SessionId; + status.PendingTransactions = Session.PendingTransactionCount(); + status.PendingEvents = Session.PendingEventCount(); + status.AcknowledgedTransactions = Session.AcknowledgedTransactionCount(); + status.AcknowledgedEvents = Session.AcknowledgedEventCount(); + status.NextTransactionId = Session.NextTransactionId(); + status.Failure = SessionFailure; + status.RetryAfter = RetryAfterUntil; + return status; +} + +bool ProducerEndpoint::IsAttempting() const noexcept +{ + return State == ConnectionState::Connecting || State == ConnectionState::Handshaking; +} + +bool ProducerEndpoint::IsOpen() const noexcept +{ + return State == ConnectionState::Handshaking || State == ConnectionState::Connected; +} + +void ProducerEndpoint::BeginAttempt(TimePoint now, TimePoint deadline, bool backoff_on_failure) +{ + State = ConnectionState::Connecting; + AttemptDeadline = std::min(deadline, now + Config.HandshakeTimeout); + BackoffOnFailure = backoff_on_failure; + Log(LogLevel::Info, "connecting to " + Address(Config)); + Actions.push_back(ConnectAction{Config.Host, Config.Port, AttemptDeadline}); +} + +void ProducerEndpoint::ResetBackoff() noexcept +{ + Reconnect.Reset(); + BackoffUntil = {}; + BackoffOnFailure = false; +} + +void ProducerEndpoint::Close(DisconnectReason reason) +{ + if (State != ConnectionState::Connecting && !IsOpen()) + { + return; + } + // The host ending a published connection says Quit first. + if (State == ConnectionState::Connected && + (reason == DisconnectReason::Stopped || reason == DisconnectReason::Cancelled)) + { + Actions.push_back(SendAction{BuildQuitFrame()}); + } + EndConnection(reason); + // An attempt the host has not taken ends here, so no close can follow it. + if (State == ConnectionState::Connecting && Discard(Actions)) + { + BackoffOnFailure = false; + State = ConnectionState::Idle; + return; + } + Actions.push_back(CloseAction{reason}); + State = ConnectionState::Closing; +} + +void ProducerEndpoint::EndConnection(DisconnectReason reason) +{ + if (State == ConnectionState::Connected) + { + Notify(Disconnected{reason}); + } + if (IsOpen()) + { + [[maybe_unused]] const ProducerResult ended = Session.Disconnect(ConnectionGeneration); + assert(ended == ProducerResult::Accepted); + } +} + +void ProducerEndpoint::HandleFrame(const std::vector& frame) +{ + EnvelopeView envelope; + const ProtocolResult decoded = DecodeEnvelope(frame.data(), frame.size(), envelope); + if (decoded != ProtocolResult::Success) + { + Log(LogLevel::Error, detail::DescribeDecodeFailure(decoded)); + Close(DisconnectReason::ProtocolError); + return; + } + if (State == ConnectionState::Handshaking) + { + HandleHandshake(envelope); + } + else + { + HandleMessage(envelope); + } +} + +void ProducerEndpoint::HandleHandshake(EnvelopeView envelope) +{ + const detail::HandshakeOutcome outcome = detail::ClassifyHandshake(envelope); + if (outcome.Accepted) + { + AcceptHello(*outcome.Accepted); + } + else if (outcome.Rejection) + { + Reject(*outcome.Rejection); + } + else + { + Log(LogLevel::Error, "unexpected handshake response"); + Close(DisconnectReason::ProtocolError); + } +} + +void ProducerEndpoint::AcceptHello(const OpenUSDConnect::HelloOk& hello) +{ + ServerInstance = Text(hello.server_instance()); + Checkpoint.reset(); + if (hello.layer_mode() != Config.LayerMode) + { + Reject({false, OpenUSDConnect::HelloRejectionCode::LayerModeMismatch, + std::string("server negotiated ") + .append(LayerModeName(hello.layer_mode())) + .append(" instead of ") + .append(LayerModeName(Config.LayerMode))}); + return; + } + LayerModeActive = hello.layer_mode(); + const std::uint64_t committed_through = hello.committed_through(); + const ProducerResult accepted = Session.AcceptHello(ConnectionGeneration, committed_through); + if (accepted != ProducerResult::Accepted) + { + Fail({committed_through, kUnexpectedId, + HighwaterFailureReason(accepted, committed_through)}); + return; + } + detail::NotifyHelloFields(hello, IssuedToken, Metadata, Notifications, Actions); + Publish(); +} + +void ProducerEndpoint::Reject(HandshakeRejected rejection) +{ + Log(LogLevel::Error, detail::DescribeRejection(rejection)); + Notify(rejection); + Rejection = std::move(rejection); + Close(DisconnectReason::HandshakeRejected); +} + +void ProducerEndpoint::Publish() +{ + State = ConnectionState::Connected; + ResetBackoff(); + Notify(Connected{}); + const std::size_t replayed = SendUnsent(); + Log(LogLevel::Info, "connected to " + Address(Config) + " (session=" + SessionId + + ", pending=" + std::to_string(replayed) + ")"); +} + +std::size_t ProducerEndpoint::SendUnsent() +{ + std::size_t sent = 0; + ProducerSessionEntry entry{}; + while (Session.ClaimNextUnsent(ConnectionGeneration, entry) == ProducerResult::Accepted) + { + Actions.emplace_back(SendAction{std::move(entry.Payload)}); + ++sent; + } + return sent; +} + +void ProducerEndpoint::HandleMessage(EnvelopeView envelope) +{ + // Other messages, such as playback replies to control frames, need no action. + const ControlMessageView message(envelope); + if (message.Kind() == ControlMessageKind::TransactionResult) + { + AcceptResult(*message.TransactionResult()); + } + else if (message.Kind() == ControlMessageKind::RateLimited) + { + AcceptRateLimit(*message.RateLimit()); + } +} + +void ProducerEndpoint::AcceptResult(const OpenUSDConnect::TransactionResult& result) +{ + const std::uint64_t transaction_id = result.txn_id(); + if (result.status() == OpenUSDConnect::TransactionStatus::Acknowledged) + { + const ProducerResult accepted = + Session.AcknowledgeThrough(ConnectionGeneration, transaction_id); + if (accepted != ProducerResult::Accepted) + { + Fail({transaction_id, kUnexpectedId, HighwaterFailureReason(accepted, transaction_id)}); + return; + } + const OpenUSDConnect::TransactionCheckpoint* checkpoint = result.checkpoint(); + if (checkpoint && !ServerInstance.empty()) + { + Checkpoint = + MirrorCheckpoint{ServerInstance, checkpoint->epoch(), checkpoint->head_seq()}; + } + else + { + Checkpoint.reset(); + } + return; + } + const auto code = static_cast(result.rejection_code()); + TransactionFailure failure{transaction_id, code, Text(result.reason()), + result.expected_txn_id()}; + const ProducerResult rejected = + Session.Reject(ConnectionGeneration, transaction_id, RejectionDisposition(code)); + assert(rejected == ProducerResult::Accepted || rejected == ProducerResult::TransactionMissing); + if (rejected == ProducerResult::TransactionMissing) + { + failure = {transaction_id, kUnexpectedId, + "server rejected unknown transaction " + std::to_string(transaction_id)}; + } + Fail(std::move(failure)); +} + +void ProducerEndpoint::AcceptRateLimit(const OpenUSDConnect::RateLimited& limited) +{ + PendingRetryAfter = RetryAfter(limited.retry_after()); + const auto wait = std::chrono::duration_cast(*PendingRetryAfter); + Log(LogLevel::Warning, + "rate limited by the server; retrying after " + std::to_string(wait.count()) + " ms"); + Close(DisconnectReason::RateLimited); +} + +void ProducerEndpoint::Fail(TransactionFailure failure) +{ + assert(Session.RecoveryRequired()); + Log(LogLevel::Error, failure.Describe()); + SessionFailure = std::move(failure); + Close(DisconnectReason::RecoveryRequired); +} + +std::string ProducerEndpoint::HighwaterFailureReason(ProducerResult result, + std::uint64_t transaction_id) const +{ + assert(result == ProducerResult::HighwaterAhead || + result == ProducerResult::HighwaterRegressed); + if (result == ProducerResult::HighwaterAhead) + { + return "server producer highwater " + std::to_string(transaction_id) + + " is ahead of local transaction " + std::to_string(Session.NextTransactionId() - 1); + } + return "server producer highwater regressed from " + + std::to_string(Session.LastAcknowledgedTransactionId()) + " to " + + std::to_string(transaction_id); +} + +RecoveryArtifact ProducerEndpoint::CaptureArtifact() const +{ + return {SessionId, *SessionFailure, Session.Entries()}; +} + +void ProducerEndpoint::Notify(Notification notification) +{ + Notifications.Push(std::move(notification)); +} + +void ProducerEndpoint::Log(LogLevel level, std::string message) +{ + Actions.push_back(LogAction{level, std::move(message)}); +} + +} // namespace openusdconnect::client diff --git a/native/client_core/src/engine/receiver_endpoint.cpp b/native/client_core/src/engine/receiver_endpoint.cpp new file mode 100644 index 0000000..76bde04 --- /dev/null +++ b/native/client_core/src/engine/receiver_endpoint.cpp @@ -0,0 +1,580 @@ +#include "openusdconnect/client/engine/receiver_endpoint.h" + +#include "endpoint_common.h" + +#include +#include +#include + +namespace openusdconnect::client +{ +namespace +{ + +using detail::Address; +using detail::Discard; +using detail::Text; +using detail::Value; + +// The consumer drains on its own thread, so an overflowed receiver polls for +// the empty queue. +constexpr std::chrono::milliseconds kDrainPollInterval{100}; + +[[nodiscard]] std::vector BuildResyncPayload() +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto resync = OpenUSDConnect::CreateResync(builder); + OpenUSDConnect::FinishEnvelopeBuffer( + builder, OpenUSDConnect::CreateEnvelope(builder, OpenUSDConnect::Payload::Resync, + resync.Union(), kSchemaVersion)); + const std::uint8_t* bytes = builder.GetBufferPointer(); + return {bytes, bytes + builder.GetSize()}; +} + +[[nodiscard]] bool IsKnownPayload(OpenUSDConnect::Payload type) noexcept +{ + return type != OpenUSDConnect::Payload::NONE && type <= OpenUSDConnect::Payload::MAX; +} + +} // namespace + +ReceiverEndpoint::ReceiverEndpoint(ReceiverConfig config, NotificationQueue& notifications) + : Config(std::move(config)) + , Notifications(notifications) + , Inbox(Config.SyncFrom, Config.MaxQueue, true) + , Reconnect(Config.Reconnect, Config.ReconnectBaseDelay, Config.ReconnectMaxDelay) +{ + assert(IsValidConfiguration(Config)); +} + +bool ReceiverEndpoint::IsValidConfiguration(const ReceiverConfig& config) noexcept +{ + const bool valid_mode = config.LayerMode == OpenUSDConnect::LayerMode::Managed || + (config.LayerMode == OpenUSDConnect::LayerMode::SharedStage && + !config.LayeredReplay && config.Department.empty()); + return valid_mode && !config.Host.empty() && config.Port != 0 && + ReceiverInbox::IsValidConfiguration(config.SyncFrom, config.MaxQueue) && + config.SocketTimeout.count() > 0 && config.MaxConsecutiveTimeouts != 0 && + ReconnectPolicy::IsValidConfiguration(config.ReconnectBaseDelay, + config.ReconnectMaxDelay); +} + +const ReceiverConfig& ReceiverEndpoint::Configuration() const noexcept +{ + return Config; +} + +void ReceiverEndpoint::SetReconnect(bool enabled) +{ + std::lock_guard lock(Mutex); + Reconnect.SetEnabled(enabled); +} + +bool ReceiverEndpoint::Start(TimePoint now) +{ + std::lock_guard lock(Mutex); + if (State != ConnectionState::Idle) + { + return false; + } + BeginAttempt(now); + return true; +} + +void ReceiverEndpoint::OnConnected(std::string_view token) +{ + std::lock_guard lock(Mutex); + if (State == ConnectionState::Stopped) + { + Actions.push_back(CloseAction{DisconnectReason::Stopped}); + return; + } + if (State != ConnectionState::Connecting) + { + return; + } + const ConnectionStart connection = Inbox.BeginConnection(); + ConnectionGeneration = connection.Generation; + ConnectionSyncFrom = connection.SyncFrom; + ConsecutiveTimeouts = 0; + Decoder.Reset(); + Rejection.reset(); + State = ConnectionState::Handshaking; + + HelloParameters hello = detail::CommonHello("receiver", Config, token); + hello.SyncFrom = ConnectionSyncFrom; + hello.LayeredReplay = Config.LayeredReplay; + hello.ReplayPrefix = Identity.BeginConnection(); + if (!detail::QueueHello(hello, Actions)) + { + Close(DisconnectReason::ProtocolError); + } +} + +void ReceiverEndpoint::OnBytes(const std::uint8_t* data, std::size_t size) +{ + std::lock_guard lock(Mutex); + if (!IsOpen()) + { + return; + } + if (size != 0) + { + ConsecutiveTimeouts = 0; + } + const bool framed = detail::HandleFrames(Decoder, data, size, + [this](std::vector& frame) + { + HandleFrame(std::move(frame)); + return IsOpen(); + }); + if (!framed) + { + Log(LogLevel::Warning, "invalid frame header"); + Close(DisconnectReason::ProtocolError); + } +} + +void ReceiverEndpoint::OnReadTimeout() +{ + std::lock_guard lock(Mutex); + if (!IsOpen()) + { + return; + } + ++ConsecutiveTimeouts; + const std::string count = std::to_string(ConsecutiveTimeouts); + if (ConsecutiveTimeouts < Config.MaxConsecutiveTimeouts) + { + Log(LogLevel::Debug, + "read timeout " + count + "/" + std::to_string(Config.MaxConsecutiveTimeouts)); + return; + } + Log(LogLevel::Warning, count + " consecutive read timeouts, reconnecting"); + Close(DisconnectReason::ReadTimeout); +} + +void ReceiverEndpoint::OnDisconnected(DisconnectReason reason, TimePoint now) +{ + std::lock_guard lock(Mutex); + if (IsOpen()) + { + Log(LogLevel::Warning, "connection to " + Address(Config) + " lost"); + EndConnection(reason); + } + else if (State != ConnectionState::Connecting && State != ConnectionState::Closing) + { + return; + } + ScheduleNextAttempt(now); +} + +void ReceiverEndpoint::OnTick(TimePoint now) +{ + std::lock_guard lock(Mutex); + if (State == ConnectionState::Backoff && now >= WakeTime) + { + BeginAttempt(now); + } + else if (State == ConnectionState::DrainWait) + { + PollDrain(now); + } +} + +void ReceiverEndpoint::Stop() +{ + std::lock_guard lock(Mutex); + if (State == ConnectionState::Stopped) + { + return; + } + if (State == ConnectionState::Connecting || IsOpen()) + { + Close(DisconnectReason::Stopped); + } + State = ConnectionState::Stopped; + Log(LogLevel::Info, "stopped"); +} + +std::vector ReceiverEndpoint::TakeActions() +{ + std::lock_guard lock(Mutex); + return std::exchange(Actions, {}); +} + +std::optional ReceiverEndpoint::TakeIssuedToken() +{ + std::lock_guard lock(Mutex); + return std::exchange(IssuedToken, std::nullopt); +} + +std::optional ReceiverEndpoint::NextWake() const +{ + std::lock_guard lock(Mutex); + if (State == ConnectionState::Backoff || State == ConnectionState::DrainWait) + { + return WakeTime; + } + return std::nullopt; +} + +std::vector> +ReceiverEndpoint::DrainFrames(std::optional max_frames) +{ + return Inbox.Drain(max_frames); +} + +std::uint64_t ReceiverEndpoint::Generation() const noexcept +{ + return Inbox.Generation(); +} + +bool ReceiverEndpoint::MarkAppliedThrough(std::uint64_t generation, std::int32_t sequence) +{ + return Inbox.MarkAppliedThrough(generation, sequence); +} + +void ReceiverEndpoint::ResetAppliedProgress() noexcept +{ + Inbox.ResetAppliedProgress(); +} + +bool ReceiverEndpoint::MarkReplayApplied() +{ + std::lock_guard lock(Mutex); + // Every drained frame applied, so the replay head counts as applied even + // when a reconnect rejected the batch's MarkAppliedThrough. + if (!Inbox.MarkReplayApplied()) + { + return false; + } + Identity.MarkReplayApplied(); + return true; +} + +bool ReceiverEndpoint::RequestReplayFrom(std::int32_t sequence) +{ + if (sequence < 1) + { + return false; + } + std::lock_guard lock(Mutex); + RequestReplay(sequence, DisconnectReason::ReplayRequested); + return true; +} + +std::uint64_t ReceiverEndpoint::FreezeMarker() const noexcept +{ + return Inbox.FreezeMarker(); +} + +bool ReceiverEndpoint::DrainedThrough(std::uint64_t marker) const noexcept +{ + return Inbox.DrainedThrough(marker); +} + +ReceiverStatus ReceiverEndpoint::Status() const +{ + std::lock_guard lock(Mutex); + ReceiverStatus status; + status.Connected = State == ConnectionState::Connected; + status.Synchronized = status.Connected && Inbox.Synchronized(); + status.Stopped = State == ConnectionState::Stopped; + status.ReplayHeadSequence = Inbox.ReplayHeadSequence(); + status.ReplayEpoch = Inbox.ReplayEpoch(); + if (const std::optional& applied = Identity.Applied()) + { + status.ServerInstance = applied->ServerInstance; + } + status.LayeredReplayActive = LayeredReplayActive; + status.LayerModeActive = LayerModeActive; + status.Rejection = Rejection; + status.Metadata = Metadata; + status.QueuedFrames = Inbox.Size(); + status.LastSequence = Inbox.LastSequence(); + status.LastAppliedSequence = Inbox.LastAppliedSequence(); + return status; +} + +bool ReceiverEndpoint::IsOpen() const noexcept +{ + return State == ConnectionState::Handshaking || State == ConnectionState::Connected; +} + +void ReceiverEndpoint::BeginAttempt(TimePoint now) +{ + State = ConnectionState::Connecting; + Log(LogLevel::Info, "connecting to " + Address(Config)); + Actions.push_back(ConnectAction{Config.Host, Config.Port, now + Config.SocketTimeout}); +} + +void ReceiverEndpoint::ScheduleNextAttempt(TimePoint now) +{ + if (Rejection || !Reconnect.Enabled()) + { + State = ConnectionState::Stopped; + Log(LogLevel::Info, "stopped"); + return; + } + if (Inbox.Overflowed()) + { + Inbox.ClearOverflow(); + DrainDeadline = Reconnect.DrainDeadline(now); + WakeTime = now; + State = ConnectionState::DrainWait; + Log(LogLevel::Info, "waiting for the queue to drain before reconnecting"); + PollDrain(now); + return; + } + WakeTime = Reconnect.NextAttempt(now); + State = ConnectionState::Backoff; + const auto delay = std::chrono::duration_cast(WakeTime - now); + Log(LogLevel::Info, "reconnecting in " + std::to_string(delay.count()) + " ms"); +} + +void ReceiverEndpoint::PollDrain(TimePoint now) +{ + if (Inbox.Size() == 0) + { + BeginAttempt(now); + } + else if (now >= DrainDeadline) + { + Log(LogLevel::Warning, "drain wait timed out, reconnecting anyway"); + BeginAttempt(now); + } + else if (now >= WakeTime) + { + WakeTime = std::min(now + kDrainPollInterval, DrainDeadline); + } +} + +void ReceiverEndpoint::Close(DisconnectReason reason) +{ + EndConnection(reason); + // An attempt the host has not taken ends here, so no close can follow it. + if (State == ConnectionState::Connecting && Discard(Actions)) + { + State = ConnectionState::Idle; + return; + } + Actions.push_back(CloseAction{reason}); + State = ConnectionState::Closing; +} + +void ReceiverEndpoint::EndConnection(DisconnectReason reason) +{ + if (State == ConnectionState::Connected) + { + Notify(Disconnected{reason}); + } + if (IsOpen()) + { + Inbox.Disconnect(ConnectionGeneration); + } +} + +void ReceiverEndpoint::RequestReplay(std::int32_t sequence, DisconnectReason reason) +{ + Identity.RequestReplayFrom(sequence, Inbox.ResetPending()); + [[maybe_unused]] const bool requested = Inbox.RequestReplayFrom(sequence); + assert(requested); + if (IsOpen()) + { + Close(reason); + } +} + +void ReceiverEndpoint::HandleFrame(std::vector frame) +{ + EnvelopeView envelope; + const ProtocolResult decoded = DecodeEnvelope(frame.data(), frame.size(), envelope); + if (decoded != ProtocolResult::Success || !IsKnownPayload(envelope.PayloadType())) + { + Log(LogLevel::Error, detail::DescribeDecodeFailure(decoded)); + Close(DisconnectReason::ProtocolError); + return; + } + if (State == ConnectionState::Handshaking) + { + HandleHandshake(envelope); + } + else + { + HandleMessage(envelope, frame); + } +} + +void ReceiverEndpoint::HandleHandshake(EnvelopeView envelope) +{ + const detail::HandshakeOutcome outcome = detail::ClassifyHandshake(envelope); + if (outcome.Accepted) + { + AcceptHello(*outcome.Accepted); + } + else if (outcome.Rejection) + { + Reject(*outcome.Rejection); + } + else + { + Log(LogLevel::Error, "unexpected handshake response"); + Close(DisconnectReason::ProtocolError); + } +} + +void ReceiverEndpoint::AcceptHello(const OpenUSDConnect::HelloOk& hello) +{ + LayerModeActive = hello.layer_mode(); + if (LayerModeActive != Config.LayerMode) + { + Reject({false, OpenUSDConnect::HelloRejectionCode::LayerModeMismatch, + "server did not negotiate requested layer mode"}); + return; + } + LayeredReplayActive = Config.LayeredReplay && hello.layered_replay(); + if (Config.LayeredReplay && !LayeredReplayActive) + { + Reject({false, OpenUSDConnect::HelloRejectionCode::LayeredReplayRequired, + "server did not negotiate requested layered replay"}); + return; + } + detail::NotifyHelloFields(hello, IssuedToken, Metadata, Notifications, Actions); + + Identity.AcceptHello(ConnectionSyncFrom, hello.replay_identity(), Text(hello.server_instance()), + Value(hello.replay_epoch())); + if (Identity.ResetRequired()) + { + assert(ConnectionSyncFrom == 1); + AcceptFrame(ReceiverMessageKind::Resync, 0, BuildResyncPayload()); + if (!IsOpen()) + { + return; + } + } + State = ConnectionState::Connected; + Reconnect.Reset(); + Notify(Connected{}); + Log(LogLevel::Info, "connected (sync_from=" + std::to_string(ConnectionSyncFrom) + ")"); +} + +void ReceiverEndpoint::Reject(HandshakeRejected rejection) +{ + Log(LogLevel::Error, detail::DescribeRejection(rejection)); + Notify(rejection); + Rejection = std::move(rejection); + Close(DisconnectReason::HandshakeRejected); +} + +void ReceiverEndpoint::HandleMessage(EnvelopeView envelope, std::vector& frame) +{ + const OpenUSDConnect::Envelope& message = *envelope.Get(); + ReceiverMessageKind kind = ReceiverMessageKind::Other; + std::int32_t sequence = 0; + switch (message.payload_type()) + { + case OpenUSDConnect::Payload::Ping: + return; + case OpenUSDConnect::Payload::ReplayComplete: + AcceptReplayComplete(*message.payload_as_ReplayComplete()); + return; + case OpenUSDConnect::Payload::PlaybackState: + { + const OpenUSDConnect::PlaybackState& state = *message.payload_as_PlaybackState(); + Notify(PlaybackState{state.time(), state.playing(), state.rate(), + Text(state.leader_client_id())}); + return; + } + case OpenUSDConnect::Payload::PlaybackClaimed: + Notify(PlaybackClaimed{Text(message.payload_as_PlaybackClaimed()->leader_client_id())}); + return; + case OpenUSDConnect::Payload::PlaybackRejected: + { + const OpenUSDConnect::PlaybackRejected& rejected = *message.payload_as_PlaybackRejected(); + Notify( + PlaybackRejected{Text(rejected.reason()), Text(rejected.current_leader_client_id())}); + return; + } + case OpenUSDConnect::Payload::BroadcastEvent: + kind = ReceiverMessageKind::Event; + sequence = message.payload_as_BroadcastEvent()->seq(); + break; + case OpenUSDConnect::Payload::LayerGraphState: + kind = ReceiverMessageKind::LayerGraphState; + sequence = message.payload_as_LayerGraphState()->seq(); + break; + case OpenUSDConnect::Payload::Resync: + kind = ReceiverMessageKind::Resync; + break; + default: + break; + } + AcceptFrame(kind, sequence, std::move(frame)); +} + +void ReceiverEndpoint::AcceptReplayComplete(const OpenUSDConnect::ReplayComplete& complete) +{ + const AcceptResult result = + Inbox.AcceptReplayComplete(ConnectionGeneration, complete.head_seq(), complete.epoch()); + if (result == AcceptResult::Accepted) + { + Identity.AcceptReplayComplete(complete.epoch()); + } + else if (result == AcceptResult::SequenceGap) + { + ReplayAfterGap("replay head " + std::to_string(complete.head_seq()) + " was not received"); + } + else if (result == AcceptResult::InvalidSequence) + { + Log(LogLevel::Warning, + "ignoring ReplayComplete with invalid head " + std::to_string(complete.head_seq())); + } +} + +void ReceiverEndpoint::AcceptFrame(ReceiverMessageKind kind, std::int32_t sequence, + std::vector frame) +{ + switch (Inbox.Accept(ConnectionGeneration, kind, sequence, std::move(frame))) + { + case AcceptResult::Accepted: + if (kind == ReceiverMessageKind::Resync) + { + Identity.AcceptResync(); + } + return; + case AcceptResult::StaleGeneration: + case AcceptResult::Duplicate: + return; + case AcceptResult::InvalidSequence: + Log(LogLevel::Warning, "ignoring frame with invalid sequence " + std::to_string(sequence)); + return; + case AcceptResult::SequenceGap: + ReplayAfterGap("sequence gap before " + std::to_string(sequence)); + return; + case AcceptResult::QueueFull: + Log(LogLevel::Warning, "queue full (" + std::to_string(Config.MaxQueue) + + "), disconnecting to replay from server"); + Close(DisconnectReason::QueueFull); + return; + } +} + +void ReceiverEndpoint::ReplayAfterGap(const std::string& description) +{ + const std::int32_t replay_from = Inbox.LastAppliedSequence() + 1; + Log(LogLevel::Error, description + "; replaying from applied " + std::to_string(replay_from)); + RequestReplay(replay_from, DisconnectReason::SequenceGap); +} + +void ReceiverEndpoint::Notify(Notification notification) +{ + Notifications.Push(std::move(notification)); +} + +void ReceiverEndpoint::Log(LogLevel level, std::string message) +{ + Actions.push_back(LogAction{level, std::move(message)}); +} + +} // namespace openusdconnect::client diff --git a/native/client_core/src/platform/bsd_socket.cpp b/native/client_core/src/platform/bsd_socket.cpp new file mode 100644 index 0000000..a018881 --- /dev/null +++ b/native/client_core/src/platform/bsd_socket.cpp @@ -0,0 +1,652 @@ +#include "openusdconnect/client/driver/socket.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +#ifdef _WIN32 +#ifndef WIN32_LEAN_AND_MEAN +#define WIN32_LEAN_AND_MEAN +#endif +#ifndef NOMINMAX +#define NOMINMAX +#endif +#include +#include +#else +#include +#include +#include +#include +#include +#include +#include +#include +#endif + +namespace openusdconnect::client +{ +namespace +{ + +#ifdef _WIN32 +using NativeSocket = SOCKET; +using SocketLength = int; +constexpr NativeSocket kInvalidSocket = INVALID_SOCKET; +constexpr int kShutdownBoth = SD_BOTH; +constexpr int kSendFlags = 0; + +[[nodiscard]] int LastSocketError() noexcept +{ + return WSAGetLastError(); +} + +[[nodiscard]] bool WouldBlock(int error) noexcept +{ + return error == WSAEWOULDBLOCK; +} + +[[nodiscard]] bool ConnectPending(int error) noexcept +{ + return error == WSAEWOULDBLOCK; +} + +[[nodiscard]] bool Retry(int) noexcept +{ + return false; +} + +[[nodiscard]] int ResolutionError(int result) noexcept +{ + return result; +} + +void CloseNative(NativeSocket handle) noexcept +{ + closesocket(handle); +} + +// A manual-reset event the blocking calls wait on beside the socket. +class Signal final +{ +public: + Signal() noexcept + : Event(WSACreateEvent()) + , CreationError(Event == WSA_INVALID_EVENT ? WSAGetLastError() : 0) + { + } + + ~Signal() + { + if (Event != WSA_INVALID_EVENT) + { + WSACloseEvent(Event); + } + } + + Signal(const Signal&) = delete; + Signal& operator=(const Signal&) = delete; + + [[nodiscard]] int Error() const noexcept + { + return CreationError; + } + + void Raise() noexcept + { + WSASetEvent(Event); + } + + void Clear() noexcept + { + WSAResetEvent(Event); + } + + [[nodiscard]] WSAEVENT Handle() const noexcept + { + return Event; + } + +private: + const WSAEVENT Event; + const int CreationError; +}; +#else +using NativeSocket = int; +using SocketLength = socklen_t; +constexpr NativeSocket kInvalidSocket = -1; +constexpr int kShutdownBoth = SHUT_RDWR; +#ifdef MSG_NOSIGNAL +constexpr int kSendFlags = MSG_NOSIGNAL; +#else +constexpr int kSendFlags = 0; +#endif + +[[nodiscard]] int LastSocketError() noexcept +{ + return errno; +} + +[[nodiscard]] bool WouldBlock(int error) noexcept +{ + return error == EAGAIN || error == EWOULDBLOCK; +} + +[[nodiscard]] bool ConnectPending(int error) noexcept +{ + return error == EINPROGRESS; +} + +[[nodiscard]] bool Retry(int error) noexcept +{ + return error == EINTR; +} + +[[nodiscard]] int ResolutionError(int result) noexcept +{ + return result == EAI_SYSTEM ? errno : EHOSTUNREACH; +} + +void CloseNative(NativeSocket handle) noexcept +{ + close(handle); +} + +[[nodiscard]] bool ConfigureDescriptor(int descriptor) noexcept +{ + return fcntl(descriptor, F_SETFD, FD_CLOEXEC) == 0 && + fcntl(descriptor, F_SETFL, fcntl(descriptor, F_GETFL) | O_NONBLOCK) == 0; +} + +// A self-pipe the blocking calls poll beside the socket. +class Signal final +{ +public: + Signal() noexcept + { + int descriptors[2] = {-1, -1}; + if (pipe(descriptors) != 0) + { + CreationError = errno; + return; + } + Read = descriptors[0]; + Write = descriptors[1]; + if (!ConfigureDescriptor(Read) || !ConfigureDescriptor(Write)) + { + CreationError = errno; + } + } + + ~Signal() + { + for (const int descriptor : {Read, Write}) + { + if (descriptor >= 0) + { + close(descriptor); + } + } + } + + Signal(const Signal&) = delete; + Signal& operator=(const Signal&) = delete; + + [[nodiscard]] int Error() const noexcept + { + return CreationError; + } + + void Raise() noexcept + { + const char byte = 0; + // A full pipe is already raised. + [[maybe_unused]] const ssize_t written = write(Write, &byte, 1); + } + + void Clear() noexcept + { + char bytes[64]; + while (read(Read, bytes, sizeof(bytes)) > 0) + { + } + } + + [[nodiscard]] int Handle() const noexcept + { + return Read; + } + +private: + int Read = -1; + int Write = -1; + int CreationError = 0; +}; +#endif + +enum class Readiness : std::uint8_t +{ + Readable, + Writable, +}; + +enum class WaitResult : std::uint8_t +{ + Ready, + Signaled, + TimedOut, + Failed, +}; + +[[nodiscard]] int RemainingMilliseconds(std::optional deadline) noexcept +{ + if (!deadline) + { + return -1; + } + const auto remaining = + std::chrono::ceil(*deadline - std::chrono::steady_clock::now()); + return static_cast( + std::clamp(remaining.count(), 0, INT_MAX - 1)); +} + +// Blocking calls wait on the socket and a Signal, so Interrupt and Wake end +// them on every platform; shutdown() does not wake a blocked Winsock call. +class BsdSocket final : public Socket +{ +public: + BsdSocket() noexcept +#ifdef _WIN32 + : NetworkEvent(WSACreateEvent()) +#endif + { + } + + ~BsdSocket() override + { + CloseNativeSocket(); +#ifdef _WIN32 + if (NetworkEvent != WSA_INVALID_EVENT) + { + WSACloseEvent(NetworkEvent); + } +#endif + } + + SocketResult Connect(const std::string& host, std::uint16_t port, TimePoint deadline) override + { + if (const int error = SetupError(); error != 0) + { + return Fail(error); + } + addrinfo hints{}; + hints.ai_family = AF_UNSPEC; + hints.ai_socktype = SOCK_STREAM; + hints.ai_protocol = IPPROTO_TCP; + addrinfo* addresses = nullptr; + const std::string service = std::to_string(port); + if (const int result = getaddrinfo(host.c_str(), service.c_str(), &hints, &addresses); + result != 0) + { + return Fail(ResolutionError(result)); + } + SocketResult result = SocketResult::Failed; + for (const addrinfo* address = addresses; address != nullptr; address = address->ai_next) + { + result = ConnectTo(*address, deadline); + if (result != SocketResult::Failed) + { + break; + } + } + freeaddrinfo(addresses); + return result; + } + + SocketResult SendAll(const std::uint8_t* data, std::size_t size, TimePoint deadline) override + { + while (size != 0) + { + if (Interrupted.load()) + { + return SocketResult::Interrupted; + } + const int chunk = static_cast(std::min(size, INT_MAX)); + const auto sent = send(Handle, reinterpret_cast(data), chunk, kSendFlags); + if (sent > 0) + { + data += sent; + size -= static_cast(sent); + continue; + } + const int error = LastSocketError(); + if (Retry(error)) + { + continue; + } + if (!WouldBlock(error)) + { + return Fail(error); + } + switch (Await(Readiness::Writable, deadline)) + { + case WaitResult::Ready: + continue; + case WaitResult::Signaled: + return SocketResult::Interrupted; + case WaitResult::TimedOut: + return SocketResult::Timeout; + case WaitResult::Failed: + return SocketResult::Failed; + } + } + return SocketResult::Success; + } + + SocketResult Receive(std::uint8_t* buffer, std::size_t capacity, + std::optional deadline, std::size_t& received) override + { + received = 0; + for (;;) + { + if (Interrupted.load() || WakePending.exchange(false)) + { + return SocketResult::Interrupted; + } + const int chunk = static_cast(std::min(capacity, INT_MAX)); + const auto count = recv(Handle, reinterpret_cast(buffer), chunk, 0); + if (count > 0) + { + received = static_cast(count); + return SocketResult::Success; + } + if (count == 0) + { + return SocketResult::Closed; + } + const int error = LastSocketError(); + if (Retry(error)) + { + continue; + } + if (!WouldBlock(error)) + { + return Fail(error); + } + switch (Await(Readiness::Readable, deadline)) + { + case WaitResult::Ready: + case WaitResult::Signaled: + continue; + case WaitResult::TimedOut: + return SocketResult::Timeout; + case WaitResult::Failed: + return SocketResult::Failed; + } + } + } + + int SystemError() const noexcept override + { + return Error; + } + + void Interrupt() noexcept override + { + Interrupted.store(true); + Wakeup.Raise(); + } + + void Wake() noexcept override + { + WakePending.store(true); + Wakeup.Raise(); + } + +private: + [[nodiscard]] SocketResult Fail(int error) noexcept + { + Error = error; + return SocketResult::Failed; + } + + [[nodiscard]] int SetupError() const noexcept + { +#ifdef _WIN32 + if (NetworkEvent == WSA_INVALID_EVENT) + { + return WSA_INVALID_HANDLE; + } +#endif + return Wakeup.Error(); + } + + [[nodiscard]] SocketResult ConnectTo(const addrinfo& address, TimePoint deadline) + { + CloseNativeSocket(); + if (const int error = Open(address.ai_family); error != 0) + { + return Fail(error); + } + if (connect(Handle, address.ai_addr, static_cast(address.ai_addrlen)) == 0) + { + return SocketResult::Success; + } + if (const int error = LastSocketError(); !ConnectPending(error)) + { + return Fail(error); + } + for (;;) + { + switch (Await(Readiness::Writable, deadline)) + { + case WaitResult::Ready: + if (const std::optional error = ConnectResult()) + { + return *error == 0 ? SocketResult::Success : Fail(*error); + } + continue; + case WaitResult::Signaled: + return SocketResult::Interrupted; + case WaitResult::TimedOut: + return SocketResult::Timeout; + case WaitResult::Failed: + return SocketResult::Failed; + } + } + } + + // Signaled means Interrupt, or a Wake while reading; a Wake during another + // call is left in WakePending for the next Receive. + [[nodiscard]] WaitResult Await(Readiness readiness, std::optional deadline) + { + for (;;) + { + const WaitResult result = Wait(readiness, deadline); + if (result != WaitResult::Signaled) + { + return result; + } + // The flags are set before raising, so reading them after clearing + // cannot miss a raise that the clear discarded. + Wakeup.Clear(); + if (Interrupted.load() || (readiness == Readiness::Readable && WakePending.load())) + { + return WaitResult::Signaled; + } + } + } + +#ifdef _WIN32 + [[nodiscard]] int Open(int family) noexcept + { + Handle = WSASocketW(family, SOCK_STREAM, IPPROTO_TCP, nullptr, 0, + WSA_FLAG_OVERLAPPED | WSA_FLAG_NO_HANDLE_INHERIT); + if (Handle == kInvalidSocket) + { + return WSAGetLastError(); + } + const BOOL enabled = TRUE; + if (setsockopt(Handle, IPPROTO_TCP, TCP_NODELAY, reinterpret_cast(&enabled), + sizeof(enabled)) != 0 || + WSAEventSelect(Handle, NetworkEvent, FD_CONNECT | FD_READ | FD_WRITE | FD_CLOSE) != 0) + { + return WSAGetLastError(); + } + return 0; + } + + // WSAEventSelect records every network event; callers retry their call. + [[nodiscard]] WaitResult Wait(Readiness, std::optional deadline) + { + const WSAEVENT events[] = {NetworkEvent, Wakeup.Handle()}; + const int timeout = RemainingMilliseconds(deadline); + const DWORD waited = WSAWaitForMultipleEvents( + 2, events, FALSE, timeout < 0 ? WSA_INFINITE : static_cast(timeout), FALSE); + if (waited == WSA_WAIT_TIMEOUT) + { + return WaitResult::TimedOut; + } + if (waited == WSA_WAIT_EVENT_0 + 1) + { + return WaitResult::Signaled; + } + WSANETWORKEVENTS network{}; + if (waited != WSA_WAIT_EVENT_0 || WSAEnumNetworkEvents(Handle, NetworkEvent, &network) != 0) + { + Error = WSAGetLastError(); + return WaitResult::Failed; + } + if ((network.lNetworkEvents & FD_CONNECT) != 0) + { + PendingConnectError = network.iErrorCode[FD_CONNECT_BIT]; + } + return WaitResult::Ready; + } + + [[nodiscard]] std::optional ConnectResult() noexcept + { + return std::exchange(PendingConnectError, std::nullopt); + } +#else + [[nodiscard]] int Open(int family) noexcept + { + Handle = socket(family, SOCK_STREAM, IPPROTO_TCP); + if (Handle == kInvalidSocket || !ConfigureDescriptor(Handle)) + { + return errno; + } + const int enabled = 1; + if (setsockopt(Handle, IPPROTO_TCP, TCP_NODELAY, &enabled, sizeof(enabled)) != 0) + { + return errno; + } +#ifdef SO_NOSIGPIPE + if (setsockopt(Handle, SOL_SOCKET, SO_NOSIGPIPE, &enabled, sizeof(enabled)) != 0) + { + return errno; + } +#endif + return 0; + } + + [[nodiscard]] WaitResult Wait(Readiness readiness, std::optional deadline) + { + pollfd descriptors[] = { + {Handle, static_cast(readiness == Readiness::Readable ? POLLIN : POLLOUT), 0}, + {Wakeup.Handle(), POLLIN, 0}, + }; + for (;;) + { + const int ready = poll(descriptors, 2, RemainingMilliseconds(deadline)); + if (ready > 0) + { + return descriptors[1].revents != 0 ? WaitResult::Signaled : WaitResult::Ready; + } + if (ready == 0) + { + return WaitResult::TimedOut; + } + if (errno != EINTR) + { + Error = errno; + return WaitResult::Failed; + } + } + } + + [[nodiscard]] std::optional ConnectResult() noexcept + { + int error = 0; + socklen_t length = sizeof(error); + if (getsockopt(Handle, SOL_SOCKET, SO_ERROR, &error, &length) != 0) + { + return errno; + } + return error; + } +#endif + + void CloseNativeSocket() noexcept + { + if (Handle == kInvalidSocket) + { + return; + } + shutdown(Handle, kShutdownBoth); + CloseNative(Handle); + Handle = kInvalidSocket; +#ifdef _WIN32 + PendingConnectError.reset(); + WSAResetEvent(NetworkEvent); +#endif + } + +#ifdef _WIN32 + const WSAEVENT NetworkEvent; + std::optional PendingConnectError; +#endif + Signal Wakeup; + NativeSocket Handle = kInvalidSocket; + int Error = 0; + std::atomic Interrupted{false}; + std::atomic WakePending{false}; +}; + +} // namespace + +std::string DescribeSystemError(int system_error) +{ + return std::system_category().message(system_error); +} + +std::string Describe(const TransportFailure& failure) +{ + return failure.Result == SocketResult::Timeout ? std::string("timed out") + : DescribeSystemError(failure.SystemError); +} + +TcpSocketFactory::TcpSocketFactory() +{ +#ifdef _WIN32 + // Winsock stays initialized for the life of the process. + static const int startup = [] + { + WSADATA data{}; + return WSAStartup(MAKEWORD(2, 2), &data); + }(); + static_cast(startup); +#endif +} + +std::unique_ptr TcpSocketFactory::Create() +{ + return std::make_unique(); +} + +} // namespace openusdconnect::client diff --git a/native/python/client_module.cpp b/native/python/client_module.cpp index cad29f3..8335ff9 100644 --- a/native/python/client_module.cpp +++ b/native/python/client_module.cpp @@ -1,122 +1,27 @@ -#include "openusdconnect/client/frame_codec.h" +#include "driver_bindings.h" + +#include "openusdconnect/client/engine/status.h" +#include "openusdconnect/client/producer_recovery.h" #include "openusdconnect/client/producer_session.h" -#include "openusdconnect/client/receiver_session.h" #include #include -#include - -#include -#include -#include -#include -#include -#include -#include +#include namespace nb = nanobind; using namespace nb::literals; -using openusdconnect::client::AcceptResult; -using openusdconnect::client::ConnectionStart; -using openusdconnect::client::FrameDecoder; -using openusdconnect::client::FrameResult; -using openusdconnect::client::ProducerConnectionStart; -using openusdconnect::client::ProducerPhase; +using openusdconnect::client::ClientPhase; +using openusdconnect::client::PhaseInputs; using openusdconnect::client::ProducerRecoveryDisposition; using openusdconnect::client::ProducerResult; -using openusdconnect::client::ReceiverMessageKind; - -namespace -{ -using PythonPayload = nb::object; -using PythonProducerSession = openusdconnect::client::OrderedProducerSession; -using PythonReceiverInbox = openusdconnect::client::OrderedReceiverSession; - -class PythonFrameError final : public std::runtime_error -{ -public: - using std::runtime_error::runtime_error; -}; - -[[noreturn]] void RaiseFrameError(FrameResult result) -{ - switch (result) - { - case FrameResult::InvalidMaxFrameSize: - throw PythonFrameError("max frame size must fit in a non-zero uint32"); - case FrameResult::EmptyPayload: - throw PythonFrameError("frame payload must not be empty"); - case FrameResult::PayloadTooLarge: - throw PythonFrameError("frame payload exceeds the configured limit"); - case FrameResult::InvalidHeader: - throw PythonFrameError("frame payload size is invalid"); - case FrameResult::Success: - break; - } - throw std::logic_error("invalid frame result"); -} - -[[nodiscard]] nb::list ToPythonBytes(std::vector> values) -{ - nb::list result; - for (const auto& value : values) - { - result.append(nb::bytes(value.data(), value.size())); - } - return result; -} - -[[nodiscard]] nb::list ToPythonBytes(std::vector values) -{ - nb::list result; - for (const auto& value : values) - { - result.append(value); - } - return result; -} - -void ValidateBytes(nb::handle value, const char* empty_message) -{ - if (!PyBytes_Check(value.ptr())) - { - throw nb::type_error("payload must be bytes"); - } - if (PyBytes_Size(value.ptr()) == 0) - { - throw std::invalid_argument(empty_message); - } -} - -} // namespace +// Defined in receiver_bindings.cpp and producer_bindings.cpp. +void BindReceiver(nb::module_& module); +void BindProducer(nb::module_& module); NB_MODULE(_native_client, module) { - module.doc() = "Native OpenUSDConnect client primitives"; - module.attr("DEFAULT_MAX_FRAME_SIZE") = openusdconnect::client::kDefaultMaxFrameSize; - - nb::exception(module, "FrameError", PyExc_ValueError); - - nb::enum_(module, "ReceiverMessageKind") - .value("EVENT", ReceiverMessageKind::Event) - .value("LAYER_GRAPH_STATE", ReceiverMessageKind::LayerGraphState) - .value("RESYNC", ReceiverMessageKind::Resync) - .value("OTHER", ReceiverMessageKind::Other); - - nb::enum_(module, "AcceptResult") - .value("ACCEPTED", AcceptResult::Accepted) - .value("STALE_GENERATION", AcceptResult::StaleGeneration) - .value("QUEUE_FULL", AcceptResult::QueueFull) - .value("DUPLICATE", AcceptResult::Duplicate) - .value("SEQUENCE_GAP", AcceptResult::SequenceGap) - .value("INVALID_SEQUENCE", AcceptResult::InvalidSequence); - - nb::enum_(module, "ProducerPhase") - .value("DISCONNECTED", ProducerPhase::Disconnected) - .value("AWAITING_HELLO", ProducerPhase::AwaitingHello) - .value("READY", ProducerPhase::Ready) - .value("RECOVERY_REQUIRED", ProducerPhase::RecoveryRequired); + module.doc() = "Native OpenUSDConnect client engine"; nb::enum_(module, "ProducerResult") .value("ACCEPTED", ProducerResult::Accepted) @@ -138,236 +43,31 @@ NB_MODULE(_native_client, module) .value("INVALID_OPERATION", ProducerRecoveryDisposition::InvalidOperation) .value("SESSION_FATAL", ProducerRecoveryDisposition::SessionFatal); - nb::class_(module, "ConnectionStart") - .def_ro("generation", &ConnectionStart::Generation) - .def_ro("sync_from", &ConnectionStart::SyncFrom); - - nb::class_(module, "ProducerConnectionStart") - .def_ro("generation", &ProducerConnectionStart::Generation); + module.def("rejection_code_name", &openusdconnect::client::RejectionCodeName, "code"_a); + module.def("rejection_disposition", &openusdconnect::client::RejectionDisposition, "code"_a); - nb::class_(module, "FrameDecoder") - .def( - "__init__", - [](FrameDecoder* decoder, std::size_t max_frame_size) - { - if (!openusdconnect::client::IsValidMaxFrameSize(max_frame_size)) - { - throw nb::value_error("max_frame_size must fit in a non-zero uint32"); - } - new (decoder) FrameDecoder(max_frame_size); - }, - "max_frame_size"_a = openusdconnect::client::kDefaultMaxFrameSize) - .def( - "feed", - [](FrameDecoder& decoder, const nb::bytes& chunk) - { - std::vector> frames; - const FrameResult result = decoder.Feed( - static_cast(chunk.data()), chunk.size(), frames); - if (result != FrameResult::Success) - { - RaiseFrameError(result); - } - return ToPythonBytes(std::move(frames)); - }, - "chunk"_a) - .def("reset", &FrameDecoder::Reset) - .def_prop_ro("buffered_bytes", &FrameDecoder::BufferedBytes) - .def_prop_ro("max_frame_size", &FrameDecoder::MaxFrameSize); + nb::enum_(module, "ClientPhase") + .value("OFFLINE", ClientPhase::Offline) + .value("CONNECTING", ClientPhase::Connecting) + .value("REPLAYING", ClientPhase::Replaying) + .value("READY", ClientPhase::Ready) + .value("RECOVERY_REQUIRED", ClientPhase::RecoveryRequired) + .value("REJECTED", ClientPhase::Rejected) + .value("CLOSED", ClientPhase::Closed) + .value("PARKED", ClientPhase::Parked); module.def( - "encode_frame", - [](const nb::bytes& payload, std::size_t max_frame_size) + "compute_phase", + [](bool closed, bool recovery_required, bool rejected, bool parked, bool replaying, + bool ready, bool connecting) { - std::vector framed; - const FrameResult result = openusdconnect::client::EncodeFrame( - static_cast(payload.data()), payload.size(), framed, - max_frame_size); - if (result != FrameResult::Success) - { - RaiseFrameError(result); - } - return nb::bytes(framed.data(), framed.size()); + return openusdconnect::client::ComputePhase( + {closed, recovery_required, rejected, parked, replaying, ready, connecting}); }, - "payload"_a, "max_frame_size"_a = openusdconnect::client::kDefaultMaxFrameSize); - - nb::class_(module, "ReceiverInbox") - .def( - "__init__", - [](PythonReceiverInbox* inbox, std::int32_t initial_sync_from, std::size_t max_messages, - bool require_contiguous) - { - if (!PythonReceiverInbox::IsValidConfiguration(initial_sync_from, max_messages)) - { - throw nb::value_error( - "initial_sync_from must be positive and max_messages must be non-zero"); - } - new (inbox) - PythonReceiverInbox(initial_sync_from, max_messages, require_contiguous); - }, - "initial_sync_from"_a, "max_messages"_a, "require_contiguous"_a = false) - .def("begin_connection", &PythonReceiverInbox::BeginConnection) - .def("disconnect", &PythonReceiverInbox::Disconnect, "generation"_a) - .def( - "accept", - [](PythonReceiverInbox& inbox, std::uint64_t generation, ReceiverMessageKind kind, - std::int32_t sequence, nb::handle frame) - { - ValidateBytes(frame, "frame must not be empty"); - const AcceptResult result = - inbox.Accept(generation, kind, sequence, nb::borrow(frame)); - if (result == AcceptResult::InvalidSequence) - { - throw nb::value_error("sequenced messages require a positive sequence"); - } - return result; - }, - "generation"_a, "kind"_a, "sequence"_a, "frame"_a) - .def( - "accept_replay_complete", - [](PythonReceiverInbox& inbox, std::uint64_t generation, std::int32_t head_seq, - std::uint64_t epoch) - { - const AcceptResult result = inbox.AcceptReplayComplete(generation, head_seq, epoch); - if (result == AcceptResult::InvalidSequence) - { - throw nb::value_error("replay head must not be negative"); - } - return result; - }, - "generation"_a, "head_seq"_a, "epoch"_a) - .def( - "drain", - [](PythonReceiverInbox& inbox, std::optional max_messages) - { - if (max_messages.has_value() && *max_messages == 0) - { - throw nb::value_error("max_messages must be non-zero when specified"); - } - return ToPythonBytes(inbox.Drain(max_messages)); - }, - "max_messages"_a = nb::none()) - .def("mark_replay_applied", &PythonReceiverInbox::MarkReplayApplied) - .def("mark_applied_through", &PythonReceiverInbox::MarkAppliedThrough, "generation"_a, - "sequence"_a) - .def( - "request_replay_from", - [](PythonReceiverInbox& inbox, std::int32_t sequence) - { - if (!inbox.RequestReplayFrom(sequence)) - { - throw nb::value_error("replay sequence must be at least one"); - } - }, - "sequence"_a) - .def("reset_applied_progress", &PythonReceiverInbox::ResetAppliedProgress) - .def("clear_overflow", &PythonReceiverInbox::ClearOverflow) - .def_prop_ro("generation", &PythonReceiverInbox::Generation) - .def_prop_ro("last_sequence", &PythonReceiverInbox::LastSequence) - .def_prop_ro("last_applied_sequence", &PythonReceiverInbox::LastAppliedSequence) - .def_prop_ro("size", &PythonReceiverInbox::Size) - .def_prop_ro("synchronized", &PythonReceiverInbox::Synchronized) - .def_prop_ro("overflowed", &PythonReceiverInbox::Overflowed) - .def_prop_ro("replay_head_sequence", &PythonReceiverInbox::ReplayHeadSequence) - .def_prop_ro("replay_epoch", &PythonReceiverInbox::ReplayEpoch); + nb::kw_only(), "closed"_a, "recovery_required"_a, "rejected"_a, "parked"_a, "replaying"_a, + "ready"_a, "connecting"_a); - nb::class_(module, "ProducerSession") - .def( - "__init__", - [](PythonProducerSession* session, std::size_t capacity) - { - if (!PythonProducerSession::IsValidConfiguration(capacity)) - { - throw nb::value_error("capacity must be non-zero"); - } - new (session) PythonProducerSession(capacity); - }, - "capacity"_a) - .def("begin_connection", &PythonProducerSession::BeginConnection) - .def("accept_hello", &PythonProducerSession::AcceptHello, "generation"_a, - "committed_through"_a) - .def("disconnect", &PythonProducerSession::Disconnect, "generation"_a) - .def( - "append", - [](PythonProducerSession& session, std::uint64_t generation, - std::uint64_t transaction_id, nb::handle payload, std::size_t event_count, - std::string layer_key) - { - ValidateBytes(payload, "transaction payload must not be empty"); - if (event_count == 0) - { - throw nb::value_error("event_count must be non-zero"); - } - return session.Append(generation, transaction_id, - nb::borrow(payload), event_count, - std::move(layer_key)); - }, - "generation"_a, "transaction_id"_a, "payload"_a, "event_count"_a, "layer_key"_a = "") - .def( - "claim_next_unsent", - [](PythonProducerSession& session, std::uint64_t generation) -> nb::object - { - PythonProducerSession::Entry entry; - const ProducerResult result = session.ClaimNextUnsent(generation, entry); - if (result == ProducerResult::NoPendingTransaction) - { - return nb::none(); - } - if (result != ProducerResult::Accepted) - { - throw std::logic_error("producer session is not ready to replay"); - } - return nb::make_tuple(entry.TransactionId, entry.Payload, entry.EventCount, - entry.LayerKey); - }, - "generation"_a) - .def("acknowledge_through", &PythonProducerSession::AcknowledgeThrough, "generation"_a, - "transaction_id"_a) - .def("reject", &PythonProducerSession::Reject, "generation"_a, "transaction_id"_a, - "disposition"_a) - .def( - "repair_rejected", - [](PythonProducerSession& session, nb::handle payload, std::size_t event_count, - std::string layer_key) - { - ValidateBytes(payload, "transaction payload must not be empty"); - if (event_count == 0) - { - throw nb::value_error("event_count must be non-zero"); - } - return session.RepairRejected(nb::borrow(payload), event_count, - std::move(layer_key)); - }, - "payload"_a, "event_count"_a, "layer_key"_a = "") - .def("reset_session", &PythonProducerSession::ResetSession) - .def("contains", &PythonProducerSession::Contains, "transaction_id"_a) - .def("drain_acknowledged_event_count", &PythonProducerSession::DrainAcknowledgedEventCount) - .def("entries", - [](const PythonProducerSession& session) - { - nb::list result; - for (const auto& entry : session.Entries()) - { - result.append(nb::make_tuple(entry.TransactionId, entry.Payload, - entry.EventCount, entry.LayerKey)); - } - return result; - }) - .def_prop_ro("phase", &PythonProducerSession::Phase) - .def_prop_ro("generation", &PythonProducerSession::Generation) - .def_prop_ro("can_append", &PythonProducerSession::CanAppend) - .def_prop_ro("empty", &PythonProducerSession::Empty) - .def_prop_ro("recovery_required", &PythonProducerSession::RecoveryRequired) - .def_prop_ro("recovery_disposition", &PythonProducerSession::RecoveryDisposition) - .def_prop_ro("rejected_transaction_id", &PythonProducerSession::RejectedTransactionId) - .def_prop_ro("pending_transaction_count", &PythonProducerSession::PendingTransactionCount) - .def_prop_ro("pending_event_count", &PythonProducerSession::PendingEventCount) - .def_prop_ro("next_transaction_id", &PythonProducerSession::NextTransactionId) - .def_prop_ro("acknowledged_transaction_count", - &PythonProducerSession::AcknowledgedTransactionCount) - .def_prop_ro("acknowledged_event_count", &PythonProducerSession::AcknowledgedEventCount) - .def_prop_ro("submitted_transaction_count", - &PythonProducerSession::SubmittedTransactionCount) - .def_prop_ro("last_acknowledged_transaction_id", - &PythonProducerSession::LastAcknowledgedTransactionId); + openusdconnect::python::BindDriverTypes(module); + BindReceiver(module); + BindProducer(module); } diff --git a/native/python/driver_bindings.cpp b/native/python/driver_bindings.cpp new file mode 100644 index 0000000..0905d34 --- /dev/null +++ b/native/python/driver_bindings.cpp @@ -0,0 +1,175 @@ +#include "driver_bindings.h" + +#include +#include + +#include +#include +#include +#include + +namespace openusdconnect::python +{ +namespace +{ + +using namespace client; +using OpenUSDConnect::LayerMode; + +// Bounds every wait so a caller-supplied number cannot overflow a deadline. +constexpr std::chrono::hours kLongestWait{24 * 365}; + +void BindEnums(nb::module_& module) +{ + nb::enum_(module, "LogLevel") + .value("DEBUG", LogLevel::Debug) + .value("INFO", LogLevel::Info) + .value("WARNING", LogLevel::Warning) + .value("ERROR", LogLevel::Error); + + nb::enum_(module, "LayerMode") + .value("MANAGED", LayerMode::Managed) + .value("SHARED_STAGE", LayerMode::SharedStage); + + nb::enum_(module, "DisconnectReason") + .value("CONNECT_FAILED", DisconnectReason::ConnectFailed) + .value("PEER_CLOSED", DisconnectReason::PeerClosed) + .value("TRANSPORT_ERROR", DisconnectReason::TransportError) + .value("STOPPED", DisconnectReason::Stopped) + .value("HANDSHAKE_REJECTED", DisconnectReason::HandshakeRejected) + .value("REPLAY_REQUESTED", DisconnectReason::ReplayRequested) + .value("SEQUENCE_GAP", DisconnectReason::SequenceGap) + .value("QUEUE_FULL", DisconnectReason::QueueFull) + .value("READ_TIMEOUT", DisconnectReason::ReadTimeout) + .value("PROTOCOL_ERROR", DisconnectReason::ProtocolError) + .value("CANCELLED", DisconnectReason::Cancelled) + .value("HANDSHAKE_TIMEOUT", DisconnectReason::HandshakeTimeout) + .value("RECOVERY_REQUIRED", DisconnectReason::RecoveryRequired) + .value("RATE_LIMITED", DisconnectReason::RateLimited); + + nb::enum_(module, "SocketResult") + .value("SUCCESS", SocketResult::Success) + .value("TIMEOUT", SocketResult::Timeout) + .value("INTERRUPTED", SocketResult::Interrupted) + .value("CLOSED", SocketResult::Closed) + .value("FAILED", SocketResult::Failed); + + nb::enum_(module, "SocketOperation") + .value("CONNECT", SocketOperation::Connect) + .value("SEND", SocketOperation::Send) + .value("RECEIVE", SocketOperation::Receive); +} + +void BindNotifications(nb::module_& module) +{ + nb::class_(module, "Connected"); + + nb::class_(module, "Disconnected").def_ro("reason", &Disconnected::Reason); + + nb::class_(module, "HandshakeRejected") + .def_ro("authentication", &HandshakeRejected::Authentication) + .def_prop_ro("code", + [](const HandshakeRejected& rejected) + { + return static_cast(rejected.Code); + }) + .def_ro("reason", &HandshakeRejected::Reason); + + nb::class_(module, "TokenIssued").def_ro("token", &TokenIssued::Token); + + nb::class_(module, "StageMetadata") + .def_ro("time_codes_per_second", &StageMetadata::TimeCodesPerSecond) + .def_ro("frames_per_second", &StageMetadata::FramesPerSecond) + .def_ro("start_time_code", &StageMetadata::StartTimeCode) + .def_ro("end_time_code", &StageMetadata::EndTimeCode) + .def_ro("meters_per_unit", &StageMetadata::MetersPerUnit) + .def_ro("up_axis", &StageMetadata::UpAxis); + + nb::class_(module, "PlaybackState") + .def_ro("time", &PlaybackState::Time) + .def_ro("playing", &PlaybackState::Playing) + .def_ro("rate", &PlaybackState::Rate) + .def_ro("leader_client_id", &PlaybackState::LeaderClientId); + + nb::class_(module, "PlaybackClaimed") + .def_ro("leader_client_id", &PlaybackClaimed::LeaderClientId); + + nb::class_(module, "PlaybackRejected") + .def_ro("reason", &PlaybackRejected::Reason) + .def_ro("current_leader_client_id", &PlaybackRejected::CurrentLeaderClientId); + + nb::class_(module, "NotificationQueue") + .def(nb::init<>()) + .def("drain", + [](NotificationQueue& queue) + { + nb::list result; + for (Notification& notification : queue.Drain()) + { + result.append(ToPython(std::move(notification))); + } + return result; + }); +} + +void BindTransport(nb::module_& module) +{ + nb::class_(module, "SocketFactory"); + + nb::class_(module, "TcpSocketFactory").def(nb::init<>()); + + nb::class_(module, "TransportFailure") + .def_ro("operation", &TransportFailure::Operation) + .def_ro("result", &TransportFailure::Result) + .def_ro("system_error", &TransportFailure::SystemError) + .def_prop_ro("description", + [](const TransportFailure& failure) + { + return Describe(failure); + }); +} + +} // namespace + +void BindDriverTypes(nb::module_& module) +{ + BindEnums(module); + BindNotifications(module); + BindTransport(module); +} + +std::chrono::milliseconds Milliseconds(double seconds) +{ + if (!(seconds > 0.0)) + { + return std::chrono::milliseconds::zero(); + } + const std::chrono::duration requested(seconds); + if (requested >= kLongestWait) + { + return kLongestWait; + } + return std::chrono::ceil(requested); +} + +double Seconds(std::chrono::milliseconds duration) +{ + return std::chrono::duration(duration).count(); +} + +std::optional Timeout(std::optional seconds) +{ + return seconds ? std::optional(Milliseconds(*seconds)) : std::nullopt; +} + +nb::object ToPython(Notification notification) +{ + return std::visit( + [](auto&& value) + { + return nb::cast(std::move(value)); + }, + std::move(notification)); +} + +} // namespace openusdconnect::python diff --git a/native/python/driver_bindings.h b/native/python/driver_bindings.h new file mode 100644 index 0000000..4c61e35 --- /dev/null +++ b/native/python/driver_bindings.h @@ -0,0 +1,272 @@ +#pragma once + +#include "openusdconnect/client/driver/socket.h" +#include "openusdconnect/client/engine/notification.h" + +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +// What the receiver and producer bindings share: durations in seconds, the +// notification types, and the Python side of a reference driver. +namespace openusdconnect::python +{ + +namespace nb = nanobind; + +// Binds the types both roles use. Call before either role's bindings. +void BindDriverTypes(nb::module_& module); + +// Seconds from Python, clamped to zero and bounded so no deadline overflows. +[[nodiscard]] std::chrono::milliseconds Milliseconds(double seconds); +[[nodiscard]] double Seconds(std::chrono::milliseconds duration); +[[nodiscard]] std::optional Timeout(std::optional seconds); + +// Binds a duration field as a property in seconds. +template +void DurationProperty(nb::class_& cls, const char* name, + std::chrono::milliseconds Config::* member) +{ + cls.def_prop_rw( + name, + [member](const Config& config) + { + return Seconds(config.*member); + }, + [member](Config& config, double seconds) + { + config.*member = Milliseconds(seconds); + }); +} + +[[nodiscard]] nb::object ToPython(client::Notification notification); + +// Owns a reference driver whose thread calls Python, and closes it when +// destroyed, so the loop writes what was queued. Every Python object here is +// touched only with the GIL, and the GIL is released around every wait. +template +class PythonDriver final +{ +public: + using Endpoint = typename Driver::EndpointType; + + // role names the callbacks in unraisable-exception reports. + PythonDriver(const std::string& role, Endpoint& endpoint, + client::NotificationQueue& notifications, + std::shared_ptr sockets, nb::object token_provider, + nb::object token_issued, nb::object notification_sink, nb::object log) + : TokenContext(role + " token provider") + , IssuedContext(role + " token issued hook") + , SinkContext(role + " notification sink") + , LogContext(role + " log") + , TokenProvider(std::move(token_provider)) + , IssuedHook(std::move(token_issued)) + , Sink(std::move(notification_sink)) + , LogCallback(std::move(log)) + , Native(std::make_unique(endpoint, notifications, std::move(sockets), Callbacks())) + { + std::lock_guard lock(RegistryMutex()); + Registry().insert(this); + } + + ~PythonDriver() + { + Destroying = true; + { + std::lock_guard lock(RegistryMutex()); + Registry().erase(this); + } + nb::gil_scoped_release release; + static_cast(Native->Close(std::nullopt)); + } + + PythonDriver(const PythonDriver&) = delete; + PythonDriver& operator=(const PythonDriver&) = delete; + + [[nodiscard]] Driver& Get() noexcept + { + return *Native; + } + + // Runs at interpreter exit, before threads may no longer take the GIL. + static void StopAll() + { + std::vector> drivers; + { + std::lock_guard lock(RegistryMutex()); + for (PythonDriver* driver : Registry()) + { + drivers.emplace_back(nb::find(driver), driver->Native.get()); + } + } + nb::gil_scoped_release release; + for (const auto& [object, driver] : drivers) + { + static_cast(driver->Close(std::nullopt)); + } + } + +private: + [[nodiscard]] static std::mutex& RegistryMutex() + { + static std::mutex mutex; + return mutex; + } + + [[nodiscard]] static std::set& Registry() + { + static std::set drivers; + return drivers; + } + + [[nodiscard]] client::DriverCallbacks Callbacks() + { + client::DriverCallbacks callbacks; + if (!TokenProvider.is_none()) + { + // Anything but a str abandons the connection attempt. + callbacks.Token = [this] + { + std::optional token; + CallPython(TokenContext, + [&] + { + std::string text; + if (nb::try_cast(TokenProvider(), text)) + { + token = std::move(text); + } + }); + return token; + }; + } + if (!IssuedHook.is_none()) + { + callbacks.TokenIssued = [this](const std::string& token) + { + CallPython(IssuedContext, + [&] + { + IssuedHook(token); + }); + }; + } + if (!Sink.is_none()) + { + callbacks.Notifications = [this](client::Notification notification) + { + CallPython(SinkContext, + [&] + { + Sink(ToPython(std::move(notification))); + }); + }; + } + if (!LogCallback.is_none()) + { + callbacks.Log = [this](client::LogLevel level, const std::string& message) + { + CallPython(LogContext, + [&] + { + LogCallback(level, message); + }); + }; + } + return callbacks; + } + + // Runs call, which calls into Python, on the driver thread with the GIL + // held; a Python exception is reported as unraisable under context. + template + void CallPython(const std::string& context, Call call) + { + nb::gil_scoped_acquire gil; + // Python may drop every other reference to this driver meanwhile, and + // destroying it on its own thread would join that thread from itself. + nb::object self = Destroying ? nb::object() : nb::find(this); + try + { + call(); + } + catch (nb::python_error& error) + { + error.discard_as_unraisable(context.c_str()); + } + if (self.is_valid() && Py_REFCNT(self.ptr()) == 1) + { + // Its owner is gone: the loop writes what is queued and exits, and the main + // thread destroys the driver. A full pending-call queue leaks it instead. + Native->StopAfterQueued(); + static_cast(Py_AddPendingCall( + [](void* object) + { + nb::handle(static_cast(object)).dec_ref(); + return 0; + }, + self.release().ptr())); + } + } + + const std::string TokenContext; + const std::string IssuedContext; + const std::string SinkContext; + const std::string LogContext; + nb::object TokenProvider; + nb::object IssuedHook; + nb::object Sink; + nb::object LogCallback; + bool Destroying = false; + // Destroyed first, so its thread has exited before the objects above go. + const std::unique_ptr Native; +}; + +// Binds the lifecycle every reference driver shares; a role adds its own methods. +template +void BindThreadedDriver(nb::class_>& cls) +{ + using Bound = PythonDriver; + cls.def("start", + [](Bound& driver) + { + return driver.Get().Start(); + }) + .def("wake", + [](Bound& driver) + { + driver.Get().Wake(); + }) + .def( + "close", + [](Bound& driver, std::optional timeout) + { + return driver.Get().Close(Timeout(timeout)); + }, + nb::arg("timeout") = nb::none(), nb::call_guard()) + .def_prop_ro("running", + [](Bound& driver) + { + return driver.Get().Running(); + }) + .def_prop_ro("stopped", + [](Bound& driver) + { + return driver.Get().Stopped(); + }) + .def_prop_ro("last_failure", + [](Bound& driver) + { + return driver.Get().LastFailure(); + }); + nb::module_::import_("atexit").attr("register")(nb::cpp_function(&Bound::StopAll)); +} + +} // namespace openusdconnect::python diff --git a/native/python/producer_bindings.cpp b/native/python/producer_bindings.cpp new file mode 100644 index 0000000..626d0ce --- /dev/null +++ b/native/python/producer_bindings.cpp @@ -0,0 +1,213 @@ +#include "driver_bindings.h" + +#include "openusdconnect/client/driver/socket.h" +#include "openusdconnect/client/driver/threaded_producer_driver.h" +#include "openusdconnect/client/engine/producer_endpoint.h" +#include "openusdconnect/client/frame_codec.h" + +#include +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace nb = nanobind; +using namespace nb::literals; +using namespace openusdconnect::client; +using openusdconnect::python::DurationProperty; +using openusdconnect::python::Milliseconds; +using openusdconnect::python::Timeout; + +namespace +{ + +using PythonProducerDriver = openusdconnect::python::PythonDriver; + +// Python passes and receives bare envelopes; the endpoint keeps them framed. +// An envelope that cannot be framed leaves the frame empty, which the endpoint +// refuses like any incomplete frame. +[[nodiscard]] std::vector Frame(const nb::bytes& envelope) +{ + std::vector frame; + static_cast( + EncodeFrame(static_cast(envelope.data()), envelope.size(), frame)); + return frame; +} + +[[nodiscard]] nb::bytes Envelope(const SharedByteBuffer& frame) +{ + return nb::bytes(frame->data() + kFrameHeaderSize, frame->size() - kFrameHeaderSize); +} + +void BindTypes(nb::module_& module) +{ + nb::class_ config_class(module, "ProducerConfig"); + config_class.def(nb::init<>()) + .def_rw("host", &ProducerConfig::Host) + .def_rw("port", &ProducerConfig::Port) + .def_rw("client_id", &ProducerConfig::ClientId) + .def_rw("origin", &ProducerConfig::Origin) + .def_rw("department", &ProducerConfig::Department) + .def_rw("layer_mode", &ProducerConfig::LayerMode) + .def_rw("session_id", &ProducerConfig::SessionId) + .def_rw("max_pending_transactions", &ProducerConfig::MaxPendingTransactions); + DurationProperty(config_class, "handshake_timeout", &ProducerConfig::HandshakeTimeout); + + nb::class_(module, "ProducerStatus") + .def_ro("connected", &ProducerStatus::Connected) + .def_ro("rejection", &ProducerStatus::Rejection) + .def_ro("layer_mode_active", &ProducerStatus::LayerModeActive) + .def_ro("metadata", &ProducerStatus::Metadata) + .def_ro("session_id", &ProducerStatus::SessionId) + .def_ro("pending_transactions", &ProducerStatus::PendingTransactions) + .def_ro("pending_events", &ProducerStatus::PendingEvents) + .def_ro("acknowledged_transactions", &ProducerStatus::AcknowledgedTransactions) + .def_ro("acknowledged_events", &ProducerStatus::AcknowledgedEvents); + + nb::class_(module, "MirrorCheckpoint") + .def_ro("server_instance", &MirrorCheckpoint::ServerInstance) + .def_ro("epoch", &MirrorCheckpoint::Epoch) + .def_ro("head_sequence", &MirrorCheckpoint::HeadSequence); + + nb::class_(module, "TransactionFailure") + .def_ro("transaction_id", &TransactionFailure::TransactionId) + .def_ro("code", &TransactionFailure::Code) + .def_ro("reason", &TransactionFailure::Reason) + .def_ro("expected_transaction_id", &TransactionFailure::ExpectedTransactionId); + + nb::class_(module, "ProducerSessionEntry") + .def_ro("transaction_id", &ProducerSessionEntry::TransactionId) + .def_prop_ro("payload", + [](const ProducerSessionEntry& entry) + { + return Envelope(entry.Payload); + }) + .def_ro("event_count", &ProducerSessionEntry::EventCount) + .def_ro("layer_key", &ProducerSessionEntry::LayerKey); + + nb::class_(module, "RecoveryArtifact") + .def_ro("session_id", &RecoveryArtifact::SessionId) + .def_ro("failure", &RecoveryArtifact::Failure) + .def_ro("transactions", &RecoveryArtifact::Transactions); + + nb::enum_(module, "FlushResult") + .value("FLUSHED", FlushResult::Flushed) + .value("RECOVERY_REQUIRED", FlushResult::RecoveryRequired) + .value("UNFINISHED", FlushResult::Unfinished); +} + +void BindEndpoint(nb::module_& module) +{ + nb::class_(module, "ProducerEndpoint") + .def( + "__init__", + [](ProducerEndpoint* endpoint, const ProducerConfig& config, + NotificationQueue& notifications) + { + if (!ProducerEndpoint::IsValidConfiguration(config)) + { + throw nb::value_error("invalid producer configuration"); + } + new (endpoint) ProducerEndpoint(config, notifications); + }, + "config"_a, "notifications"_a, nb::keep_alive<1, 3>()) + .def("status", &ProducerEndpoint::Status) + .def( + "request_connect", + [](ProducerEndpoint& endpoint, std::optional timeout) + { + // The endpoint caps the attempt at its handshake timeout. + const std::chrono::milliseconds budget = + timeout ? Milliseconds(*timeout) : endpoint.Configuration().HandshakeTimeout; + const TimePoint now = std::chrono::steady_clock::now(); + return endpoint.RequestConnect(now, now + budget); + }, + "timeout"_a = nb::none()) + .def("cancel_connect", &ProducerEndpoint::CancelConnect) + .def("disconnect", &ProducerEndpoint::Disconnect) + .def("next_transaction_id", &ProducerEndpoint::NextTransactionId) + .def( + "append", + [](ProducerEndpoint& endpoint, std::uint64_t transaction_id, const nb::bytes& envelope, + std::size_t event_count, std::string layer_key) + { + return endpoint.Append(transaction_id, Frame(envelope), event_count, + std::move(layer_key)); + }, + "transaction_id"_a, "envelope"_a, "event_count"_a, "layer_key"_a = "") + .def( + "queue_control", + [](ProducerEndpoint& endpoint, const nb::bytes& envelope) + { + return endpoint.QueueControl(Frame(envelope)); + }, + "envelope"_a) + .def("drain_acknowledged_event_count", &ProducerEndpoint::DrainAcknowledgedEventCount) + .def("acknowledged_checkpoint", &ProducerEndpoint::AcknowledgedCheckpoint) + .def("failure", &ProducerEndpoint::Failure) + .def("artifact", &ProducerEndpoint::Artifact) + .def( + "repair_rejected", + [](ProducerEndpoint& endpoint, const nb::bytes& envelope, std::size_t event_count, + std::string layer_key) + { + return endpoint.RepairRejected(Frame(envelope), event_count, std::move(layer_key)); + }, + "envelope"_a, "event_count"_a, "layer_key"_a = "") + .def("abandon_rejected_session", &ProducerEndpoint::AbandonRejectedSession, "session_id"_a); +} + +void BindDriver(nb::module_& module) +{ + nb::class_ cls(module, "ProducerDriver"); + cls.def( + "__init__", + [](PythonProducerDriver* driver, ProducerEndpoint& endpoint, + NotificationQueue& notifications, std::shared_ptr sockets, + nb::object token_provider, nb::object token_issued, nb::object notification_sink, + nb::object log) + { + new (driver) + PythonProducerDriver("producer", endpoint, notifications, std::move(sockets), + std::move(token_provider), std::move(token_issued), + std::move(notification_sink), std::move(log)); + }, + "endpoint"_a, "notifications"_a, "sockets"_a, nb::kw_only(), + "token_provider"_a = nb::none(), "token_issued"_a = nb::none(), + "notification_sink"_a = nb::none(), "log"_a = nb::none(), nb::keep_alive<1, 2>(), + nb::keep_alive<1, 3>()) + .def( + "connect", + [](PythonProducerDriver& driver, std::optional timeout) + { + return driver.Get().Connect(Timeout(timeout)); + }, + "timeout"_a = nb::none(), nb::call_guard()) + .def( + "flush", + [](PythonProducerDriver& driver, std::optional timeout) + { + return driver.Get().Flush(Timeout(timeout)); + }, + "timeout"_a = nb::none(), nb::call_guard()); + openusdconnect::python::BindThreadedDriver(cls); +} + +} // namespace + +void BindProducer(nb::module_& module) +{ + BindTypes(module); + BindEndpoint(module); + BindDriver(module); +} diff --git a/native/python/receiver_bindings.cpp b/native/python/receiver_bindings.cpp new file mode 100644 index 0000000..1dc65e7 --- /dev/null +++ b/native/python/receiver_bindings.cpp @@ -0,0 +1,154 @@ +#include "driver_bindings.h" + +#include "openusdconnect/client/driver/socket.h" +#include "openusdconnect/client/driver/threaded_receiver_driver.h" +#include "openusdconnect/client/engine/receiver_endpoint.h" + +#include +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace nb = nanobind; +using namespace nb::literals; +using namespace openusdconnect::client; +using openusdconnect::python::DurationProperty; +using openusdconnect::python::Timeout; + +namespace +{ + +using PythonReceiverDriver = openusdconnect::python::PythonDriver; + +[[nodiscard]] nb::list ToPythonBytes(const std::vector>& frames) +{ + nb::list result; + for (const std::vector& frame : frames) + { + result.append(nb::bytes(frame.data(), frame.size())); + } + return result; +} + +void BindEndpoint(nb::module_& module) +{ + nb::class_ config_class(module, "ReceiverConfig"); + config_class.def(nb::init<>()) + .def_rw("host", &ReceiverConfig::Host) + .def_rw("port", &ReceiverConfig::Port) + .def_rw("client_id", &ReceiverConfig::ClientId) + .def_rw("origin", &ReceiverConfig::Origin) + .def_rw("department", &ReceiverConfig::Department) + .def_rw("layered_replay", &ReceiverConfig::LayeredReplay) + .def_rw("layer_mode", &ReceiverConfig::LayerMode) + .def_rw("sync_from", &ReceiverConfig::SyncFrom) + .def_rw("max_queue", &ReceiverConfig::MaxQueue) + .def_rw("max_consecutive_timeouts", &ReceiverConfig::MaxConsecutiveTimeouts) + .def_rw("reconnect", &ReceiverConfig::Reconnect); + DurationProperty(config_class, "socket_timeout", &ReceiverConfig::SocketTimeout); + DurationProperty(config_class, "reconnect_base_delay", &ReceiverConfig::ReconnectBaseDelay); + DurationProperty(config_class, "reconnect_max_delay", &ReceiverConfig::ReconnectMaxDelay); + + nb::class_(module, "ReceiverStatus") + .def_ro("connected", &ReceiverStatus::Connected) + .def_ro("synchronized", &ReceiverStatus::Synchronized) + .def_ro("stopped", &ReceiverStatus::Stopped) + .def_ro("replay_head_sequence", &ReceiverStatus::ReplayHeadSequence) + .def_ro("replay_epoch", &ReceiverStatus::ReplayEpoch) + .def_ro("server_instance", &ReceiverStatus::ServerInstance) + .def_ro("layered_replay_active", &ReceiverStatus::LayeredReplayActive) + .def_ro("layer_mode_active", &ReceiverStatus::LayerModeActive) + .def_ro("rejection", &ReceiverStatus::Rejection) + .def_ro("metadata", &ReceiverStatus::Metadata) + .def_ro("queued_frames", &ReceiverStatus::QueuedFrames) + .def_ro("last_sequence", &ReceiverStatus::LastSequence) + .def_ro("last_applied_sequence", &ReceiverStatus::LastAppliedSequence); + + nb::class_(module, "ReceiverEndpoint") + .def( + "__init__", + [](ReceiverEndpoint* endpoint, const ReceiverConfig& config, + NotificationQueue& notifications) + { + if (!ReceiverEndpoint::IsValidConfiguration(config)) + { + throw nb::value_error("invalid receiver configuration"); + } + new (endpoint) ReceiverEndpoint(config, notifications); + }, + "config"_a, "notifications"_a, nb::keep_alive<1, 3>()) + .def("status", &ReceiverEndpoint::Status) + .def("stop", &ReceiverEndpoint::Stop) + .def("set_reconnect", &ReceiverEndpoint::SetReconnect, "enabled"_a) + .def( + "drain_frames", + [](ReceiverEndpoint& endpoint, std::optional max_frames) + { + if (max_frames && *max_frames == 0) + { + throw nb::value_error("max_frames must be non-zero when specified"); + } + return ToPythonBytes(endpoint.DrainFrames(max_frames)); + }, + "max_frames"_a = nb::none()) + .def_prop_ro("generation", &ReceiverEndpoint::Generation) + .def("mark_applied_through", &ReceiverEndpoint::MarkAppliedThrough, "generation"_a, + "sequence"_a) + .def("reset_applied_progress", &ReceiverEndpoint::ResetAppliedProgress) + .def("mark_replay_applied", &ReceiverEndpoint::MarkReplayApplied) + .def("request_replay_from", &ReceiverEndpoint::RequestReplayFrom, "sequence"_a) + .def("freeze_marker", &ReceiverEndpoint::FreezeMarker) + .def("drained_through", &ReceiverEndpoint::DrainedThrough, "marker"_a); +} + +void BindDriver(nb::module_& module) +{ + nb::class_ cls(module, "ReceiverDriver"); + cls.def( + "__init__", + [](PythonReceiverDriver* driver, ReceiverEndpoint& endpoint, + NotificationQueue& notifications, std::shared_ptr sockets, + nb::object token_provider, nb::object token_issued, nb::object notification_sink, + nb::object log) + { + new (driver) + PythonReceiverDriver("receiver", endpoint, notifications, std::move(sockets), + std::move(token_provider), std::move(token_issued), + std::move(notification_sink), std::move(log)); + }, + "endpoint"_a, "notifications"_a, "sockets"_a, nb::kw_only(), + "token_provider"_a = nb::none(), "token_issued"_a = nb::none(), + "notification_sink"_a = nb::none(), "log"_a = nb::none(), nb::keep_alive<1, 2>(), + nb::keep_alive<1, 3>()) + .def( + "wait_connected", + [](PythonReceiverDriver& driver, std::optional timeout) + { + return driver.Get().WaitConnected(Timeout(timeout)); + }, + "timeout"_a = nb::none(), nb::call_guard()) + .def( + "wait_synchronized", + [](PythonReceiverDriver& driver, std::optional timeout) + { + return driver.Get().WaitSynchronized(Timeout(timeout)); + }, + "timeout"_a = nb::none(), nb::call_guard()); + openusdconnect::python::BindThreadedDriver(cls); +} + +} // namespace + +void BindReceiver(nb::module_& module) +{ + BindEndpoint(module); + BindDriver(module); +} diff --git a/openusdconnect/__init__.py b/openusdconnect/__init__.py index e1a3f56..ce9fe9c 100644 --- a/openusdconnect/__init__.py +++ b/openusdconnect/__init__.py @@ -55,7 +55,7 @@ ) from .protocol import make_hello, make_quit, make_txn from .protocol_constants import LayerMode -from .receiver import ReceiverThread +from .receiver import EventReceiver from .recovery import ( QuarantinedTransaction, RecoveryArtifact, @@ -82,6 +82,7 @@ "ClientStatus", "DecodeResult", "Event", + "EventReceiver", "EventSender", "HelloRejectionCode", "LayerKeyRouter", @@ -95,7 +96,6 @@ "PlaybackState", "PluginEnvironmentError", "PluginEnvironmentResult", - "ReceiverThread", "QuarantinedTransaction", "RecoveryArtifact", "RecoveryError", diff --git a/openusdconnect/_client_backend.py b/openusdconnect/_client_backend.py index a0a0ffe..4ddf89a 100644 --- a/openusdconnect/_client_backend.py +++ b/openusdconnect/_client_backend.py @@ -1,21 +1,143 @@ -"""Native client-core API used by the Python integration.""" +"""Native client-core API used by the Python integration, and its value conversions.""" + +import logging +import weakref +from collections.abc import Callable from ._native_client import ( # type: ignore[import-not-found] - AcceptResult, - ProducerPhase, + ClientPhase, + FlushResult, + LayerMode, + LogLevel, + NotificationQueue, + PlaybackClaimed, + PlaybackRejected, + PlaybackState, + ProducerConfig, + ProducerDriver, + ProducerEndpoint, ProducerRecoveryDisposition, ProducerResult, - ProducerSession, - ReceiverInbox, - ReceiverMessageKind, + ProducerStatus, + ReceiverConfig, + ReceiverDriver, + ReceiverEndpoint, + ReceiverStatus, + SocketResult, + StageMetadata, + TcpSocketFactory, + TokenIssued, + compute_phase, + rejection_code_name, + rejection_disposition, ) +from .protocol_constants import LayerMode as _LayerMode + +NATIVE_LAYER_MODES = { + _LayerMode.MANAGED: LayerMode.MANAGED, + _LayerMode.SHARED_STAGE: LayerMode.SHARED_STAGE, +} +LAYER_MODES = {native: mode for mode, native in NATIVE_LAYER_MODES.items()} +LOG_LEVELS = { + LogLevel.DEBUG: logging.DEBUG, + LogLevel.INFO: logging.INFO, + LogLevel.WARNING: logging.WARNING, + LogLevel.ERROR: logging.ERROR, +} + + +def stage_metadata_fields(metadata) -> dict: + """The authored fields of native stage metadata, keyed as in a ``set_stage_metadata`` event.""" + fields = { + "timeCodesPerSecond": metadata.time_codes_per_second, + "framesPerSecond": metadata.frames_per_second, + "startTimeCode": metadata.start_time_code, + "endTimeCode": metadata.end_time_code, + "metersPerUnit": metadata.meters_per_unit, + "upAxis": metadata.up_axis, + } + return {key: value for key, value in fields.items() if value is not None} + + +def rejection_reason(rejection) -> str: + """Why a native handshake rejection refused the connection.""" + if rejection.authentication: + return rejection.reason + return rejection.reason or "connection rejected" + + +def notification_queue( + notifications: NotificationQueue | None, **callbacks: Callable | None +) -> tuple[NotificationQueue, bool]: + """The queue an endpoint pushes to, and whether its wrapper delivers it to *callbacks*. + + The owner of a given *notifications* drains it, so none of *callbacks* may be set. + """ + if notifications is None: + return NotificationQueue(), True + named = [name for name, callback in callbacks.items() if callback is not None] + if named: + raise ValueError(f"notifications= cannot be combined with {', '.join(named)}") + return notifications, False + + +def driver_callbacks(owner, logger: logging.Logger, *, sink: bool) -> dict: + """Driver callbacks that hold *owner* weakly, so collecting it stops its driver thread. + + *owner* supplies ``_connection_token()``, ``_token_issued(token)``, and with + *sink* ``_deliver(notification)``. + """ + reference = weakref.ref(owner) + + def method(name: str) -> Callable: + def call(*args): + alive = reference() + return None if alive is None else getattr(alive, name)(*args) + + return call + + def log(level, message: str) -> None: + logger.log(LOG_LEVELS[level], "%s", message) + + return { + "token_provider": method("_connection_token"), + "token_issued": method("_token_issued"), + "notification_sink": method("_deliver") if sink else None, + "log": log, + } + __all__ = [ - "AcceptResult", - "ProducerPhase", + "LAYER_MODES", + "LOG_LEVELS", + "NATIVE_LAYER_MODES", + "ClientPhase", + "FlushResult", + "LayerMode", + "LogLevel", + "NotificationQueue", + "PlaybackClaimed", + "PlaybackRejected", + "PlaybackState", + "ProducerConfig", + "ProducerDriver", + "ProducerEndpoint", "ProducerRecoveryDisposition", "ProducerResult", - "ProducerSession", - "ReceiverInbox", - "ReceiverMessageKind", + "ProducerStatus", + "ReceiverConfig", + "ReceiverDriver", + "ReceiverEndpoint", + "ReceiverStatus", + "SocketResult", + "StageMetadata", + "TcpSocketFactory", + "TokenIssued", + "compute_phase", + "driver_callbacks", + "notification_queue", + "rejection_code_name", + "rejection_disposition", + "rejection_reason", + "stage_metadata_fields", ] diff --git a/openusdconnect/_client_base.py b/openusdconnect/_client_base.py index 336c78f..730d10c 100644 --- a/openusdconnect/_client_base.py +++ b/openusdconnect/_client_base.py @@ -1,6 +1,6 @@ """Lifecycle shells the high-level clients build on. -``ClientBase`` owns the lifecycle, status, observer hooks, and credential of +``ClientBase`` owns the lifecycle, status, observer delivery, and credential of every client. ``PublishingClientBase`` adds the sender role: recovery, playback, and durable completion. ``EmitterClientBase`` adds stage-edit capture through a ``NoticeEmitter`` with optional transform coalescing. @@ -8,43 +8,88 @@ from __future__ import annotations +import logging +from collections import deque from collections.abc import Callable, Sequence +from dataclasses import fields from pxr import Usd +from . import _client_backend from ._client_lifecycle import ( DEFAULT_WAIT_TIMEOUT_S, - ClientCallbackQueue, _pause_before_poll, + close_endpoint, compute_phase, deadline_after, raise_if_rejected, remaining_time, - stop_receiver, submit_and_wait, wait_until_ready, ) from ._client_utils import ClientCredential -from ._observer_hooks import observer_hooks, stage_metadata_from_message -from .client_observer import ClientObserver, StageMetadata +from .client_observer import ( + AppliedBatch, + ClientObserver, + PlaybackClaim, + PlaybackState, + StageMetadata, +) from .client_types import ClientStatus, SyncUpdate from .coalescing import TransformCoalescingWindow +from .dispatcher import EventDispatcher from .emitter import NoticeEmitter, PrimChannel -from .receiver import ReceiverThread +from .receiver import EventReceiver from .recovery import RecoveryArtifact, RecoveryError, RejectionDisposition, TransactionFailure from .sender import EventSender +LOG = logging.getLogger(__name__) + + +def _stage_metadata(native) -> StageMetadata: + """Native stage metadata, whose attributes share the dataclass field names.""" + names = (field.name for field in fields(StageMetadata)) + return StageMetadata(**{name: getattr(native, name) for name in names}) + + +# The observer method each native notification reaches, and its argument. +_NOTIFICATIONS = { + _client_backend.TokenIssued: ("on_token_issued", lambda native: native.token), + _client_backend.StageMetadata: ("on_stage_metadata", _stage_metadata), + _client_backend.PlaybackState: ( + "on_playback_state", + lambda native: PlaybackState( + native.playing, native.time, native.rate, native.leader_client_id, + ), + ), + _client_backend.PlaybackClaimed: ( + "on_playback_claim", + lambda native: PlaybackClaim(True, native.leader_client_id), + ), + _client_backend.PlaybackRejected: ( + "on_playback_claim", + lambda native: PlaybackClaim(False, native.current_leader_client_id, native.reason), + ), +} + + +def _overridden(observer: ClientObserver | None, name: str) -> Callable | None: + """The observer's method *name* when its class overrides it, else ``None``.""" + if observer is None or getattr(type(observer), name) is getattr(ClientObserver, name): + return None + return getattr(observer, name) + class ClientBase: - """Lifecycle, status, observer hooks, and credential shared by every client. + """Lifecycle, status, observer, and credential shared by every client. A subclass builds ``_sender`` and/or ``_receiver`` from ``_credential`` and - ``_hooks``, implements ``_is_synchronized``, and overrides the other status - hooks it needs. Call every method from the stage-owning thread. + ``_notifications`` and overrides the status hooks it needs. Call every + method from the stage-owning thread. """ _sender: EventSender | None = None - _receiver: ReceiverThread | None = None + _receiver: EventReceiver | None = None # Reconnection deliberately suspended by the host (UsdPublisher.disconnect). _paused = False @@ -63,11 +108,21 @@ def __init__( self._closed = False self._apply_depth = 0 self._close_requested = False - self._callbacks = ClientCallbackQueue() - self._hooks = observer_hooks(observer, self._callbacks.wrap) - self._credential = ClientCredential( - host, port, token, persist_token, self._hooks.on_token_issued, - ) + if observer is not None and not isinstance(observer, ClientObserver): + raise TypeError("observer must be a ClientObserver") + self._on_applied = _overridden(observer, "on_applied") + self._on_resync = _overridden(observer, "on_resync") + self._notification_methods = { + kind: (method, convert) + for kind, (name, convert) in _NOTIFICATIONS.items() + if (method := _overridden(observer, name)) is not None + } + # Both roles push here; update() and close() deliver to the observer. + self._notifications = _client_backend.NotificationQueue() + self._undelivered = deque() + # Both roles' handshakes carry the stage metadata. + self._delivered_metadata: StageMetadata | None = None + self._credential = ClientCredential(host, port, token, persist_token) @property def client_id(self) -> str: @@ -76,50 +131,53 @@ def client_id(self) -> str: @property def stage_metadata(self) -> StageMetadata: - return stage_metadata_from_message(self._endpoints[0].stage_metadata) + return _stage_metadata(self._endpoints[0].snapshot().metadata) @property def status(self) -> ClientStatus: """Current state as one immutable value; read it on the stage-owning thread.""" receiver, sender = self._receiver, self._sender - endpoints = self._endpoints + received = None if receiver is None else receiver.snapshot() + sent = None if sender is None else sender.snapshot() + snapshots = [state for state in (sent, received) if state is not None] failure = None if sender is None else sender.transaction_failure rebuild_reason = self._rebuild_reason() - auth_rejected = any(endpoint.auth_rejected for endpoint in endpoints) - rejected = auth_rejected or any(endpoint.hello_rejected for endpoint in endpoints) - connected = not self._closed and all(endpoint.connected for endpoint in endpoints) - synchronized = not self._closed and self._is_synchronized() + rejections = [state.rejection for state in snapshots if state.rejection is not None] + connected = not self._closed and all(state.connected for state in snapshots) + synchronized = not self._closed and self._synchronized( + sent.connected if received is None else received.synchronized + ) phase = compute_phase( closed=self._closed, recovery_required=failure is not None or bool(rebuild_reason), - rejected=rejected, + rejected=bool(rejections), parked=self._is_parked(), - replaying=receiver is not None and receiver.connected and not receiver.synchronized, + replaying=received is not None and received.connected and not received.synchronized, ready=connected and synchronized, connecting=( self._started and not self._paused - and not (receiver is not None and receiver.stopped) + and not (received is not None and received.stopped) ), ) if failure is not None: reason = str(failure) else: - reasons = [e.rejection_reason for e in (sender, receiver) if e is not None] + reasons = map(_client_backend.rejection_reason, rejections) reason = rebuild_reason or next((text for text in reasons if text), "") return ClientStatus( phase=phase, connected=connected, synchronized=synchronized, - receiver_connected=None if receiver is None else receiver.connected, - sender_connected=None if sender is None else sender.connected, + receiver_connected=None if received is None else received.connected, + sender_connected=None if sent is None else sent.connected, prepared_events=self._prepared_events(), - pending_events=0 if sender is None else sender.pending_event_count, - acknowledged_events_total=0 if sender is None else sender.acknowledged_event_count, + pending_events=0 if sent is None else sent.pending_events, + acknowledged_events_total=0 if sent is None else sent.acknowledged_events, failure=failure, - recovery=None if sender is None else sender.recovery_incident, + recovery=None if failure is None else sender.recovery_incident, reason=reason, - auth_rejected=auth_rejected, + auth_rejected=any(rejection.authentication for rejection in rejections), has_unsent_changes=not self._closed and self._has_unsent_changes(), **self._role_status(), ) @@ -169,11 +227,11 @@ def close(self) -> None: self._closed = True try: if self._sender is not None: - self._sender.disconnect() + close_endpoint(self._sender) if self._receiver is not None: - stop_receiver(self._receiver) + close_endpoint(self._receiver) # A token issued by the last handshake must still reach the host. - self._callbacks.close() + self._notify_observer(closing=True) finally: self._release() @@ -189,8 +247,18 @@ def _endpoints(self) -> tuple: return tuple(endpoint for endpoint in (self._receiver, self._sender) if endpoint) def _is_synchronized(self) -> bool: - """Status hook: the local state has applied the server's replay.""" - raise NotImplementedError + """Whether the local state has applied the server's replay.""" + receiver = self._receiver + return self._synchronized( + self._sender.connected if receiver is None else receiver.synchronized + ) + + def _synchronized(self, replayed: bool) -> bool: + """Status hook: synchronization, given whether the receiver applied the replay. + + Without a receiver, *replayed* is whether the sender is connected. + """ + return replayed def _is_parked(self) -> bool: """Status hook: no stage is bound.""" @@ -235,9 +303,51 @@ def _dispatch(self, operation: Callable, /, *args, **kwargs): def _begin_update(self) -> bool: """Deliver queued notifications; ``False`` when one of them closed the client.""" self._require_started() - self._callbacks.drain() + self._notify_observer() return not self._closed + def _observe_dispatcher(self, dispatcher: EventDispatcher) -> None: + """Report *dispatcher*'s deliveries to the observer.""" + on_applied = self._on_applied + if on_applied is not None: + dispatcher.on_applied_events = lambda events: on_applied( + AppliedBatch(dispatcher.applying_seq, events) + ) + dispatcher.on_resync = self._on_resync + + def _notify_observer(self, *, closing: bool = False) -> None: + """Deliver the notifications queued so far. + + A failure propagates and the rest wait for the next call; while + closing, every one runs and the first failure is re-raised after them. + """ + undelivered = self._undelivered + undelivered.extend(self._notifications.drain()) + error = None + while undelivered: + try: + self._deliver(undelivered.popleft()) + except Exception as exc: + if not closing: + raise + if error is not None: + LOG.exception("Observer notification failed while closing") + error = error or exc + if error is not None: + raise error + + def _deliver(self, notification) -> None: + delivery = self._notification_methods.get(type(notification)) + if delivery is None: + return + method, convert = delivery + value = convert(notification) + if isinstance(value, StageMetadata): + if value == self._delivered_metadata: + return + self._delivered_metadata = value + method(value) + def _connect_sender(self, timeout: float | None = None) -> bool: if self._sender.connected: return True diff --git a/openusdconnect/_client_lifecycle.py b/openusdconnect/_client_lifecycle.py index 5780657..164b4a8 100644 --- a/openusdconnect/_client_lifecycle.py +++ b/openusdconnect/_client_lifecycle.py @@ -3,23 +3,33 @@ from __future__ import annotations import logging -import queue -import threading import time -from collections.abc import Callable from typing import TYPE_CHECKING +from . import _client_backend from .client_types import ClientPhase, ClientStatus from .sender import TransactionRejectedError if TYPE_CHECKING: - from .receiver import ReceiverThread + from .receiver import EventReceiver + from .sender import EventSender LOG = logging.getLogger(__name__) DEFAULT_WAIT_TIMEOUT_S = 10.0 _POLL_INTERVAL_S = 0.01 +_PHASES = { + _client_backend.ClientPhase.OFFLINE: ClientPhase.OFFLINE, + _client_backend.ClientPhase.CONNECTING: ClientPhase.CONNECTING, + _client_backend.ClientPhase.REPLAYING: ClientPhase.REPLAYING, + _client_backend.ClientPhase.READY: ClientPhase.READY, + _client_backend.ClientPhase.RECOVERY_REQUIRED: ClientPhase.RECOVERY_REQUIRED, + _client_backend.ClientPhase.REJECTED: ClientPhase.REJECTED, + _client_backend.ClientPhase.CLOSED: ClientPhase.CLOSED, + _client_backend.ClientPhase.PARKED: ClientPhase.PARKED, +} + def deadline_after(timeout: float | None) -> float | None: return None if timeout is None else time.monotonic() + max(timeout, 0.0) @@ -48,21 +58,17 @@ def compute_phase( connecting: bool, ) -> ClientPhase: """The one precedence order every client uses for ``ClientStatus.phase``.""" - if closed: - return ClientPhase.CLOSED - if recovery_required: - return ClientPhase.RECOVERY_REQUIRED - if rejected: - return ClientPhase.REJECTED - if parked: - return ClientPhase.PARKED - if replaying: - return ClientPhase.REPLAYING - if ready: - return ClientPhase.READY - if connecting: - return ClientPhase.CONNECTING - return ClientPhase.OFFLINE + return _PHASES[ + _client_backend.compute_phase( + closed=closed, + recovery_required=recovery_required, + rejected=rejected, + parked=parked, + replaying=replaying, + ready=ready, + connecting=connecting, + ) + ] def raise_if_blocked(client, status: ClientStatus) -> None: @@ -115,82 +121,6 @@ def submit_and_wait(client, timeout: float | None) -> bool: return False -class BacklogHold: - """Hold a local batch until the messages queued before it have been drained. - - Counting those messages, rather than checking whether a drain used its - whole budget, keeps sustained inbound traffic from holding edits forever. - """ - - __slots__ = ("_ahead",) - - def __init__(self): - self._ahead = 0 - - @property - def holding(self) -> bool: - return self._ahead > 0 - - def freeze(self, queued: int) -> None: - """Record the queue depth when a new local batch is frozen.""" - self._ahead = queued - - def drained(self, count: int, queued: int) -> None: - # The messages ahead of the batch are at the front of the queue, so a - # replay request that discards the queue also bounds them. - self._ahead = min(max(0, self._ahead - count), queued) - - -class ClientCallbackQueue: - """Deliver notifications raised on network threads during update() or close().""" - - def __init__(self): - self._queue = queue.SimpleQueue() - self._lock = threading.Lock() - self._closed = False - - def wrap(self, callback: Callable) -> Callable: - def enqueue(value): - with self._lock: - if not self._closed: - self._queue.put((callback, value)) - - return enqueue - - def drain(self) -> None: - # Only notifications queued before this tick, so a busy receiver cannot - # starve update(). - for _ in range(self._queue.qsize()): - try: - callback, value = self._queue.get_nowait() - except queue.Empty: - break - callback(value) - - def close(self) -> None: - """Refuse new notifications and deliver the queued ones. - - Every queued notification runs even if one raises; the first error is - re-raised after the rest have been delivered. - """ - with self._lock: - self._closed = True - error = None - while True: - try: - callback, value = self._queue.get_nowait() - except queue.Empty: - break - try: - callback(value) - except Exception as exc: - if error is not None: - LOG.exception("Observer notification failed while closing") - error = error or exc - if error is not None: - raise error - - def raise_if_rejected(endpoint, role: str) -> None: if endpoint.auth_rejected: raise PermissionError(f"{role} authentication rejected") @@ -198,9 +128,6 @@ def raise_if_rejected(endpoint, role: str) -> None: raise ConnectionError(endpoint.rejection_reason or f"{role} connection rejected") -def stop_receiver(receiver: ReceiverThread) -> None: - receiver.stop() - if receiver.is_alive() and receiver is not threading.current_thread(): - receiver.join(timeout=2.0) - if receiver.is_alive(): - LOG.warning("Receiver thread did not stop within 2 seconds") +def close_endpoint(endpoint: EventReceiver | EventSender) -> None: + if not endpoint.close(timeout=2.0): + LOG.warning("%s did not stop within 2 seconds", type(endpoint).__name__) diff --git a/openusdconnect/_client_utils.py b/openusdconnect/_client_utils.py index a15cd0c..36f149a 100644 --- a/openusdconnect/_client_utils.py +++ b/openusdconnect/_client_utils.py @@ -4,7 +4,6 @@ import threading import uuid -from collections.abc import Callable from .client_types import ClientPhase, ClientStatus, SyncUpdate from .token_client import load_token, save_token @@ -53,18 +52,10 @@ def resolve_client_token( class ClientCredential: """The one token both roles of a client present, persisted when enabled.""" - def __init__( - self, - host: str, - port: int, - token: str | None, - persist: bool, - on_issued: Callable[[str], None] | None = None, - ): + def __init__(self, host: str, port: int, token: str | None, persist: bool): self._host = host self._port = port self._persist = persist - self._on_issued = on_issued # Both connection threads read and replace the token. self._lock = threading.Lock() self.token = resolve_client_token(host, port, token, persist) @@ -77,13 +68,11 @@ def current(self) -> str | None: return self.token def issued(self, token: str) -> None: - """Adopt a server-issued token, persist it, then notify the host.""" + """Adopt a server-issued token and persist it.""" with self._lock: self.token = token if self._persist: save_token(self._host, self._port, token) - if self._on_issued is not None: - self._on_issued(token) def endpoint_kwargs(self) -> dict: """Keyword arguments that make an endpoint present and report this token.""" diff --git a/openusdconnect/_observer_hooks.py b/openusdconnect/_observer_hooks.py deleted file mode 100644 index 8f59b4c..0000000 --- a/openusdconnect/_observer_hooks.py +++ /dev/null @@ -1,121 +0,0 @@ -"""Translate a ClientObserver into the callables the low-level endpoints take.""" - -from __future__ import annotations - -from collections.abc import Callable -from dataclasses import dataclass - -from .client_observer import ( - AppliedBatch, - ClientObserver, - PlaybackClaim, - PlaybackState, - StageMetadata, -) - -_STAGE_METADATA_FIELDS = { - "timeCodesPerSecond": "time_codes_per_second", - "framesPerSecond": "frames_per_second", - "startTimeCode": "start_time_code", - "endTimeCode": "end_time_code", - "metersPerUnit": "meters_per_unit", - "upAxis": "up_axis", -} - - -def stage_metadata_from_message(message: dict) -> StageMetadata: - return StageMetadata(**{ - field: message[key] for key, field in _STAGE_METADATA_FIELDS.items() if key in message - }) - - -def _playback_state(message: dict) -> PlaybackState: - return PlaybackState( - playing=bool(message["playing"]), - time=float(message["time"]), - rate=float(message["rate"]), - leader_client_id=message["leader_client_id"], - ) - - -def _claim_granted(message: dict) -> PlaybackClaim: - return PlaybackClaim(granted=True, leader_client_id=message["leader_client_id"]) - - -def _claim_rejected(message: dict) -> PlaybackClaim: - return PlaybackClaim( - granted=False, - leader_client_id=message["current_leader_client_id"], - reason=message["reason"], - ) - - -@dataclass(frozen=True, slots=True) -class ObserverHooks: - """Callables for the observer methods a subclass overrides, else ``None``.""" - - on_applied: Callable[[int, list], None] | None = None - on_resync: Callable[[], None] | None = None - on_token_issued: Callable[[str], None] | None = None - on_stage_metadata: Callable[[dict], None] | None = None - on_playback_state: Callable[[dict], None] | None = None - on_playback_claimed: Callable[[dict], None] | None = None - on_playback_rejected: Callable[[dict], None] | None = None - - def receiver_callbacks(self) -> dict[str, Callable | None]: - """Keyword arguments for ``ReceiverThread``.""" - return { - "on_stage_metadata": self.on_stage_metadata, - "on_playback_state": self.on_playback_state, - "on_playback_claimed": self.on_playback_claimed, - "on_playback_rejected": self.on_playback_rejected, - } - - def applied_events_for(self, dispatcher) -> Callable[[list], None] | None: - """The ``on_applied_events`` callback reporting batches of *dispatcher*.""" - on_applied = self.on_applied - if on_applied is None: - return None - return lambda events: on_applied(dispatcher.applying_seq, events) - - -def observer_hooks( - observer: ClientObserver | None, - notify: Callable[[Callable], Callable], -) -> ObserverHooks: - """Hooks for *observer*; *notify* defers a notification until ``update()``.""" - if observer is None: - return ObserverHooks() - if not isinstance(observer, ClientObserver): - raise TypeError("observer must be a ClientObserver") - - def overrides(name: str) -> bool: - return getattr(type(observer), name) is not getattr(ClientObserver, name) - - claims = overrides("on_playback_claim") - # Delivery methods already run inside update(); notifications arrive on - # network threads and go through notify. - return ObserverHooks( - on_applied=( - (lambda seq, events: observer.on_applied(AppliedBatch(seq, events))) - if overrides("on_applied") else None - ), - on_resync=observer.on_resync if overrides("on_resync") else None, - on_token_issued=( - notify(observer.on_token_issued) if overrides("on_token_issued") else None - ), - on_stage_metadata=( - notify(lambda m: observer.on_stage_metadata(stage_metadata_from_message(m))) - if overrides("on_stage_metadata") else None - ), - on_playback_state=( - notify(lambda m: observer.on_playback_state(_playback_state(m))) - if overrides("on_playback_state") else None - ), - on_playback_claimed=( - notify(lambda m: observer.on_playback_claim(_claim_granted(m))) if claims else None - ), - on_playback_rejected=( - notify(lambda m: observer.on_playback_claim(_claim_rejected(m))) if claims else None - ), - ) diff --git a/openusdconnect/client_observer.py b/openusdconnect/client_observer.py index 1aade23..9aa3d07 100644 --- a/openusdconnect/client_observer.py +++ b/openusdconnect/client_observer.py @@ -89,7 +89,7 @@ def on_resync(self) -> None: """Delivery (ManagedClient, UsdReceiver): the stream restarted; reset host state.""" def on_stage_metadata(self, metadata: StageMetadata) -> None: - """Notification: stage settings received with a handshake.""" + """Notification: the server's stage metadata, delivered when it changes.""" def on_playback_state(self, state: PlaybackState) -> None: """Notification (receiving clients): the shared playhead changed.""" diff --git a/openusdconnect/dispatcher.py b/openusdconnect/dispatcher.py index edb8b3d..535670f 100644 --- a/openusdconnect/dispatcher.py +++ b/openusdconnect/dispatcher.py @@ -39,7 +39,7 @@ from .codec import ReceivedEvent from .emitter import NoticeEmitter from .layer_key_router import LayerKeyRouter - from .receiver import ReceiverThread + from .receiver import EventReceiver LOG = logging.getLogger(__name__) @@ -380,7 +380,7 @@ class EventDispatcher: def __init__( self, *, - receiver: ReceiverThread, + receiver: EventReceiver, adapter: DCCAdapter, mirror_stage: Usd.Stage | None = None, emitter: NoticeEmitter | None = None, @@ -455,6 +455,7 @@ def drain_and_apply(self, *, max_messages: int | None = None) -> int: """ if self._projection_state is not None: self._projection_state.ensure_native_projection_safe() + generation = self.receiver.generation bufs = ( self.receiver.drain_queue() if max_messages is None @@ -516,6 +517,9 @@ def drain_and_apply(self, *, max_messages: int | None = None) -> int: if result.errors: self.receiver.request_replay_from(self._last_seq + 1) else: + if result.resync_requested: + self.receiver.reset_applied_progress() + self.receiver.mark_applied_through(generation, self._last_seq) self.receiver.mark_replay_applied() return applied diff --git a/openusdconnect/managed_client.py b/openusdconnect/managed_client.py index faa78af..7330eda 100644 --- a/openusdconnect/managed_client.py +++ b/openusdconnect/managed_client.py @@ -12,12 +12,7 @@ from pxr import Sdf, Usd from ._client_base import EmitterClientBase -from ._client_lifecycle import ( - DEFAULT_WAIT_TIMEOUT_S, - BacklogHold, - deadline_after, - remaining_time, -) +from ._client_lifecycle import DEFAULT_WAIT_TIMEOUT_S, deadline_after, remaining_time from ._client_utils import client_origin, require_app_name, validate_layered_source from .adapters import UsdStageAdapter from .client_id import make_stable_client_id @@ -26,7 +21,7 @@ from .defaults import DEFAULT_HOST, DEFAULT_SYNC_PORT from .dispatcher import AssetDependencyRefreshResult, EventDispatcher from .emitter import PrimChannel -from .receiver import ReceiverThread +from .receiver import EventReceiver from .recovery import RecoveryArtifact, RecoveryError from .sender import EventSender @@ -65,7 +60,6 @@ def __init__( replicated_api_schemas: set[str] | None = None, extra_channels: Sequence[PrimChannel] | None = None, transform_coalesce_seconds: float = 0.0, - background_send: bool = False, ): app_name = require_app_name(app_name) if not isinstance(stage, Usd.Stage): @@ -78,7 +72,7 @@ def __init__( self._app_name = app_name self._authoring_layer: Sdf.Layer | None = None self._last_recovery_result: ManagedRecoveryResult | None = None - self._backlog = BacklogHold() + self._backlog_marker = 0 self._init_emitter( stage, attr_filter=attr_filter, @@ -92,20 +86,19 @@ def __init__( } credential = self._credential.endpoint_kwargs() self._sender = EventSender( - host, port, department=department, background_send=background_send, + host, port, department=department, notifications=self._notifications, **identity, **credential, ) - self._receiver = ReceiverThread( + self._receiver = EventReceiver( host=host, port=port, sync_from=1, reconnect=reconnect, layered_replay=True, - **identity, **credential, **self._hooks.receiver_callbacks(), + notifications=self._notifications, **identity, **credential, ) self._dispatcher = EventDispatcher( receiver=self._receiver, adapter=UsdStageAdapter(stage), emitter=self._emitter, - on_resync=self._hooks.on_resync, ) - self._dispatcher.on_applied_events = self._hooks.applied_events_for(self._dispatcher) + self._observe_dispatcher(self._dispatcher) # Modify the stage last so a failed construction leaves it untouched. with self._emitter.suppressed(): self._authoring_layer = self._create_authoring_layer(stage, app_name) @@ -121,8 +114,8 @@ def authoring_layer(self) -> Sdf.Layer | None: return self._authoring_layer @property - def receiver(self) -> ReceiverThread: - """The underlying :class:`ReceiverThread`; a diagnostic handle.""" + def receiver(self) -> EventReceiver: + """The underlying :class:`EventReceiver`; a diagnostic handle.""" return self._receiver @property @@ -165,18 +158,19 @@ def update(self, *, max_messages: int | None = None) -> SyncUpdate: had_batch = bool(self._emitter.prepared_event_count) outgoing = self._prepare_outgoing_events() if not had_batch and self._emitter.prepared_event_count: - self._backlog.freeze(self._receiver.queued_message_count) + self._backlog_marker = self._receiver.freeze_marker() received = self._apply_queued(max_messages) if self._closed: return self._progress(received) - self._backlog.drained( - self._dispatcher.drained_message_count, self._receiver.queued_message_count, - ) sent = 0 if self._receiver.connected and not self._sender.connected: self._sender.request_connect() - if self._sender.connected and self._is_synchronized() and not self._backlog.holding: + if ( + self._sender.connected + and self._is_synchronized() + and self._receiver.drained_through(self._backlog_marker) + ): sent = self._send(outgoing) return self._progress(received, sent) @@ -272,12 +266,8 @@ def recover_use_server( self._resume_sender_after_recovery(remaining_time(deadline)) return result - def _is_synchronized(self) -> bool: - return ( - self._stage is not None - and self._receiver.synchronized - and not self._sender.recovery_required - ) + def _synchronized(self, replayed: bool) -> bool: + return replayed and self._stage is not None and not self._sender.recovery_required def _is_parked(self) -> bool: return self._stage is None diff --git a/openusdconnect/receiver.py b/openusdconnect/receiver.py index 49303b0..2e13588 100644 --- a/openusdconnect/receiver.py +++ b/openusdconnect/receiver.py @@ -1,30 +1,20 @@ -"""Background TCP receiver with replay-aware reconnection.""" +"""Receive wire messages on a native thread for a stage-owning consumer to drain.""" from __future__ import annotations import logging -import socket -import threading -import time from collections import deque from collections.abc import Callable from . import _client_backend -from .codec import ( - HelloRejectionCode, - PayloadType, - _decode_stage_metadata_table, - decode_envelope, - encode_message, - message_to_dict, - payload_type_and_sequence, - resolve_payload, -) +from .codec import HelloRejectionCode from .defaults import DEFAULT_HOST, DEFAULT_SYNC_PORT -from .framing import IncompleteRead, MessageTooLarge, recv_framed -from .protocol import make_hello -from .protocol_constants import LayerMode -from .transport import send_msg +from .protocol_constants import ( + MSG_PLAYBACK_CLAIMED, + MSG_PLAYBACK_REJECTED, + MSG_PLAYBACK_STATE, + LayerMode, +) LOG = logging.getLogger(__name__) @@ -32,28 +22,22 @@ _RECONNECT_MAX_DELAY = 30.0 _SOCKET_TIMEOUT = 30.0 _MAX_QUEUE_DEPTH = 50_000 -_MAX_CONSECUTIVE_TIMEOUTS = 10 - -_MESSAGE_KIND_BY_PAYLOAD = { - PayloadType.BroadcastEvent: _client_backend.ReceiverMessageKind.EVENT, - PayloadType.LayerGraphState: _client_backend.ReceiverMessageKind.LAYER_GRAPH_STATE, - PayloadType.Resync: _client_backend.ReceiverMessageKind.RESYNC, -} -_PLAYBACK_PAYLOAD_TYPES = frozenset( - { - PayloadType.PlaybackState, - PayloadType.PlaybackClaimed, - PayloadType.PlaybackRejected, - } -) -class ReceiverThread(threading.Thread): - """Receive wire messages off-thread for a stage-owning consumer to drain. +def _transport_error(failure) -> OSError: + if failure.result == _client_backend.SocketResult.TIMEOUT: + return TimeoutError(failure.description) + return OSError(failure.system_error, failure.description) + - Scene events are queued as raw FlatBuffers for consumer-thread decoding - and USD mutation. Handshake and control messages, including their callbacks, - are processed on the receiver thread. Overflow closes the connection and +class EventReceiver: + """Receive wire messages for a stage-owning consumer to drain. + + A native thread runs the connection until :meth:`close` and queues scene + messages as raw FlatBuffers for the consumer thread to decode and apply. + Handshake and control messages, including their callbacks, are handled on + that thread; given ``notifications``, the owner drains that queue instead + and only ``on_token_issued`` runs there. Overflow closes the connection and resumes by replay after the queue drains or the drain wait expires. """ @@ -79,507 +63,241 @@ def __init__( on_playback_rejected: Callable[[dict], None] | None = None, layered_replay: bool = True, layer_mode: LayerMode | str = LayerMode.MANAGED, + notifications: _client_backend.NotificationQueue | None = None, ): - super().__init__(daemon=True) - self.host = host - self.port = port - self.sync_from = sync_from - self.reconnect = reconnect - self.max_queue = max_queue - self.socket_timeout = socket_timeout - self.client_id = client_id - self.origin = origin - self.department = department + self._host = host + self._port = port + self._sync_from = sync_from + self._max_queue = max_queue + self._socket_timeout = socket_timeout + self._client_id = client_id + self._origin = origin + self._department = department + self._layered_replay = bool(layered_replay) + self._layer_mode = LayerMode(layer_mode) + self._reconnect = bool(reconnect) self.token = token - self.layered_replay = bool(layered_replay) - self.layered_replay_active = False - self.layer_mode = LayerMode(layer_mode) - self.layer_mode_active = LayerMode.MANAGED self._token_provider = token_provider + self._token_error: Exception | None = None self._on_token_issued = on_token_issued self._on_stage_metadata = on_stage_metadata self._on_playback_state = on_playback_state self._on_playback_claimed = on_playback_claimed self._on_playback_rejected = on_playback_rejected - self._reconnect_base_delay = reconnect_base_delay - self._reconnect_max_delay = reconnect_max_delay - self._stop_event = threading.Event() - self.sock: socket.socket | None = None - self._socket_lock = threading.Lock() - self._inbox = _client_backend.ReceiverInbox(sync_from, max_queue) - self._connected_event = threading.Event() - self._synchronized_event = threading.Event() - self._handshake_event = threading.Event() - self._replay_lock = threading.RLock() - self._received_replay_identity: tuple[str, int] | None = None - self._initial_replay_identity: tuple[str, int] | None = None - self._connection_server_instance = "" - self._replay_identity_supported = False - self._hello_sent = False - self._prefix_validation_requested = False - self._connection_prefix_proven = False - self._reset_on_hello = False - self.replay_head_seq = 0 - self.replay_epoch = 0 - self.server_instance = "" - self.auth_rejected = False - self.hello_rejected = False - self.rejection_code = HelloRejectionCode.Unspecified - self.rejection_reason = "" - self.connection_error: Exception | None = None - self.stage_metadata: dict = {} + + config = _client_backend.ReceiverConfig() + config.host = host + config.port = port + config.client_id = client_id or "" + config.origin = origin or "" + config.department = department or "" + config.layered_replay = self._layered_replay + config.layer_mode = _client_backend.NATIVE_LAYER_MODES[self._layer_mode] + config.sync_from = sync_from + config.max_queue = max_queue + config.socket_timeout = socket_timeout + config.reconnect = self._reconnect + config.reconnect_base_delay = reconnect_base_delay + config.reconnect_max_delay = reconnect_max_delay + self._notifications, self._owns_notifications = _client_backend.notification_queue( + notifications, + on_stage_metadata=on_stage_metadata, + on_playback_state=on_playback_state, + on_playback_claimed=on_playback_claimed, + on_playback_rejected=on_playback_rejected, + ) + self._endpoint = _client_backend.ReceiverEndpoint(config, self._notifications) + self._driver = None + + @property + def host(self) -> str: + return self._host + + @property + def port(self) -> int: + return self._port + + @property + def sync_from(self) -> int: + """The first sequence requested; the consumer already holds every earlier one.""" + return self._sync_from + + @property + def max_queue(self) -> int: + return self._max_queue + + @property + def socket_timeout(self) -> float: + return self._socket_timeout + + @property + def client_id(self) -> str | None: + return self._client_id + + @property + def origin(self) -> str | None: + return self._origin + + @property + def department(self) -> str | None: + return self._department + + @property + def layered_replay(self) -> bool: + return self._layered_replay + + @property + def layer_mode(self) -> LayerMode: + return self._layer_mode + + @property + def reconnect(self) -> bool: + """Whether a lost connection is retried; a change applies when it ends.""" + return self._reconnect + + @reconnect.setter + def reconnect(self, enabled: bool) -> None: + self._reconnect = bool(enabled) + self._endpoint.set_reconnect(self._reconnect) @property def connected(self) -> bool: - return self._connected_event.is_set() - - @connected.setter - def connected(self, value: bool) -> None: - with self._replay_lock: - if value: - self._connected_event.set() - else: - self._connected_event.clear() - self._synchronized_event.clear() - self._inbox.disconnect(self._inbox.generation) + return self._endpoint.status().connected @property def synchronized(self) -> bool: """Whether replay through the server's advertised head was applied.""" - return self.connected and self._synchronized_event.is_set() + return self._endpoint.status().synchronized @property def last_seq(self) -> int: - return self._inbox.last_sequence + return self._endpoint.status().last_sequence + + @property + def running(self) -> bool: + """Whether the connection thread has started and not yet exited.""" + return self._driver is not None and self._driver.running @property def stopped(self) -> bool: - """Whether the thread ran and exited, so it will not reconnect.""" - return self.ident is not None and not self.is_alive() + """Whether the connection thread ran and exited, so it will not reconnect.""" + return self._driver is not None and self._driver.stopped @property def queued_message_count(self) -> int: """Number of received messages waiting for the owning thread to drain.""" + return self._endpoint.status().queued_frames - return self._inbox.size + @property + def generation(self) -> int: + """Connection generation; read it before draining for :meth:`mark_applied_through`.""" + return self._endpoint.generation - def wait_connected(self, timeout: float | None = None) -> bool: - """Wait for the current handshake result, not for replay completion.""" - if self.connected: - return True - self._handshake_event.wait(timeout=timeout) - return self.connected + @property + def layered_replay_active(self) -> bool: + return self._endpoint.status().layered_replay_active - def wait_synchronized(self, timeout: float | None = None) -> bool: - """Wait for replay to be applied by the stage-owning consumer thread.""" - if self.synchronized: - return True - self._synchronized_event.wait(timeout=timeout) - return self.synchronized + @property + def layer_mode_active(self) -> LayerMode: + return _client_backend.LAYER_MODES[self._endpoint.status().layer_mode_active] - def mark_replay_applied(self) -> bool: - """Publish READY after a successful drain applied the replay prefix.""" - with self._replay_lock: - if not self._inbox.mark_replay_applied(): - return False - self.replay_head_seq = self._inbox.replay_head_sequence - self.replay_epoch = self._inbox.replay_epoch - # Publish the applied identity, not a new socket's un-applied Hello. - self.server_instance = ( - self._received_replay_identity[0] if self._received_replay_identity else "" - ) - self._synchronized_event.set() - return True + @property + def auth_rejected(self) -> bool: + rejection = self._endpoint.status().rejection + return rejection is not None and rejection.authentication - def run(self) -> None: - delay = self._reconnect_base_delay - while not self._stop_event.is_set(): - self.connection_error = None - was_connected = False - try: - self._connect_and_recv() - except Exception as exc: - self.connection_error = exc - if not self._stop_event.is_set(): - LOG.exception("ReceiverThread: connection error") - finally: - was_connected = self.connected - self.connected = False - self._close_socket() - - if self._should_stop_reconnecting(): - # Always release waiters when this thread terminates. - self._handshake_event.set() - break - - self._handshake_event.clear() - if was_connected: - delay = self._reconnect_base_delay - - if self._inbox.overflowed: - self._inbox.clear_overflow() - delay = self._reconnect_base_delay - self._wait_for_queue_drain() - continue - - LOG.info("ReceiverThread: reconnecting in %.1fs", delay) - if self._stop_event.wait(timeout=delay): - break - delay = min(delay * 2, self._reconnect_max_delay) - - LOG.info("ReceiverThread stopped") - - def _should_stop_reconnecting(self) -> bool: - return ( - not self.reconnect - or self._stop_event.is_set() - or self.auth_rejected - or self.hello_rejected - ) + @property + def hello_rejected(self) -> bool: + rejection = self._endpoint.status().rejection + return rejection is not None and not rejection.authentication - def _wait_for_queue_drain(self) -> None: - LOG.info("ReceiverThread: waiting for queue to drain before reconnect") - deadline = time.monotonic() + self._reconnect_max_delay - while self._inbox.size and not self._stop_event.wait(timeout=0.1): - if time.monotonic() >= deadline: - LOG.warning("ReceiverThread: drain wait timed out, reconnecting anyway") - return - - def _connect_and_recv(self) -> None: - """Single connection attempt: connect, handshake, read until EOF/error.""" - LOG.info("ReceiverThread connecting to %s:%s", self.host, self.port) - sock = socket.create_connection( - (self.host, self.port), - timeout=self.socket_timeout, - ) - with self._socket_lock: - self.sock = sock - sock.settimeout(self.socket_timeout) - # Send small handshake messages promptly. - sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) - - if self._stop_event.is_set(): - self._close_socket(sock) - return + @property + def rejection_code(self) -> int: + rejection = self._endpoint.status().rejection + return HelloRejectionCode.Unspecified if rejection is None else rejection.code - with self._replay_lock: - connection = self._inbox.begin_connection() - connection_generation = connection.generation - sync_from = connection.sync_from - prefix_identity = self._received_replay_identity - self._synchronized_event.clear() + @property + def rejection_reason(self) -> str: + rejection = self._endpoint.status().rejection + return "" if rejection is None else _client_backend.rejection_reason(rejection) - if self._token_provider is not None: - self.token = self._token_provider() - hello = make_hello( - "receiver", - sync_from=sync_from, - client_id=self.client_id, - origin=self.origin, - department=self.department, - token=self.token, - layered_replay=self.layered_replay, - layer_mode=self.layer_mode, - ) - # Initial explicit cursors may refer to an externally supplied snapshot. - # Preserve that contract, but do not claim its unvalidated prefix as proof. - self._prefix_validation_requested = self._hello_sent - if self._prefix_validation_requested: - hello["replay_server_instance"] = prefix_identity[0] if prefix_identity else "" - if prefix_identity is not None: - hello["replay_epoch"] = prefix_identity[1] - self._hello_sent = True - send_msg(sock, hello) - self._handshake_event.clear() - self.auth_rejected = False - self.hello_rejected = False - self.rejection_code = HelloRejectionCode.Unspecified - self.rejection_reason = "" - - consecutive_timeouts = 0 - while not self._stop_event.is_set(): - try: - buf = recv_framed(sock) - except TimeoutError: - consecutive_timeouts += 1 - if consecutive_timeouts >= _MAX_CONSECUTIVE_TIMEOUTS: - LOG.warning( - "ReceiverThread: %d consecutive timeouts, reconnecting", - consecutive_timeouts, - ) - break - LOG.debug( - "ReceiverThread: recv timeout (%d/%d)", - consecutive_timeouts, - _MAX_CONSECUTIVE_TIMEOUTS, - ) - continue - except (IncompleteRead, MessageTooLarge): - current_generation = connection_generation == self._inbox.generation - if not self._stop_event.is_set() and current_generation: - LOG.warning("ReceiverThread: framing error during read") - break - except OSError: - current_generation = connection_generation == self._inbox.generation - if not self._stop_event.is_set() and current_generation: - LOG.warning("ReceiverThread: socket error during read") - break - - consecutive_timeouts = 0 - - if connection_generation != self._inbox.generation: - return - - if not self.connected: - if not self._handle_handshake_message(buf, sync_from, connection_generation): - return - continue - - if not self._handle_data_message(buf, connection_generation): - return + @property + def server_instance(self) -> str: + """The server whose replay the consumer applied; empty when unproven.""" + return self._endpoint.status().server_instance - @staticmethod - def _decode_text(value: str | bytes | None) -> str: - if isinstance(value, bytes): - return value.decode("utf-8") - return value or "" + @property + def replay_epoch(self) -> int: + return self._endpoint.status().replay_epoch - @staticmethod - def _invoke_callback(callback: Callable | None, value, name: str) -> None: - if callback is None: - return - try: - callback(value) - except Exception: - LOG.exception("ReceiverThread: %s callback failed", name) - - def _reject_hello(self, code: int, reason: str) -> None: - self.hello_rejected = True - self.rejection_code = code - self.rejection_reason = reason - LOG.error("ReceiverThread: connection rejected (%s): %s", code, reason) - self._handshake_event.set() - - def _connection_replay_identity(self, epoch: int | None) -> tuple[str, int] | None: - """Return an identity only when the peer supplied all negotiated fields.""" - if not self._replay_identity_supported or not self._connection_server_instance: - return None - if epoch is None: - return None - return self._connection_server_instance, int(epoch) - - def _handle_handshake_message( - self, buf: bytes, sync_from: int, connection_generation: int | None = None, - ) -> bool: - if connection_generation is None: - connection_generation = self._inbox.generation - env = decode_envelope(buf) - payload_type = env.PayloadType() - - if payload_type == PayloadType.AuthRejected: - _, rejection = resolve_payload(env) - reason = self._decode_text(rejection.Reason()) - LOG.error("ReceiverThread: auth rejected %s", reason) - self.auth_rejected = True - self.rejection_reason = reason - self._handshake_event.set() - return False - - if payload_type == PayloadType.HelloRejected: - _, rejection = resolve_payload(env) - self._reject_hello( - int(rejection.Code()), - self._decode_text(rejection.Reason()), - ) - return False + @property + def replay_head_seq(self) -> int: + return self._endpoint.status().replay_head_sequence - if payload_type != PayloadType.HelloOk: - return True + @property + def stage_metadata(self) -> dict: + """The latest stage metadata the server authored, keyed as on the wire.""" + return _client_backend.stage_metadata_fields(self._endpoint.status().metadata) - _, hello = resolve_payload(env) - self._connection_server_instance = self._decode_text(hello.ServerInstance()) or "" - self._replay_identity_supported = bool(hello.ReplayIdentity()) - self._connection_prefix_proven = sync_from == 1 or self._prefix_validation_requested - self.layer_mode_active = LayerMode.SHARED_STAGE if hello.LayerMode() else LayerMode.MANAGED - if self.layer_mode_active is not self.layer_mode: - self._reject_hello( - HelloRejectionCode.LayerModeMismatch, - "server did not negotiate requested layer mode", - ) - return False + @property + def connection_error(self) -> Exception | None: + """Why the latest connection attempt failed, or ``None``.""" + failure = None if self._driver is None else self._driver.last_failure + return self._token_error if failure is None else _transport_error(failure) - self.layered_replay_active = bool(self.layered_replay and hello.LayeredReplay()) - if self.layered_replay and not self.layered_replay_active: - self._reject_hello( - HelloRejectionCode.LayeredReplayRequired, - "server did not negotiate requested layered replay", - ) - return False - - issued_token = self._decode_text(hello.Token()) - if issued_token: - self.token = issued_token - self._invoke_callback(self._on_token_issued, issued_token, "on_token_issued") - LOG.info("ReceiverThread: token issued by server") - - metadata_table = hello.StageMetadata() - if metadata_table is not None: - metadata = _decode_stage_metadata_table(metadata_table) - if metadata: - self.stage_metadata = metadata - self._invoke_callback(self._on_stage_metadata, metadata, "on_stage_metadata") - - with self._replay_lock: - if connection_generation != self._inbox.generation: - return False - identity = self._connection_replay_identity(hello.ReplayEpoch()) - # This proof belongs only to the handshake's replay, not later live resets. - self._initial_replay_identity = identity - if identity is None: - self._received_replay_identity = None - if self._reset_on_hello: - if not self._handle_data_message( - encode_message({"type": "resync"}), connection_generation, - ): - return False - self._reset_on_hello = False - prefix_matches = ( - self._prefix_validation_requested and self._received_replay_identity == identity - ) - if identity is not None and (sync_from == 1 or prefix_matches): - self._received_replay_identity = identity - # A changed/unknown prefix keeps its old identity until Resync is accepted. - self.connected = True - self._handshake_event.set() - LOG.info("ReceiverThread connected (sync_from=%d)", sync_from) - return True - - def _handle_control_message( - self, - payload_type: int, - buf: bytes, - connection_generation: int, - ) -> bool | None: - if payload_type == PayloadType.Ping: - return True + def snapshot(self) -> _client_backend.ReceiverStatus: + """The native status in one call, for reading several fields together.""" + return self._endpoint.status() - if payload_type == PayloadType.ReplayComplete: - _, complete = resolve_payload(decode_envelope(buf)) - head_seq = int(complete.HeadSeq()) - epoch = int(complete.Epoch()) - with self._replay_lock: - result = self._inbox.accept_replay_complete( - connection_generation, - head_seq, - epoch, - ) - if result == _client_backend.AcceptResult.ACCEPTED: - self._initial_replay_identity = None - # This covers received/queued frames, even before their consumer - # drains them, including resets without a handshake epoch. - self._received_replay_identity = None - if self._connection_prefix_proven: - self._received_replay_identity = self._connection_replay_identity(epoch) - if result == _client_backend.AcceptResult.STALE_GENERATION: - return False - return True + def mark_applied_through(self, generation: int, sequence: int) -> bool: + """Advance the applied cursor for frames drained after reading ``generation``. - if payload_type not in _PLAYBACK_PAYLOAD_TYPES: - return None - - if payload_type == PayloadType.PlaybackState: - callback = self._on_playback_state - elif payload_type == PayloadType.PlaybackClaimed: - callback = self._on_playback_claimed - else: - callback = self._on_playback_rejected - self._invoke_callback(callback, message_to_dict(buf), "playback") - return True - - def _handle_data_message(self, buf: bytes, connection_generation: int) -> bool: - payload_type, sequence = payload_type_and_sequence(buf) - handled = self._handle_control_message(payload_type, buf, connection_generation) - if handled is not None: - return handled - - kind = _MESSAGE_KIND_BY_PAYLOAD.get( - payload_type, - _client_backend.ReceiverMessageKind.OTHER, - ) - with self._replay_lock: - result = self._inbox.accept(connection_generation, kind, sequence, buf) - if ( - kind == _client_backend.ReceiverMessageKind.RESYNC - and result == _client_backend.AcceptResult.ACCEPTED - ): - self._synchronized_event.clear() - self._received_replay_identity = self._initial_replay_identity - self._initial_replay_identity = None - self._connection_prefix_proven = True - if result == _client_backend.AcceptResult.STALE_GENERATION: - return False - if result == _client_backend.AcceptResult.DUPLICATE: - return True - if result == _client_backend.AcceptResult.SEQUENCE_GAP: - replay_from = self._inbox.last_applied_sequence + 1 - LOG.error( - "ReceiverThread: sequence gap before %d; replaying from applied %d", - sequence, - replay_from, - ) - self.request_replay_from(replay_from) - return False - if result == _client_backend.AcceptResult.QUEUE_FULL: - LOG.warning( - "ReceiverThread: queue full (%d), disconnecting to replay from server", - self.max_queue, - ) - return False - return True + A live sequence gap replays from this cursor, so report a batch only + after its whole apply pipeline succeeded. + """ + return self._endpoint.mark_applied_through(generation, sequence) + + def reset_applied_progress(self) -> None: + """Restart the applied cursor once the consumer has applied a queued resync.""" + self._endpoint.reset_applied_progress() + + def freeze_marker(self) -> int: + """Return a marker covering every message queued now; see :meth:`drained_through`.""" + return self._endpoint.freeze_marker() + + def drained_through(self, marker: int) -> bool: + """Whether every message queued when ``marker`` was taken has been drained.""" + return self._endpoint.drained_through(marker) + + def wait_connected(self, timeout: float | None = None) -> bool: + """Wait for the current handshake result, not for replay completion.""" + if self._driver is None: + return self.connected + return self._driver.wait_connected(timeout) + + def wait_synchronized(self, timeout: float | None = None) -> bool: + """Wait for replay to be applied by the stage-owning consumer thread.""" + if self._driver is None: + return self.synchronized + return self._driver.wait_synchronized(timeout) + + def mark_replay_applied(self) -> bool: + """Publish READY after a successful drain applied the replay prefix.""" + applied = self._endpoint.mark_replay_applied() + if applied: + self._wake() + return applied def request_replay_from(self, seq_start: int) -> None: """Reconnect and request replay beginning at ``seq_start``. Consumers call this after a queued frame fails to decode. Frames queued - after the failed frame are discarded, and the connection generation - prevents an in-flight read from adding more stale frames. + after the failed frame are discarded, and the current connection is + replaced so it cannot queue more stale frames. """ - seq_start = int(seq_start) - if seq_start < 1: + if not self._endpoint.request_replay_from(int(seq_start)): raise ValueError("replay sequence must be at least 1") - - with self._replay_lock: - self._inbox.request_replay_from(seq_start) - self._synchronized_event.clear() - # Discarded frames may include an unapplied reset. Their received - # identity cannot prove the consumer's retained prefix, at any cursor. - self._received_replay_identity = None - self._initial_replay_identity = None - if seq_start == 1: - self._reset_on_hello = True - self._close_socket() - - def _close_socket(self, sock: socket.socket | None = None) -> None: - """Detach and close a socket, logging cleanup errors at debug level.""" - with self._socket_lock: - if sock is None: - sock = self.sock - if self.sock is sock: - self.sock = None - - if sock is None: - return - try: - sock.shutdown(socket.SHUT_RDWR) - except OSError: - LOG.debug( - "ReceiverThread: socket shutdown failed during close", - exc_info=True, - ) - try: - sock.close() - except OSError: - LOG.debug("ReceiverThread: socket close failed", exc_info=True) + self._wake() def drain_queue(self, max_messages: int | None = None) -> deque: """Drain queued wire messages, optionally limiting work for this tick.""" @@ -587,10 +305,97 @@ def drain_queue(self, max_messages: int | None = None) -> deque: isinstance(max_messages, bool) or not isinstance(max_messages, int) or max_messages < 1 ): raise ValueError("max_messages must be a positive integer or None") - return deque(self._inbox.drain(max_messages)) + return deque(self._endpoint.drain_frames(max_messages)) + + def start(self) -> None: + """Start connecting on a native thread that runs until closed or collected; once only.""" + if self._driver is not None: + raise RuntimeError("a receiver can only be started once") + self._driver = _client_backend.ReceiverDriver( + self._endpoint, + self._notifications, + _client_backend.TcpSocketFactory(), + **_client_backend.driver_callbacks(self, LOG, sink=self._owns_notifications), + ) + if not self._driver.start(): + raise RuntimeError("could not start the receiver's connection thread") + + def close(self, timeout: float | None = None) -> bool: + """Stop the connection thread and close its connection; repeated calls are harmless. + + Waits up to ``timeout`` seconds for the thread to exit, without limit + for ``None``, and returns whether it has exited: ``True`` before + :meth:`start`, ``False`` on timeout or from a callback, which runs on + that thread. + """ + if self._driver is None: + self._endpoint.stop() + return True + return self._driver.close(timeout) + + def __enter__(self) -> EventReceiver: + return self - def stop(self) -> None: - """Request clean shutdown.""" - self._stop_event.set() - self._handshake_event.set() - self._close_socket() + def __exit__(self, exc_type, exc, traceback) -> None: + self.close() + + def _wake(self) -> None: + if self._driver is not None: + self._driver.wake() + + def _connection_token(self) -> str | None: + """The token for the next handshake; ``None`` abandons the attempt.""" + self._token_error = None + if self._token_provider is not None: + try: + self.token = self._token_provider() + except Exception as exc: + LOG.exception("EventReceiver: token provider failed") + self._token_error = exc + return None + return self.token or "" + + def _token_issued(self, token: str) -> None: + """Adopt a token the server issued, on the connection thread.""" + self.token = token + self._notify(self._on_token_issued, token, "on_token_issued") + + def _deliver(self, notification) -> None: + """Run the callback for a notification on the connection thread.""" + if isinstance(notification, _client_backend.StageMetadata): + self._notify( + self._on_stage_metadata, + _client_backend.stage_metadata_fields(notification), + "on_stage_metadata", + ) + elif isinstance(notification, _client_backend.PlaybackState): + message = { + "type": MSG_PLAYBACK_STATE, + "time": notification.time, + "playing": notification.playing, + "rate": notification.rate, + "leader_client_id": notification.leader_client_id, + } + self._notify(self._on_playback_state, message, "playback") + elif isinstance(notification, _client_backend.PlaybackClaimed): + message = { + "type": MSG_PLAYBACK_CLAIMED, + "leader_client_id": notification.leader_client_id, + } + self._notify(self._on_playback_claimed, message, "playback") + elif isinstance(notification, _client_backend.PlaybackRejected): + message = { + "type": MSG_PLAYBACK_REJECTED, + "reason": notification.reason, + "current_leader_client_id": notification.current_leader_client_id, + } + self._notify(self._on_playback_rejected, message, "playback") + + @staticmethod + def _notify(callback: Callable | None, value, name: str) -> None: + if callback is None: + return + try: + callback(value) + except Exception: + LOG.exception("EventReceiver: %s callback failed", name) diff --git a/openusdconnect/recovery.py b/openusdconnect/recovery.py index 6293065..b14bf09 100644 --- a/openusdconnect/recovery.py +++ b/openusdconnect/recovery.py @@ -5,7 +5,7 @@ from dataclasses import dataclass from enum import StrEnum -from .codec import TransactionRejectionCode +from . import _client_backend class RecoveryError(RuntimeError): @@ -30,16 +30,14 @@ class RecoveryKind(StrEnum): TRANSACTION_REJECTED = "transaction_rejected" -_CODE_NAMES = { - TransactionRejectionCode.InvalidIdentity: "invalid_identity", - TransactionRejectionCode.UnexpectedId: "unexpected_id", - TransactionRejectionCode.StaleLayerGraph: "stale_layer_graph", - TransactionRejectionCode.InvalidTransaction: "invalid_transaction", -} - -_CODE_DISPOSITIONS = { - TransactionRejectionCode.StaleLayerGraph: RejectionDisposition.RECOVERABLE_CONFLICT, - TransactionRejectionCode.InvalidTransaction: RejectionDisposition.INVALID_OPERATION, +_DISPOSITIONS = { + _client_backend.ProducerRecoveryDisposition.SESSION_FATAL: RejectionDisposition.SESSION_FATAL, + _client_backend.ProducerRecoveryDisposition.RECOVERABLE_CONFLICT: ( + RejectionDisposition.RECOVERABLE_CONFLICT + ), + _client_backend.ProducerRecoveryDisposition.INVALID_OPERATION: ( + RejectionDisposition.INVALID_OPERATION + ), } @@ -54,12 +52,11 @@ class TransactionFailure: @property def code_name(self) -> str: - return _CODE_NAMES.get(self.code, f"unknown_{self.code}") + return _client_backend.rejection_code_name(self.code) or f"unknown_{self.code}" @property def disposition(self) -> RejectionDisposition: - # Unknown rejection codes fail closed for forward compatibility. - return _CODE_DISPOSITIONS.get(self.code, RejectionDisposition.SESSION_FATAL) + return _DISPOSITIONS[_client_backend.rejection_disposition(self.code)] def __str__(self) -> str: expected = f", expected transaction {self.expected_txn_id}" if self.expected_txn_id else "" diff --git a/openusdconnect/send.py b/openusdconnect/send.py index 6bcddef..f01a3c8 100644 --- a/openusdconnect/send.py +++ b/openusdconnect/send.py @@ -102,7 +102,7 @@ def main(argv: list[str] | None = None): sys.exit(1) print(f"Sent {msg.get('type', '?')} message") finally: - sender.disconnect() + sender.close() if __name__ == "__main__": diff --git a/openusdconnect/sender.py b/openusdconnect/sender.py index 7e46bc5..3278b4e 100644 --- a/openusdconnect/sender.py +++ b/openusdconnect/sender.py @@ -1,33 +1,16 @@ -"""TCP producer with durable transaction acknowledgement and reconnect replay.""" +"""Durable transaction producer whose connection runs on a native thread.""" from __future__ import annotations import logging -import socket import threading -import time import uuid from collections.abc import Callable from . import _client_backend from .checkpoints import MirrorCheckpoint -from .codec import ( - PayloadType, - TransactionRejectionCode, - TransactionStatus, - _decode_stage_metadata_table, - decode_envelope, - encode_message, - resolve_payload, -) -from .framing import IncompleteRead, MessageTooLarge, recv_framed -from .protocol import ( - make_claim_playback, - make_hello, - make_playback_control, - make_quit, - make_txn, -) +from .codec import encode_message +from .protocol import make_claim_playback, make_playback_control, make_txn from .protocol_constants import LayerMode from .protocol_validation import validate_events from .recovery import ( @@ -38,7 +21,6 @@ TransactionFailure, make_recovery_incident, ) -from .transport import send_msg, send_raw LOG = logging.getLogger(__name__) @@ -54,16 +36,49 @@ def __init__(self, failure: TransactionFailure): self.failure = failure -class EventSender: - """Pipelined producer whose outbox survives socket reconnects. +def _session_id(session_id: str | None) -> str: + session_id = session_id or uuid.uuid4().hex + if len(session_id) > 128: + raise ValueError("session_id must contain 1-128 characters") + return session_id + + +def _failure(native) -> TransactionFailure: + return TransactionFailure( + txn_id=native.transaction_id, + code=native.code, + reason=native.reason, + expected_txn_id=native.expected_transaction_id, + ) + + +def _artifact(native) -> RecoveryArtifact: + return RecoveryArtifact( + producer_session_id=native.session_id, + failure=_failure(native.failure), + transactions=tuple( + QuarantinedTransaction( + txn_id=entry.transaction_id, + payload=entry.payload, + event_count=entry.event_count, + layer_key=entry.layer_key, + ) + for entry in native.transactions + ), + ) - ``send_events`` returns once this object owns the encoded transaction. A - background reader removes it only after the server's cumulative durable - acknowledgement covers it. The same encoded bytes and Hello-bound producer - identity are replayed after reconnect, so an ACK lost after commit cannot - apply the USD edits twice. - ``background_send=True`` moves transaction writes and replay to a worker. +class EventSender: + """Pipelined producer whose outbox survives reconnects. + + ``send_events`` returns once this object owns the encoded transaction. The + native outbox removes it only after the server's cumulative durable + acknowledgement covers it, and replays the same bytes under the same + Hello-bound producer identity after a reconnect, so an acknowledgement lost + after commit cannot apply the USD edits twice. A native thread, started by + the first connection request, writes, reads, and runs the callbacks until + :meth:`close`; given ``notifications``, the owner drains that queue instead + and only ``on_token_issued`` runs there. """ def __init__( @@ -83,549 +98,276 @@ def __init__( layer_mode: LayerMode | str = LayerMode.MANAGED, session_id: str | None = None, max_pending_transactions: int = _MAX_PENDING_TRANSACTIONS, - background_send: bool = False, + notifications: _client_backend.NotificationQueue | None = None, ): if role != "emitter": raise ValueError("EventSender role must be 'emitter'") if max_pending_transactions < 1: raise ValueError("max_pending_transactions must be positive") - self.host = host - self.port = port - self.client_id = client_id - self.role = role - self.origin = origin - self.department = department + self._host = host + self._port = port + self._client_id = client_id + self._origin = origin + self._department = department + self._layer_mode = LayerMode(layer_mode) + self._handshake_timeout = handshake_timeout + self._max_pending_transactions = max_pending_transactions self.token = token - self.layer_mode = LayerMode(layer_mode) - self.layer_mode_active = LayerMode.MANAGED - self.handshake_timeout = handshake_timeout - self.session_id = session_id or uuid.uuid4().hex - if not self.session_id or len(self.session_id) > 128: - raise ValueError("session_id must contain 1-128 characters") - self.max_pending_transactions = max_pending_transactions - self._background_send = background_send + self._token_provider = token_provider self._on_token_issued = on_token_issued self._on_stage_metadata = on_stage_metadata - self._token_provider = token_provider - self.sock: socket.socket | None = None - self.auth_rejected = False - self.hello_rejected = False - self.rejection_reason = "" - self.stage_metadata: dict = {} - - # Reentrant: background submission holds it while appending to the outbox. - self._condition = threading.Condition(threading.RLock()) - self._connect_lock = threading.Lock() - self._connect_epoch = 0 - self._connect_thread: threading.Thread | None = None - self._connecting_socket: socket.socket | None = None - self._connect_active = 0 - self._connect_retry_at = 0.0 - self._connect_retry_delay = 1.0 - self._send_lock = threading.Lock() - self._reader_thread: threading.Thread | None = None - self._writer_thread: threading.Thread | None = None - self._socket_generation = 0 - self._session = _client_backend.ProducerSession(max_pending_transactions) - self._failure: TransactionFailure | None = None - self._recovery_artifact: RecoveryArtifact | None = None - self._recovery_incident: RecoveryIncident | None = None - self._retry_after_until = 0.0 - self._server_instance = "" - self._acknowledged_checkpoint: MirrorCheckpoint | None = None + config = _client_backend.ProducerConfig() + config.host = host + config.port = port + config.client_id = client_id + config.origin = origin or "" + config.department = department or "" + config.layer_mode = _client_backend.NATIVE_LAYER_MODES[self._layer_mode] + config.session_id = _session_id(session_id) + config.handshake_timeout = handshake_timeout + config.max_pending_transactions = max_pending_transactions + notifications, owns_notifications = _client_backend.notification_queue( + notifications, on_stage_metadata=on_stage_metadata + ) + self._endpoint = _client_backend.ProducerEndpoint(config, notifications) + self._driver = _client_backend.ProducerDriver( + self._endpoint, + notifications, + _client_backend.TcpSocketFactory(), + **_client_backend.driver_callbacks(self, LOG, sink=owns_notifications), + ) + self._started = False + # Pairs each transaction ID with the frame that encodes it. + self._submit_lock = threading.Lock() + # Builds the recovery objects once per failure, so their identity holds. + self._recovery_lock = threading.Lock() + self._recovery: tuple[RecoveryArtifact, RecoveryIncident] | None = None @property - def acknowledged_checkpoint(self) -> MirrorCheckpoint | None: - """Post-commit mirror position when all sends are acknowledged. + def host(self) -> str: + return self._host - None means pending/rejected work or a peer without checkpoint support. - A Hello highwater alone does not establish a mirror checkpoint. - """ - with self._condition: - if self._failure is not None or not self._session.empty: - return None - return self._acknowledged_checkpoint + @property + def port(self) -> int: + return self._port + + @property + def client_id(self) -> str: + return self._client_id + + @property + def role(self) -> str: + return "emitter" + + @property + def origin(self) -> str | None: + return self._origin + + @property + def department(self) -> str | None: + return self._department + + @property + def layer_mode(self) -> LayerMode: + return self._layer_mode + + @property + def handshake_timeout(self) -> float: + return self._handshake_timeout + + @property + def max_pending_transactions(self) -> int: + return self._max_pending_transactions + + @property + def session_id(self) -> str: + """The producer session; :meth:`abandon_rejected_session` starts a new one.""" + return self._endpoint.status().session_id @property def is_connected(self) -> bool: - with self._condition: - return self.sock is not None + return self._endpoint.status().connected @property def connected(self) -> bool: return self.is_connected + @property + def layer_mode_active(self) -> LayerMode: + return _client_backend.LAYER_MODES[self._endpoint.status().layer_mode_active] + + @property + def stage_metadata(self) -> dict: + """The latest stage metadata the server authored, keyed as on the wire.""" + return _client_backend.stage_metadata_fields(self._endpoint.status().metadata) + + @property + def auth_rejected(self) -> bool: + rejection = self._endpoint.status().rejection + return rejection is not None and rejection.authentication + + @property + def hello_rejected(self) -> bool: + rejection = self._endpoint.status().rejection + return rejection is not None and not rejection.authentication + + @property + def rejection_reason(self) -> str: + """Why the latest handshake was refused, or the failure that refuses new ones.""" + rejection = self._endpoint.status().rejection + if rejection is not None: + return _client_backend.rejection_reason(rejection) + failure = self.transaction_failure + return "" if failure is None else failure.reason + @property def pending_transaction_count(self) -> int: - return self._session.pending_transaction_count + return self._endpoint.status().pending_transactions @property def pending_event_count(self) -> int: - return self._session.pending_event_count + return self._endpoint.status().pending_events @property def acknowledged_transaction_count(self) -> int: - return self._session.acknowledged_transaction_count + return self._endpoint.status().acknowledged_transactions @property def acknowledged_event_count(self) -> int: - return self._session.acknowledged_event_count + return self._endpoint.status().acknowledged_events @property - def _next_txn_id(self) -> int: - """Compatibility view of the native outbox sequence cursor.""" - return self._session.next_transaction_id + def acknowledged_checkpoint(self) -> MirrorCheckpoint | None: + """Post-commit mirror position when all sends are acknowledged. - @property - def transaction_error(self) -> str: - with self._condition: - return str(self._failure) if self._failure is not None else "" + None means pending/rejected work or a peer without checkpoint support. + A Hello highwater alone does not establish a mirror checkpoint. + """ + checkpoint = self._endpoint.acknowledged_checkpoint() + if checkpoint is None: + return None + return MirrorCheckpoint( + server_instance=checkpoint.server_instance, + epoch=checkpoint.epoch, + head_seq=checkpoint.head_sequence, + ) @property def transaction_failure(self) -> TransactionFailure | None: """Structured terminal result for UI and recovery policy.""" - with self._condition: - return self._failure + recovery = self._current_recovery() + return None if recovery is None else recovery[0].failure + + @property + def transaction_error(self) -> str: + failure = self.transaction_failure + return "" if failure is None else str(failure) @property def recovery_incident(self) -> RecoveryIncident | None: """Immutable summary suitable for status polling and host UI.""" - with self._condition: - return self._recovery_incident + recovery = self._current_recovery() + return None if recovery is None else recovery[1] @property def recovery_artifact(self) -> RecoveryArtifact | None: """Exact quarantined bytes for inspection or application-owned export.""" - with self._condition: - return self._recovery_artifact + recovery = self._current_recovery() + return None if recovery is None else recovery[0] @property def recovery_disposition(self) -> RejectionDisposition | None: """Recommended response category for the current rejection.""" - with self._condition: - return self._failure.disposition if self._failure is not None else None + failure = self.transaction_failure + return None if failure is None else failure.disposition @property def recovery_required(self) -> bool: """Whether a deterministic rejection quarantined this producer session.""" - with self._condition: - return self._session.recovery_required + return self._current_recovery() is not None + + def snapshot(self) -> _client_backend.ProducerStatus: + """The native status in one call, for reading several fields together.""" + return self._endpoint.status() def request_connect(self, timeout: float | None = 2.0) -> bool: """Start one background attempt, returning whether it was scheduled. Call again from an update loop to retry transient failures. Attempts back off from one to eight seconds; rejection requires explicit connect. - Callbacks run on the handshake thread, just as for synchronous connect. + Callbacks run on the connection thread, just as for synchronous connect. """ - with self._condition: - if ( - self.sock is not None - or self._connect_active - or self._connect_thread is not None - or self.auth_rejected - or self.hello_rejected - or self._session.recovery_required - or time.monotonic() < max(self._connect_retry_at, self._retry_after_until) - ): - return False - epoch = self._connect_epoch - thread = threading.Thread( - target=self._connect_worker, - args=(timeout, epoch), - name=f"openusdconnect-reconnect-{self.client_id}", - daemon=True, - ) - self._connect_thread = thread - try: - thread.start() - except Exception: - self._connect_thread = None - raise - return True - - def _connect_worker(self, timeout: float | None, epoch: int) -> None: - connected = False - try: - connected = self._connect_attempt(timeout, epoch) - except Exception: - LOG.exception("EventSender: background connect failed") - finally: - with self._condition: - if epoch == self._connect_epoch: - if connected: - self._connect_retry_at = 0.0 - self._connect_retry_delay = 1.0 - else: - self._connect_retry_at = time.monotonic() + self._connect_retry_delay - self._connect_retry_delay = min(8.0, self._connect_retry_delay * 2) - self._connect_thread = None - self._condition.notify_all() + if not self._endpoint.request_connect(timeout): + return False + self._start() + self._driver.wake() + return True def cancel_connect(self) -> bool: """Invalidate pending handshakes without waiting; report completion. An already published connection is left intact. Use disconnect to close - it as well. A blocked connection creation may finish later, but cannot - publish its socket after cancellation. + it as well. A cancelled attempt can no longer publish its connection. """ - with self._condition: - self._connect_epoch += 1 - self._connect_retry_at = 0.0 - self._connect_retry_delay = 1.0 - sock = self._connecting_socket - finished = not self._connect_active and self._connect_thread is None - self._close_socket_object(sock) + finished = self._endpoint.cancel_connect() + self._driver.wake() return finished def connect(self, timeout: float | None = None) -> bool: - """Handshake, start the result reader, and replay the exact outbox. + """Handshake and replay the exact outbox; return whether connected. - ``timeout`` bounds this attempt and never extends the configured - handshake timeout. + An attempt already in flight is waited for first. ``timeout`` bounds the + whole call and never extends the configured handshake timeout. """ - with self._condition: - epoch = self._connect_epoch - return self._connect_attempt(timeout, epoch) - - def _connect_attempt(self, timeout: float | None, epoch: int) -> bool: - budget = self.handshake_timeout - if timeout is not None: - budget = min(budget, max(timeout, 0.0)) - deadline = time.monotonic() + budget - with self._condition: - if epoch != self._connect_epoch: - return False - self._connect_active += 1 - acquired = False - try: - acquired = self._connect_lock.acquire(timeout=max(0.0, budget)) - if not acquired: - return False - return self._connect_locked(deadline, epoch) - finally: - with self._condition: - self._connect_active -= 1 - if acquired: - self._connecting_socket = None - self._condition.notify_all() - if acquired: - self._connect_lock.release() - - def _connect_locked(self, deadline: float, epoch: int) -> bool: - # The connect lock is held throughout handshake and replay. - with self._condition: - if epoch != self._connect_epoch: - return False - if self.sock is not None: - return True - if self._session.recovery_required or time.monotonic() < self._retry_after_until: - return False - - connect_timeout = deadline - time.monotonic() - if connect_timeout <= 0.0: - return False - if self._token_provider is not None: - self.token = self._token_provider() + self._start() + return self._driver.connect(timeout) - connection = self._session.begin_connection() - if connection is None: - return False - generation = connection.generation - - self.auth_rejected = False - self.hello_rejected = False - self.rejection_reason = "" - sock: socket.socket | None = None - published = False - acquired_send = False - try: - sock = socket.create_connection((self.host, self.port), timeout=connect_timeout) - with self._condition: - if epoch != self._connect_epoch: - return False - self._connecting_socket = sock - sock.settimeout(max(0.001, deadline - time.monotonic())) - sock.setsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY, 1) - send_msg( - sock, - make_hello( - self.role, - client_id=self.client_id, - origin=self.origin, - department=self.department, - token=self.token, - layer_mode=self.layer_mode, - producer_session_id=self.session_id, - ), - ) - sock.settimeout(max(0.001, deadline - time.monotonic())) - buf = recv_framed(sock) - with self._condition: - if epoch != self._connect_epoch: - return False - env = decode_envelope(buf) - pt = env.PayloadType() - if not self._accept_handshake_response(sock, env, pt, generation): - return False - - # Serialize publication of the socket with outbox replay. A new - # send cannot overtake an older pending transaction here. - acquired_send = self._send_lock.acquire(timeout=max(0.0, deadline - time.monotonic())) - if not acquired_send: - return False - with self._condition: - if epoch != self._connect_epoch or time.monotonic() >= deadline: - return False - self.sock = sock - self._socket_generation = generation - self._connecting_socket = None - published = True - sock.settimeout(max(0.001, deadline - time.monotonic())) - replayed = 0 - if not self._background_send: - while pending := self._session.claim_next_unsent(generation): - remaining = deadline - time.monotonic() - if remaining <= 0: - raise TimeoutError("reconnect replay timed out") - sock.settimeout(remaining) - send_raw(sock, pending[1]) - replayed += 1 - sock.settimeout(None) - except Exception as exc: - if not published: - if isinstance(exc, OSError): - # An unreachable server is expected while retrying. - LOG.info("EventSender: connect to %s:%d failed: %s", self.host, self.port, exc) - else: - LOG.exception("EventSender: handshake failed") - return False - self._close(expected=sock) - if not isinstance(exc, OSError): - raise - LOG.info("EventSender: reconnect replay failed", exc_info=True) - return False - finally: - if acquired_send: - self._send_lock.release() - # Until publication this attempt owns both resources. Afterwards, - # _close(expected=sock) handles failures without closing a newer socket. - if not published: - self._session.disconnect(generation) - self._close_socket_object(sock) - - reader = threading.Thread( - target=self._read_results, - args=(sock, generation), - name=f"openusdconnect-ack-{self.client_id}", - daemon=True, - ) - try: - with self._condition: - if self.sock is not sock or generation != self._socket_generation: - return False - self._reader_thread = reader - reader.start() - if self._background_send: - writer = threading.Thread( - target=self._write_pending, - args=(sock, generation), - name=f"openusdconnect-send-{self.client_id}", - daemon=True, - ) - writer.start() - self._writer_thread = writer - except Exception: - self._close(expected=sock) - raise - LOG.info( - "EventSender connected to %s:%d (session=%s, pending=%d)", - self.host, - self.port, - self.session_id, - replayed, - ) - return True - - def _accept_handshake_response( - self, sock: socket.socket, env, payload_type: int, generation: int - ) -> bool: - """Validate one server hello and initialize connection metadata. - - Application callbacks are observers. Their failure is logged but does - not invalidate an otherwise completed protocol handshake. + def disconnect(self) -> None: + """Close the connection while retaining unacknowledged transactions.""" + self._endpoint.disconnect() + self._driver.wake() + + def close(self, timeout: float | None = None) -> bool: + """Write the queued transactions and the Quit message, then close the connection. + + Then stops the connection thread. The writes happen only while the socket + is open and within one second, or ``timeout`` when shorter, so ``close(0)`` + does not wait for them. Closing does not wait for acknowledgements; + :meth:`flush` does. Waits up to ``timeout`` seconds in all, without limit + for ``None``, and returns whether the thread has exited: ``True`` before + the first connection request, ``False`` on timeout or at once from a + callback, which runs on that thread; the thread then exits once the + callback returns. Afterwards :meth:`connect` and :meth:`request_connect` + return ``False``; repeated calls are harmless. """ - if payload_type == PayloadType.AuthRejected: - _, rejected = resolve_payload(env) - self.rejection_reason = self._decode_string(rejected.Reason()) - self.auth_rejected = True - return False - if payload_type == PayloadType.HelloRejected: - _, rejected = resolve_payload(env) - self.rejection_reason = self._decode_string(rejected.Reason()) or "connection rejected" - self.hello_rejected = True - return False - if payload_type != PayloadType.HelloOk: - LOG.error("EventSender: unexpected handshake response %s", payload_type) - return False - - _, hello_ok = resolve_payload(env) - with self._condition: - self._server_instance = self._decode_string(hello_ok.ServerInstance()) or "" - self._acknowledged_checkpoint = None - active_mode = LayerMode("shared_stage" if hello_ok.LayerMode() else "managed") - if active_mode is not self.layer_mode: - self.rejection_reason = ( - f"server negotiated {active_mode.value} instead of {self.layer_mode.value}" - ) - self.hello_rejected = True - return False - self.layer_mode_active = active_mode + return self._driver.close(timeout) - committed_through = int(hello_ok.CommittedThrough()) - result = self._session.accept_hello(generation, committed_through) - if result != _client_backend.ProducerResult.ACCEPTED: - self.rejection_reason = self._highwater_failure_reason(result, committed_through) - self._record_session_failure( - txn_id=committed_through, - code=int(TransactionRejectionCode.UnexpectedId), - reason=self.rejection_reason, - ) - return False - - issued = self._decode_string(hello_ok.Token()) - if issued: - self.token = issued - self._notify_handshake_callback( - self._on_token_issued, - issued, - name="on_token_issued", - ) - - metadata = hello_ok.StageMetadata() - if metadata is not None: - decoded = _decode_stage_metadata_table(metadata) - if decoded: - self.stage_metadata = decoded - self._notify_handshake_callback( - self._on_stage_metadata, - decoded, - name="on_stage_metadata", - ) - sock.settimeout(None) - return True - - @staticmethod - def _notify_handshake_callback(callback, value, *, name: str) -> None: - if callback is None: - return - try: - callback(value) - except Exception: - LOG.exception("EventSender: %s callback failed", name) + def __enter__(self) -> EventSender: + return self - def disconnect(self) -> None: - """Close the socket while retaining unacknowledged transactions.""" - self.cancel_connect() - if self._background_send: - # Shutdown must interrupt a blocked writer, not wait for its lock. - self._close() - return - with self._condition: - if self.sock is None: - return - with self._send_lock: - with self._condition: - sock = self.sock - if sock is not None: - try: - send_msg(sock, make_quit()) - except OSError: - pass - if sock is not None: - self._close(expected=sock) + def __exit__(self, exc_type, exc, traceback) -> None: + self.close() def send_events(self, events: list, *, layer_key: str = "") -> bool: """Submit a transaction without waiting for its durable result. ``True`` means the encoded bytes are owned by the bounded outbox, even - if the socket fails during this call. ``False`` means no ownership was - taken (disconnected before submission, empty input, full outbox, or a - terminal rejection). + if the connection fails afterwards. ``False`` means no ownership was + taken (disconnected, empty input, full outbox, or a terminal rejection). """ if not events: return False - validate_events(events, layer_mode=self.layer_mode) - submission_lock = self._condition if self._background_send else self._send_lock - try: - # Synchronous callers allocate IDs in socket order. Background - # callers only append; the writer claims that same ordered outbox. - with submission_lock: - with self._condition: - if self.sock is None or self._failure is not None: - return False - if not self._session.can_append: - return False - generation = self._socket_generation - txn_id = self._session.next_transaction_id - payload = encode_message( - make_txn( - events, - layer_key=layer_key, - txn_id=txn_id, - ) - ) - result = self._session.append( - generation, txn_id, payload, len(events), layer_key - ) - if result != _client_backend.ProducerResult.ACCEPTED: - return False - if self._background_send: - self._condition.notify_all() - return True - sock = self.sock - if sock is not None: - send_raw(sock, payload) - except OSError: - LOG.info( - "EventSender: send became ambiguous; retaining transaction", - exc_info=True, - ) - self._close(expected=sock) + validate_events(events, layer_mode=self._layer_mode) + with self._submit_lock: + txn_id = self._endpoint.next_transaction_id() + payload = encode_message(make_txn(events, layer_key=layer_key, txn_id=txn_id)) + result = self._endpoint.append(txn_id, payload, len(events), layer_key) + if result != _client_backend.ProducerResult.ACCEPTED: + return False + self._driver.wake() return True - def _write_pending(self, sock: socket.socket, generation: int) -> None: - """Send the native outbox in order for exactly one socket generation.""" - try: - while True: - with self._condition: - if ( - self.sock is not sock - or self._socket_generation != generation - or self._failure is not None - ): - return - pending = self._session.claim_next_unsent(generation) - if pending is None: - self._condition.wait() - continue - with self._send_lock: - with self._condition: - if self.sock is not sock or self._socket_generation != generation: - return - send_raw(sock, pending[1]) - except OSError: - LOG.info("EventSender: background send failed; retaining transaction", exc_info=True) - except Exception: - LOG.exception("EventSender: background writer failed") - finally: - self._close(expected=sock) - with self._condition: - if self._writer_thread is threading.current_thread(): - self._writer_thread = None - self._condition.notify_all() - def repair_rejected_transaction(self, events: list, *, layer_key: str = "") -> int: """Replace a recoverable rejected transaction at the same ordered ID. @@ -637,33 +379,21 @@ def repair_rejected_transaction(self, events: list, *, layer_key: str = "") -> i """ if not events: raise ValueError("repair events must not be empty") - validate_events(events, layer_mode=self.layer_mode) - with self._condition: - failure = self._failure - if failure is None: + validate_events(events, layer_mode=self._layer_mode) + with self._recovery_lock: + recovery = self._current_recovery_locked() + if recovery is None: raise RuntimeError("there is no rejected transaction to retry") + failure = recovery[0].failure if failure.disposition is not RejectionDisposition.RECOVERABLE_CONFLICT: raise RuntimeError( f"{failure.code_name} is {failure.disposition.value}, not recoverable" ) - - # A rejection normally already closes this socket. Make the boundary - # explicit so a racing reader cannot leave a repaired outbox attached to - # the connection that delivered the rejection. - self.disconnect() - payload = encode_message(make_txn(events, layer_key=layer_key, txn_id=failure.txn_id)) - with self._send_lock: - with self._condition: - if self._failure is not failure: - raise RuntimeError("transaction rejection changed during recovery") - result = self._session.repair_rejected(payload, len(events), layer_key) - if result != _client_backend.ProducerResult.ACCEPTED: - raise RuntimeError(f"native recovery rejected repair: {result}") - self._failure = None - self._recovery_artifact = None - self._recovery_incident = None - self._retry_after_until = 0.0 - self._condition.notify_all() + payload = encode_message(make_txn(events, layer_key=layer_key, txn_id=failure.txn_id)) + result = self._endpoint.repair_rejected(payload, len(events), layer_key) + if result != _client_backend.ProducerResult.ACCEPTED: + raise RuntimeError(f"native recovery rejected repair: {result}") + self._recovery = None return failure.txn_id def abandon_rejected_session(self, *, session_id: str | None = None) -> RecoveryArtifact: @@ -673,54 +403,30 @@ def abandon_rejected_session(self, *, session_id: str | None = None) -> Recovery its USD stage before reconnecting or submitting rebuilt intent with the new producer session. """ - replacement = session_id or uuid.uuid4().hex - if not replacement or len(replacement) > 128: - raise ValueError("session_id must contain 1-128 characters") - - with self._condition: - failure = self._failure - artifact = self._recovery_artifact - previous_session_id = self.session_id - if failure is None or artifact is None: + replacement = _session_id(session_id) + with self._recovery_lock: + recovery = self._current_recovery_locked() + if recovery is None: raise RuntimeError("there is no rejected producer session to abandon") - if replacement == previous_session_id: + artifact = recovery[0] + if replacement == artifact.producer_session_id: raise ValueError("replacement session_id must differ from rejected session") - - self.disconnect() - with self._send_lock: - with self._condition: - if self._failure is not failure or self._recovery_artifact is not artifact: - raise RuntimeError("transaction rejection changed during recovery") - self._session.reset_session() - self.session_id = replacement - self._failure = None - self._recovery_artifact = None - self._recovery_incident = None - self._retry_after_until = 0.0 - self.rejection_reason = "" - self._condition.notify_all() + # The checks above are the endpoint's preconditions. + abandoned = self._endpoint.abandon_rejected_session(replacement) + assert abandoned is not None + self._recovery = None return artifact def send_message(self, msg: dict) -> bool: """Send a non-transaction protocol message (not retained for replay).""" - with self._condition: - sock = self.sock - if sock is None: - return False - try: - with self._send_lock: - with self._condition: - if self.sock is not sock: - return False - send_msg(sock, msg) - return True - except OSError: - self._close(expected=sock) + if not self._endpoint.queue_control(encode_message(msg)): return False + self._driver.wake() + return True def drain_acknowledged_event_count(self) -> int: """Return successful event acknowledgements received since the last drain.""" - return self._session.drain_acknowledged_event_count() + return self._endpoint.drain_acknowledged_event_count() def flush(self, timeout: float | None = None) -> bool: """Wait for all submitted transactions to reach a terminal result. @@ -728,39 +434,16 @@ def flush(self, timeout: float | None = None) -> bool: Reconnects and replays while time remains. Returns ``False`` on timeout; a deterministic server rejection raises ``TransactionRejectedError``. """ - deadline = None if timeout is None else time.monotonic() + max(timeout, 0.0) - while True: - with self._condition: - if self._failure is not None: - raise TransactionRejectedError(self._failure) - if self._session.empty: - return True - connected = self.sock is not None - retry_at = self._retry_after_until - remaining = None if deadline is None else deadline - time.monotonic() - if remaining is not None and remaining <= 0: - return False - if connected: - self._condition.wait( - timeout=0.25 if remaining is None else min(0.25, remaining) - ) - continue - delay = max(0.0, retry_at - time.monotonic()) - if delay: - if deadline is not None and time.monotonic() + delay >= deadline: - return False - time.sleep(min(delay, 0.25)) - continue - remaining = None if deadline is None else max(0.0, deadline - time.monotonic()) - if not self.connect(timeout=remaining): - with self._condition: - remaining = None if deadline is None else deadline - time.monotonic() - if remaining is not None and remaining <= 0: - return False - self._condition.wait(timeout=0.1 if remaining is None else min(0.1, remaining)) + result = self._driver.flush(timeout) + if result == _client_backend.FlushResult.RECOVERY_REQUIRED: + failure = self.transaction_failure + # None only when another thread resolved the failure meanwhile. + if failure is not None: + raise TransactionRejectedError(failure) + return result == _client_backend.FlushResult.FLUSHED def claim_playback(self, time: float | None = None) -> bool: - return self.send_message(make_claim_playback(self.client_id, time=time)) + return self.send_message(make_claim_playback(self._client_id, time=time)) def send_playback_control( self, @@ -771,169 +454,57 @@ def send_playback_control( ) -> bool: return self.send_message(make_playback_control(action, time=time, rate=rate)) - def _read_results(self, sock: socket.socket, generation: int) -> None: - try: - while True: - buf = recv_framed(sock) - env = decode_envelope(buf) - if env.PayloadType() == PayloadType.TransactionResult: - _, result = resolve_payload(env) - self._accept_result(result, generation) - elif env.PayloadType() == PayloadType.RateLimited: - _, limited = resolve_payload(env) - with self._condition: - self._retry_after_until = max( - self._retry_after_until, - time.monotonic() + float(limited.RetryAfter()), - ) - break - except (OSError, IncompleteRead, MessageTooLarge, ValueError): - pass - finally: - with self._condition: - is_current = generation == self._socket_generation and self.sock is sock - if is_current: - self._close(expected=sock) - - def _accept_result(self, result, generation: int) -> None: - txn_id = int(result.TxnId()) - rejected_socket = None - with self._condition: - if generation != self._socket_generation or self.sock is None: - return - status_value = int(result.Status()) - if status_value == TransactionStatus.Acknowledged: - accepted = self._session.acknowledge_through(generation, txn_id) - if accepted == _client_backend.ProducerResult.STALE_GENERATION: - return - if accepted == _client_backend.ProducerResult.ACCEPTED: - checkpoint = result.Checkpoint() - self._acknowledged_checkpoint = ( - MirrorCheckpoint( - server_instance=self._server_instance, - epoch=int(checkpoint.Epoch()), - head_seq=int(checkpoint.HeadSeq()), - ) - if self._server_instance and checkpoint is not None - else None - ) - if accepted != _client_backend.ProducerResult.ACCEPTED: - failure = TransactionFailure( - txn_id=txn_id, - code=int(TransactionRejectionCode.UnexpectedId), - reason=self._highwater_failure_reason(accepted, txn_id), - ) - self._record_session_failure_locked(failure) - rejected_socket = self.sock - else: - code = int(result.RejectionCode()) - reason = self._decode_string(result.Reason()) - failure = TransactionFailure( - txn_id=txn_id, - code=code, - reason=reason, - expected_txn_id=int(result.ExpectedTxnId()), - ) - native_disposition = { - RejectionDisposition.RECOVERABLE_CONFLICT: ( - _client_backend.ProducerRecoveryDisposition.RECOVERABLE_CONFLICT - ), - RejectionDisposition.INVALID_OPERATION: ( - _client_backend.ProducerRecoveryDisposition.INVALID_OPERATION - ), - RejectionDisposition.SESSION_FATAL: ( - _client_backend.ProducerRecoveryDisposition.SESSION_FATAL - ), - }[failure.disposition] - accepted = self._session.reject(generation, txn_id, native_disposition) - if accepted == _client_backend.ProducerResult.STALE_GENERATION: - return - if accepted == _client_backend.ProducerResult.TRANSACTION_MISSING: - failure = TransactionFailure( - txn_id=txn_id, - code=int(TransactionRejectionCode.UnexpectedId), - reason=f"server rejected unknown transaction {txn_id}", - ) - self._record_session_failure_locked(failure) - rejected_socket = self.sock - self._condition.notify_all() - if rejected_socket is not None: - self._close(expected=rejected_socket) - - def _record_session_failure( - self, *, txn_id: int, code: int, reason: str, expected_txn_id: int = 0 - ) -> None: - with self._condition: - self._record_session_failure_locked( - TransactionFailure( - txn_id=txn_id, - code=code, - reason=reason, - expected_txn_id=expected_txn_id, - ) - ) - self._condition.notify_all() - - def _record_session_failure_locked(self, failure: TransactionFailure) -> None: - self._failure = failure - self._recovery_artifact = RecoveryArtifact( - producer_session_id=self.session_id, - failure=failure, - transactions=tuple( - QuarantinedTransaction( - txn_id=pending_txn_id, - payload=payload, - event_count=event_count, - layer_key=layer_key, - ) - for pending_txn_id, payload, event_count, layer_key in self._session.entries() - ), - ) - self._recovery_incident = make_recovery_incident(self._recovery_artifact) + def _start(self) -> None: + """Start the connection thread once; it runs until closed or collected.""" + if not self._started: + self._driver.start() + self._started = True + + def _current_recovery(self) -> tuple[RecoveryArtifact, RecoveryIncident] | None: + with self._recovery_lock: + return self._current_recovery_locked() + + def _current_recovery_locked(self) -> tuple[RecoveryArtifact, RecoveryIncident] | None: + # Only repair and abandon clear a failure, and both reset this cache. + if self._recovery is None: + native = self._endpoint.artifact() + if native is not None: + artifact = _artifact(native) + self._recovery = (artifact, make_recovery_incident(artifact)) + return self._recovery + + def _connection_token(self) -> str | None: + """The token for the next handshake; ``None`` abandons the attempt.""" + if self._token_provider is not None: + try: + self.token = self._token_provider() + except Exception: + LOG.exception("EventSender: token provider failed") + return None + return self.token or "" - def _highwater_failure_reason(self, result, transaction_id: int) -> str: - if result == _client_backend.ProducerResult.HIGHWATER_AHEAD: - return ( - f"server producer highwater {transaction_id} is ahead of local " - f"transaction {self._session.next_transaction_id - 1}" - ) - if result == _client_backend.ProducerResult.HIGHWATER_REGRESSED: - return ( - f"server producer highwater regressed from " - f"{self._session.last_acknowledged_transaction_id} to {transaction_id}" + def _token_issued(self, token: str) -> None: + """Adopt a token the server issued, on the connection thread.""" + self.token = token + self._notify(self._on_token_issued, token, "on_token_issued") + + def _deliver(self, notification) -> None: + """Run the callback for a notification on the connection thread.""" + if isinstance(notification, _client_backend.StageMetadata): + self._notify( + self._on_stage_metadata, + _client_backend.stage_metadata_fields(notification), + "on_stage_metadata", ) - return f"invalid producer result for transaction {transaction_id}: {result}" - - def _close(self, *, expected: socket.socket | None = None) -> None: - with self._condition: - sock = self.sock - if sock is None or (expected is not None and sock is not expected): - return - self.sock = None - # An unpublished handshake still owns its native connection and - # ends it itself; only the published socket's generation ends here. - self._session.disconnect(self._socket_generation) - self._condition.notify_all() - self._close_socket_object(sock) @staticmethod - def _close_socket_object(sock: socket.socket | None) -> None: - if sock is None: + def _notify(callback: Callable | None, value, name: str) -> None: + if callback is None: return try: - sock.shutdown(socket.SHUT_RDWR) - except OSError: - pass - try: - sock.close() - except OSError: - pass - - @staticmethod - def _decode_string(value) -> str: - if isinstance(value, bytes): - return value.decode("utf-8") - return value or "" + callback(value) + except Exception: + LOG.exception("EventSender: %s callback failed", name) __all__ = ["EventSender", "TransactionRejectedError"] diff --git a/openusdconnect/shared_stage_client.py b/openusdconnect/shared_stage_client.py index e9835fd..25c42d1 100644 --- a/openusdconnect/shared_stage_client.py +++ b/openusdconnect/shared_stage_client.py @@ -10,12 +10,7 @@ from pxr import Sdf, Usd from ._client_base import PublishingClientBase -from ._client_lifecycle import ( - DEFAULT_WAIT_TIMEOUT_S, - BacklogHold, - deadline_after, - remaining_time, -) +from ._client_lifecycle import DEFAULT_WAIT_TIMEOUT_S, deadline_after, remaining_time from ._client_utils import client_origin, require_app_name from .client_id import make_stable_client_id from .client_observer import ClientObserver @@ -30,7 +25,7 @@ K_SET_SUBLAYERS, LayerMode, ) -from .receiver import ReceiverThread +from .receiver import EventReceiver from .recovery import RecoveryArtifact, RecoveryError from .sdf_layer_tracker import SdfLayerChangeTracker from .sender import EventSender @@ -128,7 +123,6 @@ def __init__( token: str | None = None, persist_token: bool = True, reconnect: bool = True, - background_send: bool = False, observer: ClientObserver | None = None, delegate_bridge_path: str | Path | None = None, ): @@ -153,17 +147,17 @@ def __init__( "origin": origin or client_origin(app_name, "shared"), } credential = self._credential.endpoint_kwargs() - self._receiver = ReceiverThread( + self._receiver = EventReceiver( host=host, port=port, sync_from=1, reconnect=reconnect, layered_replay=False, layer_mode=LayerMode.SHARED_STAGE, - **identity, **credential, **self._hooks.receiver_callbacks(), + notifications=self._notifications, **identity, **credential, ) self._sender = EventSender( - host, port, layer_mode=LayerMode.SHARED_STAGE, background_send=background_send, + host, port, layer_mode=LayerMode.SHARED_STAGE, notifications=self._notifications, **identity, **credential, ) self._last_seq = 0 - self._backlog = BacklogHold() + self._backlog_marker = 0 self._set_deferred([]) self._last_recovery_assessment: SharedRecoveryAssessment | None = None self._recovery_rebind_artifact: RecoveryArtifact | None = None @@ -489,7 +483,11 @@ def update(self, *, max_messages: int | None = None) -> SyncUpdate: sent = 0 if self._graph.ready and self._receiver.connected and not self._sender.connected: self._sender.request_connect() - if self._sender.connected and self._is_synchronized() and not self._backlog.holding: + if ( + self._sender.connected + and self._is_synchronized() + and self._receiver.drained_through(self._backlog_marker) + ): while routed := self._tracker.next_routed_batch(): batch, layer_key, events = routed if not self._sender.send_events(events, layer_key=layer_key): @@ -503,15 +501,15 @@ def _apply_queued(self, max_messages: int | None = None) -> int: had_batch = bool(self._tracker.prepared_event_count) self._tracker.prepare_local_changes() if not had_batch and self._tracker.prepared_event_count: - self._backlog.freeze(self._receiver.queued_message_count) + self._backlog_marker = self._receiver.freeze_marker() try: return self._apply_incoming(max_messages) finally: self._tracker.restore_prepared() def _apply_incoming(self, max_messages: int | None = None) -> int: + generation = self._receiver.generation buffers = self._receiver.drain_queue(max_messages) - self._backlog.drained(len(buffers), self._receiver.queued_message_count) if not buffers: self._receiver.mark_replay_applied() return 0 @@ -557,6 +555,9 @@ def _apply_incoming(self, max_messages: int | None = None) -> int: self._receiver.request_replay_from(self._last_seq + 1) LOG.warning("Shared-stage decode failed: %s", result.errors[0]) else: + if result.resync_requested: + self._receiver.reset_applied_progress() + self._receiver.mark_applied_through(generation, self._last_seq) self._receiver.mark_replay_applied() return applied @@ -704,12 +705,8 @@ def _rebind_stage_for_recovery(self, stage: Usd.Stage) -> None: self._last_seq = 0 old_tracker.close() - def _is_synchronized(self) -> bool: - return ( - self._graph.ready - and self._receiver.synchronized - and not self._sender.recovery_required - ) + def _synchronized(self, replayed: bool) -> bool: + return replayed and self._graph.ready and not self._sender.recovery_required def _prepared_events(self) -> int: return self._tracker.prepared_event_count diff --git a/openusdconnect/usd_client.py b/openusdconnect/usd_client.py index d4be293..4f0b560 100644 --- a/openusdconnect/usd_client.py +++ b/openusdconnect/usd_client.py @@ -21,7 +21,7 @@ from .defaults import DEFAULT_HOST, DEFAULT_SYNC_PORT from .dispatcher import AssetDependencyRefreshResult, EventDispatcher from .emitter import PrimChannel -from .receiver import ReceiverThread +from .receiver import EventReceiver from .sender import EventSender @@ -60,7 +60,7 @@ def __init__( self._stage: Usd.Stage | None = stage self._owns_stage_adapter = adapter is None destination = adapter or UsdStageAdapter(stage) - self._receiver = ReceiverThread( + self._receiver = EventReceiver( host=host, port=port, sync_from=1, @@ -68,16 +68,15 @@ def __init__( client_id=client_id or make_stable_client_id(app_name), origin=origin or client_origin(app_name, "recv"), layered_replay=True, + notifications=self._notifications, **self._credential.endpoint_kwargs(), - **self._hooks.receiver_callbacks(), ) self._dispatcher = EventDispatcher( receiver=self._receiver, adapter=destination, mirror_stage=None if destination.targets_stage() is stage else stage, - on_resync=self._hooks.on_resync, ) - self._dispatcher.on_applied_events = self._hooks.applied_events_for(self._dispatcher) + self._observe_dispatcher(self._dispatcher) @property def stage(self) -> Usd.Stage | None: @@ -89,8 +88,8 @@ def stage(self) -> Usd.Stage | None: return self._stage @property - def receiver(self) -> ReceiverThread: - """The underlying :class:`ReceiverThread`; a diagnostic handle.""" + def receiver(self) -> EventReceiver: + """The underlying :class:`EventReceiver`; a diagnostic handle.""" return self._receiver @property @@ -166,8 +165,8 @@ def acknowledge_native_scene_rebuilt(self) -> None: self._require_open() self._dispatcher.acknowledge_native_scene_rebuilt() - def _is_synchronized(self) -> bool: - return self._stage is not None and self._receiver.synchronized + def _synchronized(self, replayed: bool) -> bool: + return replayed and self._stage is not None def _is_parked(self) -> bool: return self._stage is None @@ -201,7 +200,6 @@ def __init__( replicated_api_schemas: set[str] | None = None, extra_channels: Sequence[PrimChannel] | None = None, transform_coalesce_seconds: float = 0.0, - background_send: bool = False, ): app_name = require_app_name(app_name) if not isinstance(stage, Usd.Stage): @@ -223,8 +221,7 @@ def __init__( client_id=client_id or make_stable_client_id(app_name), origin=origin or client_origin(app_name, "emit"), department=department, - on_stage_metadata=self._hooks.on_stage_metadata, - background_send=background_send, + notifications=self._notifications, **self._credential.endpoint_kwargs(), ) @@ -254,9 +251,6 @@ def update(self, *, max_messages: int | None = None) -> SyncUpdate: self._sender.request_connect() return self._progress(submitted=sent) - def _is_synchronized(self) -> bool: - return self._sender.connected - def _connect_sender(self, timeout: float | None = None) -> bool: self._paused = False return super()._connect_sender(timeout) diff --git a/packaging/blender_native/CMakeLists.txt b/packaging/blender_native/CMakeLists.txt index f3e216a..abc61d7 100644 --- a/packaging/blender_native/CMakeLists.txt +++ b/packaging/blender_native/CMakeLists.txt @@ -22,8 +22,11 @@ add_subdirectory( # support applications that still embed Python 3.11. nanobind_add_module(_native_client "${OPENUSDCONNECT_SOURCE_DIR}/native/python/client_module.cpp" + "${OPENUSDCONNECT_SOURCE_DIR}/native/python/driver_bindings.cpp" + "${OPENUSDCONNECT_SOURCE_DIR}/native/python/producer_bindings.cpp" + "${OPENUSDCONNECT_SOURCE_DIR}/native/python/receiver_bindings.cpp" ) -target_link_libraries(_native_client PRIVATE OpenUSDConnectClientCore) +target_link_libraries(_native_client PRIVATE OpenUSDConnect::ClientDriver) target_compile_features(_native_client PRIVATE cxx_std_17) if(MSVC) diff --git a/scripts/check_versions.py b/scripts/check_versions.py index c4332a8..a054895 100644 --- a/scripts/check_versions.py +++ b/scripts/check_versions.py @@ -57,6 +57,12 @@ def _generated_flatbuffers_version() -> str | None: return ".".join(parts) +def _fetched_flatbuffers_version() -> str | None: + cmake = _read("native/client_core/CMakeLists.txt") + match = re.search(r"google/flatbuffers/archive/refs/tags/v(\d+\.\d+\.\d+)\.tar\.gz", cmake) + return match.group(1) if match else None + + def _docker_instructions(dockerfile: str) -> list[str]: instructions: list[str] = [] current = "" @@ -143,9 +149,11 @@ def collect_errors() -> list[str]: setup_flatbuffers = str( _assignment("integrations/unreal/OpenUSDConnect/setup_flatbuffers.py", "DEFAULT_VERSION") ) - generated_flatbuffers = _generated_flatbuffers_version() - if not flatbuffers or flatbuffers != setup_flatbuffers or flatbuffers != generated_flatbuffers: - errors.append("FlatBuffers Python, Unreal setup, and generated-header versions must match") + mirrors = {setup_flatbuffers, _generated_flatbuffers_version(), _fetched_flatbuffers_version()} + if not flatbuffers or mirrors != {flatbuffers}: + errors.append( + "FlatBuffers Python, Unreal setup, CMake fetch, and generated-header versions must match" + ) vendored_version = ROOT / ( "integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/ThirdParty/flatbuffers/VERSION" ) diff --git a/scripts/material_zoo_blender_viewer.py b/scripts/material_zoo_blender_viewer.py index 9f2d46f..b27b462 100644 --- a/scripts/material_zoo_blender_viewer.py +++ b/scripts/material_zoo_blender_viewer.py @@ -117,7 +117,7 @@ def poll() -> float | None: print("[Material Zoo Viewer] receiver authentication rejected", flush=True) return None sender = capture.get_emitter_sender() - emitter_connected = sender is not None and sender.sock is not None + emitter_connected = sender is not None and sender.connected receiver_connected = receiver is not None and receiver.connected if last_seq >= args.expected_seq and emitter_connected and receiver_connected: _present(args.camera) diff --git a/scripts/run_material_zoo.py b/scripts/run_material_zoo.py index 70703f0..bb01977 100644 --- a/scripts/run_material_zoo.py +++ b/scripts/run_material_zoo.py @@ -477,8 +477,10 @@ def _publish(port: int, fixture_events: list[dict], presentation_events: list[di raise RuntimeError("Material Zoo fixture transaction failed") if presentation_events and not sender.send_events(presentation_events): raise RuntimeError("Material Zoo presentation transaction failed") + if not sender.flush(timeout=10): + raise RuntimeError("Material Zoo transactions were not acknowledged within 10 seconds") finally: - sender.disconnect() + sender.close() def _raise_if_viewer_failed(processes: list[subprocess.Popen]) -> None: diff --git a/scripts/run_tla_models.py b/scripts/run_tla_models.py index cacefaa..65c5fe2 100644 --- a/scripts/run_tla_models.py +++ b/scripts/run_tla_models.py @@ -29,12 +29,24 @@ ("TransactionRecoveryFirst.cfg", "TransactionRecovery.tla", "recovery: reject 1"), ("TransactionRecovery.cfg", "TransactionRecovery.tla", "recovery: reject 3"), ("RecoverySessionRollover.cfg", "RecoverySessionRollover.tla", "session rollover"), + ("ProducerConnection.cfg", "ProducerConnection.tla", "producer connection: honest"), + ( + "ProducerConnectionDivergence.cfg", + "ProducerConnection.tla", + "producer connection: divergence", + ), ("ReceiverSynchronization.cfg", "ReceiverSynchronization.tla", "receiver: queue 3"), ( "ReceiverSynchronizationTight.cfg", "ReceiverSynchronization.tla", "receiver: queue 1", ), + ("ReceiverReplayIdentity.cfg", "ReceiverReplayIdentity.tla", "replay identity: fresh"), + ( + "ReceiverReplayIdentitySnapshot.cfg", + "ReceiverReplayIdentity.tla", + "replay identity: snapshot", + ), ("TransactionCoordinator.cfg", "TransactionCoordinator.tla", "coordinator: valid"), ( "TransactionCoordinatorInvalid.cfg", @@ -53,6 +65,13 @@ # These are the adversarial or split-boundary actions most likely to become # accidentally unreachable while the models are edited. REQUIRED_ACTIONS = { + "ProducerConnection.tla": { + "AbandonAttempt", + "ConnectInterrupted", + "PeerCloses", + "ServerLosesProgress", + "ServerRunsAhead", + }, "RecoverySessionRollover.tla": { "ConcurrentAuthoritativeCommit", "RefreshCheckpoint", @@ -64,6 +83,13 @@ "InjectStaleComplete", "DiscardStaleFrame", }, + "ReceiverReplayIdentity.tla": { + "ReceiveResync", + "DropEvent", + "LiveReset", + "Restart", + "ConsumerFail", + }, "TransactionCoordinator.tla": { "GroupApply", "GroupPersist", diff --git a/tests/helpers.py b/tests/helpers.py index 079dc59..c07b7e3 100644 --- a/tests/helpers.py +++ b/tests/helpers.py @@ -6,11 +6,9 @@ import sys import threading import time -from collections import deque from contextlib import contextmanager from openusdconnect.client_observer import ClientObserver -from openusdconnect.codec import encode_message from openusdconnect.protocol_constants import ( K_SET_REFERENCE, K_SET_XFORM_TRS, @@ -150,105 +148,128 @@ def ensure_prim_event(path): @contextmanager -def in_process_server(): - """Run an isolated TCP server on an ephemeral port.""" +def serving(state, port=0): + """Serve *state* on a loopback port; leaving closes the listener and its connections.""" from openusdconnect.server.connection import ConnectionHandler, ThreadedTCPServer - from openusdconnect.server.state import UsdSyncServer - state = UsdSyncServer(log_path=":memory:", txn_batch_size=1) - tcp = ThreadedTCPServer(("127.0.0.1", 0), ConnectionHandler, state, max_workers=8) - thread = threading.Thread(target=tcp.serve_forever, daemon=True) + tcp = ThreadedTCPServer(("127.0.0.1", port), ConnectionHandler, state, max_workers=8) + # Shutdown waits for the next poll. + thread = threading.Thread(target=tcp.serve_forever, args=(0.05,), daemon=True) thread.start() try: - yield state, tcp.server_address[1] + yield tcp.server_address[1] finally: tcp.shutdown() tcp.server_close() thread.join(5) + assert not thread.is_alive() + + +@contextmanager +def server_state(): + """An isolated server state with an in-memory event log.""" + from openusdconnect.server.state import UsdSyncServer + + state = UsdSyncServer(log_path=":memory:", txn_batch_size=1) + try: + yield state + finally: state.shutdown() state.store.close() - assert not thread.is_alive() -def wait_until(predicate): - deadline = time.monotonic() + 5 - while time.monotonic() < deadline: - if predicate(): - return - time.sleep(0.005) - assert predicate() +@contextmanager +def in_process_server(): + """Run an isolated TCP server on an ephemeral port.""" + with server_state() as state, serving(state) as port: + yield state, port + + +def recorded_hellos(monkeypatch): + """Record every hello the in-process server decodes, in arrival order.""" + from openusdconnect.server import connection + + hellos = [] + decode = connection.decode_hello + + def record(table): + hellos.append(decode(table)) + return hellos[-1] + + monkeypatch.setattr(connection, "decode_hello", record) + return hellos @contextmanager -def receiver_connection(receiver): - """Run one connection attempt and close its socket before returning.""" - thread = threading.Thread(target=receiver._connect_and_recv, daemon=True) - thread.start() +def embedded_server(**config): + """Run a ``ServerRuntime`` on an ephemeral loopback port with an in-memory event log.""" + from openusdconnect.server import ServerConfig, ServerRuntime + + runtime = ServerRuntime( + ServerConfig( + host="127.0.0.1", port=0, log_path=":memory:", preflight_plugins=False, **config + ) + ) try: - wait_until(lambda: receiver.connected) - yield + with runtime: + yield runtime finally: - receiver._close_socket() - thread.join(5) - assert not thread.is_alive() - receiver.connected = False + if runtime.sync_server is not None and runtime.sync_server.token_store is not None: + runtime.sync_server.token_store.close() -def mcp_session_with_receiver(port): - """Build an MCP mirror whose connection timing is controlled by the test.""" - from pxr import Usd +def client_registered(runtime, client_id): + """Whether the server of *runtime* holds a connection from *client_id*.""" + state = runtime.sync_server + with state.clients_lock: + return any(info.client_id == client_id for info in state.clients.values()) - from integrations.mcp.config import McpConfig - from integrations.mcp.session import ConnectionSession - from openusdconnect.usd_client import UsdReceiver - session = ConnectionSession(McpConfig(read_after_write_timeout_s=1)) - session.mirror_stage = Usd.Stage.CreateInMemory() - session.receiver = UsdReceiver( - session.mirror_stage, - app_name="replay-identity-test", - host="127.0.0.1", - port=port, - persist_token=False, - ) - # These tests drive one connection attempt directly to control reconnect timing. - session.receiver._started = True - return session +def wait_until(predicate, timeout=5.0): + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if predicate(): + return + time.sleep(0.005) + assert predicate() + + +def connect_client(client, *, synchronized=True): + """Start a high-level client and complete its receiver's handshake with the server. + + ``synchronized`` also applies the server's replay the way ``update()`` does, + without publishing local edits. + """ + client.start() + receiver = client._receiver + assert receiver.wait_connected(5), receiver.connection_error + if synchronized: + apply = client.update if client._sender is None else client._apply_queued + wait_until(lambda: apply() is not None and receiver.synchronized) class PeerTraffic: - """Replaces a receiver's queue with ping messages that peers keep sending.""" - - def __init__(self, receiver, monkeypatch, *, queued=0): - self._ping = encode_message({"type": "ping"}) - self._frames = deque() - self.arrive(queued) - monkeypatch.setattr(receiver, "drain_queue", self._drain) - monkeypatch.setattr( - type(receiver), "queued_message_count", - property(lambda _receiver: len(self._frames)), - ) + """Records a peer producer commits, each queued by *receiver* when ``arrive`` returns.""" + + def __init__(self, state, receiver, make_event, *, layer_key=""): + self._state = state + self._receiver = receiver + self._make_event = make_event + self._layer_key = layer_key + self._txn_id = 0 def arrive(self, count): - self._frames.extend([self._ping] * count) - - def _drain(self, max_messages=None): - count = len(self._frames) if max_messages is None else min(max_messages, len(self._frames)) - return deque(self._frames.popleft() for _ in range(count)) - - -def force_handshake(client, *, synchronized=False): - """Mark a high-level client started with a completed receiver handshake.""" - client._started = True - receiver = getattr(client, "_receiver", None) - if receiver is not None: - receiver.connected = True - receiver.layered_replay_active = receiver.layered_replay - if synchronized: - receiver._synchronized_event.set() - graph = getattr(client, "_graph", None) - if graph is not None: - graph._ready = True + queued = self._receiver.queued_message_count + count + for _ in range(count): + self._txn_id += 1 + self._state.process_idempotent_txn( + [self._make_event(f"/Peer{self._txn_id}")], + session_id="peer", + txn_id=self._txn_id, + client_id="peer", + layer_key=self._layer_key, + ) + wait_until(lambda: self._receiver.queued_message_count == queued) class RecordingObserver(ClientObserver): diff --git a/tests/integration/asset_tests/test_assets.py b/tests/integration/asset_tests/test_assets.py index d540dda..6c177cf 100644 --- a/tests/integration/asset_tests/test_assets.py +++ b/tests/integration/asset_tests/test_assets.py @@ -290,12 +290,12 @@ def _verify_material_zoo_stage_receiver(base_path, port, event_count): from openusdconnect.adapters import UsdStageAdapter from openusdconnect.dispatcher import EventDispatcher - from openusdconnect.receiver import ReceiverThread + from openusdconnect.receiver import EventReceiver stage = Usd.Stage.Open(base_path) assert stage is not None, f"Could not open Material Zoo base stage: {base_path}" stage.SetEditTarget(stage.GetSessionLayer()) - receiver = ReceiverThread( + receiver = EventReceiver( host="127.0.0.1", port=port, sync_from=1, @@ -339,8 +339,7 @@ def _bound_material(prim_path): ).Get() assert tuple(sphere_translate) == (0.0, 1.5, 0.0) finally: - receiver.stop() - receiver.join(timeout=2.0) + receiver.close(timeout=2.0) dispatcher.close() diff --git a/tests/integration/scripts/blender_emitter_reconnect_script.py b/tests/integration/scripts/blender_emitter_reconnect_script.py index 7c6e132..e9a44e6 100644 --- a/tests/integration/scripts/blender_emitter_reconnect_script.py +++ b/tests/integration/scripts/blender_emitter_reconnect_script.py @@ -70,7 +70,7 @@ def _tick(): _result( status="FAIL", reason=f"timeout in {_phase}", - connected=bool(sender and sender.sock), + connected=bool(sender and sender.connected), pending=sender.pending_transaction_count if sender else -1, ) bpy.ops.wm.quit_blender() @@ -92,7 +92,7 @@ def _tick(): sender = capture.get_emitter_sender() if _phase == "wait_for_outage": - if not _exists("edit-now") or sender is None or sender.sock is not None: + if not _exists("edit-now") or sender is None or sender.connected: return 0.1 cube = _find_cube() assert cube is not None @@ -119,7 +119,7 @@ def _tick(): clean = not emitter or (not emitter.dirty and not emitter.prepared_event_count) if ( sender is None - or sender.sock is None + or not sender.connected or sender.pending_transaction_count or not clean or sender.acknowledged_event_count < 1 diff --git a/tests/integration/scripts/blender_receiver_script.py b/tests/integration/scripts/blender_receiver_script.py index 37a50ea..c6c1bdc 100644 --- a/tests/integration/scripts/blender_receiver_script.py +++ b/tests/integration/scripts/blender_receiver_script.py @@ -35,7 +35,7 @@ from openusdconnect.protocol_constants import ( MSG_EVENT, ) -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver def main(): @@ -55,7 +55,7 @@ def main(): sys.exit(1) print(f"[Receiver] Connecting to 127.0.0.1:{port}") - receiver = ReceiverThread(host="127.0.0.1", port=port, sync_from=1) + receiver = EventReceiver(host="127.0.0.1", port=port, sync_from=1) receiver.start() # Wait for connection + replay to complete @@ -84,11 +84,7 @@ def main(): except Exception as e: print(f"[Receiver] Error processing: {e}") - receiver.stop() - try: - receiver.join(timeout=2.0) - except Exception: - pass + receiver.close(timeout=2.0) # --- Verify results --- results = {} diff --git a/tests/integration/scripts/mtlx_ref_receiver_script.py b/tests/integration/scripts/mtlx_ref_receiver_script.py index 13504e8..545d78f 100644 --- a/tests/integration/scripts/mtlx_ref_receiver_script.py +++ b/tests/integration/scripts/mtlx_ref_receiver_script.py @@ -35,7 +35,7 @@ from openusdconnect.protocol_constants import ( MSG_EVENT, ) -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver def _process_event(adapter, ev): @@ -60,7 +60,7 @@ def main(): sys.exit(1) print(f"[MtlxRefReceiver] Connecting to 127.0.0.1:{port}") - receiver = ReceiverThread(host="127.0.0.1", port=port, sync_from=1) + receiver = EventReceiver(host="127.0.0.1", port=port, sync_from=1) receiver.start() # Poll until we receive events (the emitter may still be sending @@ -87,11 +87,7 @@ def main(): print(f"[MtlxRefReceiver] Processing: {k} {prim_path}") _process_event(adapter, ev) - receiver.stop() - try: - receiver.join(timeout=2.0) - except Exception: - pass + receiver.close(timeout=2.0) # ------------------------------------------------------------------ # Verify hierarchy diff --git a/tests/integration/scripts/payload_axis_test_script.py b/tests/integration/scripts/payload_axis_test_script.py index 9f6bdc8..04f4036 100644 --- a/tests/integration/scripts/payload_axis_test_script.py +++ b/tests/integration/scripts/payload_axis_test_script.py @@ -165,8 +165,8 @@ def main(): # Step 3: Start receiver # ================================================================== print("[Test] Step 3: Starting receiver") - from openusdconnect.receiver import ReceiverThread - receiver_addon._RECEIVER = ReceiverThread( + from openusdconnect.receiver import EventReceiver + receiver_addon._RECEIVER = EventReceiver( host="127.0.0.1", port=port, sync_from=1, ) receiver_addon._RECEIVER.start() @@ -248,8 +248,7 @@ def main(): # Cleanup if receiver_addon._RECEIVER is not None: - receiver_addon._RECEIVER.stop() - receiver_addon._RECEIVER.join(timeout=2) + receiver_addon._RECEIVER.close(timeout=2) receiver_addon._RECEIVER = None with open(out_path, "w") as f: diff --git a/tests/integration/scripts/ref_loopback_script.py b/tests/integration/scripts/ref_loopback_script.py index 0d634d0..494679f 100644 --- a/tests/integration/scripts/ref_loopback_script.py +++ b/tests/integration/scripts/ref_loopback_script.py @@ -97,10 +97,10 @@ def main(): K_SET_REFERENCE, MSG_EVENT, ) - from openusdconnect.receiver import ReceiverThread + from openusdconnect.receiver import EventReceiver adapter = BlenderAdapter() - receiver = ReceiverThread(host="127.0.0.1", port=port, sync_from=1) + receiver = EventReceiver(host="127.0.0.1", port=port, sync_from=1) receiver.start() time.sleep(2.0) @@ -122,11 +122,7 @@ def main(): adapter.apply_event(ev) - receiver.stop() - try: - receiver.join(timeout=2.0) - except Exception: - pass + receiver.close(timeout=2.0) # ================================================================== # Step 3: Inspect scene for duplicates diff --git a/tests/integration/scripts/ref_manual_then_move_script.py b/tests/integration/scripts/ref_manual_then_move_script.py index 6259f1e..759a96c 100644 --- a/tests/integration/scripts/ref_manual_then_move_script.py +++ b/tests/integration/scripts/ref_manual_then_move_script.py @@ -66,7 +66,7 @@ def main(): K_SET_XFORM_TRS, MSG_EVENT, ) - from openusdconnect.receiver import ReceiverThread + from openusdconnect.receiver import EventReceiver from openusdconnect.transport import send_line # ================================================================== @@ -108,7 +108,7 @@ def main(): # ================================================================== print("[ManualThenMove] Starting receiver...") adapter = BlenderAdapter() - receiver = ReceiverThread(host="127.0.0.1", port=port, sync_from=1) + receiver = EventReceiver(host="127.0.0.1", port=port, sync_from=1) receiver.start() time.sleep(1.5) @@ -198,11 +198,7 @@ def main(): adapter.apply_event(ev) - receiver.stop() - try: - receiver.join(timeout=2.0) - except Exception: - pass + receiver.close(timeout=2.0) # ================================================================== # Check for duplicates diff --git a/tests/integration/scripts/ref_receiver_script.py b/tests/integration/scripts/ref_receiver_script.py index 05c8975..0ef43e4 100644 --- a/tests/integration/scripts/ref_receiver_script.py +++ b/tests/integration/scripts/ref_receiver_script.py @@ -36,7 +36,7 @@ from openusdconnect.protocol_constants import ( MSG_EVENT, ) -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver def _process_event(adapter, ev): @@ -74,7 +74,7 @@ def main(): print(f"[RefReceiver] Asset root: {asset_root}") print(f"[RefReceiver] Connecting to 127.0.0.1:{port}") - receiver = ReceiverThread(host="127.0.0.1", port=port, sync_from=1) + receiver = EventReceiver(host="127.0.0.1", port=port, sync_from=1) receiver.start() # Wait for events to arrive via replay @@ -94,11 +94,7 @@ def main(): print(f"[RefReceiver] Processing: {k} {prim_path}") _process_event(adapter, ev) - receiver.stop() - try: - receiver.join(timeout=2.0) - except Exception: - pass + receiver.close(timeout=2.0) # ------------------------------------------------------------------ # Inspect scene diff --git a/tests/integration/test_department_projection.py b/tests/integration/test_department_projection.py index 4491a69..618416d 100644 --- a/tests/integration/test_department_projection.py +++ b/tests/integration/test_department_projection.py @@ -9,7 +9,7 @@ from openusdconnect.codec import HelloRejectionCode, message_to_dict from openusdconnect.protocol_constants import MSG_RESYNC -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver from openusdconnect.sender import EventSender from openusdconnect.server import UsdSyncServer from openusdconnect.server.connection import ConnectionHandler, ThreadedTCPServer @@ -68,7 +68,7 @@ def test_flat_receiver_is_admitted_only_for_single_layer( accepted, ): sync_server, port = server_factory(departments) - receiver = ReceiverThread( + receiver = EventReceiver( port=port, reconnect=False, client_id="flat-observer", @@ -78,8 +78,7 @@ def test_flat_receiver_is_admitted_only_for_single_layer( receiver.start() try: if not accepted: - receiver.join(timeout=5) - assert not receiver.is_alive() + assert _wait_until(lambda: receiver.stopped) assert not receiver.connected assert receiver.hello_rejected assert receiver.rejection_code == HelloRejectionCode.LayeredReplayRequired @@ -92,14 +91,13 @@ def test_flat_receiver_is_admitted_only_for_single_layer( assert not receiver.layered_replay_active assert sync_server._collaboration.flat_receiver_count == 1 finally: - receiver.stop() - receiver.join(timeout=2) + receiver.close(timeout=2) @pytest.mark.parametrize("departments", [[], ["animation", "layout"]]) def test_layered_receiver_is_admitted_for_both_server_modes(server_factory, departments): sync_server, port = server_factory(departments) - receiver = ReceiverThread( + receiver = EventReceiver( port=port, reconnect=False, client_id="layered-observer", @@ -111,14 +109,13 @@ def test_layered_receiver_is_admitted_for_both_server_modes(server_factory, depa assert receiver.layered_replay_active assert sync_server._collaboration.flat_receiver_count == 0 finally: - receiver.stop() - receiver.join(timeout=2) + receiver.close(timeout=2) def test_single_layer_flat_receiver_gets_live_and_replayed_records(server_factory): sync_server, port = server_factory([]) - live = ReceiverThread( + live = EventReceiver( port=port, reconnect=False, client_id="live-flat", @@ -142,11 +139,10 @@ def test_single_layer_flat_receiver_gets_live_and_replayed_records(server_factor live_records = [message_to_dict(raw) for raw in live.drain_queue()] assert [record["event"] for record in live_records] == [event] - live.stop() - live.join(timeout=2) + live.close(timeout=2) assert _wait_until(lambda: sync_server._collaboration.flat_receiver_count == 0) - late = ReceiverThread( + late = EventReceiver( port=port, reconnect=False, sync_from=1, @@ -160,16 +156,14 @@ def test_single_layer_flat_receiver_gets_live_and_replayed_records(server_factor assert [record["event"] for record in replay_records] == [event] finally: sender.disconnect() - live.stop() - live.join(timeout=2) + live.close(timeout=2) if late is not None: - late.stop() - late.join(timeout=2) + late.close(timeout=2) def test_stale_cursor_resyncs_against_an_empty_log(server_factory): _sync_server, port = server_factory([]) - receiver = ReceiverThread( + receiver = EventReceiver( port=port, reconnect=False, sync_from=5, @@ -189,14 +183,13 @@ def _received_resync(): assert _wait_until(_received_resync) finally: - receiver.stop() - receiver.join(timeout=2) + receiver.close(timeout=2) def test_flat_receiver_blocks_enabling_department_policy(server_factory): sync_server, port = server_factory([]) - receiver = ReceiverThread( + receiver = EventReceiver( port=port, reconnect=False, client_id="flat-policy-guard", @@ -208,8 +201,7 @@ def test_flat_receiver_blocks_enabling_department_policy(server_factory): with pytest.raises(ReplayModeConflictError, match="layer-stack changes"): sync_server.set_department_priority(["animation", "layout"]) finally: - receiver.stop() - receiver.join(timeout=2) + receiver.close(timeout=2) assert _wait_until(lambda: sync_server._collaboration.flat_receiver_count == 0) sync_server.set_department_priority(["animation", "layout"]) @@ -223,26 +215,24 @@ def _fail_replay(_handler, _records): raise OSError("injected replay failure") monkeypatch.setattr(sync_server, "replay_records", _fail_replay) - receiver = ReceiverThread( + receiver = EventReceiver( port=port, reconnect=False, client_id="failing-replay", origin="failing-replay-origin", ) receiver.start() - receiver.join(timeout=5) try: - assert not receiver.is_alive() + assert _wait_until(lambda: receiver.stopped) assert _wait_until(lambda: not sync_server.receivers) assert _wait_until(lambda: not sync_server.clients) finally: - receiver.stop() - receiver.join(timeout=2) + receiver.close(timeout=2) def test_compaction_replay_failure_releases_flat_reservation(server_factory, monkeypatch): sync_server, port = server_factory([]) - receiver = ReceiverThread( + receiver = EventReceiver( port=port, reconnect=False, client_id="failing-flat-compaction", @@ -274,8 +264,7 @@ def _fail_replay(_handler, _seq_start, *, seq_end=None): assert sync_server._collaboration.flat_receiver_count == 0 finally: sender.disconnect() - receiver.stop() - receiver.join(timeout=2) + receiver.close(timeout=2) def test_realtime_receiver_boundary_waits_for_pending_persistence( @@ -300,7 +289,7 @@ def _blocked_append(records, **kwargs): client_id="realtime-author", origin="realtime-author-origin", ) - receiver = ReceiverThread( + receiver = EventReceiver( port=port, reconnect=False, client_id="realtime-observer", @@ -324,6 +313,4 @@ def _blocked_append(records, **kwargs): finally: allow_persist.set() sender.disconnect() - receiver.stop() - if receiver.ident is not None: - receiver.join(timeout=2) + receiver.close(timeout=2) diff --git a/tests/integration/test_logical_layer_replay.py b/tests/integration/test_logical_layer_replay.py index 2984cc0..71976d8 100644 --- a/tests/integration/test_logical_layer_replay.py +++ b/tests/integration/test_logical_layer_replay.py @@ -12,7 +12,7 @@ from openusdconnect.adapters import MockAdapter, UsdStageAdapter from openusdconnect.dispatcher import EventDispatcher from openusdconnect.emitter import NoticeEmitter -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver from openusdconnect.sender import EventSender from openusdconnect.server import UsdSyncServer from openusdconnect.server.connection import ConnectionHandler, ThreadedTCPServer @@ -65,7 +65,7 @@ def close(self): class _LayeredClient: def __init__(self, port, client_id): self.stage = Usd.Stage.CreateInMemory() - self.receiver = ReceiverThread( + self.receiver = EventReceiver( port=port, reconnect=False, client_id=client_id, @@ -79,8 +79,7 @@ def __init__(self, port, client_id): self.receiver.start() def close(self): - self.receiver.stop() - self.receiver.join(timeout=2) + self.receiver.close(timeout=2) self.dispatcher.close() def pump_until(self, predicate, timeout=5.0): @@ -125,7 +124,7 @@ class _NativeLayeredClient: def __init__(self, port, client_id, department): self.stage = Usd.Stage.CreateInMemory() self.adapter = MockAdapter() - self.receiver = ReceiverThread( + self.receiver = EventReceiver( port=port, reconnect=False, client_id=client_id, @@ -141,8 +140,7 @@ def __init__(self, port, client_id, department): self.receiver.start() def close(self): - self.receiver.stop() - self.receiver.join(timeout=2) + self.receiver.close(timeout=2) self.dispatcher.close() def pump_until(self, predicate, timeout=5.0): diff --git a/tests/integration/test_managed_client.py b/tests/integration/test_managed_client.py index 3473d1e..179b3d8 100644 --- a/tests/integration/test_managed_client.py +++ b/tests/integration/test_managed_client.py @@ -406,10 +406,8 @@ def test_managed_client_shares_reissued_tokens(tmp_path, background, first_recon stage, app_name="token-refresh", client_id="token-refresh", port=runtime.server_address[1], persist_token=False, ) - sender_readers = [] try: assert client.connect(timeout=5) - sender_readers.append(client.sender._reader_thread) assert _drain_until(client, lambda: client.status.synchronized) old_token = client.sender.token assert old_token == client.receiver.token @@ -437,7 +435,6 @@ def test_managed_client_shares_reissued_tokens(tmp_path, background, first_recon assert client.status.connected else: assert client.connect(timeout=3) - sender_readers.append(client.sender._reader_thread) sender_tokens.append(client.sender.token) assert not client.sender.auth_rejected assert sender_tokens[0] != old_token @@ -454,9 +451,6 @@ def test_managed_client_shares_reissued_tokens(tmp_path, background, first_recon assert client.sender.token == client.receiver.token == sender_tokens[0] finally: client.close() - for worker in (*sender_readers, client.sender._connect_thread): - if worker is not None: - worker.join(timeout=3) runtime.sync_server.token_store.close() diff --git a/tests/integration/test_mcp_roundtrip.py b/tests/integration/test_mcp_roundtrip.py index b3a6682..f1290e4 100644 --- a/tests/integration/test_mcp_roundtrip.py +++ b/tests/integration/test_mcp_roundtrip.py @@ -1,6 +1,6 @@ """E2E: the MCP session authors over real TCP; mirror + other clients reflect it. -Exercises the full networked path (EventSender -> server -> ReceiverThread -> +Exercises the full networked path (EventSender -> server -> EventReceiver -> UsdStageAdapter mirror), the read-after-write drain, ancestor auto-create, and fan-out to an independent client. Headless, no DCC. """ @@ -20,7 +20,7 @@ from openusdconnect.adapters import UsdStageAdapter from openusdconnect.checkpoints import MirrorCheckpoint from openusdconnect.dispatcher import EventDispatcher -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver from openusdconnect.sender import EventSender from tests.helpers import ensure_prim_event, in_process_server, start_server, stop_server @@ -148,15 +148,15 @@ def test_reconnect_and_disconnect_join_mirror_threads(server): session = _connect(server) first = session.receiver.receiver try: - assert first.is_alive() + assert first.running session.sender.disconnect() session.connect() second = session.receiver.receiver assert second is not first - assert not first.is_alive() - assert second.is_alive() + assert not first.running + assert second.running session.disconnect() - assert not second.is_alive() + assert not second.running session.disconnect() finally: session.disconnect() @@ -167,7 +167,7 @@ def test_mesh_roundtrip_and_fanout(server): other = None try: other_stage = Usd.Stage.CreateInMemory() - other = ReceiverThread( + other = EventReceiver( host="127.0.0.1", port=server, sync_from=1, client_id="other", origin="other-recv" ) other.start() @@ -204,7 +204,7 @@ def test_mesh_roundtrip_and_fanout(server): assert UsdGeom.Mesh(om).GetPointsAttr().Get() is not None finally: if other is not None: - other.stop() + other.close() session.disconnect() diff --git a/tests/integration/test_receiver_overflow_identity.py b/tests/integration/test_receiver_overflow_identity.py index d14df32..824abb9 100644 --- a/tests/integration/test_receiver_overflow_identity.py +++ b/tests/integration/test_receiver_overflow_identity.py @@ -1,210 +1,224 @@ """A partial replay must retain enough identity to resume after queue overflow.""" -import threading +from typing import NamedTuple import pytest from pxr import Usd from openusdconnect.adapters import UsdStageAdapter -from openusdconnect.codec import ( - PayloadType, - payload_type_and_sequence, -) from openusdconnect.dispatcher import EventDispatcher -from openusdconnect.receiver import ReceiverThread -from tests.helpers import ensure_prim_event, in_process_server - - -@pytest.fixture -def replay_server(): - with in_process_server() as server: - yield server - - -def _receive_until_boundary(receiver, dispatcher, monkeypatch, *, after_replay=None): - """Do not drain until overflow or ReplayComplete, independent of scheduling.""" - boundary = threading.Event() - replay_complete = threading.Event() - errors = [] - handle_data = receiver._handle_data_message - - def observe_data(buf, generation): - accepted = handle_data(buf, generation) - payload_type, _sequence = payload_type_and_sequence(buf) - if accepted and payload_type == PayloadType.ReplayComplete: - replay_complete.set() - boundary.set() - elif not accepted: - boundary.set() - return accepted - - def receive(): - try: - receiver._connect_and_recv() - except Exception as exc: - errors.append(exc) - finally: - boundary.set() +from openusdconnect.receiver import EventReceiver +from tests.helpers import ensure_prim_event, recorded_hellos, server_state, serving, wait_until - with monkeypatch.context() as patch: - patch.setattr(receiver, "_handle_data_message", observe_data) - worker = threading.Thread(target=receive, daemon=True) - worker.start() - try: - assert boundary.wait(5), "receiver reached neither overflow nor ReplayComplete" - assert not errors, errors - if after_replay is not None: - assert replay_complete.is_set() - dispatcher.drain_and_apply() - assert receiver.synchronized - boundary.clear() - replay_complete.clear() - after_replay() - assert boundary.wait(5), "live reset reached neither overflow nor ReplayComplete" - assert not errors, errors - completed = replay_complete.is_set() - assert completed or receiver._inbox.overflowed, "receiver stopped before replay" - if not completed: - assert receiver.queued_message_count == receiver.max_queue - dispatcher.drain_and_apply() - if completed: - assert receiver.synchronized - return completed - finally: - receiver._close_socket() - worker.join(5) - receiver.connected = False - receiver._inbox.clear_overflow() - assert not worker.is_alive() - assert not errors, errors +# Covers a reconnect after the default base delay, including a refused attempt. +BOUNDARY_TIMEOUT = 15 -@pytest.mark.parametrize("after_purge", [False, True], ids=["initial", "after-purge"]) -def test_partial_replay_advances_across_queue_overflow(replay_server, monkeypatch, after_purge): - state, port = replay_server - receiver = ReceiverThread(host="127.0.0.1", port=port, max_queue=3) - stage = Usd.Stage.CreateInMemory() - dispatcher = EventDispatcher(receiver=receiver, adapter=UsdStageAdapter(stage)) - try: - if after_purge: - state._commit_events([ensure_prim_event("/Before")]) - assert _receive_until_boundary(receiver, dispatcher, monkeypatch) - assert stage.GetPrimAtPath("/Before") - state.purge() - - paths = [f"/P{index}" for index in range(8)] - state._commit_events([ensure_prim_event(path) for path in paths]) - head = state.store.get_max_seq() - assert head == len(paths) - - progress = [] - completed = False - for _ in range(len(paths)): - completed = _receive_until_boundary(receiver, dispatcher, monkeypatch) - progress.append(dispatcher.last_seq) - if completed: - break - - assert len(progress) > 1, "fixture must force at least one replay overflow" - assert all(after > before for before, after in zip(progress, progress[1:], strict=False)), ( - f"partial replay restarted instead of advancing: {progress}" +class Boundary(NamedTuple): + completed: bool + # The hello that opened the connection. + hello: dict + + +class _Replica: + """A stage whose consumer applies the receiver queue only at replay boundaries. + + Until the consumer drains, a connection either overflows the queue or + holds the whole replay, so each boundary is independent of scheduling. + After an overflow the receiver reconnects only once the queue is drained. + """ + + def __init__(self, state, port, hellos, *, stage=None, **options): + self.stage = stage or Usd.Stage.CreateInMemory() + self.receiver = EventReceiver(host="127.0.0.1", port=port, **options) + self.dispatcher = EventDispatcher( + receiver=self.receiver, adapter=UsdStageAdapter(self.stage) + ) + self._state = state + self._hellos = hellos + self._next_hello = len(hellos) + + def close(self): + self.receiver.close(timeout=5) + self.dispatcher.close() + + def next_boundary(self, *, after_replay=None): + """Wait for overflow or the replay's last record, then apply the queue.""" + if not (self.receiver.running or self.receiver.stopped): + self.receiver.start() + hello = self._next_hello + completed = self._reach_boundary() + if after_replay is not None: + assert completed + self._apply(completed) + after_replay() + completed = self._reach_boundary() + self._apply(completed) + return Boundary(completed, self._hellos[hello]) + + def _overflowed(self): + receiver = self.receiver + return not receiver.connected and receiver.queued_message_count == receiver.max_queue + + def _reach_boundary(self): + receiver = self.receiver + head = self._state.store.get_max_seq() + wait_until( + lambda: self._overflowed() or (receiver.connected and receiver.last_seq == head), + timeout=BOUNDARY_TIMEOUT, ) - assert completed, f"replay never completed: {progress}" - assert dispatcher.last_seq == head - assert receiver.server_instance == state.server_instance - assert receiver.replay_epoch == state.get_replay_token()[0] - assert all(stage.GetPrimAtPath(path) for path in paths) - assert not stage.GetPrimAtPath("/Before") - finally: - receiver.stop() - dispatcher.close() + return not self._overflowed() + + def _apply(self, completed): + # The receiver reconnects only after this drain, so its next hello follows. + self._next_hello = len(self._hellos) + if completed: + wait_until( + lambda: self.dispatcher.drain_and_apply() is not None and self.receiver.synchronized + ) + else: + self.dispatcher.drain_and_apply() + + +@pytest.mark.parametrize("after_purge", [False, True], ids=["initial", "after-purge"]) +def test_partial_replay_advances_across_queue_overflow(monkeypatch, after_purge): + hellos = recorded_hellos(monkeypatch) + with server_state() as state: + port = 0 + replica = None + try: + if after_purge: + state._commit_events([ensure_prim_event("/Before")]) + with serving(state) as port: + replica = _Replica(state, port, hellos, max_queue=3) + assert replica.next_boundary().completed + assert replica.stage.GetPrimAtPath("/Before") + # The receiver finds the purge when it reconnects. + state.purge() + + paths = [f"/P{index}" for index in range(8)] + state._commit_events([ensure_prim_event(path) for path in paths]) + head = state.store.get_max_seq() + assert head == len(paths) + + with serving(state, port) as port: + replica = replica or _Replica(state, port, hellos, max_queue=3) + progress = [] + completed = False + for _ in range(len(paths)): + completed = replica.next_boundary().completed + progress.append(replica.dispatcher.last_seq) + if completed: + break + + assert len(progress) > 1, "fixture must force at least one replay overflow" + assert all( + after > before for before, after in zip(progress, progress[1:], strict=False) + ), f"partial replay restarted instead of advancing: {progress}" + assert completed, f"replay never completed: {progress}" + assert replica.dispatcher.last_seq == head + assert replica.receiver.server_instance == state.server_instance + assert replica.receiver.replay_epoch == state.get_replay_token()[0] + assert all(replica.stage.GetPrimAtPath(path) for path in paths) + assert not replica.stage.GetPrimAtPath("/Before") + finally: + if replica is not None: + replica.close() @pytest.mark.parametrize("layered_replay", [False, True], ids=["flat", "layered"]) @pytest.mark.parametrize("explicit_replay", [False, True], ids=["initial-overflow", "full-replay"]) def test_snapshot_cursor_overflow_before_first_reset_event_recovers( - replay_server, monkeypatch, layered_replay, explicit_replay, + monkeypatch, + layered_replay, + explicit_replay, ): - state, port = replay_server + hellos = recorded_hellos(monkeypatch) snapshot_paths = [f"/P{index}" for index in range(1, 4)] remaining_paths = [f"/P{index}" for index in range(4, 7)] - state._commit_events([ensure_prim_event(path) for path in snapshot_paths]) stage = Usd.Stage.CreateInMemory() for path in snapshot_paths: stage.DefinePrim(path, "Xform") - receiver = ReceiverThread( - host="127.0.0.1", port=port, sync_from=4, - max_queue=2 if layered_replay else 1, layered_replay=layered_replay, - ) - dispatcher = EventDispatcher(receiver=receiver, adapter=UsdStageAdapter(stage)) - dispatcher.last_seq = 3 - try: - if explicit_replay: - assert _receive_until_boundary(receiver, dispatcher, monkeypatch) - assert receiver.server_instance == "" - receiver.request_replay_from(1) - state._commit_events([ensure_prim_event(path) for path in remaining_paths]) - - # The reset (plus layer stack when negotiated) fills the queue before - # event 1. Its next connection must start at 1, not the snapshot cursor 4. - progress = [] - completed = False - for _ in range(10): - completed = _receive_until_boundary(receiver, dispatcher, monkeypatch) - progress.append(dispatcher.last_seq) - if completed: - break - assert 0 in progress, f"fixture did not overflow before event 1: {progress}" - assert completed, f"snapshot replay never completed: {progress}" - reset_index = progress.index(0) - assert progress[reset_index:] == list(range(7)) - assert receiver.last_seq == dispatcher.last_seq == state.store.get_max_seq() == 6 - assert receiver.server_instance == state.server_instance - assert receiver.replay_epoch == state.get_replay_token()[0] - assert all(stage.GetPrimAtPath(path) for path in snapshot_paths + remaining_paths) - finally: - receiver.stop() - dispatcher.close() - - -def test_live_compaction_overflow_resumes_in_new_epoch(replay_server, monkeypatch): - state, port = replay_server + with server_state() as state, serving(state) as port: + state._commit_events([ensure_prim_event(path) for path in snapshot_paths]) + replica = _Replica( + state, + port, + hellos, + stage=stage, + sync_from=4, + max_queue=2 if layered_replay else 1, + layered_replay=layered_replay, + ) + replica.dispatcher.last_seq = 3 + try: + if explicit_replay: + assert replica.next_boundary().completed + assert replica.receiver.server_instance == "" + replica.receiver.request_replay_from(1) + state._commit_events([ensure_prim_event(path) for path in remaining_paths]) + + # The reset (plus layer stack when negotiated) fills the queue before + # event 1. Its next connection must start at 1, not the snapshot cursor 4. + progress = [] + completed = False + for _ in range(10): + completed = replica.next_boundary().completed + progress.append(replica.dispatcher.last_seq) + if completed: + break + assert 0 in progress, f"fixture did not overflow before event 1: {progress}" + assert completed, f"snapshot replay never completed: {progress}" + reset_index = progress.index(0) + assert progress[reset_index:] == list(range(7)) + receiver = replica.receiver + assert receiver.last_seq == replica.dispatcher.last_seq == 6 + assert state.store.get_max_seq() == 6 + assert receiver.server_instance == state.server_instance + assert receiver.replay_epoch == state.get_replay_token()[0] + assert all(stage.GetPrimAtPath(path) for path in snapshot_paths + remaining_paths) + finally: + replica.close() + + +def test_live_compaction_overflow_resumes_in_new_epoch(monkeypatch): + hellos = recorded_hellos(monkeypatch) paths = [f"/P{index}" for index in range(8)] - state._commit_events([ensure_prim_event(paths[0])]) - receiver = ReceiverThread(host="127.0.0.1", port=port, max_queue=3) - stage = Usd.Stage.CreateInMemory() - dispatcher = EventDispatcher(receiver=receiver, adapter=UsdStageAdapter(stage)) + with server_state() as state, serving(state) as port: + state._commit_events([ensure_prim_event(paths[0])]) + replica = _Replica(state, port, hellos, max_queue=3) - def compact_live(): - # Keep the initial replay small, then force a large reset on the same socket. - state._commit_events([ensure_prim_event(path) for path in paths[1:]]) - state.compact_log() + def compact_live(): + # Keep the initial replay small, then force a large reset on the same socket. + state._commit_events([ensure_prim_event(path) for path in paths[1:]]) + state.compact_log() - try: - assert not _receive_until_boundary( - receiver, dispatcher, monkeypatch, after_replay=compact_live, - ) - assert state.get_replay_token()[0] == 1 - assert receiver._received_replay_identity is None - assert receiver.replay_epoch == 0 - - # The epoch-less live Resync needs one fresh handshake/reset. Subsequent - # overflows must retain that handshake's epoch and advance normally. - progress = [] - completed = False - for _ in range(len(paths)): - completed = _receive_until_boundary(receiver, dispatcher, monkeypatch) - progress.append(dispatcher.last_seq) - if completed: - break - assert all(after > before for before, after in zip(progress, progress[1:], strict=False)), ( - f"live reset replay restarted instead of advancing: {progress}" - ) - assert completed, f"live reset replay never completed: {progress}" - assert dispatcher.last_seq == state.store.get_max_seq() - assert receiver.replay_epoch == 1 - assert receiver.server_instance == state.server_instance - assert all(stage.GetPrimAtPath(path) for path in paths) - finally: - receiver.stop() - dispatcher.close() + try: + assert not replica.next_boundary(after_replay=compact_live).completed + assert state.get_replay_token()[0] == 1 + assert replica.receiver.replay_epoch == 0 + + # The epoch-less live Resync needs one fresh handshake/reset. Subsequent + # overflows must retain that handshake's epoch and advance normally. + progress = [] + completed = False + for _ in range(len(paths)): + boundary = replica.next_boundary() + if not progress: + # The live reset left no proven identity, so the hello claims none. + assert boundary.hello["replay_server_instance"] == "" + assert "replay_epoch" not in boundary.hello + progress.append(replica.dispatcher.last_seq) + completed = boundary.completed + if completed: + break + assert all( + after > before for before, after in zip(progress, progress[1:], strict=False) + ), f"live reset replay restarted instead of advancing: {progress}" + assert completed, f"live reset replay never completed: {progress}" + assert replica.dispatcher.last_seq == state.store.get_max_seq() + assert replica.receiver.replay_epoch == 1 + assert replica.receiver.server_instance == state.server_instance + assert all(replica.stage.GetPrimAtPath(path) for path in paths) + finally: + replica.close() diff --git a/tests/integration/test_receiver_replay_identity.py b/tests/integration/test_receiver_replay_identity.py index 32160cf..6645be9 100644 --- a/tests/integration/test_receiver_replay_identity.py +++ b/tests/integration/test_receiver_replay_identity.py @@ -1,54 +1,78 @@ """Reconnect must not reuse a cursor from another server or sequence epoch.""" import socket +from contextlib import ExitStack import pytest +from pxr import Usd +from integrations.mcp.config import McpConfig +from integrations.mcp.session import ConnectionSession from openusdconnect.codec import encode_message, message_to_dict from openusdconnect.framing import recv_framed, send_framed from openusdconnect.protocol import make_hello -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver from openusdconnect.sender import EventSender +from openusdconnect.usd_client import UsdReceiver from tests.helpers import ( ensure_prim_event, in_process_server, - mcp_session_with_receiver, - receiver_connection, + recorded_hellos, + server_state, + serving, wait_until, ) +# A receiver retries a lost connection after its reconnect delay, and a +# refused attempt takes seconds on some platforms. +RECONNECT_TIMEOUT = 15 + + +def _mcp_session(port): + """An MCP mirror session whose receiver is connecting to *port*.""" + session = ConnectionSession(McpConfig(read_after_write_timeout_s=1)) + session.mirror_stage = Usd.Stage.CreateInMemory() + session.receiver = UsdReceiver( + session.mirror_stage, + app_name="replay-identity-test", + host="127.0.0.1", + port=port, + persist_token=False, + ) + session.receiver.start() + return session + def _drain_ready(session): def ready(): session.receiver.update() return session.receiver.status.synchronized - wait_until(ready) + + wait_until(ready, timeout=RECONNECT_TIMEOUT) @pytest.mark.parametrize("reset", ["compact", "purge", "restart"]) def test_colliding_reconnect_cannot_confirm_missing_own_write(reset): - with ( - in_process_server() as (state, port), - in_process_server() as (replacement, replacement_port), - ): + with server_state() as state, server_state() as replacement, ExitStack() as cleanup: state._commit_events([ensure_prim_event("/Before")] * 3) - session = mcp_session_with_receiver(port) - try: - with receiver_connection(session.receiver.receiver): - _drain_ready(session) - assert session.receiver.last_seq == 3 - assert session.mirror_stage.GetPrimAtPath("/Before") - wait_until(lambda: not state.receivers) - - if reset == "compact": - state.compact_log() - assert state.store.get_max_seq() == 1 - elif reset == "purge": - state.purge() - else: - state, port = replacement, replacement_port - session.receiver.receiver.port = port - session.sender = EventSender("127.0.0.1", port, client_id="own") + with serving(state) as port: + session = _mcp_session(port) + cleanup.callback(session.disconnect) + _drain_ready(session) + assert session.receiver.last_seq == 3 + assert session.mirror_stage.GetPrimAtPath("/Before") + # Closing the listener dropped the receiver; it reconnects once the port serves again. + assert not state.receivers + + if reset == "compact": + state.compact_log() + assert state.store.get_max_seq() == 1 + elif reset == "purge": + state.purge() + else: + state = replacement + with serving(state) as producer_port: + session.sender = EventSender("127.0.0.1", producer_port, client_id="own") assert session.sender.connect() assert session.sender.send_events([ensure_prim_event("/Own")]) assert session.sender.flush(5) @@ -56,15 +80,14 @@ def test_colliding_reconnect_cannot_confirm_missing_own_write(reset): state._commit_events([ensure_prim_event("/Foreign")]) assert not session.mirror_stage.GetPrimAtPath("/Own") - with receiver_connection(session.receiver.receiver): + with serving(state, port): + wait_until(lambda: session.receiver.receiver.connected, timeout=RECONNECT_TIMEOUT) assert session._drain_after_write() assert session.mirror_stage.GetPrimAtPath("/Own") assert session.receiver.last_seq == 3 assert session.receiver.server_instance == state.server_instance if reset != "compact": assert not session.mirror_stage.GetPrimAtPath("/Before") - finally: - session.disconnect() @pytest.mark.parametrize("prefix", ["matching", "epoch", "instance", "unknown", "legacy", "fresh"]) @@ -77,8 +100,11 @@ def test_server_validates_prefix_inside_replay_window(prefix): hello["sync_from"] = 1 if prefix != "legacy": hello["replay_server_instance"] = ( - "" if prefix in ("unknown", "fresh") else - "another-server" if prefix == "instance" else state.server_instance + "" + if prefix in ("unknown", "fresh") + else "another-server" + if prefix == "instance" + else state.server_instance ) if prefix not in ("unknown", "fresh"): hello["replay_epoch"] = epoch + (prefix == "epoch") @@ -95,70 +121,34 @@ def test_server_validates_prefix_inside_replay_window(prefix): ) -@pytest.mark.parametrize("instance", ["", "older-checkpoint-server"]) -def test_old_server_retains_cursor_without_publishing_confirmation_identity(instance): - with socket.socket() as listener: - listener.bind(("127.0.0.1", 0)) - listener.listen(2) - listener.settimeout(5) - receiver = ReceiverThread( - host="127.0.0.1", port=listener.getsockname()[1], - sync_from=4, reconnect=False, layered_replay=False, - ) - receiver.start() - try: - with listener.accept()[0] as second: - second.settimeout(5) - hello = message_to_dict(recv_framed(second)) - assert hello["sync_from"] == 4 - send_framed(second, encode_message({ - "type": "hello_ok", "server_instance": instance, - })) - assert receiver.wait_connected(5) - send_framed(second, encode_message({ - "type": "event", "seq": 4, "event": ensure_prim_event("/Own"), - })) - send_framed(second, encode_message({ - "type": "replay_complete", "head_seq": 4, "epoch": 0, - })) - wait_until(lambda: receiver.last_seq == 4) - assert [message_to_dict(raw)["type"] for raw in receiver.drain_queue()] == [ - "event", - ] - wait_until(receiver.mark_replay_applied) - assert receiver.synchronized - assert receiver.server_instance == "" - assert receiver._received_replay_identity is None - receiver.join(5) - assert not receiver.is_alive() - finally: - receiver.stop() - receiver.join(5) - - -def test_initial_snapshot_cursor_is_preserved_without_claiming_prefix_proof(): - with in_process_server() as (state, port): +def test_initial_snapshot_cursor_is_preserved_without_claiming_prefix_proof(monkeypatch): + hellos = recorded_hellos(monkeypatch) + with server_state() as state, ExitStack() as cleanup: state._commit_events([ensure_prim_event("/Snapshot"), ensure_prim_event("/PostSnapshot")]) - receiver = ReceiverThread(host="127.0.0.1", port=port, sync_from=2) received = [] def drain_ready(): received.extend(message_to_dict(raw) for raw in receiver.drain_queue()) return receiver.mark_replay_applied() - with receiver_connection(receiver): + with serving(state) as port: + receiver = EventReceiver(host="127.0.0.1", port=port, sync_from=2) + cleanup.callback(receiver.close, 5) + receiver.start() wait_until(drain_ready) assert not any(msg["type"] == "resync" for msg in received) assert [msg["seq"] for msg in received if msg["type"] == "event"] == [2] assert receiver.synchronized assert receiver.server_instance == "" - assert receiver._received_replay_identity is None # The next connection must validate that externally supplied prefix. # With no identity for it, a complete reset/replay establishes proof. received.clear() - with receiver_connection(receiver): - wait_until(drain_ready) + with serving(state, port): + wait_until(drain_ready, timeout=RECONNECT_TIMEOUT) + assert len(hellos) == 2 + assert hellos[1]["replay_server_instance"] == "" + assert "replay_epoch" not in hellos[1] assert any(msg["type"] == "resync" for msg in received) assert [msg["seq"] for msg in received if msg["type"] == "event"] == [1, 2] assert receiver.server_instance == state.server_instance @@ -167,52 +157,53 @@ def drain_ready(): def test_apply_failure_discards_unapplied_replay_identity(monkeypatch): with in_process_server() as (state, port): state._commit_events([ensure_prim_event("/Before")] * 3) - session = mcp_session_with_receiver(port) + session = _mcp_session(port) try: - with receiver_connection(session.receiver.receiver): - _drain_ready(session) - session.sender = EventSender("127.0.0.1", port, client_id="own") - assert session.sender.connect() + _drain_ready(session) + session.sender = EventSender("127.0.0.1", port, client_id="own") + assert session.sender.connect() + state.process_idempotent_txn( + [ensure_prim_event("/ApplyFailure")], + session_id="foreign", + txn_id=1, + client_id="foreign", + ) + wait_until(lambda: session.receiver.receiver.last_seq == 4) + + def fail_while_reset_arrives(*args, **kwargs): + # The consumer owns an old-epoch batch while the receiver + # queues a reset and new-epoch records with colliding IDs. + state.purge() + assert session.sender.send_events([ensure_prim_event("/Own")]) + assert session.sender.flush(5) state.process_idempotent_txn( - [ensure_prim_event("/ApplyFailure")], session_id="foreign", txn_id=1, + [ensure_prim_event("/Foreign")] * 2, + session_id="foreign", + txn_id=2, client_id="foreign", ) - wait_until(lambda: session.receiver.receiver.last_seq == 4) - - def fail_while_reset_arrives(*args, **kwargs): - # The consumer owns an old-epoch batch while the receiver - # queues a reset and new-epoch records with colliding IDs. - state.purge() - assert session.sender.send_events([ensure_prim_event("/Own")]) - assert session.sender.flush(5) - state.process_idempotent_txn( - [ensure_prim_event("/Foreign")] * 2, session_id="foreign", txn_id=2, - client_id="foreign", - ) - wait_until(lambda: ( - session.receiver.receiver._received_replay_identity - == (state.server_instance, 1) - and session.receiver.receiver.last_seq == 3 - )) - assert session.receiver.receiver.replay_epoch == 0 - raise RuntimeError("injected apply failure") - - with monkeypatch.context() as patch: - patch.setattr( - session.receiver.dispatcher, "_apply_layered", fail_while_reset_arrives - ) - with pytest.raises(RuntimeError, match="injected apply failure"): - session.receiver.update() - assert session.receiver.last_seq == 3 - assert session.mirror_stage.GetPrimAtPath("/Before") - assert not session.mirror_stage.GetPrimAtPath("/Own") + # The purge's reset and new-epoch marker precede these records. + wait_until(lambda: session.receiver.receiver.last_seq == 3) + assert session.receiver.receiver.replay_epoch == 0 + raise RuntimeError("injected apply failure") + + with monkeypatch.context() as patch: + patch.setattr( + session.receiver.dispatcher, "_apply_layered", fail_while_reset_arrives + ) + with pytest.raises(RuntimeError, match="injected apply failure"): + session.receiver.update() + assert session.receiver.last_seq == 3 + assert session.mirror_stage.GetPrimAtPath("/Before") + assert not session.mirror_stage.GetPrimAtPath("/Own") - with receiver_connection(session.receiver.receiver): - assert session._drain_after_write() - assert session.mirror_stage.GetPrimAtPath("/Own") - assert not session.mirror_stage.GetPrimAtPath("/Before") - assert session.receiver.last_seq == 3 - assert session.receiver.replay_epoch == 1 - assert session.receiver.server_instance == state.server_instance + # The failed batch requested a replay, which reconnects. + wait_until(lambda: session.receiver.receiver.connected, timeout=RECONNECT_TIMEOUT) + assert session._drain_after_write() + assert session.mirror_stage.GetPrimAtPath("/Own") + assert not session.mirror_stage.GetPrimAtPath("/Before") + assert session.receiver.last_seq == 3 + assert session.receiver.replay_epoch == 1 + assert session.receiver.server_instance == state.server_instance finally: session.disconnect() diff --git a/tests/integration/test_replay_capture_identity.py b/tests/integration/test_replay_capture_identity.py index 72a25dc..1024268 100644 --- a/tests/integration/test_replay_capture_identity.py +++ b/tests/integration/test_replay_capture_identity.py @@ -7,48 +7,45 @@ from pxr import Usd from openusdconnect.checkpoints import MirrorCheckpoint -from openusdconnect.codec import PayloadType, encode_message, message_to_dict +from openusdconnect.codec import encode_message, message_to_dict from openusdconnect.framing import recv_framed, send_framed from openusdconnect.protocol import make_hello from openusdconnect.sender import EventSender from openusdconnect.server import connection as connection_mod -from tests.helpers import ( - ReceiverStub, - ensure_prim_event, - in_process_server, - mcp_session_with_receiver, - receiver_connection, -) +from tests.helpers import ReceiverStub, ensure_prim_event, in_process_server -def test_snapshot_replacement_after_capture_cannot_confirm_unapplied_write(monkeypatch): +def _receive_replay(sock): + """Messages through the next replay completion marker.""" + messages = [] + while not messages or messages[-1]["type"] != "replay_complete": + messages.append(message_to_dict(recv_framed(sock))) + return messages + + +def test_snapshot_replacement_after_capture_follows_the_captured_replay(monkeypatch): with in_process_server() as (state, port): - session = mcp_session_with_receiver(port) - session.config.read_after_write_timeout_s = 0.1 - session.sender = EventSender("127.0.0.1", port, client_id="own") - replay_complete = threading.Event() - resume_receiver = threading.Event() + sender = EventSender("127.0.0.1", port, client_id="own") replacement_done = threading.Event() replacement_errors = [] replacement_worker = None try: - assert session.sender.connect() - assert session.sender.send_events([ensure_prim_event("/Own")]) - assert session.sender.flush(5) - assert session.sender.acknowledged_checkpoint == MirrorCheckpoint( - state.server_instance, 0, 1 - ) + assert sender.connect() + assert sender.send_events([ensure_prim_event("/Own")]) + assert sender.flush(5) + assert sender.acknowledged_checkpoint == MirrorCheckpoint(state.server_instance, 0, 1) state._broadcast_queue.join() replacement = Usd.Stage.CreateInMemory() replacement.DefinePrim("/Replacement", "Xform") epoch, head = state.get_snapshot_token() replacement.GetRootLayer().customLayerData = { "openusdconnect": { - "scene_id": state.scene_id, "epoch": epoch, "snapshot_seq": head, + "scene_id": state.scene_id, + "epoch": epoch, + "snapshot_seq": head, }, } send = connection_mod.send_msg - control = session.receiver.receiver._handle_control_message def replace_snapshot(): try: @@ -70,32 +67,29 @@ def replace_before_hello(sock, message): assert state.get_replay_token()[0] == 1 send(sock, message) - def pause_after_initial_complete(payload_type, buf, generation): - result = control(payload_type, buf, generation) - if payload_type == PayloadType.ReplayComplete and not replay_complete.is_set(): - replay_complete.set() - assert resume_receiver.wait(5) - return result - monkeypatch.setattr(connection_mod, "send_msg", replace_before_hello) - monkeypatch.setattr( - session.receiver.receiver, - "_handle_control_message", - pause_after_initial_complete, - ) - with receiver_connection(session.receiver.receiver): - try: - assert replay_complete.wait(5) - assert session._drain_after_write() - assert session.mirror_stage.GetPrimAtPath("/Own") - assert not session.mirror_stage.GetPrimAtPath("/Replacement") - assert session.receiver.last_seq == 1 - assert session.receiver.replay_epoch == 0 - finally: - resume_receiver.set() + with socket.create_connection(("127.0.0.1", port), timeout=5) as sock: + send_framed(sock, encode_message(make_hello("receiver", layered_replay=True))) + captured = _receive_replay(sock) + replaced = _receive_replay(sock) + assert replacement_done.is_set() + + # The captured epoch replays its own write, so a mirror can confirm + # it before the replacement that follows resets the stream. + assert captured[0]["type"] == "hello_ok" and captured[0]["replay_epoch"] == 0 + assert [ + (message["seq"], message["event"]["prim"]) + for message in captured + if message["type"] == "event" + ] == [(1, "/Own")] + assert (captured[-1]["head_seq"], captured[-1]["epoch"]) == (1, 0) + assert replaced[0]["type"] == "resync" + assert "/Replacement" in { + message["event"]["prim"] for message in replaced if message["type"] == "event" + } + assert replaced[-1]["epoch"] == 1 finally: - resume_receiver.set() - session.disconnect() + sender.disconnect() if replacement_worker is not None: replacement_worker.join(5) assert not replacement_worker.is_alive() diff --git a/tests/integration/test_shared_stage_client.py b/tests/integration/test_shared_stage_client.py index 1251ee3..0e05d25 100644 --- a/tests/integration/test_shared_stage_client.py +++ b/tests/integration/test_shared_stage_client.py @@ -88,10 +88,8 @@ def test_shared_client_shares_reissued_tokens(tmp_path, background, first_reconn stage, app_name="token-refresh", client_id="token-refresh", port=runtime.server_address[1], persist_token=False, ) - sender_readers = [] try: assert client.connect(timeout=5) - sender_readers.append(client._sender._reader_thread) assert _pump_until([client], lambda: client.status.synchronized) old_token = client._sender.token assert old_token == client._receiver.token @@ -119,7 +117,6 @@ def test_shared_client_shares_reissued_tokens(tmp_path, background, first_reconn assert client.status.connected else: assert client.connect(timeout=3) - sender_readers.append(client._sender._reader_thread) sender_tokens.append(client._sender.token) assert not client._sender.auth_rejected assert sender_tokens[0] != old_token @@ -136,9 +133,6 @@ def test_shared_client_shares_reissued_tokens(tmp_path, background, first_reconn assert client._sender.token == client._receiver.token == sender_tokens[0] finally: client.close() - for worker in (*sender_readers, client._sender._connect_thread): - if worker is not None: - worker.join(timeout=3) runtime.sync_server.token_store.close() diff --git a/tests/integration/test_stage_first_receiver.py b/tests/integration/test_stage_first_receiver.py index 884a668..e2db17e 100644 --- a/tests/integration/test_stage_first_receiver.py +++ b/tests/integration/test_stage_first_receiver.py @@ -1,7 +1,7 @@ """Integration test for the Blender receiver's stage-first architecture. Verifies the full flow: server → receiver → stage commit (atomic) → adapter. -Uses a real server (subprocess), real ReceiverThread, and MockAdapter. +Uses a real server (subprocess), real EventReceiver, and MockAdapter. No Blender required headless, runs in CI. """ @@ -26,7 +26,7 @@ MSG_EVENT, MSG_RESYNC, ) -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver from openusdconnect.transport import recv_msg, send_msg from tests.helpers import start_server, stop_server @@ -85,7 +85,7 @@ def _receive_events(min_events=1, timeout=30.0): return as soon as the events land. Fails here, at the wait, rather than letting a short drain confuse downstream assertions. """ - rt = ReceiverThread( + rt = EventReceiver( host="127.0.0.1", port=PORT, sync_from=1, client_id="test-receiver", origin="test-recv" ) rt.start() @@ -96,7 +96,7 @@ def _receive_events(min_events=1, timeout=30.0): if len(_parse_events_from_bufs(lines)) >= min_events: break time.sleep(0.05) - rt.stop() + rt.close() got = len(_parse_events_from_bufs(lines)) assert got >= min_events, f"receiver drained {got}/{min_events} events within {timeout}s" return lines diff --git a/tests/integration/test_transaction_reconnect.py b/tests/integration/test_transaction_reconnect.py index 272aa95..962ba57 100644 --- a/tests/integration/test_transaction_reconnect.py +++ b/tests/integration/test_transaction_reconnect.py @@ -17,7 +17,7 @@ MSG_TRANSACTION_RESULT, PROTOCOL_VERSION, ) -from openusdconnect.receiver import ReceiverThread +from openusdconnect.receiver import EventReceiver from openusdconnect.sender import EventSender from openusdconnect.server import UsdSyncServer from openusdconnect.server.connection import ConnectionHandler, ThreadedTCPServer @@ -439,7 +439,7 @@ def observe_group(records, *, producer_progress=()): return append_batch(records, producer_progress=producer_progress) state.store.append_batch = observe_group - receiver = ReceiverThread( + receiver = EventReceiver( host="127.0.0.1", port=port, reconnect=False, @@ -493,8 +493,7 @@ def received_everything(): assert sequences == list(range(1, len(senders) + 1)) assert group_sizes == [len(senders)] finally: - receiver.stop() - receiver.join(timeout=5) + receiver.close(timeout=5) for sender in senders: sender.disconnect() @@ -536,7 +535,7 @@ def observe_group_publication(transactions): monkeypatch.setattr( state, "broadcast_transaction_group_views", observe_group_publication ) - receiver = ReceiverThread( + receiver = EventReceiver( host="127.0.0.1", port=port, reconnect=False, @@ -575,8 +574,7 @@ def received_both(): assert [record["seq"] for record in records] == [1, 2] finally: second_published.set() - receiver.stop() - receiver.join(timeout=5) + receiver.close(timeout=5) first.disconnect() second.disconnect() diff --git a/tests/integration/test_vfs_webdav.py b/tests/integration/test_vfs_webdav.py index f2ca16e..39bfe0a 100644 --- a/tests/integration/test_vfs_webdav.py +++ b/tests/integration/test_vfs_webdav.py @@ -23,7 +23,7 @@ from openusdconnect.codec import message_to_dict # noqa: E402 from openusdconnect.managed_client import ManagedClient # noqa: E402 -from openusdconnect.receiver import ReceiverThread # noqa: E402 +from openusdconnect.receiver import EventReceiver # noqa: E402 from openusdconnect.sender import EventSender # noqa: E402 from openusdconnect.server import UsdSyncServer # noqa: E402 from openusdconnect.server.connection import ConnectionHandler, ThreadedTCPServer # noqa: E402 @@ -878,7 +878,7 @@ def test_receiver_from_snapshot_seq_gets_only_post_snapshot_events(self, tmp_pat meta = stage.GetRootLayer().customLayerData["openusdconnect"] assert meta["snapshot_seq"] == 1 - receiver = ReceiverThread( + receiver = EventReceiver( host="127.0.0.1", port=sync_port, sync_from=meta["snapshot_seq"] + 1, @@ -921,8 +921,7 @@ def _poll_events(): if sender is not None: sender.disconnect() if receiver is not None: - receiver.stop() - receiver.join(timeout=2) + receiver.close(timeout=2) tcp_server.shutdown() tcp_server.server_close() srv.shutdown() diff --git a/tests/native/CMakeLists.txt b/tests/native/CMakeLists.txt index e7b342e..9aa853d 100644 --- a/tests/native/CMakeLists.txt +++ b/tests/native/CMakeLists.txt @@ -12,25 +12,68 @@ endif() add_subdirectory(../../native/client_core client_core) -set(OPENUSDCONNECT_FLATBUFFERS_INCLUDE_DIR - "${CMAKE_CURRENT_SOURCE_DIR}/../../integrations/unreal/OpenUSDConnect/Source/OpenUSDConnectPXR/ThirdParty/flatbuffers/include" - CACHE PATH "Directory containing the pinned FlatBuffers C++ headers" +add_executable(OpenUSDConnectProtocolTests test_protocol_codec.cpp) +target_link_libraries(OpenUSDConnectProtocolTests PRIVATE OpenUSDConnect::ClientProtocol) +add_test(NAME OpenUSDConnectProtocolTests COMMAND OpenUSDConnectProtocolTests) + +add_executable(OpenUSDConnectClientEngineTests test_client_engine.cpp) +target_link_libraries(OpenUSDConnectClientEngineTests PRIVATE OpenUSDConnect::ClientProtocol) +add_test(NAME OpenUSDConnectClientEngineTests COMMAND OpenUSDConnectClientEngineTests) + +add_executable(OpenUSDConnectReceiverEndpointTests test_receiver_endpoint.cpp) +target_link_libraries(OpenUSDConnectReceiverEndpointTests PRIVATE OpenUSDConnect::ClientEngine) +if(MSVC) + target_compile_options(OpenUSDConnectReceiverEndpointTests PRIVATE /W4 /permissive-) +else() + target_compile_options(OpenUSDConnectReceiverEndpointTests PRIVATE -Wall -Wextra -Wpedantic) +endif() +add_test(NAME OpenUSDConnectReceiverEndpointTests COMMAND OpenUSDConnectReceiverEndpointTests) + +find_package(Threads REQUIRED) +add_executable(OpenUSDConnectProducerEndpointTests test_producer_endpoint.cpp) +target_link_libraries(OpenUSDConnectProducerEndpointTests + PRIVATE OpenUSDConnect::ClientEngine Threads::Threads ) -if(NOT EXISTS "${OPENUSDCONNECT_FLATBUFFERS_INCLUDE_DIR}/flatbuffers/flatbuffers.h") - message(FATAL_ERROR - "FlatBuffers headers are missing. Run " - "integrations/unreal/OpenUSDConnect/setup_flatbuffers.py before configuring native tests, " - "or set OPENUSDCONNECT_FLATBUFFERS_INCLUDE_DIR." - ) +if(MSVC) + target_compile_options(OpenUSDConnectProducerEndpointTests PRIVATE /W4 /permissive-) +else() + target_compile_options(OpenUSDConnectProducerEndpointTests PRIVATE -Wall -Wextra -Wpedantic) endif() +add_test(NAME OpenUSDConnectProducerEndpointTests COMMAND OpenUSDConnectProducerEndpointTests) -add_executable(OpenUSDConnectProtocolTests test_protocol_codec.cpp) -target_link_libraries(OpenUSDConnectProtocolTests PRIVATE OpenUSDConnect::ClientProtocol) -target_include_directories(OpenUSDConnectProtocolTests PRIVATE - "${OPENUSDCONNECT_FLATBUFFERS_INCLUDE_DIR}" +add_executable(OpenUSDConnectThreadedReceiverDriverTests test_threaded_receiver_driver.cpp) +target_link_libraries(OpenUSDConnectThreadedReceiverDriverTests + PRIVATE OpenUSDConnect::ClientDriver OpenUSDConnect::ClientDriverTesting +) +if(MSVC) + target_compile_options(OpenUSDConnectThreadedReceiverDriverTests PRIVATE /W4 /permissive-) +else() + target_compile_options(OpenUSDConnectThreadedReceiverDriverTests PRIVATE -Wall -Wextra -Wpedantic) +endif() +add_test(NAME OpenUSDConnectThreadedReceiverDriverTests + COMMAND OpenUSDConnectThreadedReceiverDriverTests +) + +add_executable(OpenUSDConnectThreadedProducerDriverTests test_threaded_producer_driver.cpp) +target_link_libraries(OpenUSDConnectThreadedProducerDriverTests + PRIVATE OpenUSDConnect::ClientDriver OpenUSDConnect::ClientDriverTesting +) +if(MSVC) + target_compile_options(OpenUSDConnectThreadedProducerDriverTests PRIVATE /W4 /permissive-) +else() + target_compile_options(OpenUSDConnectThreadedProducerDriverTests PRIVATE -Wall -Wextra -Wpedantic) +endif() +add_test(NAME OpenUSDConnectThreadedProducerDriverTests + COMMAND OpenUSDConnectThreadedProducerDriverTests ) -add_test(NAME OpenUSDConnectProtocolTests COMMAND OpenUSDConnectProtocolTests) -add_executable(OpenUSDConnectReplayIdentityTests test_replay_identity.cpp) -target_link_libraries(OpenUSDConnectReplayIdentityTests PRIVATE OpenUSDConnect::ClientCore) -add_test(NAME OpenUSDConnectReplayIdentityTests COMMAND OpenUSDConnectReplayIdentityTests) +add_executable(OpenUSDConnectTcpSocketTests test_tcp_socket.cpp) +target_link_libraries(OpenUSDConnectTcpSocketTests + PRIVATE OpenUSDConnect::ClientDriver $<$:ws2_32> +) +if(MSVC) + target_compile_options(OpenUSDConnectTcpSocketTests PRIVATE /W4 /permissive-) +else() + target_compile_options(OpenUSDConnectTcpSocketTests PRIVATE -Wall -Wextra -Wpedantic) +endif() +add_test(NAME OpenUSDConnectTcpSocketTests COMMAND OpenUSDConnectTcpSocketTests) diff --git a/tests/native/driver_recorder.h b/tests/native/driver_recorder.h new file mode 100644 index 0000000..b65f884 --- /dev/null +++ b/tests/native/driver_recorder.h @@ -0,0 +1,203 @@ +#pragma once + +#include "openusdconnect/client/driver/socket.h" +#include "openusdconnect/client/driver/testing/scripted_socket.h" +#include "openusdconnect/client/engine/notification.h" + +#include "frames.h" +#include "test_check.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +// What the driver tests share: patience, a recorder, and a harness. +namespace driver_test +{ + +using namespace openusdconnect::client; + +// Bounds every wait for the driver thread, which a passing test never reaches. +inline constexpr std::chrono::milliseconds kPatience{5'000}; +// A scripted refusal's system error. +inline constexpr int kRefused = 10061; + +// Polls for state the driver thread reports after the event a test waited for. +template +[[nodiscard]] bool Eventually(Ready ready) +{ + const auto deadline = std::chrono::steady_clock::now() + kPatience; + while (!ready()) + { + if (std::chrono::steady_clock::now() >= deadline) + { + return false; + } + std::this_thread::sleep_for(std::chrono::milliseconds(1)); + } + return true; +} + +// Records what the driver thread reports. +class Recorder final +{ +public: + void Add(Notification notification) + { + std::lock_guard lock(Mutex); + if (const TokenIssued* issued = std::get_if(¬ification)) + { + Token = issued->Token; + } + Notices.push_back(std::move(notification)); + } + + void Log(const std::string& message) + { + std::lock_guard lock(Mutex); + Logs.push_back(message); + } + + // Callbacks that record notifications, unless callbacks has its own, and logs. + [[nodiscard]] DriverCallbacks Recording(DriverCallbacks callbacks) + { + if (!callbacks.Notifications) + { + callbacks.Notifications = [this](Notification notification) + { + Add(std::move(notification)); + }; + } + callbacks.Log = [this](LogLevel, const std::string& message) + { + Log(message); + }; + return callbacks; + } + + [[nodiscard]] std::string IssuedToken() const + { + std::lock_guard lock(Mutex); + return Token; + } + + [[nodiscard]] bool Logged(std::string_view text) const + { + std::lock_guard lock(Mutex); + for (const std::string& log : Logs) + { + if (log.find(text) != std::string::npos) + { + return true; + } + } + return false; + } + + template + [[nodiscard]] std::size_t Count() const + { + std::lock_guard lock(Mutex); + std::size_t count = 0; + for (const Notification& notice : Notices) + { + count += std::holds_alternative(notice) ? 1 : 0; + } + return count; + } + + template + [[nodiscard]] std::vector All() const + { + std::lock_guard lock(Mutex); + std::vector matching; + for (const Notification& notice : Notices) + { + if (const T* value = std::get_if(¬ice)) + { + matching.push_back(*value); + } + } + return matching; + } + +private: + mutable std::mutex Mutex; + std::vector Notices; + std::vector Logs; + std::string Token; +}; + +// One endpoint driven by its reference driver over scripted sockets. +template +class DriverHarness +{ +public: + using EndpointType = typename DriverType::EndpointType; + + template + explicit DriverHarness(const Config& config, DriverCallbacks callbacks = {}) + : Endpoint(config, Notifications) + , Sockets(std::make_shared()) + , Driver(std::make_unique(Endpoint, Notifications, Sockets, + Record.Recording(std::move(callbacks)))) + { + } + + ~DriverHarness() + { + Driver->Stop(); + CHECK(Driver->Join(kPatience)); + } + + DriverHarness(const DriverHarness&) = delete; + DriverHarness& operator=(const DriverHarness&) = delete; + + // Accepts the pending connect and returns once the client sent its Hello. + [[nodiscard]] std::shared_ptr Accept() + { + std::shared_ptr connection = Sockets->Accept(kPatience); + CHECK(connection != nullptr); + CHECK(connection->WaitIdle(kPatience)); + return connection; + } + + void Deliver(ScriptedConnection& connection, const endpoint_test::Bytes& frame) + { + CHECK(connection.Deliver(frame)); + CHECK(connection.WaitIdle(kPatience)); + } + + // Starts the loop if needed, then accepts the pending attempt with a HelloOk. + [[nodiscard]] std::shared_ptr + Handshake(const endpoint_test::server::Hello& hello = {}) + { + if (!Driver->Running()) + { + CHECK(Driver->Start()); + } + std::shared_ptr connection = Accept(); + Deliver(*connection, endpoint_test::server::HelloOk(hello)); + CHECK(Eventually( + [this] + { + return Endpoint.Status().Connected; + })); + return connection; + } + + NotificationQueue Notifications; + EndpointType Endpoint; + Recorder Record; + const std::shared_ptr Sockets; + const std::unique_ptr Driver; +}; + +} // namespace driver_test diff --git a/tests/native/endpoint_host.h b/tests/native/endpoint_host.h new file mode 100644 index 0000000..4d4e2cd --- /dev/null +++ b/tests/native/endpoint_host.h @@ -0,0 +1,144 @@ +#pragma once + +#include "openusdconnect/client/engine/actions.h" +#include "openusdconnect/client/engine/notification.h" + +#include "frames.h" +#include "test_check.h" + +#include +#include +#include +#include +#include +#include + +namespace endpoint_test +{ + +template +[[nodiscard]] Config With(Config config, Field Config::* field, std::common_type_t value) +{ + config.*field = value; + return config; +} + +template +[[nodiscard]] const T& As(const Notification& notification) +{ + CHECK(std::holds_alternative(notification)); + return std::get(notification); +} + +// Plays the host around one endpoint, on a clock the test advances. +template +class Host +{ +public: + template + explicit Host(const Config& config) + : Endpoint(config, Notifications) + { + } + + // Actions since the last call, without log lines. + [[nodiscard]] std::vector Commands() + { + Collect(); + return std::exchange(Pending, {}); + } + + template + [[nodiscard]] T Single() + { + std::vector commands = Commands(); + CHECK(commands.size() == 1); + CHECK(std::holds_alternative(commands.front())); + return std::get(std::move(commands.front())); + } + + // The oldest action not yet taken, which must be a T. + template + [[nodiscard]] T Next() + { + Collect(); + CHECK(!Pending.empty() && std::holds_alternative(Pending.front())); + T next = std::get(std::move(Pending.front())); + Pending.erase(Pending.begin()); + return next; + } + + // Whether the endpoint logged at level since the last call. + [[nodiscard]] bool Logged(LogLevel level) + { + Collect(); + const bool logged = std::any_of(Logs.begin(), Logs.end(), + [level](const LogAction& log) + { + return log.Level == level; + }); + Logs.clear(); + return logged; + } + + [[nodiscard]] std::vector Notices() + { + return Notifications.Drain(); + } + + template + [[nodiscard]] T Notice() + { + std::vector notices = Notices(); + CHECK(notices.size() == 1); + CHECK(std::holds_alternative(notices.front())); + return std::get(std::move(notices.front())); + } + + [[nodiscard]] auto Status() const + { + return Endpoint.Status(); + } + + void Feed(const Bytes& bytes) + { + Endpoint.OnBytes(bytes.data(), bytes.size()); + } + + void Disconnect(DisconnectReason reason = DisconnectReason::PeerClosed) + { + Endpoint.OnDisconnected(reason, Now); + } + + void Advance(std::chrono::milliseconds elapsed) + { + Now += elapsed; + Endpoint.OnTick(Now); + } + + NotificationQueue Notifications; + EndpointType Endpoint; + // Away from the clock's epoch, which no deadline may depend on. + TimePoint Now = TimePoint{} + std::chrono::hours(1); + +private: + void Collect() + { + for (Action& action : Endpoint.TakeActions()) + { + if (LogAction* log = std::get_if(&action)) + { + Logs.push_back(std::move(*log)); + } + else + { + Pending.push_back(std::move(action)); + } + } + } + + std::vector Pending; + std::vector Logs; +}; + +} // namespace endpoint_test diff --git a/tests/native/frames.h b/tests/native/frames.h new file mode 100644 index 0000000..9479cfc --- /dev/null +++ b/tests/native/frames.h @@ -0,0 +1,302 @@ +#pragma once + +#include "openusdconnect/client/engine/notification.h" +#include "openusdconnect/client/protocol_codec.h" + +#include "test_check.h" + +#include +#include +#include +#include +#include +#include + +// Wire frames shared by the endpoint and driver tests. +namespace endpoint_test +{ + +using namespace openusdconnect::client; +using OpenUSDConnect::HelloRejectionCode; +using OpenUSDConnect::LayerMode; +using OpenUSDConnect::Payload; +using OpenUSDConnect::TransactionRejectionCode; + +using Bytes = std::vector; + +[[nodiscard]] inline std::string Text(const flatbuffers::String* value) +{ + return value ? value->str() : std::string(); +} + +[[nodiscard]] inline const OpenUSDConnect::Envelope& Decode(const Bytes& payload) +{ + EnvelopeView view; + CHECK(DecodeEnvelope(payload.data(), payload.size(), view) == ProtocolResult::Success); + return *view.Get(); +} + +// A frame a client sent, whose envelope borrows from it. +[[nodiscard]] inline const OpenUSDConnect::Envelope& DecodeSent(const Bytes& frame) +{ + std::size_t size = 0; + CHECK(TryReadFrameHeader(frame.data(), kDefaultMaxFrameSize, size)); + CHECK(size + kFrameHeaderSize == frame.size()); + EnvelopeView view; + CHECK(DecodeEnvelope(frame.data() + kFrameHeaderSize, size, view) == ProtocolResult::Success); + return *view.Get(); +} + +// The payload type of each frame, decoded with Decoder. +template +[[nodiscard]] std::vector Kinds(const std::vector& frames) +{ + std::vector kinds; + for (const Bytes& frame : frames) + { + kinds.push_back(Decoder(frame).payload_type()); + } + return kinds; +} + +// Server-to-client frames, length-prefixed as they arrive on the socket. +namespace server +{ + +struct Hello final +{ + std::string ServerInstance = "server"; + bool ReplayIdentity = true; + std::optional ReplayEpoch = 0; + bool LayeredReplay = true; + LayerMode Mode = LayerMode::Managed; + std::string Token; + std::optional Metadata; + std::uint64_t CommittedThrough = 0; +}; + +struct Checkpoint final +{ + std::uint64_t Epoch = 0; + std::int32_t HeadSequence = 0; +}; + +[[nodiscard]] inline Bytes Frame(flatbuffers::FlatBufferBuilder& builder, Payload type, + flatbuffers::Offset payload, + std::uint16_t schema_version = kSchemaVersion) +{ + OpenUSDConnect::FinishEnvelopeBuffer( + builder, OpenUSDConnect::CreateEnvelope(builder, type, payload, schema_version)); + Bytes frame; + CHECK(EncodeFrame(builder.GetBufferPointer(), builder.GetSize(), frame) == + FrameResult::Success); + return frame; +} + +[[nodiscard]] inline flatbuffers::Offset +OptionalString(flatbuffers::FlatBufferBuilder& builder, std::string_view text) +{ + return text.empty() ? flatbuffers::Offset() : CreateString(builder, text); +} + +[[nodiscard]] inline flatbuffers::Optional Wire(std::optional value) +{ + return value ? flatbuffers::Optional(*value) : flatbuffers::nullopt; +} + +[[nodiscard]] inline Bytes HelloOk(const Hello& hello = {}) +{ + flatbuffers::FlatBufferBuilder builder(256); + flatbuffers::Offset metadata; + if (hello.Metadata) + { + const StageMetadata& fields = *hello.Metadata; + const auto up_axis = OptionalString(builder, fields.UpAxis.value_or("")); + metadata = OpenUSDConnect::CreateSetStageMetadata( + builder, Wire(fields.TimeCodesPerSecond), Wire(fields.FramesPerSecond), + Wire(fields.StartTimeCode), Wire(fields.EndTimeCode), Wire(fields.MetersPerUnit), + up_axis); + } + const auto token = OptionalString(builder, hello.Token); + const auto instance = OptionalString(builder, hello.ServerInstance); + const auto epoch = hello.ReplayEpoch ? flatbuffers::Optional(*hello.ReplayEpoch) + : flatbuffers::nullopt; + const auto accepted = OpenUSDConnect::CreateHelloOk( + builder, token, metadata, hello.LayeredReplay, hello.Mode, hello.CommittedThrough, instance, + hello.ReplayIdentity, epoch); + return Frame(builder, Payload::HelloOk, accepted.Union()); +} + +[[nodiscard]] inline Bytes AuthRejected(std::string_view reason) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto rejected = + OpenUSDConnect::CreateAuthRejected(builder, CreateString(builder, reason)); + return Frame(builder, Payload::AuthRejected, rejected.Union()); +} + +[[nodiscard]] inline Bytes HelloRejected(HelloRejectionCode code, std::string_view reason) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto rejected = + OpenUSDConnect::CreateHelloRejected(builder, code, OptionalString(builder, reason)); + return Frame(builder, Payload::HelloRejected, rejected.Union()); +} + +[[nodiscard]] inline Bytes Ping() +{ + flatbuffers::FlatBufferBuilder builder(32); + return Frame(builder, Payload::Ping, OpenUSDConnect::CreatePing(builder).Union()); +} + +[[nodiscard]] inline Bytes Resync() +{ + flatbuffers::FlatBufferBuilder builder(32); + return Frame(builder, Payload::Resync, OpenUSDConnect::CreateResync(builder).Union()); +} + +[[nodiscard]] inline Bytes ReplayComplete(std::int32_t head, std::uint64_t epoch) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto complete = OpenUSDConnect::CreateReplayComplete(builder, head, epoch); + return Frame(builder, Payload::ReplayComplete, complete.Union()); +} + +[[nodiscard]] inline Bytes Event(std::int32_t sequence) +{ + flatbuffers::FlatBufferBuilder builder(128); + const auto prim = OpenUSDConnect::CreateEnsurePrim( + builder, CreateString(builder, "/World/P" + std::to_string(sequence))); + const auto event = OpenUSDConnect::CreateEventWrapper( + builder, OpenUSDConnect::EventPayload::EnsurePrim, prim.Union()); + const auto broadcast = OpenUSDConnect::CreateBroadcastEvent(builder, sequence, event); + return Frame(builder, Payload::BroadcastEvent, broadcast.Union()); +} + +[[nodiscard]] inline Bytes LayerGraph(std::int32_t sequence) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto state = OpenUSDConnect::CreateLayerGraphState(builder, sequence); + return Frame(builder, Payload::LayerGraphState, state.Union()); +} + +[[nodiscard]] inline Bytes LayerStack() +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto layer = OpenUSDConnect::CreateLogicalLayerState(builder, CreateString(builder, "a")); + const auto stack = OpenUSDConnect::CreateLayerStackState( + builder, CreateString(builder, "generation"), 1, builder.CreateVector(&layer, 1)); + return Frame(builder, Payload::LayerStackState, stack.Union()); +} + +[[nodiscard]] inline Bytes Playback(double time, bool playing, double rate, std::string_view leader) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto state = OpenUSDConnect::CreatePlaybackState(builder, time, playing, rate, + CreateString(builder, leader)); + return Frame(builder, Payload::PlaybackState, state.Union()); +} + +[[nodiscard]] inline Bytes Claimed(std::string_view leader) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto claimed = + OpenUSDConnect::CreatePlaybackClaimed(builder, CreateString(builder, leader)); + return Frame(builder, Payload::PlaybackClaimed, claimed.Union()); +} + +[[nodiscard]] inline Bytes ClaimRejected(std::string_view reason, std::string_view leader) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto rejected = OpenUSDConnect::CreatePlaybackRejected( + builder, CreateString(builder, reason), CreateString(builder, leader)); + return Frame(builder, Payload::PlaybackRejected, rejected.Union()); +} + +[[nodiscard]] inline Bytes Acknowledged(std::uint64_t transaction_id, + std::optional checkpoint = std::nullopt) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto wire_checkpoint = + checkpoint ? OpenUSDConnect::CreateTransactionCheckpoint(builder, checkpoint->Epoch, + checkpoint->HeadSequence) + : flatbuffers::Offset(); + const auto result = OpenUSDConnect::CreateTransactionResult( + builder, transaction_id, OpenUSDConnect::TransactionStatus::Acknowledged, 0, + TransactionRejectionCode::None, 0, wire_checkpoint); + return Frame(builder, Payload::TransactionResult, result.Union()); +} + +[[nodiscard]] inline Bytes Rejected(std::uint64_t transaction_id, TransactionRejectionCode code, + std::string_view reason, std::uint64_t expected = 0) +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto result = OpenUSDConnect::CreateTransactionResult( + builder, transaction_id, OpenUSDConnect::TransactionStatus::Rejected, expected, code, + OptionalString(builder, reason)); + return Frame(builder, Payload::TransactionResult, result.Union()); +} + +[[nodiscard]] inline Bytes RateLimited(float retry_after) +{ + flatbuffers::FlatBufferBuilder builder(32); + const auto limited = OpenUSDConnect::CreateRateLimited(builder, retry_after); + return Frame(builder, Payload::RateLimited, limited.Union()); +} + +} // namespace server + +struct SentHello final +{ + std::string Role; + std::int32_t ProtocolVersion = 0; + std::int32_t SyncFrom = 0; + std::string ClientId; + std::string Origin; + std::string Department; + std::string Token; + bool LayeredReplay = false; + LayerMode Mode = LayerMode::Managed; + std::string ProducerSessionId; + // Absent without a claim; empty when the claimed prefix is unknown. + std::optional ReplayServerInstance; + std::optional ReplayEpoch; + + [[nodiscard]] bool Claims(std::string_view instance, std::uint64_t epoch) const + { + return ReplayServerInstance == instance && ReplayEpoch == epoch; + } + + [[nodiscard]] bool ClaimsUnknownPrefix() const + { + return ReplayServerInstance == "" && !ReplayEpoch; + } +}; + +[[nodiscard]] inline SentHello DecodeHello(const Bytes& frame) +{ + const OpenUSDConnect::Hello* hello = DecodeSent(frame).payload_as_Hello(); + CHECK(hello != nullptr); + SentHello sent; + sent.Role = Text(hello->role()); + sent.ProtocolVersion = hello->protocol_version(); + sent.SyncFrom = hello->sync_from(); + sent.ClientId = Text(hello->client_id()); + sent.Origin = Text(hello->origin()); + sent.Department = Text(hello->department()); + sent.Token = Text(hello->token()); + sent.LayeredReplay = hello->layered_replay(); + sent.Mode = hello->layer_mode(); + sent.ProducerSessionId = Text(hello->producer_session_id()); + if (hello->replay_server_instance()) + { + sent.ReplayServerInstance = hello->replay_server_instance()->str(); + } + if (hello->replay_epoch().has_value()) + { + sent.ReplayEpoch = *hello->replay_epoch(); + } + return sent; +} + +} // namespace endpoint_test diff --git a/tests/native/producer_frames.h b/tests/native/producer_frames.h new file mode 100644 index 0000000..80498e5 --- /dev/null +++ b/tests/native/producer_frames.h @@ -0,0 +1,97 @@ +#pragma once + +#include "openusdconnect/client/engine/producer_endpoint.h" + +#include "frames.h" +#include "test_check.h" + +#include +#include +#include +#include + +// Producer test fixtures shared by the endpoint and driver tests. +namespace producer_test +{ + +using namespace endpoint_test; + +[[nodiscard]] inline ProducerConfig TestConfig() +{ + ProducerConfig config; + config.Host = "127.0.0.1"; + config.Port = 7200; + config.ClientId = "client"; + config.Origin = "origin"; + config.SessionId = "session"; + return config; +} + +// A one-event transaction frame, encoded as a host encodes it. +[[nodiscard]] inline Bytes TransactionFrame(std::uint64_t transaction_id, std::string_view prim, + std::string_view layer_key = {}) +{ + flatbuffers::FlatBufferBuilder builder(128); + flatbuffers::Offset event; + CHECK(BuildVisibilityEvent(builder, VisibilityEventView{prim, true}, event) == + ProtocolResult::Success); + CHECK(FinishTransactionFrame(builder, transaction_id, &event, 1, layer_key) == + ProtocolResult::Success); + const std::uint8_t* bytes = builder.GetBufferPointer(); + return {bytes, bytes + builder.GetSize()}; +} + +[[nodiscard]] inline Bytes ClaimFrame() +{ + flatbuffers::FlatBufferBuilder builder(64); + const auto claim = + OpenUSDConnect::CreateClaimPlayback(builder, CreateString(builder, "client")); + CHECK(FinishEnvelopeFrame(builder, OpenUSDConnect::CreateEnvelope( + builder, Payload::ClaimPlayback, claim.Union(), + kSchemaVersion)) == ProtocolResult::Success); + const std::uint8_t* bytes = builder.GetBufferPointer(); + return {bytes, bytes + builder.GetSize()}; +} + +[[nodiscard]] inline std::vector TransactionIds(const std::vector& frames) +{ + std::vector ids; + for (const Bytes& frame : frames) + { + const OpenUSDConnect::Txn* transaction = DecodeSent(frame).payload_as_Txn(); + CHECK(transaction != nullptr); + ids.push_back(transaction->txn_id()); + } + return ids; +} + +[[nodiscard]] inline Bytes Concatenate(const std::vector& frames) +{ + Bytes bytes; + for (const Bytes& frame : frames) + { + bytes.insert(bytes.end(), frame.begin(), frame.end()); + } + return bytes; +} + +// Splits what a client sent into its length-prefixed frames. +[[nodiscard]] inline std::vector SplitFrames(const Bytes& stream) +{ + std::vector frames; + std::size_t offset = 0; + while (offset < stream.size()) + { + std::size_t size = 0; + CHECK(stream.size() - offset >= kFrameHeaderSize); + CHECK(TryReadFrameHeader(stream.data() + offset, kDefaultMaxFrameSize, size)); + const std::size_t end = offset + kFrameHeaderSize + size; + CHECK(end <= stream.size()); + frames.emplace_back(stream.begin() + static_cast(offset), + stream.begin() + static_cast(end)); + offset = end; + } + return frames; +} + +} // namespace producer_test diff --git a/tests/native/receiver_frames.h b/tests/native/receiver_frames.h new file mode 100644 index 0000000..0e71685 --- /dev/null +++ b/tests/native/receiver_frames.h @@ -0,0 +1,40 @@ +#pragma once + +#include "openusdconnect/client/engine/receiver_endpoint.h" + +#include "frames.h" +#include "test_check.h" + +#include +#include + +// Receiver test fixtures shared by the endpoint and driver tests. +namespace receiver_test +{ + +using namespace endpoint_test; + +[[nodiscard]] inline ReceiverConfig TestConfig() +{ + ReceiverConfig config; + config.Host = "127.0.0.1"; + config.Port = 7200; + config.ClientId = "client"; + config.Origin = "origin"; + return config; +} + +[[nodiscard]] inline std::vector Sequences(const std::vector& frames) +{ + std::vector sequences; + for (const Bytes& frame : frames) + { + if (const auto* event = Decode(frame).payload_as_BroadcastEvent()) + { + sequences.push_back(event->seq()); + } + } + return sequences; +} + +} // namespace receiver_test diff --git a/tests/native/test_client_engine.cpp b/tests/native/test_client_engine.cpp new file mode 100644 index 0000000..711c782 --- /dev/null +++ b/tests/native/test_client_engine.cpp @@ -0,0 +1,326 @@ +#include "openusdconnect/client/engine/status.h" +#include "openusdconnect/client/producer_recovery.h" +#include "openusdconnect/client/receiver_session.h" +#include "openusdconnect/client/schema/messages_generated.h" + +#include "test_check.h" + +#include +#include +#include +#include + +using namespace openusdconnect::client; + +[[nodiscard]] constexpr bool MatchesWire(RejectionCode code, + OpenUSDConnect::TransactionRejectionCode wire) noexcept +{ + return static_cast(code) == static_cast(wire); +} + +static_assert(MatchesWire(RejectionCode::None, OpenUSDConnect::TransactionRejectionCode::None)); +static_assert(MatchesWire(RejectionCode::InvalidIdentity, + OpenUSDConnect::TransactionRejectionCode::InvalidIdentity)); +static_assert(MatchesWire(RejectionCode::UnexpectedId, + OpenUSDConnect::TransactionRejectionCode::UnexpectedId)); +static_assert(MatchesWire(RejectionCode::StaleLayerGraph, + OpenUSDConnect::TransactionRejectionCode::StaleLayerGraph)); +static_assert(MatchesWire(RejectionCode::InvalidTransaction, + OpenUSDConnect::TransactionRejectionCode::InvalidTransaction)); +// A new wire code needs a rejection policy entry. +static_assert(MatchesWire(RejectionCode::InvalidTransaction, + OpenUSDConnect::TransactionRejectionCode::MAX)); + +struct PhaseCase final +{ + bool PhaseInputs::* Input; + ClientPhase Phase; +}; + +constexpr PhaseCase kPhasePrecedence[] = { + {&PhaseInputs::Closed, ClientPhase::Closed}, + {&PhaseInputs::RecoveryRequired, ClientPhase::RecoveryRequired}, + {&PhaseInputs::Rejected, ClientPhase::Rejected}, + {&PhaseInputs::Parked, ClientPhase::Parked}, + {&PhaseInputs::Replaying, ClientPhase::Replaying}, + {&PhaseInputs::Ready, ClientPhase::Ready}, + {&PhaseInputs::Connecting, ClientPhase::Connecting}, +}; + +static void TestEachPhaseOutranksThePhasesAfterIt() +{ + const std::size_t count = std::size(kPhasePrecedence); + for (std::size_t index = 0; index < count; ++index) + { + PhaseInputs inputs; + for (std::size_t lower = index; lower < count; ++lower) + { + inputs.*kPhasePrecedence[lower].Input = true; + } + CHECK(ComputePhase(inputs) == kPhasePrecedence[index].Phase); + } + CHECK(ComputePhase(PhaseInputs{}) == ClientPhase::Offline); +} + +static void TestTransactionFailureDescription() +{ + TransactionFailure failure{3, 3, "layer was remapped", 2}; + CHECK(failure.Disposition() == ProducerRecoveryDisposition::RecoverableConflict); + CHECK(failure.Describe() == + "transaction 3 rejected (stale_layer_graph, expected transaction 2): layer was remapped"); + failure = {7, 9, "", 0}; + CHECK(failure.Disposition() == ProducerRecoveryDisposition::SessionFatal); + CHECK(failure.Describe() == "transaction 7 rejected (unknown_9): no reason supplied"); +} + +using TestInbox = OrderedReceiverSession; + +static void AcceptFrames(TestInbox& inbox, std::uint64_t generation, int count) +{ + for (int frame = 0; frame < count; ++frame) + { + CHECK(inbox.Accept(generation, ReceiverMessageKind::Other, 0, frame) == + AcceptResult::Accepted); + } +} + +static void TestHoldCoversOnlyFramesQueuedBeforeTheMarker() +{ + TestInbox inbox(1, 8); + const std::uint64_t generation = inbox.BeginConnection().Generation; + CHECK(inbox.DrainedThrough(inbox.FreezeMarker())); + AcceptFrames(inbox, generation, 3); + const std::uint64_t marker = inbox.FreezeMarker(); + AcceptFrames(inbox, generation, 2); + CHECK(inbox.Drain(2).size() == 2); + CHECK(!inbox.DrainedThrough(marker)); + AcceptFrames(inbox, generation, 2); + int frame = -1; + CHECK(inbox.TryPop(frame)); + CHECK(inbox.DrainedThrough(marker)); + CHECK(inbox.Size() == 4); +} + +static void TestRejectedFramesDoNotExtendTheHold() +{ + TestInbox inbox(1, 2); + const std::uint64_t generation = inbox.BeginConnection().Generation; + AcceptFrames(inbox, generation, 2); + CHECK(inbox.Accept(generation, ReceiverMessageKind::Other, 0, 2) == AcceptResult::QueueFull); + const std::uint64_t marker = inbox.FreezeMarker(); + CHECK(inbox.Drain().size() == 2); + CHECK(inbox.DrainedThrough(marker)); +} + +static void TestReplayRequestReleasesTheHold() +{ + TestInbox inbox(1, 8); + const std::uint64_t generation = inbox.BeginConnection().Generation; + AcceptFrames(inbox, generation, 3); + const std::uint64_t marker = inbox.FreezeMarker(); + CHECK(!inbox.DrainedThrough(marker)); + CHECK(inbox.RequestReplayFrom(1)); + CHECK(inbox.DrainedThrough(marker)); +} + +static void TestReplayMarkerRequiresItsRecords() +{ + TestInbox contiguous(1, 8, true); + const std::uint64_t generation = contiguous.BeginConnection().Generation; + CHECK(contiguous.Accept(generation, ReceiverMessageKind::Event, 1, 1) == + AcceptResult::Accepted); + CHECK(contiguous.AcceptReplayComplete(generation, 2, 0) == AcceptResult::SequenceGap); + CHECK(contiguous.AcceptReplayComplete(generation, 1, 0) == AcceptResult::Accepted); + + TestInbox unordered(1, 8); + CHECK(unordered.AcceptReplayComplete(unordered.BeginConnection().Generation, 2, 0) == + AcceptResult::Accepted); +} + +static void TestResetPendingUntilAppliedOrDiscarded() +{ + TestInbox inbox(1, 8, true); + std::uint64_t generation = inbox.BeginConnection().Generation; + const auto accept_reset = [&] + { + CHECK(inbox.Accept(generation, ReceiverMessageKind::Resync, 0, 0) == + AcceptResult::Accepted); + }; + CHECK(!inbox.ResetPending()); + accept_reset(); + CHECK(inbox.ResetPending()); + int frame = -1; + CHECK(inbox.TryPop(frame)); + CHECK(inbox.ResetPending()); + inbox.ResetAppliedProgress(); + CHECK(!inbox.ResetPending()); + + // One report covers every reset drained before it. + accept_reset(); + CHECK(inbox.Accept(generation, ReceiverMessageKind::Event, 1, 1) == AcceptResult::Accepted); + accept_reset(); + CHECK(inbox.Drain().size() == 3); + inbox.ResetAppliedProgress(); + CHECK(!inbox.ResetPending()); + + // A replay request settles queued and drained resets alike. + accept_reset(); + CHECK(inbox.TryPop(frame)); + accept_reset(); + CHECK(inbox.RequestReplayFrom(1)); + CHECK(!inbox.ResetPending()); + + // A late report for the drained reset cannot settle a newer one. + generation = inbox.BeginConnection().Generation; + accept_reset(); + inbox.ResetAppliedProgress(); + CHECK(inbox.ResetPending()); +} + +// Queues an event whose payload is its sequence. +[[nodiscard]] static AcceptResult AcceptEvent(TestInbox& inbox, std::uint64_t generation, + std::int32_t sequence) +{ + return inbox.Accept(generation, ReceiverMessageKind::Event, sequence, sequence); +} + +[[nodiscard]] static AcceptResult AcceptReset(TestInbox& inbox, std::uint64_t generation) +{ + return inbox.Accept(generation, ReceiverMessageKind::Resync, 0, 0); +} + +static void TestInboxRejectsInvalidArguments() +{ + CHECK(!TestInbox::IsValidConfiguration(0, 1)); + CHECK(!TestInbox::IsValidConfiguration(1, 0)); + TestInbox inbox(1, 1); + const std::uint64_t generation = inbox.BeginConnection().Generation; + CHECK(AcceptEvent(inbox, generation, 0) == AcceptResult::InvalidSequence); + CHECK(inbox.Accept(generation, ReceiverMessageKind::LayerGraphState, 0, 0) == + AcceptResult::InvalidSequence); + CHECK(inbox.AcceptReplayComplete(generation, -1, 0) == AcceptResult::InvalidSequence); + CHECK(!inbox.RequestReplayFrom(0)); + CHECK(inbox.Size() == 0); +} + +static void TestReplayAppliesBeforeLiveFramesDrain() +{ + TestInbox inbox(1, 8); + const ConnectionStart connection = inbox.BeginConnection(); + CHECK(connection.SyncFrom == 1); + CHECK(AcceptEvent(inbox, connection.Generation, 1) == AcceptResult::Accepted); + CHECK(inbox.AcceptReplayComplete(connection.Generation, 1, 7) == AcceptResult::Accepted); + CHECK(AcceptEvent(inbox, connection.Generation, 2) == AcceptResult::Accepted); + + CHECK(!inbox.MarkReplayApplied()); + CHECK(inbox.Drain(1) == std::vector{1}); + CHECK(inbox.MarkReplayApplied()); + CHECK(inbox.ReplayHeadSequence() == 1); + CHECK(inbox.ReplayEpoch() == 7); + CHECK(inbox.Drain() == std::vector{2}); +} + +static void TestStaleGenerationIsRejectedWithoutMutation() +{ + TestInbox inbox(4, 8); + const ConnectionStart first = inbox.BeginConnection(); + const ConnectionStart second = inbox.BeginConnection(); + CHECK(AcceptEvent(inbox, first.Generation, 4) == AcceptResult::StaleGeneration); + CHECK(inbox.Size() == 0); + CHECK(inbox.LastSequence() == 3); + CHECK(second.SyncFrom == 4); +} + +static void TestOverflowIsBoundedAndReplayable() +{ + TestInbox inbox(1, 1); + const std::uint64_t generation = inbox.BeginConnection().Generation; + CHECK(AcceptEvent(inbox, generation, 1) == AcceptResult::Accepted); + CHECK(AcceptEvent(inbox, generation, 2) == AcceptResult::QueueFull); + CHECK(inbox.Overflowed()); + CHECK(inbox.Drain() == std::vector{1}); + + CHECK(inbox.RequestReplayFrom(2)); + CHECK(inbox.BeginConnection().SyncFrom == 2); + CHECK(!inbox.Overflowed()); +} + +static void TestResetReconnectsFromOneWithoutDiscardingTheQueue(bool queued_prefix) +{ + TestInbox inbox(4, queued_prefix ? 2 : 1, true); + ConnectionStart connection = inbox.BeginConnection(); + CHECK(connection.SyncFrom == 4); + std::vector expected; + if (queued_prefix) + { + CHECK(AcceptEvent(inbox, connection.Generation, 4) == AcceptResult::Accepted); + expected.push_back(4); + } + CHECK(AcceptReset(inbox, connection.Generation) == AcceptResult::Accepted); + expected.push_back(0); + CHECK(inbox.Size() == expected.size()); + CHECK(inbox.LastSequence() == 0); + CHECK(AcceptEvent(inbox, connection.Generation, 1) == AcceptResult::QueueFull); + // Disconnects before the first new event keep the reset cursor and every + // queued frame, even before the consumer drains. + for (int attempt = 0; attempt < 2; ++attempt) + { + inbox.Disconnect(connection.Generation); + connection = inbox.BeginConnection(); + CHECK(connection.SyncFrom == 1); + CHECK(inbox.Size() == expected.size()); + } + CHECK(inbox.Drain() == expected); + inbox.ClearOverflow(); + CHECK(AcceptEvent(inbox, connection.Generation, 1) == AcceptResult::Accepted); + inbox.Disconnect(connection.Generation); + CHECK(inbox.BeginConnection().SyncFrom == 2); + CHECK(inbox.Drain() == std::vector{1}); +} + +static void TestFullReplayCursorSurvivesDisconnectBeforeAnyFrames() +{ + TestInbox inbox(4, 1); + CHECK(inbox.BeginConnection().SyncFrom == 4); + CHECK(inbox.RequestReplayFrom(1)); + for (int attempt = 0; attempt < 2; ++attempt) + { + const ConnectionStart connection = inbox.BeginConnection(); + CHECK(connection.SyncFrom == 1); + inbox.Disconnect(connection.Generation); + } +} + +static void TestRejectedResetPreservesTheSnapshotCursorAndQueue() +{ + TestInbox inbox(4, 1); + const std::uint64_t generation = inbox.BeginConnection().Generation; + CHECK(AcceptEvent(inbox, generation, 4) == AcceptResult::Accepted); + CHECK(AcceptReset(inbox, generation) == AcceptResult::QueueFull); + CHECK(inbox.LastSequence() == 4); + inbox.Disconnect(generation); + CHECK(inbox.BeginConnection().SyncFrom == 5); + CHECK(inbox.Drain() == std::vector{4}); +} + +int main() +{ + TestEachPhaseOutranksThePhasesAfterIt(); + TestTransactionFailureDescription(); + TestHoldCoversOnlyFramesQueuedBeforeTheMarker(); + TestRejectedFramesDoNotExtendTheHold(); + TestReplayRequestReleasesTheHold(); + TestReplayMarkerRequiresItsRecords(); + TestResetPendingUntilAppliedOrDiscarded(); + TestInboxRejectsInvalidArguments(); + TestReplayAppliesBeforeLiveFramesDrain(); + TestStaleGenerationIsRejectedWithoutMutation(); + TestOverflowIsBoundedAndReplayable(); + for (const bool queued_prefix : {false, true}) + { + TestResetReconnectsFromOneWithoutDiscardingTheQueue(queued_prefix); + } + TestFullReplayCursorSurvivesDisconnectBeforeAnyFrames(); + TestRejectedResetPreservesTheSnapshotCursorAndQueue(); + return 0; +} diff --git a/tests/native/test_producer_endpoint.cpp b/tests/native/test_producer_endpoint.cpp new file mode 100644 index 0000000..80e83a7 --- /dev/null +++ b/tests/native/test_producer_endpoint.cpp @@ -0,0 +1,1029 @@ +#include "openusdconnect/client/engine/producer_endpoint.h" + +#include "endpoint_host.h" +#include "producer_frames.h" +#include "test_check.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +using namespace openusdconnect::client; +using namespace std::chrono_literals; + +namespace +{ + +using namespace producer_test; + +// Plays the host around one endpoint. +class Producer final : public Host +{ +public: + explicit Producer(const ProducerConfig& config = TestConfig()) + : Host(config) + { + } + + // The frames to send since the last call; any other action fails. + [[nodiscard]] std::vector SendActions() + { + std::vector sends; + for (Action& action : Commands()) + { + CHECK(std::holds_alternative(action)); + sends.push_back(std::get(std::move(action))); + } + return sends; + } + + [[nodiscard]] std::vector Sends() + { + std::vector frames; + for (const SendAction& send : SendActions()) + { + frames.push_back(*send.Bytes); + } + return frames; + } + + // Takes a new attempt's actions and returns its deadline, when the endpoint wakes. + TimePoint Attempt() + { + const ConnectAction connect = Single(); + CHECK(connect.Host == "127.0.0.1" && connect.Port == 7200); + CHECK(Endpoint.NextWake() == connect.Deadline); + return connect.Deadline; + } + + // Opens the pending attempt's socket and returns the Hello it sent. + SentHello Open(std::string_view token = {}) + { + Endpoint.OnConnected(token); + return DecodeHello(*Single().Bytes); + } + + SentHello Request(std::string_view token = {}) + { + CHECK(Endpoint.RequestConnect(Now, Now + 2s)); + static_cast(Attempt()); + return Open(token); + } + + // Connects and returns what publication replayed. + std::vector Handshake(const server::Hello& hello = {}) + { + static_cast(Request()); + Feed(server::HelloOk(hello)); + CHECK(Status().Connected); + static_cast(Notices()); + return Sends(); + } + + ProducerResult Submit(std::string_view prim, std::string_view layer_key = {}, + std::size_t events = 1) + { + const std::uint64_t id = Endpoint.NextTransactionId(); + return Endpoint.Append(id, TransactionFrame(id, prim, layer_key), events, + std::string(layer_key)); + } + + // Expects the endpoint to close the connection, then reports the close. + void ExpectClose(DisconnectReason reason) + { + CHECK(Single().Reason == reason); + Disconnect(reason); + } +}; + +void TestConfigurationValidation() +{ + const ProducerConfig valid = TestConfig(); + const std::string two_byte = "\xC3\xA9"; + std::string longest_multibyte; + for (std::size_t index = 0; index < kMaxProducerSessionIdLength; ++index) + { + longest_multibyte += two_byte; + } + const std::pair rules[] = { + {valid, true}, + {With(valid, &ProducerConfig::LayerMode, LayerMode::SharedStage), true}, + {With(valid, &ProducerConfig::Origin, ""), true}, + {With(valid, &ProducerConfig::Port, 0), false}, + {With(valid, &ProducerConfig::SessionId, longest_multibyte), true}, + {With(valid, &ProducerConfig::SessionId, longest_multibyte + two_byte), false}, + {With(valid, &ProducerConfig::ReconnectMaxDelay, 999ms), false}, + {With(With(valid, &ProducerConfig::LayerMode, LayerMode::SharedStage), + &ProducerConfig::Department, "layout"), + false}, + }; + for (const auto& [config, expected] : rules) + { + CHECK(ProducerEndpoint::IsValidConfiguration(config) == expected); + } +} + +void TestAttemptsAreBoundedAndExclusive() +{ + Producer producer; + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + CHECK(producer.Attempt() == producer.Now + 2s); + CHECK(producer.Status().Handshaking && !producer.Status().Connected); + CHECK(!producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + CHECK(producer.Endpoint.Connect(producer.Now, producer.Now + 2s) == ConnectResult::Busy); + static_cast(producer.Open()); + CHECK(producer.Status().Handshaking); + CHECK(producer.Endpoint.Connect(producer.Now, producer.Now + 2s) == ConnectResult::Busy); + producer.Feed(server::HelloOk()); + CHECK(producer.Endpoint.Connect(producer.Now, producer.Now) == ConnectResult::Connected); + CHECK(!producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + CHECK(!producer.Endpoint.NextWake()); + + // The configured handshake timeout caps a longer deadline. + Producer patient; + CHECK(patient.Endpoint.Connect(patient.Now, patient.Now + 1h) == ConnectResult::Started); + CHECK(patient.Attempt() == patient.Now + 10s); + + // No attempt starts without time left. + Producer hurried; + CHECK(!hurried.Endpoint.RequestConnect(hurried.Now, hurried.Now)); + CHECK(hurried.Endpoint.Connect(hurried.Now, hurried.Now - 1ms) == ConnectResult::Refused); + CHECK(hurried.Commands().empty()); +} + +void TestHelloCarriesTheProducerIdentity() +{ + ProducerConfig config = TestConfig(); + config.Department = "layout"; + Producer producer(config); + const SentHello hello = producer.Request("token-1"); + CHECK(hello.Role == "emitter"); + CHECK(hello.ProtocolVersion == kProtocolVersion); + CHECK(hello.SyncFrom == 0); + CHECK(hello.ClientId == "client"); + CHECK(hello.Origin == "origin"); + CHECK(hello.Department == "layout"); + CHECK(hello.Token == "token-1"); + CHECK(!hello.LayeredReplay); + CHECK(hello.Mode == LayerMode::Managed); + CHECK(hello.ProducerSessionId == "session"); + CHECK(!hello.ReplayServerInstance && !hello.ReplayEpoch); +} + +void TestAcceptedHelloNotifiesThenPublishes() +{ + Producer producer; + static_cast(producer.Request()); + StageMetadata metadata; + metadata.MetersPerUnit = 0.01; + metadata.UpAxis = "Y"; + server::Hello hello; + hello.Token = "issued"; + hello.Metadata = metadata; + producer.Feed(server::HelloOk(hello)); + + const std::vector notices = producer.Notices(); + CHECK(notices.size() == 3); + CHECK(As(notices[0]).Token == "issued"); + CHECK(As(notices[1]).MetersPerUnit == 0.01); + static_cast(As(notices[2])); + + // Metadata the server did not author is neither notified nor kept. + Producer bare; + hello.Metadata = StageMetadata{}; + hello.Token.clear(); + static_cast(bare.Request()); + bare.Feed(server::HelloOk(hello)); + static_cast(bare.Notice()); + CHECK(!bare.Status().Metadata.UpAxis); +} + +void TestHandshakeRejectionsHoldUntilAnExplicitConnect() +{ + { + Producer producer; + static_cast(producer.Request()); + producer.Feed(server::AuthRejected("invalid token")); + CHECK(producer.Single().Reason == DisconnectReason::HandshakeRejected); + const HandshakeRejected rejected = producer.Notice(); + CHECK(rejected.Authentication && rejected.Reason == "invalid token"); + producer.Disconnect(); + const ProducerStatus status = producer.Status(); + CHECK(!status.Connected && !status.Handshaking && !status.Stopped && !status.Failure); + CHECK(status.Rejection && status.Rejection->Authentication); + + // The backoff has passed, yet only an explicit connect retries. + producer.Now += 1h; + CHECK(!producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + CHECK(producer.Endpoint.Connect(producer.Now, producer.Now + 2s) == ConnectResult::Started); + CHECK(!producer.Status().Rejection); + static_cast(producer.Attempt()); + static_cast(producer.Open()); + producer.Feed(server::HelloOk()); + CHECK(producer.Status().Connected); + } + server::Hello shared; + shared.Mode = LayerMode::SharedStage; + shared.Token = "not-issued"; + // An empty reason stays empty; hosts word their own default. + const std::pair rejections[] = { + {server::HelloRejected(HelloRejectionCode::Unspecified, ""), + {false, HelloRejectionCode::Unspecified, ""}}, + {server::HelloRejected(HelloRejectionCode::LayerModeMismatch, "server uses shared_stage"), + {false, HelloRejectionCode::LayerModeMismatch, "server uses shared_stage"}}, + {server::HelloOk(shared), + {false, HelloRejectionCode::LayerModeMismatch, + "server negotiated shared_stage instead of managed"}}, + }; + for (const auto& [frame, expected] : rejections) + { + Producer producer; + static_cast(producer.Request()); + producer.Feed(frame); + producer.ExpectClose(DisconnectReason::HandshakeRejected); + const std::optional rejection = producer.Status().Rejection; + CHECK(rejection && !rejection->Authentication); + CHECK(rejection->Code == expected.Code && rejection->Reason == expected.Reason); + CHECK(!producer.Endpoint.RequestConnect(producer.Now + 1h, producer.Now + 2h)); + } +} + +void CheckSessionFailure(Producer& producer, std::uint64_t transaction_id, std::string_view reason) +{ + const std::optional failure = producer.Endpoint.Failure(); + CHECK(failure); + CHECK(failure->TransactionId == transaction_id); + CHECK(failure->Code == static_cast(RejectionCode::UnexpectedId)); + CHECK(failure->ExpectedTransactionId == 0); + CHECK(failure->Reason == reason); + CHECK(failure->Disposition() == ProducerRecoveryDisposition::SessionFatal); +} + +struct HighwaterCase final +{ + int Submitted; + std::uint64_t AcknowledgedBefore; + Bytes Frame; + std::uint64_t TransactionId; + std::string_view Reason; + std::size_t ArtifactSize; +}; + +void TestHighwaterContradictionsRequireRecovery() +{ + const auto hello_ok = [](std::uint64_t committed_through) + { + server::Hello hello; + hello.CommittedThrough = committed_through; + hello.Token = "not-issued"; + return server::HelloOk(hello); + }; + const HighwaterCase cases[] = { + {0, 0, hello_ok(1), 1, "server producer highwater 1 is ahead of local transaction 0", 0}, + // A session ahead of its local outbox keeps the outbox as evidence. + {2, 0, hello_ok(3), 3, "server producer highwater 3 is ahead of local transaction 2", 2}, + {2, 2, hello_ok(1), 1, "server producer highwater regressed from 2 to 1", 0}, + {1, 0, server::Acknowledged(9), 9, + "server producer highwater 9 is ahead of local transaction 1", 1}, + {3, 2, server::Acknowledged(1), 1, "server producer highwater regressed from 2 to 1", 1}, + }; + for (const HighwaterCase& expected : cases) + { + Producer producer; + static_cast(producer.Handshake()); + for (int index = 0; index < expected.Submitted; ++index) + { + CHECK(producer.Submit("/P") == ProducerResult::Accepted); + } + static_cast(producer.Sends()); + if (expected.AcknowledgedBefore != 0) + { + producer.Feed(server::Acknowledged(expected.AcknowledgedBefore)); + } + const bool hello = DecodeSent(expected.Frame).payload_type() == Payload::HelloOk; + if (hello) + { + producer.Disconnect(); + static_cast(producer.Request()); + static_cast(producer.Notices()); + } + producer.Feed(expected.Frame); + producer.ExpectClose(DisconnectReason::RecoveryRequired); + CheckSessionFailure(producer, expected.TransactionId, expected.Reason); + CHECK(producer.Endpoint.Artifact()->Transactions.size() == expected.ArtifactSize); + if (hello) + { + CHECK(producer.Notices().empty()); + } + } +} + +void TestHandshakeProtocolErrorsAreFailedAttempts() +{ + for (const Bytes& bytes : {server::Ping(), Bytes{0, 0, 0, 2, 0xFF, 0xFF}}) + { + Producer producer; + static_cast(producer.Request()); + producer.Feed(bytes); + CHECK(producer.Single().Reason == DisconnectReason::ProtocolError); + producer.Disconnect(); + const ProducerStatus status = producer.Status(); + CHECK(!status.Connected && !status.Rejection && !status.Failure); + CHECK(producer.Notices().empty()); + // A failed request backs off before the next one. + CHECK(!producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + CHECK(producer.Endpoint.RequestConnect(producer.Now + 1s, producer.Now + 3s)); + } +} + +void TestHandshakeDeadline() +{ + { + Producer producer; + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + static_cast(producer.Attempt()); + producer.Advance(1999ms); + CHECK(producer.Commands().empty()); + producer.Advance(1ms); + CHECK(producer.Single().Reason == DisconnectReason::HandshakeTimeout); + CHECK(!producer.Status().Handshaking && !producer.Endpoint.NextWake()); + // The host's connect finishing late opens nothing. + producer.Endpoint.OnConnected({}); + CHECK(producer.Commands().empty()); + producer.Disconnect(DisconnectReason::ConnectFailed); + CHECK(!producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + } + { + Producer producer; + static_cast(producer.Request()); + producer.Advance(2s); + CHECK(producer.Single().Reason == DisconnectReason::HandshakeTimeout); + producer.Feed(server::HelloOk()); + CHECK(!producer.Status().Connected && producer.Notices().empty()); + producer.Disconnect(); + producer.Now += 1s; + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + } +} + +void FailRequest(Producer& producer) +{ + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + static_cast(producer.Attempt()); + producer.Disconnect(DisconnectReason::ConnectFailed); +} + +// After a failed requested attempt, what happens next decides whether +// RequestConnect is allowed at once and a second later. +void TestRequestBackoff() +{ + const std::tuple rows[] = { + {nullptr, false, true}, + // A second failure once the backoff ends doubles it. + {[](Producer& producer) + { + producer.Now += 1s; + FailRequest(producer); + }, + false, false}, + // An explicit Connect neither waits for nor extends it. + {[](Producer& producer) + { + CHECK(producer.Endpoint.Connect(producer.Now, producer.Now + 2s) == + ConnectResult::Started); + static_cast(producer.Attempt()); + producer.Disconnect(DisconnectReason::ConnectFailed); + }, + false, true}, + // Cancelling and a published connection reset it; a lost connection is + // not a failed attempt. + {[](Producer& producer) + { + CHECK(producer.Endpoint.CancelConnect()); + }, + true, true}, + {[](Producer& producer) + { + producer.Now += 1s; + static_cast(producer.Handshake()); + producer.Disconnect(); + FailRequest(producer); + }, + false, true}, + }; + for (const auto& [then, allowed_now, allowed_later] : rows) + { + for (const std::chrono::milliseconds later : {0s, 1s}) + { + Producer producer; + FailRequest(producer); + if (then) + { + then(producer); + } + CHECK(producer.Commands().empty()); + const TimePoint at = producer.Now + later; + CHECK(producer.Endpoint.RequestConnect(at, at + 2s) == + (later == 0s ? allowed_now : allowed_later)); + } + } +} + +// A cancelled handshake must not end a later attempt's session generation or +// accept its own late HelloOk, either of which quarantines a healthy session. +void TestCancelDuringHandshakeKeepsTheSessionHealthy(bool disconnect) +{ + Producer producer; + static_cast(producer.Handshake()); + CHECK(producer.Submit("/Unsent") == ProducerResult::Accepted); + const Bytes unsent = producer.Sends().at(0); + producer.Disconnect(); + CHECK(producer.Notice().Reason == DisconnectReason::PeerClosed); + + static_cast(producer.Request()); + bool finished = false; + if (disconnect) + { + producer.Endpoint.Disconnect(); + } + else + { + finished = producer.Endpoint.CancelConnect(); + } + // The HelloOk was already in flight when the host cancelled. + producer.Feed(server::HelloOk()); + const ProducerStatus status = producer.Status(); + CHECK(!status.Connected && !status.Handshaking && !status.Failure); + CHECK(status.Closing && !finished); + CHECK(producer.Single().Reason == DisconnectReason::Cancelled); + CHECK(producer.Notices().empty()); + CHECK(!producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + CHECK(producer.Endpoint.Connect(producer.Now, producer.Now + 2s) == ConnectResult::Busy); + producer.Disconnect(DisconnectReason::Cancelled); + CHECK(!producer.Status().Closing); + CHECK(producer.Endpoint.CancelConnect()); + + CHECK(producer.Handshake() == std::vector{unsent}); + CHECK(producer.Submit("/After") == ProducerResult::Accepted); + CHECK(TransactionIds(producer.Sends()) == std::vector{2}); + CHECK(!producer.Status().Failure); +} + +void TestCancelBeforeTheSocketOpens() +{ + Producer producer; + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + static_cast(producer.Attempt()); + CHECK(!producer.Endpoint.CancelConnect()); + CHECK(producer.Single().Reason == DisconnectReason::Cancelled); + // The host's connect finished before it applied the close. + producer.Endpoint.OnConnected({}); + CHECK(producer.Commands().empty()); + producer.Disconnect(DisconnectReason::Cancelled); + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + + // Cancelling leaves a published connection intact. + CHECK(producer.Endpoint.CancelConnect()); + CHECK(producer.Status().Connected); + CHECK(TransactionIds(producer.Sends()) == std::vector{1}); +} + +// The host may report a socket's end before it applies the close queued for +// it. Applying that close later would end the next attempt in its place. +void TestAReportedEndVoidsActionsQueuedForTheConnection() +{ + Producer producer; + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + static_cast(producer.Attempt()); + CHECK(!producer.Endpoint.CancelConnect()); + producer.Disconnect(DisconnectReason::ConnectFailed); + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + static_cast(producer.Attempt()); + static_cast(producer.Open()); + producer.Feed(server::HelloOk()); + CHECK(producer.Status().Connected); + + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + producer.Endpoint.Disconnect(); + producer.Disconnect(); + CHECK(producer.Commands().empty()); + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + static_cast(producer.Attempt()); +} + +// An attempt the host has not taken ends at once, so its close cannot share +// a batch with its connect and outlive it. +void TestAnUntakenAttemptIsWithdrawn() +{ + for (int end = 0; end < 4; ++end) + { + Producer producer; + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + switch (end) + { + case 0: + CHECK(producer.Endpoint.CancelConnect()); + break; + case 1: + producer.Endpoint.Disconnect(); + break; + case 2: + producer.Advance(2s); + break; + default: + producer.Endpoint.Stop(); + break; + } + CHECK(producer.Commands().empty()); + const ProducerStatus status = producer.Status(); + CHECK(!status.Handshaking && !status.Closing && !producer.Endpoint.NextWake()); + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s) == (end != 3)); + } +} + +void TestPublicationReplaysUnsentFramesInOrderWithoutCopies() +{ + Producer producer; + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A", "", 1) == ProducerResult::Accepted); + CHECK(producer.Submit("/B", "", 2) == ProducerResult::Accepted); + CHECK(producer.Submit("/C", "", 3) == ProducerResult::Accepted); + const std::vector first = producer.SendActions(); + CHECK(first.size() == 3); + producer.Feed(server::Acknowledged(1)); + producer.Disconnect(); + CHECK(producer.Notice().Reason == DisconnectReason::PeerClosed); + CHECK(producer.Status().PendingTransactions == 2); + + static_cast(producer.Request()); + server::Hello hello; + hello.CommittedThrough = 2; + producer.Feed(server::HelloOk(hello)); + const std::vector replayed = producer.SendActions(); + CHECK(replayed.size() == 1); + CHECK(replayed[0].Bytes == first[2].Bytes); + const ProducerStatus status = producer.Status(); + CHECK(status.PendingTransactions == 1 && status.AcknowledgedTransactions == 2); + + CHECK(producer.Submit("/D") == ProducerResult::Accepted); + producer.Feed(server::Acknowledged(4)); + CHECK(producer.Endpoint.OutboxEmpty()); +} + +void TestAppendRequiresAPublishedHealthyConnection() +{ + ProducerConfig config = TestConfig(); + config.MaxPendingTransactions = 2; + Producer producer(config); + CHECK(producer.Submit("/A") == ProducerResult::InvalidPhase); + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + static_cast(producer.Attempt()); + CHECK(producer.Submit("/A") == ProducerResult::InvalidPhase); + static_cast(producer.Open()); + CHECK(producer.Submit("/A") == ProducerResult::InvalidPhase); + CHECK(producer.Commands().empty()); + CHECK(producer.Endpoint.NextTransactionId() == 1 && producer.Endpoint.OutboxEmpty()); + + producer.Feed(server::HelloOk()); + static_cast(producer.Notices()); + const Bytes frame = TransactionFrame(1, "/A"); + CHECK(producer.Endpoint.Append(2, TransactionFrame(2, "/A"), 1, "") == + ProducerResult::SequenceMismatch); + CHECK(producer.Endpoint.Append(1, frame, 0, "") == ProducerResult::InvalidArgument); + CHECK(producer.Endpoint.Append(1, Bytes(frame.begin(), frame.end() - 1), 1, "") == + ProducerResult::InvalidArgument); + CHECK(producer.Commands().empty()); + + CHECK(producer.Endpoint.Append(1, frame, 1, "") == ProducerResult::Accepted); + CHECK(producer.Sends() == std::vector{frame}); + CHECK(producer.Submit("/B") == ProducerResult::Accepted); + CHECK(producer.Submit("/C") == ProducerResult::OutboxFull); + CHECK(TransactionIds(producer.Sends()) == std::vector{2}); + producer.Feed(server::Acknowledged(1)); + CHECK(producer.Submit("/C") == ProducerResult::Accepted); + static_cast(producer.Sends()); + + producer.Disconnect(); + CHECK(producer.Submit("/D") == ProducerResult::InvalidPhase); + server::Hello hello; + hello.CommittedThrough = 1; + static_cast(producer.Handshake(hello)); + producer.Feed(server::Rejected(2, TransactionRejectionCode::InvalidTransaction, "bad")); + producer.ExpectClose(DisconnectReason::RecoveryRequired); + CHECK(producer.Submit("/D") == ProducerResult::RecoveryRequired); + CHECK(producer.Status().PendingTransactions == 2); +} + +void TestAcknowledgedCheckpointNeedsAnEmptyOutbox() +{ + Producer producer; + server::Hello hello; + hello.ServerInstance = "server-a"; + static_cast(producer.Handshake(hello)); + CHECK(!producer.Endpoint.AcknowledgedCheckpoint()); + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + CHECK(producer.Submit("/B") == ProducerResult::Accepted); + producer.Feed(server::Acknowledged(1, server::Checkpoint{2, 8})); + CHECK(!producer.Endpoint.AcknowledgedCheckpoint()); + producer.Feed(server::Acknowledged(2, server::Checkpoint{3, 9})); + std::optional checkpoint = producer.Endpoint.AcknowledgedCheckpoint(); + CHECK(checkpoint && checkpoint->ServerInstance == "server-a"); + CHECK(checkpoint->Epoch == 3 && checkpoint->HeadSequence == 9); + + CHECK(producer.Submit("/C") == ProducerResult::Accepted); + CHECK(!producer.Endpoint.AcknowledgedCheckpoint()); + producer.Feed(server::Acknowledged(3)); + CHECK(producer.Endpoint.OutboxEmpty() && !producer.Endpoint.AcknowledgedCheckpoint()); + + CHECK(producer.Submit("/D") == ProducerResult::Accepted); + producer.Feed(server::Acknowledged(4, server::Checkpoint{4, 10})); + CHECK(producer.Endpoint.AcknowledgedCheckpoint()); + // A Hello acknowledges no mirror position, even with nothing pending. + producer.Disconnect(); + CHECK(producer.Endpoint.AcknowledgedCheckpoint()); + static_cast(producer.Request()); + hello.ServerInstance = "server-b"; + hello.CommittedThrough = 4; + producer.Feed(server::HelloOk(hello)); + CHECK(producer.Status().Connected && !producer.Endpoint.AcknowledgedCheckpoint()); + + // Without a server instance, a checkpoint cannot name its sequence domain. + Producer anonymous; + hello = {}; + hello.ServerInstance.clear(); + static_cast(anonymous.Handshake(hello)); + CHECK(anonymous.Submit("/A") == ProducerResult::Accepted); + anonymous.Feed(server::Acknowledged(1, server::Checkpoint{1, 1})); + CHECK(anonymous.Endpoint.OutboxEmpty() && !anonymous.Endpoint.AcknowledgedCheckpoint()); +} + +void TestRejectionQuarantinesTheOutbox() +{ + Producer producer; + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A", "layer-a", 1) == ProducerResult::Accepted); + CHECK(producer.Submit("/B", "layer-b", 2) == ProducerResult::Accepted); + const std::vector sent = producer.Sends(); + // Frames after a rejection in the same read are not handled. + producer.Feed(Concatenate( + {server::Rejected(1, TransactionRejectionCode::StaleLayerGraph, "layer was remapped", 1), + server::Acknowledged(1)})); + CHECK(producer.Single().Reason == DisconnectReason::RecoveryRequired); + CHECK(producer.Notice().Reason == DisconnectReason::RecoveryRequired); + + const std::optional failure = producer.Endpoint.Failure(); + CHECK(failure && failure->TransactionId == 1 && failure->ExpectedTransactionId == 1); + CHECK(failure->Code == static_cast(TransactionRejectionCode::StaleLayerGraph)); + CHECK(failure->Reason == "layer was remapped"); + + CHECK(producer.Submit("/C") == ProducerResult::RecoveryRequired); + CHECK(!producer.Endpoint.QueueControl(ClaimFrame())); + producer.Disconnect(); + const ProducerStatus status = producer.Status(); + CHECK(!status.Connected && status.Failure && status.PendingTransactions == 2); + CHECK(!producer.Endpoint.AcknowledgedCheckpoint()); + CHECK(!producer.Endpoint.RequestConnect(producer.Now + 1h, producer.Now + 2h)); + CHECK(producer.Endpoint.Connect(producer.Now, producer.Now + 2s) == ConnectResult::Refused); + + const std::optional artifact = producer.Endpoint.Artifact(); + CHECK(artifact && artifact->SessionId == "session"); + CHECK(artifact->Transactions.size() == 2); + const std::pair transactions[] = {{"layer-a", 1}, + {"layer-b", 2}}; + for (std::size_t index = 0; index < 2; ++index) + { + const ProducerSessionEntry& entry = artifact->Transactions[index]; + CHECK(entry.TransactionId == index + 1); + CHECK(*entry.Payload == sent[index]); + CHECK(entry.LayerKey == transactions[index].first); + CHECK(entry.EventCount == transactions[index].second); + } +} + +void TestRejectionOfAnUnknownTransaction() +{ + Producer producer; + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + CHECK(TransactionIds(producer.Sends()) == std::vector{1}); + producer.Feed(server::Rejected(7, TransactionRejectionCode::StaleLayerGraph, "stale", 1)); + producer.ExpectClose(DisconnectReason::RecoveryRequired); + const std::optional failure = producer.Endpoint.Failure(); + CHECK(failure && failure->TransactionId == 7 && failure->ExpectedTransactionId == 0); + CHECK(failure->Code == static_cast(RejectionCode::UnexpectedId)); + CHECK(failure->Reason == "server rejected unknown transaction 7"); + CHECK(failure->Disposition() == ProducerRecoveryDisposition::SessionFatal); + CHECK(producer.Endpoint.RepairRejected(TransactionFrame(7, "/A"), 1, "") == + ProducerResult::RecoveryNotRecoverable); +} + +void TestRateLimitClosesAndOpensARetryWindow() +{ + Producer producer; + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + const std::vector sent = producer.Sends(); + producer.Feed(server::RateLimited(1.5F)); + CHECK(producer.Single().Reason == DisconnectReason::RateLimited); + CHECK(producer.Notice().Reason == DisconnectReason::RateLimited); + // The window starts once the close is reported. + producer.Now += 100ms; + producer.Disconnect(); + CHECK(producer.Status().RetryAfter == producer.Now + 1500ms); + CHECK(!producer.Endpoint.RequestConnect(producer.Now + 1499ms, producer.Now + 1h)); + CHECK(producer.Endpoint.Connect(producer.Now + 1499ms, producer.Now + 1h) == + ConnectResult::Refused); + producer.Now += 1500ms; + CHECK(producer.Handshake() == sent); + CHECK(!producer.Status().Failure); + + const std::pair hostile[] = { + {std::numeric_limits::quiet_NaN(), std::chrono::steady_clock::duration::zero()}, + {std::numeric_limits::infinity(), 1h}, + }; + for (const auto& [seconds, window] : hostile) + { + Producer limited; + static_cast(limited.Handshake()); + limited.Feed(server::RateLimited(seconds)); + limited.ExpectClose(DisconnectReason::RateLimited); + CHECK(limited.Status().RetryAfter == limited.Now + window); + CHECK(limited.Endpoint.RequestConnect(limited.Now + window, limited.Now + window + 1s)); + } +} + +void TestRepairReplacesTheRejectedTransaction() +{ + Producer producer; + static_cast(producer.Handshake()); + CHECK(producer.Endpoint.RepairRejected(TransactionFrame(1, "/A"), 1, "") == + ProducerResult::InvalidPhase); + CHECK(producer.Submit("/Stale", "old-layer") == ProducerResult::Accepted); + CHECK(producer.Submit("/Later", "stable-layer") == ProducerResult::Accepted); + const std::vector sent = producer.Sends(); + producer.Feed(server::Rejected(1, TransactionRejectionCode::StaleLayerGraph, "remapped")); + producer.ExpectClose(DisconnectReason::RecoveryRequired); + + const Bytes repaired = TransactionFrame(1, "/Repaired", "new-layer"); + CHECK(producer.Endpoint.Failure()); + CHECK(producer.Endpoint.RepairRejected(repaired, 2, "new-layer") == ProducerResult::Accepted); + const ProducerStatus status = producer.Status(); + CHECK(!status.Failure && !producer.Endpoint.Artifact()); + CHECK(status.PendingTransactions == 2); + + CHECK(producer.Handshake() == (std::vector{repaired, sent[1]})); + producer.Feed(server::Acknowledged(2)); + CHECK(producer.Endpoint.OutboxEmpty()); + + Producer invalid; + static_cast(invalid.Handshake()); + CHECK(invalid.Submit("/A") == ProducerResult::Accepted); + invalid.Feed(server::Rejected(1, TransactionRejectionCode::InvalidTransaction, "bad")); + CHECK(invalid.Endpoint.RepairRejected(TransactionFrame(1, "/A"), 1, "") == + ProducerResult::RecoveryNotRecoverable); + CHECK(invalid.Endpoint.Failure()); +} + +void TestAbandonContinuesAsAFreshSession() +{ + Producer producer; + CHECK(!producer.Endpoint.AbandonRejectedSession("replacement")); + static_cast(producer.Handshake()); + CHECK(producer.Submit("/Rejected") == ProducerResult::Accepted); + CHECK(producer.Submit("/Suffix") == ProducerResult::Accepted); + static_cast(producer.Sends()); + producer.Feed(server::Rejected(1, TransactionRejectionCode::InvalidTransaction, "injected")); + producer.ExpectClose(DisconnectReason::RecoveryRequired); + + CHECK(!producer.Endpoint.AbandonRejectedSession("")); + CHECK(!producer.Endpoint.AbandonRejectedSession("session")); + CHECK(producer.Endpoint.Failure()); + + const std::optional artifact = + producer.Endpoint.AbandonRejectedSession("replacement"); + CHECK(artifact && artifact->SessionId == "session"); + CHECK(artifact->Transactions.size() == 2); + const ProducerStatus status = producer.Status(); + CHECK(status.SessionId == "replacement" && !status.Failure); + CHECK(status.PendingTransactions == 0 && status.NextTransactionId == 1); + CHECK(!producer.Endpoint.AbandonRejectedSession("another")); + + CHECK(producer.Request().ProducerSessionId == "replacement"); + producer.Feed(server::HelloOk()); + CHECK(producer.Status().Connected && producer.Sends().empty()); + CHECK(producer.Submit("/Rebuilt") == ProducerResult::Accepted); + CHECK(TransactionIds(producer.Sends()) == std::vector{1}); +} + +void TestDisconnectSaysQuitAndKeepsTheOutbox() +{ + Producer producer; + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + const std::vector sent = producer.Sends(); + producer.Endpoint.Disconnect(); + CHECK(DecodeSent(*producer.Next().Bytes).payload_type() == Payload::Quit); + CHECK(producer.Next().Reason == DisconnectReason::Cancelled); + CHECK(producer.Notice().Reason == DisconnectReason::Cancelled); + CHECK(!producer.Status().Connected); + CHECK(producer.Submit("/B") == ProducerResult::InvalidPhase); + CHECK(!producer.Endpoint.QueueControl(ClaimFrame())); + CHECK(!producer.Endpoint.CancelConnect()); + producer.Disconnect(DisconnectReason::Cancelled); + + producer.Endpoint.Disconnect(); + CHECK(producer.Commands().empty()); + CHECK(producer.Handshake() == sent); + + // An attempt in flight is closed without a Quit. + producer.Disconnect(); + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); + static_cast(producer.Attempt()); + producer.Endpoint.Disconnect(); + CHECK(producer.Single().Reason == DisconnectReason::Cancelled); +} + +void TestStopIsFinal() +{ + Producer producer; + static_cast(producer.Handshake()); + producer.Endpoint.Stop(); + CHECK(DecodeSent(*producer.Next().Bytes).payload_type() == Payload::Quit); + CHECK(producer.Next().Reason == DisconnectReason::Stopped); + CHECK(producer.Notice().Reason == DisconnectReason::Stopped); + producer.Disconnect(DisconnectReason::Stopped); + CHECK(producer.Status().Stopped && !producer.Status().Connected); + CHECK(!producer.Endpoint.RequestConnect(producer.Now + 1h, producer.Now + 2h)); + CHECK(producer.Endpoint.Connect(producer.Now, producer.Now + 2s) == ConnectResult::Refused); + CHECK(producer.Endpoint.CancelConnect()); + producer.Endpoint.OnConnected({}); + CHECK(producer.Single().Reason == DisconnectReason::Stopped); + producer.Endpoint.Stop(); + producer.Endpoint.Disconnect(); + CHECK(producer.Commands().empty()); + + Producer attempting; + CHECK(attempting.Endpoint.RequestConnect(attempting.Now, attempting.Now + 2s)); + static_cast(attempting.Attempt()); + attempting.Endpoint.Stop(); + CHECK(attempting.Single().Reason == DisconnectReason::Stopped); + CHECK(attempting.Status().Stopped && !attempting.Endpoint.NextWake()); +} + +void TestControlFramesFollowTheConnection() +{ + Producer producer; + CHECK(!producer.Endpoint.QueueControl(ClaimFrame())); + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + CHECK(producer.Endpoint.QueueControl(ClaimFrame())); + CHECK(producer.Submit("/B") == ProducerResult::Accepted); + CHECK(!producer.Endpoint.QueueControl(Bytes{1, 2, 3})); + CHECK((Kinds(producer.Sends()) == + std::vector{Payload::Txn, Payload::ClaimPlayback, Payload::Txn})); + producer.Disconnect(); + CHECK(TransactionIds(producer.Handshake()) == (std::vector{1, 2})); +} + +void TestMessagesWithoutAProducerActionAreIgnored() +{ + Producer producer; + static_cast(producer.Handshake()); + flatbuffers::FlatBufferBuilder newer(32); + const Bytes frames[] = { + server::Claimed("client"), + server::Event(1), + server::HelloOk(), + server::Frame(newer, static_cast(200), OpenUSDConnect::CreatePing(newer).Union()), + }; + for (const Bytes& frame : frames) + { + producer.Feed(frame); + } + CHECK(producer.Commands().empty() && producer.Notices().empty()); + CHECK(producer.Status().Connected); +} + +// A result read with the HelloOk applies to the connection the HelloOk published. +void TestFramesAfterTheHelloOkInOneReadFollowPublication() +{ + Producer producer; + static_cast(producer.Handshake()); + CHECK(producer.Submit("/A") == ProducerResult::Accepted); + CHECK(producer.Submit("/B") == ProducerResult::Accepted); + static_cast(producer.Sends()); + producer.Disconnect(); + + static_cast(producer.Request()); + server::Hello hello; + hello.CommittedThrough = 1; + producer.Feed(Concatenate({server::HelloOk(hello), server::Acknowledged(2)})); + CHECK(producer.Status().Connected && producer.Endpoint.OutboxEmpty()); + CHECK(TransactionIds(producer.Sends()) == std::vector{2}); +} + +void TestStaleHostReportsAreIgnored() +{ + Producer producer; + producer.Feed(server::HelloOk()); + producer.Endpoint.OnConnected({}); + producer.Disconnect(DisconnectReason::ConnectFailed); + producer.Advance(1h); + CHECK(producer.Commands().empty() && producer.Notices().empty()); + CHECK(producer.Endpoint.RequestConnect(producer.Now, producer.Now + 2s)); +} + +// The host thread appends while the loop thread reads acknowledgements. +void TestConcurrentAppendWhileTheLoopReads() +{ + constexpr std::uint64_t kTransactions = 2'000; + ProducerConfig config = TestConfig(); + config.MaxPendingTransactions = 64; + Producer producer(config); + static_cast(producer.Handshake()); + ProducerEndpoint& endpoint = producer.Endpoint; + + std::thread host( + [&endpoint] + { + for (std::uint64_t id = 1; id <= kTransactions;) + { + const ProducerResult result = + endpoint.Append(id, TransactionFrame(id, "/P"), 1, ""); + CHECK(result == ProducerResult::Accepted || result == ProducerResult::OutboxFull); + if (result == ProducerResult::Accepted) + { + ++id; + } + else + { + std::this_thread::yield(); + } + } + }); + std::uint64_t next = 1; + while (next <= kTransactions) + { + for (Action& action : endpoint.TakeActions()) + { + const SendAction* send = std::get_if(&action); + CHECK(send != nullptr); + CHECK(TransactionIds({*send->Bytes}) == std::vector{next}); + const Bytes ack = server::Acknowledged(next++); + endpoint.OnBytes(ack.data(), ack.size()); + } + CHECK(endpoint.Status().Connected); + std::this_thread::yield(); + } + host.join(); + CHECK(endpoint.OutboxEmpty()); + CHECK(endpoint.Status().AcknowledgedTransactions == kTransactions); +} + +} // namespace + +int main() +{ + TestConfigurationValidation(); + TestAttemptsAreBoundedAndExclusive(); + TestHelloCarriesTheProducerIdentity(); + TestAcceptedHelloNotifiesThenPublishes(); + TestHandshakeRejectionsHoldUntilAnExplicitConnect(); + TestHighwaterContradictionsRequireRecovery(); + TestHandshakeProtocolErrorsAreFailedAttempts(); + TestHandshakeDeadline(); + TestRequestBackoff(); + TestCancelDuringHandshakeKeepsTheSessionHealthy(false); + TestCancelDuringHandshakeKeepsTheSessionHealthy(true); + TestCancelBeforeTheSocketOpens(); + TestAReportedEndVoidsActionsQueuedForTheConnection(); + TestAnUntakenAttemptIsWithdrawn(); + TestPublicationReplaysUnsentFramesInOrderWithoutCopies(); + TestAppendRequiresAPublishedHealthyConnection(); + TestAcknowledgedCheckpointNeedsAnEmptyOutbox(); + TestRejectionQuarantinesTheOutbox(); + TestRejectionOfAnUnknownTransaction(); + TestRateLimitClosesAndOpensARetryWindow(); + TestRepairReplacesTheRejectedTransaction(); + TestAbandonContinuesAsAFreshSession(); + TestDisconnectSaysQuitAndKeepsTheOutbox(); + TestStopIsFinal(); + TestControlFramesFollowTheConnection(); + TestMessagesWithoutAProducerActionAreIgnored(); + TestFramesAfterTheHelloOkInOneReadFollowPublication(); + TestStaleHostReportsAreIgnored(); + TestConcurrentAppendWhileTheLoopReads(); + return 0; +} diff --git a/tests/native/test_protocol_codec.cpp b/tests/native/test_protocol_codec.cpp index beabae8..b100367 100644 --- a/tests/native/test_protocol_codec.cpp +++ b/tests/native/test_protocol_codec.cpp @@ -1,13 +1,69 @@ +#include "openusdconnect/client/frame_codec.h" #include "openusdconnect/client/protocol_codec.h" #include "test_check.h" #include +#include #include +#include using namespace openusdconnect::client; +using Bytes = std::vector; + +[[nodiscard]] static Bytes ToBytes(std::string_view text) +{ + return Bytes(text.begin(), text.end()); +} + +static void TestFrameDecoderHandlesFragmentedAndCoalescedInput() +{ + Bytes stream; + for (const std::string_view payload : {"alpha", "beta"}) + { + Bytes frame; + CHECK(EncodeFrame(ToBytes(payload).data(), payload.size(), frame) == FrameResult::Success); + stream.insert(stream.end(), frame.begin(), frame.end()); + } + FrameDecoder decoder; + std::vector frames; + CHECK(decoder.Feed(stream.data(), 3, frames) == FrameResult::Success); + CHECK(frames.empty()); + CHECK(decoder.BufferedBytes() == 3); + CHECK(decoder.Feed(stream.data() + 3, 5, frames) == FrameResult::Success); + CHECK(frames.empty()); + CHECK(decoder.Feed(stream.data() + 8, stream.size() - 8, frames) == FrameResult::Success); + CHECK((frames == std::vector{ToBytes("alpha"), ToBytes("beta")})); + CHECK(decoder.BufferedBytes() == 0); +} + +static void TestFrameDecoderRejectsAnInvalidSizeAtTheHeaderBoundary() +{ + FrameDecoder decoder(8); + const std::uint8_t header[kFrameHeaderSize] = {0, 0, 0, 9}; + std::vector frames; + CHECK(decoder.Feed(header, sizeof(header), frames) == FrameResult::InvalidHeader); + CHECK(frames.empty()); + CHECK(decoder.BufferedBytes() == 0); +} + +static void TestFrameSizeValidation() +{ + CHECK(!IsValidMaxFrameSize(0)); + CHECK(IsValidMaxFrameSize(kDefaultMaxFrameSize)); + const Bytes payload = ToBytes("oversized"); + Bytes frame; + CHECK(EncodeFrame(payload.data(), 0, frame) == FrameResult::EmptyPayload); + CHECK(EncodeFrame(payload.data(), payload.size(), frame, 4) == FrameResult::PayloadTooLarge); + CHECK(frame.empty()); +} + int main() { + TestFrameDecoderHandlesFragmentedAndCoalescedInput(); + TestFrameDecoderRejectsAnInvalidSizeAtTheHeaderBoundary(); + TestFrameSizeValidation(); + flatbuffers::FlatBufferBuilder hello_builder(256); const HelloParameters hello{"emitter", 0, "client", "origin", "lighting", "token", false, OpenUSDConnect::LayerMode::Managed, @@ -43,6 +99,22 @@ int main() CHECK(decoded_replay_hello->replay_epoch().has_value()); CHECK(*decoded_replay_hello->replay_epoch() == 3); + flatbuffers::FlatBufferBuilder anonymous_builder(128); + CHECK(BuildHelloFrame(anonymous_builder, HelloParameters{"receiver", 1}) == + ProtocolResult::Success); + CHECK(BuildHelloFrame(anonymous_builder, HelloParameters{"emitter", 0, "", "origin"}) == + ProtocolResult::InvalidArgument); + HelloParameters emitter{"emitter", 0, "client"}; + CHECK(!IsValidHelloParameters(emitter)); + emitter.ProducerSessionId = "producer"; + CHECK(IsValidHelloParameters(emitter)); + const std::string longest(kMaxProducerSessionIdLength, 'p'); + emitter.ProducerSessionId = longest; + CHECK(IsValidHelloParameters(emitter)); + const std::string too_long = longest + "p"; + emitter.ProducerSessionId = too_long; + CHECK(!IsValidHelloParameters(emitter)); + flatbuffers::FlatBufferBuilder transaction_builder(256); const VisibilityEventView visibility{"/World/Sphere", true}; flatbuffers::Offset event; diff --git a/tests/native/test_receiver_endpoint.cpp b/tests/native/test_receiver_endpoint.cpp new file mode 100644 index 0000000..a3e75bd --- /dev/null +++ b/tests/native/test_receiver_endpoint.cpp @@ -0,0 +1,1138 @@ +#include "openusdconnect/client/engine/receiver_endpoint.h" + +#include "endpoint_host.h" +#include "receiver_frames.h" +#include "test_check.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +using namespace openusdconnect::client; +using namespace std::chrono_literals; +using OpenUSDConnect::HelloRejectionCode; +using OpenUSDConnect::LayerMode; +using OpenUSDConnect::Payload; + +namespace +{ + +using namespace receiver_test; + +// Whether the server resumes this Hello rather than sending Resync. +[[nodiscard]] bool ServerResumes(const SentHello& hello, std::string_view instance, + std::uint64_t epoch, std::int32_t head) +{ + const bool claim_holds = !hello.ReplayServerInstance || hello.Claims(instance, epoch); + return hello.SyncFrom <= head + 1 && (hello.SyncFrom == 1 || claim_holds); +} + +// Plays both the host and the stage-owning consumer around one endpoint. +class Receiver final : public Host +{ +public: + explicit Receiver(const ReceiverConfig& config = TestConfig()) + : Host(config) + , Applied(config.SyncFrom - 1) + { + } + + // Completes the pending connect attempt and returns the Hello it sent. + SentHello Connect(std::string_view token = {}) + { + static_cast(Single()); + Endpoint.OnConnected(token); + return DecodeHello(*Single().Bytes); + } + + SentHello Start(std::string_view token = {}) + { + CHECK(Endpoint.Start(Now)); + return Connect(token); + } + + SentHello Handshake(const server::Hello& hello = {}) + { + const SentHello sent = Start(); + Feed(server::HelloOk(hello)); + CHECK(Status().Connected); + return sent; + } + + // Reports the disconnect, waits out the backoff, and returns the next Hello. + SentHello Reconnect() + { + Disconnect(); + Now = Endpoint.NextWake().value(); + Endpoint.OnTick(Now); + return Connect(); + } + + // Reports the disconnect of an overflowed connection, applies the queue, and + // returns the Hello sent at the next wake. + SentHello ReconnectAfterDrain() + { + Disconnect(); + const TimePoint poll = Endpoint.NextWake().value(); + static_cast(Apply()); + Now = poll; + Endpoint.OnTick(Now); + return Connect(); + } + + // Drains and applies every queued frame as a successful consumer batch does. + std::vector Apply() + { + const std::uint64_t generation = Endpoint.Generation(); + std::vector frames = Endpoint.DrainFrames(); + bool reset = false; + for (const Bytes& frame : frames) + { + const OpenUSDConnect::Envelope& envelope = Decode(frame); + if (envelope.payload_type() == Payload::Resync) + { + reset = true; + Applied = 0; + } + else if (envelope.payload_type() == Payload::BroadcastEvent) + { + Applied = envelope.payload_as_BroadcastEvent()->seq(); + } + } + if (reset) + { + Endpoint.ResetAppliedProgress(); + } + if (!frames.empty()) + { + static_cast(Endpoint.MarkAppliedThrough(generation, Applied)); + } + static_cast(Endpoint.MarkReplayApplied()); + return frames; + } + + // The consumer's applied cursor. + std::int32_t Applied; +}; + +void TestConfigurationValidation() +{ + const ReceiverConfig valid = TestConfig(); + const ReceiverConfig shared_stage = + With(With(valid, &ReceiverConfig::LayerMode, LayerMode::SharedStage), + &ReceiverConfig::LayeredReplay, false); + const std::pair rules[] = { + {valid, true}, + {shared_stage, true}, + {With(With(valid, &ReceiverConfig::ClientId, ""), &ReceiverConfig::Origin, ""), true}, + {With(valid, &ReceiverConfig::Port, 0), false}, + {With(valid, &ReceiverConfig::ReconnectMaxDelay, 999ms), false}, + {With(valid, &ReceiverConfig::LayerMode, LayerMode::SharedStage), false}, + {With(shared_stage, &ReceiverConfig::Department, "layout"), false}, + }; + for (const auto& [config, expected] : rules) + { + CHECK(ReceiverEndpoint::IsValidConfiguration(config) == expected); + } +} + +void TestStartConnectsOnceAndSendsTheConfiguredHello() +{ + ReceiverConfig config = TestConfig(); + config.Department = "layout"; + config.SyncFrom = 5; + Receiver receiver(config); + CHECK(receiver.Endpoint.Start(receiver.Now)); + const ConnectAction connect = receiver.Single(); + CHECK(connect.Host == "127.0.0.1" && connect.Port == 7200); + CHECK(connect.Deadline == receiver.Now + config.SocketTimeout); + CHECK(!receiver.Endpoint.Start(receiver.Now)); + receiver.Endpoint.OnConnected("token-1"); + const SentHello hello = DecodeHello(*receiver.Single().Bytes); + CHECK(hello.Role == "receiver"); + CHECK(hello.ProtocolVersion == kProtocolVersion); + CHECK(hello.SyncFrom == 5); + CHECK(hello.ClientId == "client"); + CHECK(hello.Origin == "origin"); + CHECK(hello.Department == "layout"); + CHECK(hello.Token == "token-1"); + CHECK(hello.LayeredReplay); + CHECK(hello.Mode == LayerMode::Managed); + CHECK(!hello.ReplayServerInstance && !hello.ReplayEpoch); +} + +void TestAcceptedHelloNotifiesInOrder() +{ + Receiver receiver; + static_cast(receiver.Start()); + CHECK(!receiver.Status().Connected); + StageMetadata metadata; + metadata.TimeCodesPerSecond = 24.0; + metadata.UpAxis = "Z"; + server::Hello hello; + hello.Token = "issued"; + hello.Metadata = metadata; + receiver.Feed(server::HelloOk(hello)); + + const std::vector notices = receiver.Notices(); + CHECK(notices.size() == 3); + CHECK(As(notices[0]).Token == "issued"); + const StageMetadata& received = As(notices[1]); + CHECK(received.TimeCodesPerSecond == 24.0); + CHECK(received.UpAxis == "Z"); + CHECK(!received.FramesPerSecond && !received.StartTimeCode && !received.EndTimeCode && + !received.MetersPerUnit); + static_cast(As(notices[2])); + + const ReceiverStatus status = receiver.Status(); + CHECK(status.Connected && !status.Synchronized && status.LayeredReplayActive); + CHECK(status.Metadata.TimeCodesPerSecond == 24.0); + CHECK(status.QueuedFrames == 0); +} + +void TestHandshakeRejectionsStop() +{ + server::Hello shared; + shared.Mode = LayerMode::SharedStage; + shared.Token = "not-issued"; + server::Hello flat; + flat.LayeredReplay = false; + const std::pair rejections[] = { + {server::AuthRejected("invalid token"), + {true, HelloRejectionCode::Unspecified, "invalid token"}}, + {server::HelloRejected(HelloRejectionCode::LayeredReplayRequired, "replay is required"), + {false, HelloRejectionCode::LayeredReplayRequired, "replay is required"}}, + {server::HelloOk(shared), + {false, HelloRejectionCode::LayerModeMismatch, + "server did not negotiate requested layer mode"}}, + {server::HelloOk(flat), + {false, HelloRejectionCode::LayeredReplayRequired, + "server did not negotiate requested layered replay"}}, + }; + const auto same = [](const HandshakeRejected& left, const HandshakeRejected& right) + { + return left.Authentication == right.Authentication && left.Code == right.Code && + left.Reason == right.Reason; + }; + for (const auto& [frame, expected] : rejections) + { + Receiver receiver; + static_cast(receiver.Start()); + receiver.Feed(frame); + CHECK(receiver.Single().Reason == DisconnectReason::HandshakeRejected); + CHECK(same(receiver.Notice(), expected)); + receiver.Disconnect(); + CHECK(receiver.Commands().empty() && !receiver.Endpoint.NextWake()); + const ReceiverStatus status = receiver.Status(); + CHECK(status.Stopped && status.Rejection && same(*status.Rejection, expected)); + CHECK(receiver.Notices().empty()); + } +} + +void TestNegotiatedModesConnect() +{ + { + ReceiverConfig config = TestConfig(); + config.LayeredReplay = false; + Receiver receiver(config); + CHECK(!receiver.Start().LayeredReplay); + server::Hello hello; + hello.LayeredReplay = false; + receiver.Feed(server::HelloOk(hello)); + CHECK(receiver.Status().Connected); + CHECK(!receiver.Status().LayeredReplayActive); + } + { + ReceiverConfig config = TestConfig(); + config.LayerMode = LayerMode::SharedStage; + config.LayeredReplay = false; + Receiver receiver(config); + CHECK(receiver.Start().Mode == LayerMode::SharedStage); + server::Hello hello; + hello.LayeredReplay = false; + hello.Mode = LayerMode::SharedStage; + receiver.Feed(server::HelloOk(hello)); + CHECK(receiver.Status().Connected); + CHECK(receiver.Status().LayerModeActive == LayerMode::SharedStage); + receiver.Feed(server::LayerGraph(1)); + CHECK(receiver.Status().LastSequence == 1); + CHECK(Kinds(receiver.Endpoint.DrainFrames()) == + std::vector{Payload::LayerGraphState}); + } +} + +void TestControlMessages() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + static_cast(receiver.Notices()); + receiver.Feed(server::Ping()); + receiver.Feed(server::Playback(12.5, true, 2.0, "leader")); + receiver.Feed(server::Claimed("me")); + receiver.Feed(server::ClaimRejected("already led", "other")); + receiver.Feed(server::LayerStack()); + CHECK(receiver.Commands().empty()); + + const std::vector notices = receiver.Notices(); + CHECK(notices.size() == 3); + const PlaybackState& state = As(notices[0]); + CHECK(state.Time == 12.5 && state.Playing && state.Rate == 2.0); + CHECK(state.LeaderClientId == "leader"); + CHECK(As(notices[1]).LeaderClientId == "me"); + const PlaybackRejected& rejected = As(notices[2]); + CHECK(rejected.Reason == "already led" && rejected.CurrentLeaderClientId == "other"); + CHECK(Kinds(receiver.Endpoint.DrainFrames()) == + std::vector{Payload::LayerStackState}); +} + +void TestReplayCompleteWaitsForDrainedFrames() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + for (std::int32_t sequence = 1; sequence <= 3; ++sequence) + { + receiver.Feed(server::Event(sequence)); + } + receiver.Feed(server::ReplayComplete(3, 7)); + CHECK(!receiver.Status().Synchronized); + const std::uint64_t generation = receiver.Endpoint.Generation(); + CHECK((Sequences(receiver.Endpoint.DrainFrames(2)) == std::vector{1, 2})); + CHECK(!receiver.Endpoint.MarkReplayApplied()); + CHECK(Sequences(receiver.Endpoint.DrainFrames(2)) == std::vector{3}); + CHECK(receiver.Endpoint.MarkAppliedThrough(generation, 3)); + CHECK(receiver.Endpoint.MarkReplayApplied()); + const ReceiverStatus status = receiver.Status(); + CHECK(status.Synchronized); + CHECK(status.ReplayHeadSequence == 3); + CHECK(status.ReplayEpoch == 7); + CHECK(status.ServerInstance == "server"); + + receiver.Feed(server::ReplayComplete(-1, 8)); + CHECK(receiver.Logged(LogLevel::Warning)); + CHECK(receiver.Status().Synchronized); +} + +void TestInPlaceResyncClearsReadyUntilApplied() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Feed(server::ReplayComplete(0, 1)); + CHECK(receiver.Endpoint.MarkReplayApplied()); + CHECK(receiver.Status().Synchronized); + + receiver.Feed(server::Resync()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::ReplayComplete(1, 2)); + CHECK(!receiver.Status().Synchronized); + CHECK((Kinds(receiver.Apply()) == + std::vector{Payload::Resync, Payload::BroadcastEvent})); + const ReceiverStatus status = receiver.Status(); + CHECK(status.Synchronized); + CHECK(status.ReplayHeadSequence == 1); + CHECK(status.ReplayEpoch == 2); +} + +void TestDataFrameAcceptResults() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + static_cast(receiver.Notices()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::Event(2)); + static_cast(receiver.Apply()); + receiver.Feed(server::Event(3)); + receiver.Feed(server::Event(2)); + CHECK(receiver.Status().QueuedFrames == 1); + CHECK(!receiver.Logged(LogLevel::Warning)); + receiver.Feed(server::Event(0)); + CHECK(receiver.Logged(LogLevel::Warning)); + CHECK(receiver.Status().QueuedFrames == 1); + CHECK(receiver.Commands().empty()); + + receiver.Feed(server::Event(5)); + CHECK(receiver.Single().Reason == DisconnectReason::SequenceGap); + CHECK(receiver.Notice().Reason == DisconnectReason::SequenceGap); + const ReceiverStatus status = receiver.Status(); + CHECK(status.QueuedFrames == 0); + CHECK(status.LastSequence == 2); + + // Bytes for a connection the endpoint closed are ignored. + receiver.Feed(server::Event(3)); + CHECK(receiver.Status().QueuedFrames == 0); + CHECK(receiver.Reconnect().SyncFrom == 3); + + // The applied cursor restarts with an applied reset. + receiver.Feed(server::HelloOk()); + receiver.Feed(server::Resync()); + receiver.Feed(server::Event(1)); + static_cast(receiver.Apply()); + receiver.Feed(server::Event(3)); + CHECK(receiver.Single().Reason == DisconnectReason::SequenceGap); + CHECK(receiver.Reconnect().SyncFrom == 2); +} + +enum class DrainOutcome +{ + DrainedEarlier, + DrainedDuringWait, + TimedOut, +}; + +// An overflow closes the connection and reconnects from the received cursor +// once the consumer drained the queue, or anyway at the drain deadline. +void TestOverflowWaitsForTheDrain(DrainOutcome outcome) +{ + ReceiverConfig config = TestConfig(); + config.MaxQueue = 2; + config.ReconnectMaxDelay = 2s; + Receiver receiver(config); + static_cast(receiver.Handshake()); + static_cast(receiver.Notices()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::Event(2)); + receiver.Feed(server::Event(3)); + CHECK(receiver.Single().Reason == DisconnectReason::QueueFull); + CHECK(receiver.Notice().Reason == DisconnectReason::QueueFull); + if (outcome == DrainOutcome::DrainedEarlier) + { + static_cast(receiver.Apply()); + } + receiver.Disconnect(); + if (outcome != DrainOutcome::DrainedEarlier) + { + // The endpoint polls for the drain without reconnecting. + const std::optional poll = receiver.Endpoint.NextWake(); + CHECK(poll && *poll > receiver.Now && *poll < receiver.Now + 2s); + receiver.Advance(1s); + CHECK(receiver.Commands().empty()); + } + if (outcome == DrainOutcome::DrainedDuringWait) + { + static_cast(receiver.Apply()); + receiver.Now = *receiver.Endpoint.NextWake(); + receiver.Endpoint.OnTick(receiver.Now); + } + else if (outcome == DrainOutcome::TimedOut) + { + receiver.Advance(1s); + } + CHECK(receiver.Connect().SyncFrom == 3); + CHECK(receiver.Status().QueuedFrames == (outcome == DrainOutcome::TimedOut ? 2U : 0U)); +} + +void TestConsecutiveReadTimeouts() +{ + ReceiverConfig config = TestConfig(); + config.MaxConsecutiveTimeouts = 3; + Receiver receiver(config); + static_cast(receiver.Start()); + const auto time_out = [&](int count) + { + for (int timeout = 0; timeout < count; ++timeout) + { + receiver.Endpoint.OnReadTimeout(); + } + }; + // The handshake counts too, and every received byte restarts the count. + time_out(2); + receiver.Feed(server::HelloOk()); + time_out(2); + receiver.Feed(server::Ping()); + time_out(2); + CHECK(receiver.Commands().empty()); + CHECK(receiver.Status().Connected); + time_out(1); + CHECK(receiver.Single().Reason == DisconnectReason::ReadTimeout); + receiver.Endpoint.OnReadTimeout(); + CHECK(receiver.Commands().empty()); +} + +void TestBackoffDoublesAndResetsAfterConnectedSession() +{ + ReceiverConfig config = TestConfig(); + config.ReconnectMaxDelay = 8s; + Receiver receiver(config); + CHECK(receiver.Endpoint.Start(receiver.Now)); + static_cast(receiver.Single()); + std::vector waits; + const auto next_attempt = [&](DisconnectReason reason) + { + receiver.Disconnect(reason); + const TimePoint due = receiver.Endpoint.NextWake().value(); + waits.push_back(std::chrono::duration_cast(due - receiver.Now)); + receiver.Now = due - 1ms; + receiver.Endpoint.OnTick(receiver.Now); + CHECK(receiver.Commands().empty()); + receiver.Advance(1ms); + static_cast(receiver.Single()); + }; + next_attempt(DisconnectReason::ConnectFailed); + next_attempt(DisconnectReason::ConnectFailed); + receiver.Endpoint.OnConnected({}); + static_cast(receiver.Single()); + receiver.Feed(server::HelloOk()); + next_attempt(DisconnectReason::PeerClosed); + for (int attempt = 0; attempt < 4; ++attempt) + { + next_attempt(DisconnectReason::ConnectFailed); + } + CHECK((waits == std::vector{1s, 2s, 1s, 2s, 4s, 8s, 8s})); +} + +void TestReconnectDisabledStops() +{ + ReceiverConfig config = TestConfig(); + config.Reconnect = false; + { + Receiver receiver(config); + static_cast(receiver.Handshake()); + receiver.Disconnect(); + CHECK(receiver.Commands().empty()); + CHECK(receiver.Status().Stopped); + CHECK(!receiver.Endpoint.NextWake()); + } + { + config.MaxQueue = 1; + Receiver receiver(config); + static_cast(receiver.Handshake()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::Event(2)); + CHECK(receiver.Single().Reason == DisconnectReason::QueueFull); + receiver.Disconnect(); + CHECK(receiver.Commands().empty()); + CHECK(receiver.Status().Stopped); + } +} + +// A toggle applies to the session that is open when it ends. +void TestReconnectToggleAppliesWhenTheSessionEnds() +{ + ReceiverConfig config = TestConfig(); + config.Reconnect = false; + Receiver receiver(config); + static_cast(receiver.Handshake()); + receiver.Endpoint.SetReconnect(true); + CHECK(receiver.Endpoint.RequestReplayFrom(1)); + CHECK(receiver.Single().Reason == DisconnectReason::ReplayRequested); + CHECK(receiver.Reconnect().SyncFrom == 1); + receiver.Feed(server::HelloOk()); + receiver.Endpoint.SetReconnect(false); + receiver.Disconnect(); + CHECK(receiver.Commands().empty()); + CHECK(receiver.Status().Stopped); + receiver.Endpoint.SetReconnect(true); + receiver.Advance(60s); + CHECK(receiver.Commands().empty()); +} + +void TestReplayRequests() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::Event(2)); + CHECK(receiver.Endpoint.RequestReplayFrom(2)); + CHECK(receiver.Single().Reason == DisconnectReason::ReplayRequested); + const ReceiverStatus status = receiver.Status(); + CHECK(status.LastSequence == 1 && status.QueuedFrames == 0); + CHECK(receiver.Reconnect().SyncFrom == 2); + receiver.Feed(server::HelloOk()); + receiver.Feed(server::Event(2)); + CHECK(receiver.Status().LastSequence == 2); + + // Without an open connection the next Hello carries the request. + receiver.Disconnect(); + CHECK(receiver.Endpoint.RequestReplayFrom(2)); + CHECK(receiver.Commands().empty()); + receiver.Advance(1s); + CHECK(receiver.Connect().SyncFrom == 2); + + // A request during the handshake abandons it before the Hello is accepted. + static_cast(receiver.Notices()); + CHECK(receiver.Endpoint.RequestReplayFrom(1)); + CHECK(receiver.Single().Reason == DisconnectReason::ReplayRequested); + receiver.Feed(server::HelloOk()); + CHECK(!receiver.Status().Connected); + CHECK(receiver.Notices().empty()); + CHECK(receiver.Reconnect().SyncFrom == 1); +} + +void TestProtocolErrorsCloseTheConnection() +{ + flatbuffers::FlatBufferBuilder future(32); + const Bytes future_schema = server::Frame( + future, Payload::Ping, OpenUSDConnect::CreatePing(future).Union(), kSchemaVersion + 1); + flatbuffers::FlatBufferBuilder none(32); + const Bytes no_payload = server::Frame(none, Payload::NONE, 0); + flatbuffers::FlatBufferBuilder unknown(32); + const Bytes unknown_payload = server::Frame(unknown, static_cast(200), + OpenUSDConnect::CreatePing(unknown).Union()); + Bytes garbage; + const std::uint8_t noise[] = {1, 2, 3, 4, 5, 6, 7, 8}; + CHECK(EncodeFrame(noise, sizeof(noise), garbage) == FrameResult::Success); + const Bytes empty_header{0, 0, 0, 0}; + for (const Bytes& error : {future_schema, no_payload, unknown_payload, garbage, empty_header}) + { + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Feed(error); + CHECK(receiver.Single().Reason == DisconnectReason::ProtocolError); + CHECK(receiver.Reconnect().SyncFrom == 1); + } + // The server answers a Hello before it sends anything else. + for (const Bytes& early : {server::Ping(), server::Event(1)}) + { + Receiver receiver; + static_cast(receiver.Start()); + receiver.Feed(early); + CHECK(receiver.Single().Reason == DisconnectReason::ProtocolError); + CHECK(!receiver.Status().Connected && receiver.Status().QueuedFrames == 0); + } +} + +void TestHostDisconnectEndsTheSession() +{ + Receiver receiver; + static_cast(receiver.Start()); + receiver.Disconnect(DisconnectReason::TransportError); + CHECK(receiver.Notices().empty()); + receiver.Advance(1s); + static_cast(receiver.Connect()); + receiver.Feed(server::HelloOk()); + static_cast(receiver.Notices()); + receiver.Feed(server::ReplayComplete(0, 0)); + CHECK(receiver.Endpoint.MarkReplayApplied()); + receiver.Disconnect(DisconnectReason::TransportError); + CHECK(receiver.Notice().Reason == DisconnectReason::TransportError); + CHECK(!receiver.Status().Connected); + CHECK(!receiver.Status().Synchronized); + CHECK(receiver.Endpoint.NextWake()); + // A duplicate report for the same connection changes nothing. + receiver.Disconnect(); + CHECK(receiver.Commands().empty()); +} + +void TestStop() +{ + { + Receiver receiver; + static_cast(receiver.Handshake()); + static_cast(receiver.Notices()); + receiver.Endpoint.Stop(); + CHECK(receiver.Single().Reason == DisconnectReason::Stopped); + CHECK(receiver.Notice().Reason == DisconnectReason::Stopped); + CHECK(receiver.Status().Stopped && !receiver.Status().Connected); + receiver.Disconnect(); + CHECK(!receiver.Endpoint.Start(receiver.Now)); + receiver.Endpoint.Stop(); + CHECK(receiver.Commands().empty()); + } + { + Receiver receiver; + CHECK(receiver.Endpoint.Start(receiver.Now)); + // An attempt the host has not taken is dropped, so a loop never connects after Stop. + receiver.Endpoint.Stop(); + CHECK(receiver.Commands().empty()); + CHECK(receiver.Status().Stopped); + } + { + Receiver receiver; + CHECK(receiver.Endpoint.Start(receiver.Now)); + static_cast(receiver.Single()); + receiver.Endpoint.Stop(); + CHECK(receiver.Single().Reason == DisconnectReason::Stopped); + // A connect that completes after Stop is closed again, without a Hello. + receiver.Endpoint.OnConnected({}); + CHECK(receiver.Single().Reason == DisconnectReason::Stopped); + } + { + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Disconnect(); + CHECK(receiver.Endpoint.NextWake()); + receiver.Endpoint.Stop(); + CHECK(receiver.Commands().empty()); + CHECK(!receiver.Endpoint.NextWake()); + receiver.Advance(60s); + CHECK(receiver.Commands().empty()); + } +} + +// Replay identity. Each case follows a receiver scenario from the Python +// receiver's unit and integration tests. + +void TestReceivedIdentityIsPublishedOnlyWhenApplied() +{ + Receiver receiver; + server::Hello old_server; + old_server.ServerInstance = "old"; + old_server.ReplayEpoch.reset(); + static_cast(receiver.Handshake(old_server)); + receiver.Feed(server::ReplayComplete(0, 2)); + CHECK(receiver.Status().ServerInstance.empty()); + CHECK(receiver.Endpoint.MarkReplayApplied()); + CHECK(receiver.Status().ServerInstance == "old"); + + CHECK(receiver.Reconnect().Claims("old", 2)); + server::Hello new_server; + new_server.ServerInstance = "new"; + new_server.ReplayEpoch.reset(); + receiver.Feed(server::HelloOk(new_server)); + CHECK(receiver.Status().ServerInstance == "old"); + CHECK(!receiver.Status().Synchronized); + receiver.Feed(server::Resync()); + receiver.Feed(server::ReplayComplete(0, 0)); + CHECK(!receiver.Endpoint.MarkReplayApplied()); + CHECK(receiver.Endpoint.DrainFrames().size() == 1); + CHECK(receiver.Endpoint.MarkReplayApplied()); + CHECK(receiver.Status().ServerInstance == "new"); + CHECK(receiver.Status().ReplayEpoch == 0); +} + +void TestInterruptedLiveResetClaimsUnknownPrefix() +{ + Receiver receiver; + server::Hello hello; + hello.ReplayEpoch.reset(); + static_cast(receiver.Handshake(hello)); + receiver.Feed(server::ReplayComplete(0, 3)); + CHECK(receiver.Endpoint.MarkReplayApplied()); + receiver.Feed(server::Resync()); + CHECK(receiver.Reconnect().ClaimsUnknownPrefix()); + CHECK(!receiver.Endpoint.MarkReplayApplied()); +} + +void TestFullReplayRequestQueuesItsOwnReset() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Feed(server::LayerStack()); + receiver.Feed(server::Event(1)); + static_cast(receiver.Apply()); + CHECK(receiver.Reconnect().SyncFrom == 2); + CHECK(receiver.Endpoint.RequestReplayFrom(1)); + CHECK(receiver.Single().Reason == DisconnectReason::ReplayRequested); + CHECK(receiver.Reconnect().SyncFrom == 1); + receiver.Feed(server::HelloOk()); + receiver.Feed(server::LayerStack()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::ReplayComplete(1, 0)); + CHECK( + (Kinds(receiver.Apply()) == + std::vector{Payload::Resync, Payload::LayerStackState, Payload::BroadcastEvent})); + CHECK(receiver.Status().Synchronized); + + // The reset belongs to that one replay. + CHECK(receiver.Reconnect().SyncFrom == 2); + receiver.Feed(server::HelloOk()); + CHECK(receiver.Status().QueuedFrames == 0); +} + +void TestChangedHelloIdentityWaitsForTheReset(bool queue_full) +{ + ReceiverConfig config = TestConfig(); + config.MaxQueue = 1; + Receiver receiver(config); + static_cast(receiver.Handshake()); + receiver.Feed(server::Event(1)); + CHECK(receiver.Endpoint.DrainFrames().size() == 1); + const SentHello hello = receiver.Reconnect(); + CHECK(hello.SyncFrom == 2); + CHECK(hello.Claims("server", 0)); + + server::Hello changed; + changed.ReplayEpoch = 1; + receiver.Feed(server::HelloOk(changed)); + const ReceiverStatus status = receiver.Status(); + CHECK(status.LastSequence == 1); + CHECK(status.ServerInstance.empty()); + CHECK(!status.Synchronized); + if (queue_full) + { + receiver.Feed(server::LayerStack()); + } + receiver.Feed(server::Resync()); + CHECK(receiver.Status().LastSequence == (queue_full ? 1 : 0)); + if (queue_full) + { + CHECK(receiver.Single().Reason == DisconnectReason::QueueFull); + const SentHello retry = receiver.ReconnectAfterDrain(); + CHECK(retry.SyncFrom == 2); + CHECK(retry.Claims("server", 0)); + } + else + { + const SentHello retry = receiver.Reconnect(); + CHECK(retry.SyncFrom == 1); + CHECK(retry.Claims("server", 1)); + } +} + +void TestOldServerRetainsCursorWithoutIdentity() +{ + ReceiverConfig config = TestConfig(); + config.SyncFrom = 4; + config.LayeredReplay = false; + Receiver receiver(config); + static_cast(receiver.Start()); + server::Hello old_server; + old_server.ServerInstance.clear(); + old_server.ReplayIdentity = false; + old_server.ReplayEpoch.reset(); + old_server.LayeredReplay = false; + receiver.Feed(server::HelloOk(old_server)); + receiver.Feed(server::Event(4)); + receiver.Feed(server::ReplayComplete(4, 0)); + static_cast(receiver.Apply()); + CHECK(receiver.Status().Synchronized); + CHECK(receiver.Status().ServerInstance.empty()); + const SentHello retry = receiver.Reconnect(); + CHECK(retry.SyncFrom == 5); + CHECK(retry.ClaimsUnknownPrefix()); +} + +void TestInitialSnapshotCursorIsNotProof() +{ + ReceiverConfig config = TestConfig(); + config.SyncFrom = 2; + Receiver receiver(config); + static_cast(receiver.Handshake()); + receiver.Feed(server::Event(2)); + receiver.Feed(server::ReplayComplete(2, 0)); + static_cast(receiver.Apply()); + CHECK(receiver.Status().Synchronized); + CHECK(receiver.Status().ServerInstance.empty()); + + const SentHello retry = receiver.Reconnect(); + CHECK(retry.SyncFrom == 3); + CHECK(retry.ClaimsUnknownPrefix()); + receiver.Feed(server::HelloOk()); + receiver.Feed(server::Resync()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::Event(2)); + receiver.Feed(server::ReplayComplete(2, 0)); + static_cast(receiver.Apply()); + CHECK(receiver.Status().ServerInstance == "server"); +} + +// Feeds events from first until the queue overflows or last is queued. +std::int32_t FeedEvents(Receiver& receiver, std::int32_t first, std::int32_t last) +{ + for (std::int32_t sequence = first; sequence <= last; ++sequence) + { + receiver.Feed(server::Event(sequence)); + if (!receiver.Status().Connected) + { + CHECK(receiver.Single().Reason == DisconnectReason::QueueFull); + return sequence; + } + } + return last + 1; +} + +enum class OverflowStart +{ + FirstReplay, + AfterServerReset, + AfterLiveReset, + SnapshotCursor, + FullReplayRequest, +}; + +// A replay that overflows the queue resumes where the queue stopped, claims the +// replay's identity, and advances on every reconnect until it completes. +void TestQueueOverflowResumesTheReplay(OverflowStart start) +{ + const bool snapshot = + start == OverflowStart::SnapshotCursor || start == OverflowStart::FullReplayRequest; + const bool layered = start == OverflowStart::FullReplayRequest; + ReceiverConfig config = TestConfig(); + config.SyncFrom = snapshot ? 4 : 1; + config.LayeredReplay = layered || !snapshot; + config.MaxQueue = snapshot ? (layered ? 2 : 1) : 3; + const std::int32_t last = snapshot ? 6 : 8; + Receiver receiver(config); + server::Hello hello; + hello.LayeredReplay = config.LayeredReplay; + const auto accept = [&](bool reset) + { + receiver.Feed(server::HelloOk(hello)); + if (reset) + { + receiver.Feed(server::Resync()); + } + if (layered) + { + receiver.Feed(server::LayerStack()); + } + }; + static_cast(receiver.Start()); + accept(false); + if (start == OverflowStart::AfterServerReset || start == OverflowStart::AfterLiveReset) + { + receiver.Feed(server::Event(1)); + receiver.Feed(server::ReplayComplete(1, 0)); + static_cast(receiver.Apply()); + hello.ReplayEpoch = 1; + } + if (start == OverflowStart::AfterServerReset) + { + static_cast(receiver.Reconnect()); + accept(true); + } + else if (start == OverflowStart::AfterLiveReset) + { + // A live Resync names no epoch, so overflowing it needs one fresh reset. + receiver.Feed(server::Resync()); + CHECK(FeedEvents(receiver, 1, last) == 3); + CHECK(receiver.ReconnectAfterDrain().ClaimsUnknownPrefix()); + accept(true); + } + else if (start == OverflowStart::SnapshotCursor) + { + // The snapshot prefix is never proven, so its first overflow needs a reset. + CHECK(FeedEvents(receiver, 4, last) == 5); + CHECK(receiver.ReconnectAfterDrain().ClaimsUnknownPrefix()); + accept(true); + } + else if (start == OverflowStart::FullReplayRequest) + { + receiver.Feed(server::ReplayComplete(3, 0)); + static_cast(receiver.Apply()); + CHECK(receiver.Endpoint.RequestReplayFrom(1)); + static_cast(receiver.Single()); + static_cast(receiver.Reconnect()); + accept(false); + } + + std::vector progress; + std::int32_t next = 1; + while ((next = FeedEvents(receiver, next, last)) <= last) + { + const SentHello resumed = receiver.ReconnectAfterDrain(); + progress.push_back(receiver.Applied); + CHECK(resumed.SyncFrom == next); + CHECK(resumed.Claims("server", *hello.ReplayEpoch)); + accept(false); + } + receiver.Feed(server::ReplayComplete(last, *hello.ReplayEpoch)); + const std::vector kinds = Kinds(receiver.Apply()); + CHECK(std::find(kinds.begin(), kinds.end(), Payload::Resync) == kinds.end()); + CHECK(progress.size() > 1); + CHECK(std::adjacent_find(progress.begin(), progress.end(), std::greater_equal<>()) == + progress.end()); + CHECK(receiver.Applied == last); + const ReceiverStatus status = receiver.Status(); + CHECK(status.Synchronized && status.ServerInstance == "server"); + CHECK(status.ReplayEpoch == *hello.ReplayEpoch); +} + +// After a live Resync the applied cursor may still count the old sequence +// domain. A replay from it must claim the applied identity, so the server +// resets rather than resuming the new domain from an old-domain cursor. +void TestOldDomainCursorReplayClaimsTheAppliedIdentity() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + for (std::int32_t sequence = 1; sequence <= 3; ++sequence) + { + receiver.Feed(server::Event(sequence)); + } + receiver.Feed(server::ReplayComplete(3, 0)); + static_cast(receiver.Apply()); + receiver.Feed(server::Event(4)); + receiver.Feed(server::Event(5)); + receiver.Feed(server::Resync()); + for (std::int32_t sequence = 1; sequence <= 5; ++sequence) + { + receiver.Feed(server::Event(sequence)); + } + receiver.Feed(server::ReplayComplete(5, 1)); + + const std::uint64_t generation = receiver.Endpoint.Generation(); + CHECK(Sequences(receiver.Endpoint.DrainFrames(1)) == std::vector{4}); + CHECK(receiver.Endpoint.MarkAppliedThrough(generation, 4)); + receiver.Feed(server::Event(7)); + CHECK(receiver.Single().Reason == DisconnectReason::SequenceGap); + const SentHello hello = receiver.Reconnect(); + CHECK(hello.SyncFrom == 5); + CHECK(hello.Claims("server", 0)); +} + +// The consumer applied the replay, but a reconnect during its batch rejected +// MarkAppliedThrough and the next connection delivers no newer frame. +void TestReplayAppliesAcrossAMidBatchReconnect() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::Event(2)); + receiver.Feed(server::ReplayComplete(2, 0)); + const std::uint64_t generation = receiver.Endpoint.Generation(); + CHECK(receiver.Endpoint.DrainFrames().size() == 2); + + const SentHello hello = receiver.Reconnect(); + CHECK(hello.SyncFrom == 3 && hello.Claims("server", 0)); + receiver.Feed(server::HelloOk()); + receiver.Feed(server::ReplayComplete(2, 0)); + CHECK(!receiver.Endpoint.MarkAppliedThrough(generation, 2)); + CHECK(receiver.Endpoint.MarkReplayApplied()); + const ReceiverStatus status = receiver.Status(); + CHECK(status.Synchronized); + CHECK(status.LastAppliedSequence == 2); +} + +enum class ReplayCause +{ + ConsumerFailure, + SequenceGap, + ReplayCompleteGap, +}; + +// A replay request during the connection's first replay keeps the handshake +// identity it proved, so the server resumes from the applied cursor. A marker +// beyond the received records is such a gap, or marking it would skip them. +void TestReplayRequestWithoutResetResumes(ReplayCause cause) +{ + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::Event(2)); + const std::uint64_t generation = receiver.Endpoint.Generation(); + CHECK(Sequences(receiver.Endpoint.DrainFrames(1)) == std::vector{1}); + CHECK(receiver.Endpoint.MarkAppliedThrough(generation, 1)); + if (cause == ReplayCause::ConsumerFailure) + { + CHECK(Sequences(receiver.Endpoint.DrainFrames(1)) == std::vector{2}); + CHECK(receiver.Endpoint.RequestReplayFrom(2)); + CHECK(receiver.Single().Reason == DisconnectReason::ReplayRequested); + } + else + { + receiver.Feed(cause == ReplayCause::SequenceGap ? server::Event(4) + : server::ReplayComplete(3, 0)); + CHECK(receiver.Single().Reason == DisconnectReason::SequenceGap); + } + if (cause == ReplayCause::ReplayCompleteGap) + { + CHECK(!receiver.Endpoint.MarkReplayApplied()); + } + const SentHello hello = receiver.Reconnect(); + CHECK(hello.SyncFrom == 2); + CHECK(hello.Claims("server", 0)); + CHECK(ServerResumes(hello, "server", 0, 3)); +} + +enum class ResetState +{ + Queued, + Drained, + Applied, +}; + +// A live reset that is still queued or unapplied separates the stage from the +// received identity, so the claim falls back to the applied replay. +void TestReplayRequestWithPendingResetClaimsTheAppliedReplay(ResetState state) +{ + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::ReplayComplete(1, 0)); + static_cast(receiver.Apply()); + receiver.Feed(server::Resync()); + receiver.Feed(server::Event(1)); + receiver.Feed(server::ReplayComplete(1, 1)); + if (state != ResetState::Queued) + { + CHECK((Kinds(receiver.Endpoint.DrainFrames(1)) == + std::vector{Payload::Resync})); + } + if (state == ResetState::Applied) + { + receiver.Endpoint.ResetAppliedProgress(); + } + CHECK(receiver.Endpoint.RequestReplayFrom(1)); + static_cast(receiver.Single()); + CHECK(receiver.Reconnect().Claims("server", state == ResetState::Applied ? 1 : 0)); +} + +// The reset the receiver queues ahead of a replay from one is pending too. +void TestOwnQueuedResetIsPending() +{ + Receiver receiver; + static_cast(receiver.Handshake()); + receiver.Feed(server::Event(1)); + static_cast(receiver.Apply()); + CHECK(receiver.Endpoint.RequestReplayFrom(1)); + static_cast(receiver.Single()); + CHECK(receiver.Reconnect().SyncFrom == 1); + receiver.Feed(server::HelloOk()); + CHECK(receiver.Status().QueuedFrames == 1); + CHECK(receiver.Endpoint.RequestReplayFrom(1)); + static_cast(receiver.Single()); + CHECK(receiver.Reconnect().ClaimsUnknownPrefix()); +} + +} // namespace + +int main() +{ + TestConfigurationValidation(); + TestStartConnectsOnceAndSendsTheConfiguredHello(); + TestAcceptedHelloNotifiesInOrder(); + TestHandshakeRejectionsStop(); + TestNegotiatedModesConnect(); + TestControlMessages(); + TestReplayCompleteWaitsForDrainedFrames(); + TestInPlaceResyncClearsReadyUntilApplied(); + TestDataFrameAcceptResults(); + for (const DrainOutcome outcome : + {DrainOutcome::DrainedEarlier, DrainOutcome::DrainedDuringWait, DrainOutcome::TimedOut}) + { + TestOverflowWaitsForTheDrain(outcome); + } + TestConsecutiveReadTimeouts(); + TestBackoffDoublesAndResetsAfterConnectedSession(); + TestReconnectDisabledStops(); + TestReconnectToggleAppliesWhenTheSessionEnds(); + TestReplayRequests(); + TestProtocolErrorsCloseTheConnection(); + TestHostDisconnectEndsTheSession(); + TestStop(); + TestReceivedIdentityIsPublishedOnlyWhenApplied(); + TestInterruptedLiveResetClaimsUnknownPrefix(); + TestFullReplayRequestQueuesItsOwnReset(); + for (const bool queue_full : {false, true}) + { + TestChangedHelloIdentityWaitsForTheReset(queue_full); + } + TestOldServerRetainsCursorWithoutIdentity(); + TestInitialSnapshotCursorIsNotProof(); + for (const OverflowStart start : {OverflowStart::FirstReplay, OverflowStart::AfterServerReset, + OverflowStart::AfterLiveReset, OverflowStart::SnapshotCursor, + OverflowStart::FullReplayRequest}) + { + TestQueueOverflowResumesTheReplay(start); + } + TestOldDomainCursorReplayClaimsTheAppliedIdentity(); + TestReplayAppliesAcrossAMidBatchReconnect(); + for (const ReplayCause cause : + {ReplayCause::ConsumerFailure, ReplayCause::SequenceGap, ReplayCause::ReplayCompleteGap}) + { + TestReplayRequestWithoutResetResumes(cause); + } + for (const ResetState state : {ResetState::Queued, ResetState::Drained, ResetState::Applied}) + { + TestReplayRequestWithPendingResetClaimsTheAppliedReplay(state); + } + TestOwnQueuedResetIsPending(); + return 0; +} diff --git a/tests/native/test_replay_identity.cpp b/tests/native/test_replay_identity.cpp deleted file mode 100644 index 6bd2bc1..0000000 --- a/tests/native/test_replay_identity.cpp +++ /dev/null @@ -1,157 +0,0 @@ -#include "openusdconnect/client/replay_identity.h" -#include "openusdconnect/client/receiver_session.h" - -#include "test_check.h" - -using namespace openusdconnect::client; - -static void TestQueuedResync(bool empty_replay) -{ - OrderedReceiverSession session(1, 8, true); - ReceiverReplayIdentity identity; - const auto connection = session.BeginConnection(); - CHECK(!identity.BeginConnection()); - identity.AcceptHello(1, true, "server", 0); - CHECK(session.Accept(connection.Generation, ReceiverMessageKind::Resync, 0, 0) == - AcceptResult::Accepted); - identity.AcceptResync(); - if (!empty_replay) - { - CHECK(session.Accept(connection.Generation, ReceiverMessageKind::Event, 1, 1) == - AcceptResult::Accepted); - } - const int head = empty_replay ? 0 : 1; - CHECK(session.AcceptReplayComplete(connection.Generation, head, 1) == AcceptResult::Accepted); - identity.AcceptReplayComplete(1); - CHECK(!session.TryMarkReplayApplied()); - int frame = -1; - CHECK(session.TryPop(frame) && frame == 0); - session.ResetAppliedProgress(); - CHECK(!session.Synchronized()); - if (!empty_replay) - { - CHECK(!session.TryMarkReplayApplied()); - CHECK(session.TryPop(frame) && frame == 1); - CHECK(!session.TryMarkReplayApplied()); - CHECK(session.MarkAppliedThrough(connection.Generation, 1)); - } - CHECK(session.TryMarkReplayApplied()); - identity.MarkReplayApplied(); - CHECK(session.Synchronized()); - CHECK(session.ReplayEpoch() == 1); - CHECK((identity.Applied() == ReplayIdentity{"server", 1})); -} - -static void TestInterruptedReplay(bool apply_failure) -{ - OrderedReceiverSession session(1, 8, true); - ReceiverReplayIdentity identity; - auto connection = session.BeginConnection(); - CHECK(!identity.BeginConnection()); - identity.AcceptHello(1, true, "server", 0); - CHECK(session.Accept(connection.Generation, ReceiverMessageKind::Resync, 0, 0) == - AcceptResult::Accepted); - identity.AcceptResync(); - CHECK(session.Accept(connection.Generation, ReceiverMessageKind::Event, 1, 1) == - AcceptResult::Accepted); - CHECK(session.AcceptReplayComplete(connection.Generation, 1, 1) == AcceptResult::Accepted); - identity.AcceptReplayComplete(1); - if (apply_failure) - { - int frame; - CHECK(session.TryPop(frame) && frame == 0); - session.ResetAppliedProgress(); - CHECK(session.TryPop(frame) && frame == 1); - // Applying the event failed; never advance the applied cursor. - CHECK(!session.TryMarkReplayApplied()); - CHECK(session.RequestReplayFrom(session.LastAppliedSequence() + 1)); - } - session.Disconnect(connection.Generation); - CHECK(!session.TryMarkReplayApplied()); - CHECK(!identity.Applied()); - connection = session.BeginConnection(); - const auto claim = identity.BeginConnection(); - CHECK(claim && !claim->IsKnown()); - CHECK(!identity.Pending()); - identity.AcceptHello(connection.SyncFrom, true, "server", 1); - // The server resets the unknown prefix before replaying the complete scene. - CHECK(session.Accept(connection.Generation, ReceiverMessageKind::Resync, 0, 0) == - AcceptResult::Accepted); - identity.AcceptResync(); - CHECK(session.Accept(connection.Generation, ReceiverMessageKind::Event, 1, 1) == - AcceptResult::Accepted); - CHECK(session.AcceptReplayComplete(connection.Generation, 1, 1) == AcceptResult::Accepted); - identity.AcceptReplayComplete(1); - int frame; - while (session.TryPop(frame)) - { - if (frame == 0) - session.ResetAppliedProgress(); - else - CHECK(session.MarkAppliedThrough(connection.Generation, frame)); - } - CHECK(session.TryMarkReplayApplied()); - identity.MarkReplayApplied(); - CHECK((identity.Applied() == ReplayIdentity{"server", 1})); -} - -int main() -{ - TestQueuedResync(false); - TestQueuedResync(true); - TestInterruptedReplay(false); - TestInterruptedReplay(true); - ReceiverReplayIdentity initial_replay; - CHECK(!initial_replay.BeginConnection()); - initial_replay.AcceptHello(1, true, "server-a", 0); - CHECK(initial_replay.IsConnectionPrefixProven()); - initial_replay.AcceptReplayComplete(0); - CHECK((initial_replay.Pending() == ReplayIdentity{"server-a", 0})); - initial_replay.MarkReplayApplied(); - CHECK((initial_replay.Applied() == ReplayIdentity{"server-a", 0})); - - const ReplayPrefixClaim matching_claim = initial_replay.BeginConnection(); - CHECK(matching_claim); - CHECK(matching_claim->IsKnown()); - CHECK(matching_claim->ServerInstance() == "server-a"); - CHECK(matching_claim->Epoch() == 0); - initial_replay.AcceptHello(2, true, "server-a", 0); - CHECK(initial_replay.IsConnectionPrefixProven()); - initial_replay.AcceptReplayComplete(0); - initial_replay.MarkReplayApplied(); - CHECK((initial_replay.Applied() == ReplayIdentity{"server-a", 0})); - - const ReplayPrefixClaim stale_claim = initial_replay.BeginConnection(); - CHECK(stale_claim && stale_claim->Epoch() == 0); - initial_replay.AcceptHello(2, true, "server-b", 1); - CHECK(!initial_replay.IsConnectionPrefixProven()); - initial_replay.AcceptReplayComplete(1); - CHECK(!initial_replay.Pending()); - initial_replay.AcceptResync(); - CHECK(initial_replay.IsConnectionPrefixProven()); - initial_replay.AcceptReplayComplete(1); - CHECK((initial_replay.Pending() == ReplayIdentity{"server-b", 1})); - initial_replay.MarkReplayApplied(); - CHECK((initial_replay.Applied() == ReplayIdentity{"server-b", 1})); - - ReceiverReplayIdentity external_prefix; - CHECK(!external_prefix.BeginConnection()); - external_prefix.AcceptHello(4, true, "server-a", 0); - CHECK(!external_prefix.IsConnectionPrefixProven()); - external_prefix.AcceptReplayComplete(0); - external_prefix.MarkReplayApplied(); - CHECK(!external_prefix.Applied()); - const ReplayPrefixClaim unknown_claim = external_prefix.BeginConnection(); - CHECK(unknown_claim); - CHECK(!unknown_claim->IsKnown()); - CHECK(unknown_claim->ServerInstance().empty()); - CHECK(!unknown_claim->Epoch()); - - ReceiverReplayIdentity legacy_server; - CHECK(!legacy_server.BeginConnection()); - legacy_server.AcceptHello(1, false, {}, std::nullopt); - legacy_server.AcceptReplayComplete(0); - legacy_server.MarkReplayApplied(); - CHECK(!legacy_server.Applied()); - return 0; -} diff --git a/tests/native/test_tcp_socket.cpp b/tests/native/test_tcp_socket.cpp new file mode 100644 index 0000000..cde8d37 --- /dev/null +++ b/tests/native/test_tcp_socket.cpp @@ -0,0 +1,214 @@ +#include "openusdconnect/client/driver/socket.h" + +#include "test_check.h" + +#include +#include +#include +#include +#include +#include + +#ifdef _WIN32 +#include +#include +#else +#include +#include +#include +#include +#endif + +using namespace openusdconnect::client; +using namespace std::chrono_literals; + +namespace +{ + +// Long enough that reaching it means an interrupt was missed. +constexpr std::chrono::seconds kPatience{5}; + +#ifdef _WIN32 +using NativeSocket = SOCKET; + +void CloseNative(NativeSocket socket) +{ + closesocket(socket); +} +#else +using NativeSocket = int; + +void CloseNative(NativeSocket socket) +{ + close(socket); +} +#endif + +// The server side of the test connections, on the loopback interface. Create a +// TcpSocketFactory first, which initializes Winsock. +class Listener final +{ +public: + Listener() + : Handle(socket(AF_INET, SOCK_STREAM, 0)) + { + sockaddr_in address{}; + address.sin_family = AF_INET; + address.sin_addr.s_addr = htonl(INADDR_LOOPBACK); + CHECK(bind(Handle, reinterpret_cast(&address), sizeof(address)) == 0); + CHECK(listen(Handle, 4) == 0); + socklen_t length = sizeof(address); + CHECK(getsockname(Handle, reinterpret_cast(&address), &length) == 0); + Port = ntohs(address.sin_port); + } + + ~Listener() + { + Close(); + } + + Listener(const Listener&) = delete; + Listener& operator=(const Listener&) = delete; + + [[nodiscard]] NativeSocket Accept() const + { + return accept(Handle, nullptr, nullptr); + } + + void Close() + { + if (Open) + { + CloseNative(Handle); + Open = false; + } + } + + std::uint16_t Port = 0; + +private: + NativeSocket Handle; + bool Open = true; +}; + +[[nodiscard]] TimePoint Deadline(std::chrono::milliseconds wait = kPatience) +{ + return std::chrono::steady_clock::now() + wait; +} + +void TestSendReceiveAndTimeout() +{ + TcpSocketFactory factory; + Listener listener; + const std::unique_ptr client = factory.Create(); + CHECK(client->Connect("127.0.0.1", listener.Port, Deadline()) == SocketResult::Success); + const NativeSocket server = listener.Accept(); + CHECK(send(server, "abc", 3, 0) == 3); + std::uint8_t buffer[16]; + std::size_t received = 0; + CHECK(client->Receive(buffer, sizeof(buffer), Deadline(), received) == SocketResult::Success); + CHECK(received == 3 && buffer[0] == 'a'); + const std::uint8_t reply[] = {'h', 'i'}; + CHECK(client->SendAll(reply, sizeof(reply), Deadline()) == SocketResult::Success); + char echoed[2]; + CHECK(recv(server, echoed, 2, 0) == 2); + CHECK(client->Receive(buffer, sizeof(buffer), Deadline(50ms), received) == + SocketResult::Timeout); + CloseNative(server); + CHECK(client->Receive(buffer, sizeof(buffer), Deadline(), received) == SocketResult::Closed); +} + +// Wake ends one Receive and is consumed; Interrupt ends every later call too. +void TestWakeAndInterruptEndBlockedCalls() +{ + TcpSocketFactory factory; + Listener listener; + const std::unique_ptr client = factory.Create(); + CHECK(client->Connect("127.0.0.1", listener.Port, Deadline()) == SocketResult::Success); + const NativeSocket server = listener.Accept(); + std::uint8_t buffer[16]; + std::size_t received = 0; + + std::thread waker( + [&] + { + std::this_thread::sleep_for(50ms); + client->Wake(); + }); + CHECK(client->Receive(buffer, sizeof(buffer), Deadline(), received) == + SocketResult::Interrupted); + waker.join(); + CHECK(client->Receive(buffer, sizeof(buffer), Deadline(50ms), received) == + SocketResult::Timeout); + client->Wake(); + CHECK(client->Receive(buffer, sizeof(buffer), Deadline(), received) == + SocketResult::Interrupted); + + std::thread interrupter( + [&] + { + std::this_thread::sleep_for(50ms); + client->Interrupt(); + }); + CHECK(client->Receive(buffer, sizeof(buffer), std::nullopt, received) == + SocketResult::Interrupted); + interrupter.join(); + CHECK(client->Receive(buffer, sizeof(buffer), Deadline(), received) == + SocketResult::Interrupted); + const std::uint8_t byte = 0; + CHECK(client->SendAll(&byte, 1, Deadline()) == SocketResult::Interrupted); + CloseNative(server); +} + +// A peer that stops reading stalls the writer, which gives up at its deadline. +// Winsock accepts a whole send while its buffer has room, so the writer keeps sending. +void TestStalledWriteTimesOut() +{ + TcpSocketFactory factory; + Listener listener; + const std::unique_ptr client = factory.Create(); + CHECK(client->Connect("127.0.0.1", listener.Port, Deadline()) == SocketResult::Success); + const NativeSocket server = listener.Accept(); + const std::vector chunk(1024 * 1024); + SocketResult result = SocketResult::Success; + auto started = std::chrono::steady_clock::now(); + for (int chunks = 0; chunks < 1024 && result == SocketResult::Success; ++chunks) + { + started = std::chrono::steady_clock::now(); + result = client->SendAll(chunk.data(), chunk.size(), Deadline(100ms)); + } + CHECK(result == SocketResult::Timeout); + CHECK(std::chrono::steady_clock::now() - started < kPatience); + CloseNative(server); +} + +// Winsock retries a refused loopback connect for seconds; Interrupt ends it. +void TestConnectEndsWhenRefusedOrInterrupted() +{ + TcpSocketFactory factory; + Listener listener; + const std::uint16_t port = listener.Port; + listener.Close(); + const std::unique_ptr client = factory.Create(); + std::thread interrupter( + [&] + { + std::this_thread::sleep_for(100ms); + client->Interrupt(); + }); + const SocketResult result = client->Connect("127.0.0.1", port, Deadline()); + interrupter.join(); + CHECK(result == SocketResult::Interrupted || + (result == SocketResult::Failed && client->SystemError() != 0)); +} + +} // namespace + +int main() +{ + TestSendReceiveAndTimeout(); + TestWakeAndInterruptEndBlockedCalls(); + TestStalledWriteTimesOut(); + TestConnectEndsWhenRefusedOrInterrupted(); + return 0; +} diff --git a/tests/native/test_threaded_producer_driver.cpp b/tests/native/test_threaded_producer_driver.cpp new file mode 100644 index 0000000..b50549b --- /dev/null +++ b/tests/native/test_threaded_producer_driver.cpp @@ -0,0 +1,331 @@ +#include "openusdconnect/client/driver/testing/scripted_socket.h" +#include "openusdconnect/client/driver/threaded_producer_driver.h" + +#include "driver_recorder.h" +#include "producer_frames.h" +#include "test_check.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include + +using namespace driver_test; +using namespace producer_test; +using namespace std::chrono_literals; + +namespace +{ + +[[nodiscard]] TimePoint Now() +{ + return std::chrono::steady_clock::now(); +} + +class Harness final : public DriverHarness +{ +public: + explicit Harness(const ProducerConfig& config = TestConfig(), DriverCallbacks callbacks = {}) + : DriverHarness(config, std::move(callbacks)) + { + } + + [[nodiscard]] std::future ConnectAsync() + { + return std::async(std::launch::async, + [this] + { + return Driver->Connect(kPatience); + }); + } + + [[nodiscard]] std::future FlushAsync() + { + return std::async(std::launch::async, + [this] + { + return Driver->Flush(kPatience); + }); + } + + // Starts the loop if needed and connects; returns the accepted connection. + [[nodiscard]] std::shared_ptr Handshake(const server::Hello& hello = {}) + { + if (!Driver->Running()) + { + CHECK(Driver->Start()); + } + std::future connected = ConnectAsync(); + std::shared_ptr connection = DriverHarness::Handshake(hello); + CHECK(connected.get()); + return connection; + } + + // Appends the next transaction as a host thread does, and returns its frame. + Bytes Submit(std::string_view prim) + { + const std::uint64_t id = Endpoint.NextTransactionId(); + Bytes frame = TransactionFrame(id, prim); + CHECK(Endpoint.Append(id, frame, 1, "") == ProducerResult::Accepted); + Driver->Wake(); + return frame; + } +}; + +// The frames a connection carried after its Hello. +[[nodiscard]] std::vector AfterHello(const ScriptedConnection& connection) +{ + std::vector frames = SplitFrames(connection.Sent()); + CHECK(!frames.empty() && DecodeSent(frames.front()).payload_type() == Payload::Hello); + frames.erase(frames.begin()); + return frames; +} + +void TestRequestConnectThenHandshake() +{ + DriverCallbacks callbacks; + callbacks.Token = [] + { + return std::optional("token-1"); + }; + Harness harness(TestConfig(), std::move(callbacks)); + CHECK(!harness.Driver->Running() && !harness.Driver->Stopped()); + CHECK(harness.Driver->Start()); + CHECK(harness.Endpoint.RequestConnect(Now(), Now() + kPatience)); + harness.Driver->Wake(); + + const std::shared_ptr connection = harness.Accept(); + const SentHello hello = DecodeHello(connection->Sent()); + CHECK(hello.Role == "emitter" && hello.Token == "token-1"); + CHECK(hello.ProducerSessionId == "session"); + server::Hello accepted; + accepted.Token = "issued"; + CHECK(connection->Deliver(server::HelloOk(accepted))); + // Connect waits for the requested attempt instead of starting its own. + CHECK(harness.Driver->Connect(kPatience)); + CHECK(harness.Sockets->Attempts() == 1); + CHECK(Eventually( + [&harness] + { + return harness.Record.Count() == 1; + })); + CHECK(harness.Record.IssuedToken() == "issued"); +} + +void TestConnectReturnsOnceTheAttemptEnds() +{ + { + Harness harness; + // Without a running loop nothing would apply the attempt. + CHECK(!harness.Driver->Connect(kPatience)); + CHECK(harness.Driver->Start()); + std::future refused = harness.ConnectAsync(); + CHECK(harness.Sockets->Refuse(kPatience, kRefused)); + CHECK(!refused.get()); + const std::optional failure = harness.Driver->LastFailure(); + CHECK(failure && failure->Operation == SocketOperation::Connect); + CHECK(failure->SystemError == kRefused); + + CHECK(!harness.Driver->Connect(0ms)); + static_cast(harness.Handshake()); + CHECK(harness.Driver->Connect(0ms)); + } + { + // A server that never answers ends the attempt at the handshake timeout. + ProducerConfig config = TestConfig(); + config.HandshakeTimeout = 50ms; + Harness harness(config); + CHECK(harness.Driver->Start()); + const TimePoint started = Now(); + std::future silent = harness.ConnectAsync(); + const std::shared_ptr connection = harness.Accept(); + CHECK(!silent.get()); + CHECK(Now() - started >= 50ms && Now() - started < kPatience); + CHECK(connection->WaitClosed(kPatience)); + CHECK(!harness.Endpoint.Status().Connected); + } +} + +void TestFlushReconnectsAndReplays() +{ + Harness harness; + const std::shared_ptr first = harness.Handshake(); + const std::size_t hello_size = first->Sent().size(); + const Bytes frame = harness.Submit("/A"); + CHECK(first->WaitSent(hello_size + frame.size(), kPatience)); + first->Close(); + CHECK(first->WaitClosed(kPatience)); + // A zero timeout only reports. + CHECK(harness.Driver->Flush(0ms) == FlushResult::Unfinished); + CHECK(harness.Sockets->Attempts() == 1); + + std::future flushed = harness.FlushAsync(); + const std::shared_ptr second = harness.Accept(); + CHECK(second->Deliver(server::HelloOk())); + CHECK(second->WaitSent(hello_size + frame.size(), kPatience)); + CHECK(AfterHello(*second) == std::vector{frame}); + CHECK(second->Deliver(server::Acknowledged(1))); + CHECK(flushed.get() == FlushResult::Flushed); +} + +void TestFlushWaitsOutTheRateLimit() +{ + Harness harness; + const std::shared_ptr first = harness.Handshake(); + const std::size_t hello_size = first->Sent().size(); + const Bytes frame = harness.Submit("/A"); + CHECK(first->Deliver(server::RateLimited(0.2F))); + CHECK(first->WaitClosed(kPatience)); + const TimePoint closed = Now(); + + std::future flushed = harness.FlushAsync(); + const std::shared_ptr second = harness.Accept(); + CHECK(Now() - closed >= 150ms); + CHECK(second->Deliver(server::HelloOk())); + CHECK(second->WaitSent(hello_size + frame.size(), kPatience)); + CHECK(AfterHello(*second) == std::vector{frame}); + CHECK(second->Deliver(server::Acknowledged(1))); + CHECK(flushed.get() == FlushResult::Flushed); +} + +void TestStalledWriteClosesAtTheSendDeadline() +{ + ProducerConfig config = TestConfig(); + config.HandshakeTimeout = 250ms; + Harness harness(config); + const std::shared_ptr connection = harness.Handshake(); + connection->StallSends(); + const TimePoint stalled = Now(); + static_cast(harness.Submit("/A")); + CHECK(connection->WaitClosed(kPatience)); + CHECK(Now() - stalled >= 250ms); + CHECK(Eventually( + [&harness] + { + return harness.Record.Count() == 1; + })); + const std::optional failure = harness.Driver->LastFailure(); + CHECK(failure && failure->Operation == SocketOperation::Send); + CHECK(failure->Result == SocketResult::Timeout); + CHECK(harness.Record.All().front().Reason == DisconnectReason::TransportError); + const ProducerStatus status = harness.Endpoint.Status(); + CHECK(!status.Connected && status.PendingTransactions == 1); +} + +void TestStopInterruptsBlockingCalls() +{ + { + Harness harness; + static_cast(harness.Handshake()); + static_cast(harness.Submit("/A")); + harness.Driver->Stop(); + CHECK(harness.Driver->Join(kPatience)); + CHECK(!harness.Driver->Connect(kPatience)); + CHECK(harness.Driver->Flush(kPatience) == FlushResult::Unfinished); + } + { + // A connect the server never completes. + Harness harness; + CHECK(harness.Driver->Start()); + std::future pending = harness.ConnectAsync(); + CHECK(Eventually( + [&harness] + { + return harness.Sockets->Attempts() == 1; + })); + harness.Driver->Stop(); + CHECK(harness.Driver->Join(kPatience)); + CHECK(!pending.get()); + } +} + +void TestCloseWritesWhatWasQueued() +{ + { + Harness harness; + const std::shared_ptr connection = harness.Handshake(); + const Bytes frame = harness.Submit("/A"); + CHECK(harness.Driver->Close(kPatience)); + CHECK(connection->ClosedByClient()); + const std::vector frames = AfterHello(*connection); + CHECK(frames.size() == 2 && frames.front() == frame); + CHECK(DecodeSent(frames.back()).payload_type() == Payload::Quit); + } + { + // A peer that stopped reading cannot hold Close past its bound. + ProducerConfig config = TestConfig(); + config.HandshakeTimeout = kPatience; + Harness harness(config); + const std::shared_ptr connection = harness.Handshake(); + connection->StallSends(); + static_cast(harness.Submit("/A")); + const TimePoint closing = Now(); + static_cast(harness.Driver->Close(100ms)); + CHECK(Now() - closing < kPatience / 5); + CHECK(harness.Driver->Join(kPatience)); + CHECK(connection->ClosedByClient()); + } +} + +void TestRejectedTransactionSurfacesThroughFailure() +{ + Harness harness; + const std::shared_ptr connection = harness.Handshake(); + static_cast(harness.Submit("/A")); + std::future flushed = harness.FlushAsync(); + CHECK(connection->Deliver( + server::Rejected(1, TransactionRejectionCode::InvalidTransaction, "bad edit"))); + CHECK(flushed.get() == FlushResult::RecoveryRequired); + CHECK(connection->WaitClosed(kPatience)); + const std::optional failure = harness.Endpoint.Failure(); + CHECK(failure && failure->TransactionId == 1 && failure->Reason == "bad edit"); + // Recovery comes first, so nothing reconnects. + CHECK(!harness.Driver->Connect(kPatience)); + CHECK(harness.Driver->Flush(kPatience) == FlushResult::RecoveryRequired); + CHECK(harness.Sockets->Attempts() == 1); +} + +// The loop thread cannot wait for itself, so its blocking calls only report. +void TestCallbacksCannotBlockOnTheLoop() +{ + ThreadedProducerDriver* driver = nullptr; + std::promise> reported; + DriverCallbacks callbacks; + callbacks.Notifications = [&](Notification notification) + { + if (std::holds_alternative(notification)) + { + reported.set_value({driver->Connect(kPatience), driver->Flush(kPatience)}); + } + }; + Harness harness(TestConfig(), std::move(callbacks)); + driver = harness.Driver.get(); + const std::shared_ptr connection = harness.Handshake(); + static_cast(harness.Submit("/A")); + connection->Close(); + std::future> result = reported.get_future(); + CHECK(result.wait_for(kPatience / 2) == std::future_status::ready); + const auto [connected, flushed] = result.get(); + CHECK(!connected && flushed == FlushResult::Unfinished); +} + +} // namespace + +int main() +{ + TestRequestConnectThenHandshake(); + TestConnectReturnsOnceTheAttemptEnds(); + TestFlushReconnectsAndReplays(); + TestFlushWaitsOutTheRateLimit(); + TestStalledWriteClosesAtTheSendDeadline(); + TestStopInterruptsBlockingCalls(); + TestCloseWritesWhatWasQueued(); + TestRejectedTransactionSurfacesThroughFailure(); + TestCallbacksCannotBlockOnTheLoop(); + return 0; +} diff --git a/tests/native/test_threaded_receiver_driver.cpp b/tests/native/test_threaded_receiver_driver.cpp new file mode 100644 index 0000000..708b2b0 --- /dev/null +++ b/tests/native/test_threaded_receiver_driver.cpp @@ -0,0 +1,260 @@ +#include "openusdconnect/client/driver/testing/scripted_socket.h" +#include "openusdconnect/client/driver/threaded_receiver_driver.h" + +#include "driver_recorder.h" +#include "receiver_frames.h" +#include "test_check.h" + +#include +#include +#include +#include +#include +#include +#include +#include + +using namespace driver_test; +using namespace receiver_test; +using namespace std::chrono_literals; + +namespace +{ + +[[nodiscard]] ReceiverConfig FastConfig() +{ + ReceiverConfig config = TestConfig(); + config.ReconnectBaseDelay = 10ms; + config.ReconnectMaxDelay = 40ms; + return config; +} + +using Harness = DriverHarness; + +void TestConnectSendsTheHelloWithTheProvidedToken() +{ + DriverCallbacks callbacks; + callbacks.Token = [] + { + return std::optional("token-1"); + }; + Harness harness(FastConfig(), std::move(callbacks)); + CHECK(!harness.Driver->Running() && !harness.Driver->Stopped()); + CHECK(harness.Driver->Start()); + CHECK(!harness.Driver->Start()); + CHECK(harness.Driver->Running()); + CHECK(harness.Driver->ThreadId().has_value()); + const std::shared_ptr connection = harness.Accept(); + const SentHello hello = DecodeHello(connection->Sent()); + CHECK(hello.Role == "receiver" && hello.SyncFrom == 1); + CHECK(hello.Token == "token-1"); + CHECK(!harness.Driver->WaitConnected(1ms)); + CHECK(harness.Record.Logged("connecting to 127.0.0.1:7200")); +} + +void TestFramesReachTheEndpoint() +{ + Harness harness(FastConfig()); + const std::shared_ptr connection = harness.Handshake(); + Bytes batch = server::Event(1); + const Bytes second = server::Event(2); + batch.insert(batch.end(), second.begin(), second.end()); + harness.Deliver(*connection, batch); + harness.Deliver(*connection, server::Ping()); + harness.Deliver(*connection, server::ReplayComplete(2, 3)); + CHECK(harness.Endpoint.Status().LastSequence == 2); + CHECK(!harness.Driver->WaitSynchronized(1ms)); + + const std::uint64_t generation = harness.Endpoint.Generation(); + CHECK(Sequences(harness.Endpoint.DrainFrames()) == (std::vector{1, 2})); + CHECK(harness.Endpoint.MarkAppliedThrough(generation, 2)); + CHECK(harness.Endpoint.MarkReplayApplied()); + harness.Driver->Wake(); + CHECK(harness.Driver->WaitSynchronized(kPatience)); + CHECK(harness.Endpoint.Status().ReplayEpoch == 3); + CHECK(harness.Record.Count() == 1); +} + +// A token issued by one handshake reaches the hook once, before the next +// handshake presents it. +void TestIssuedTokenReachesTheNextHandshake() +{ + std::mutex mutex; + std::vector issued; + DriverCallbacks callbacks; + callbacks.TokenIssued = [&](const std::string& token) + { + std::lock_guard lock(mutex); + issued.push_back(token); + }; + callbacks.Token = [&] + { + std::lock_guard lock(mutex); + return std::optional(issued.empty() ? std::string() : issued.back()); + }; + Harness harness(FastConfig(), std::move(callbacks)); + server::Hello hello; + hello.Token = "issued"; + const std::shared_ptr first = harness.Handshake(hello); + CHECK(DecodeHello(first->Sent()).Token.empty()); + first->Close(); + CHECK(first->WaitClosed(kPatience)); + const std::shared_ptr second = harness.Accept(); + CHECK(DecodeHello(second->Sent()).Token == "issued"); + CHECK(harness.Record.Count() == 1); + CHECK(harness.Record.IssuedToken() == "issued"); + std::lock_guard lock(mutex); + CHECK(issued == std::vector{"issued"}); +} + +void TestAbandonedTokenRetries() +{ + std::atomic calls = 0; + DriverCallbacks callbacks; + callbacks.Token = [&calls] + { + return ++calls == 1 ? std::nullopt : std::optional("later"); + }; + Harness harness(FastConfig(), std::move(callbacks)); + CHECK(harness.Driver->Start()); + const std::shared_ptr abandoned = harness.Sockets->Accept(kPatience); + CHECK(abandoned != nullptr && abandoned->WaitClosed(kPatience)); + CHECK(abandoned->Sent().empty()); + CHECK(DecodeHello(harness.Accept()->Sent()).Token == "later"); +} + +void TestReadTimeoutsReconnect() +{ + ReceiverConfig config = FastConfig(); + config.SocketTimeout = 50ms; + config.MaxConsecutiveTimeouts = 2; + Harness harness(config); + const std::shared_ptr connection = harness.Handshake(); + CHECK(connection->WaitClosed(kPatience)); + CHECK(DecodeHello(harness.Accept()->Sent()).SyncFrom == 1); +} + +void TestFailedConnectIsRecordedAndRetried() +{ + Harness harness(FastConfig()); + CHECK(harness.Driver->Start()); + CHECK(harness.Sockets->Refuse(kPatience, kRefused)); + const std::shared_ptr connection = harness.Accept(); + CHECK(harness.Sockets->Attempts() == 2); + CHECK(!harness.Driver->LastFailure()); + + ReceiverConfig config = FastConfig(); + config.Reconnect = false; + Harness once(config); + CHECK(once.Driver->Start()); + CHECK(once.Sockets->Refuse(kPatience, kRefused)); + CHECK(once.Driver->Join(kPatience)); + CHECK(once.Driver->Stopped() && !once.Driver->Running()); + CHECK(!once.Driver->WaitConnected(kPatience)); + const std::optional failure = once.Driver->LastFailure(); + CHECK(failure && failure->Operation == SocketOperation::Connect); + CHECK(failure->Result == SocketResult::Failed && failure->SystemError == kRefused); + CHECK(once.Endpoint.Status().Stopped); +} + +void TestStopInterruptsBlockingCalls() +{ + { + Harness harness(FastConfig()); + const std::shared_ptr connection = harness.Handshake(); + harness.Driver->Stop(); + CHECK(harness.Driver->Join(kPatience)); + CHECK(connection->ClosedByClient()); + CHECK(harness.Driver->Stopped()); + CHECK(harness.Endpoint.Status().Stopped); + } + { + Harness harness(FastConfig()); + CHECK(harness.Driver->Start()); + harness.Driver->Stop(); + CHECK(harness.Driver->Join(kPatience)); + CHECK(!harness.Driver->WaitConnected(std::nullopt)); + } + { + // A host may stop the endpoint before the driver starts. + Harness harness(FastConfig()); + harness.Driver->Stop(); + CHECK(harness.Driver->Start()); + CHECK(harness.Driver->Join(kPatience)); + CHECK(harness.Sockets->Attempts() == 0); + } +} + +void TestCloseEndsTheConnection() +{ + Harness harness(FastConfig()); + const std::shared_ptr connection = harness.Handshake(); + CHECK(harness.Driver->Close(kPatience)); + CHECK(connection->ClosedByClient()); +} + +void TestWakeAppliesAReplayRequest() +{ + Harness harness(FastConfig()); + const std::shared_ptr connection = harness.Handshake(); + harness.Deliver(*connection, server::Event(1)); + CHECK(harness.Endpoint.RequestReplayFrom(1)); + harness.Driver->Wake(); + CHECK(connection->WaitClosed(kPatience)); + CHECK(DecodeHello(harness.Accept()->Sent()).SyncFrom == 1); +} + +// Callbacks may call back into the endpoint and the driver. +void TestCallbacksReenterTheEndpointAndDriver() +{ + ThreadedReceiverDriver* driver = nullptr; + ReceiverEndpoint* endpoint = nullptr; + std::atomic joined_self = true; + DriverCallbacks callbacks; + callbacks.Notifications = [&](Notification notification) + { + if (std::holds_alternative(notification)) + { + CHECK(endpoint->Status().Connected); + CHECK(endpoint->RequestReplayFrom(2)); + driver->Wake(); + joined_self = driver->Join(1ms); + } + else if (std::holds_alternative(notification)) + { + driver->Stop(); + } + }; + Harness harness(FastConfig(), std::move(callbacks)); + driver = harness.Driver.get(); + endpoint = &harness.Endpoint; + server::Hello hello; + hello.Token = "issued"; + CHECK(harness.Driver->Start()); + const std::shared_ptr first = harness.Accept(); + CHECK(first->Deliver(server::HelloOk(hello))); + CHECK(first->WaitClosed(kPatience)); + CHECK(!joined_self); + const std::shared_ptr second = harness.Accept(); + CHECK(DecodeHello(second->Sent()).SyncFrom == 2); + CHECK(second->Deliver(server::HelloOk())); + CHECK(second->Deliver(server::Claimed("me"))); + CHECK(harness.Driver->Join(kPatience)); +} + +} // namespace + +int main() +{ + TestConnectSendsTheHelloWithTheProvidedToken(); + TestFramesReachTheEndpoint(); + TestIssuedTokenReachesTheNextHandshake(); + TestAbandonedTokenRetries(); + TestReadTimeoutsReconnect(); + TestFailedConnectIsRecordedAndRetried(); + TestStopInterruptsBlockingCalls(); + TestCloseEndsTheConnection(); + TestWakeAppliesAReplayRequest(); + TestCallbacksReenterTheEndpointAndDriver(); + return 0; +} diff --git a/tests/unit/layered_replay_test_support.py b/tests/unit/layered_replay_test_support.py index e2584ba..5ae7399 100644 --- a/tests/unit/layered_replay_test_support.py +++ b/tests/unit/layered_replay_test_support.py @@ -37,6 +37,7 @@ class _LayeredQueue: layered_replay_active = True sync_from = 1 origin = None + generation = 0 def __init__(self, messages): self.messages = list(messages) @@ -50,6 +51,12 @@ def drain_queue(self): def request_replay_from(self, seq_start): self.replay_requests.append(seq_start) + def reset_applied_progress(self): + pass + + def mark_applied_through(self, _generation, _sequence): + return True + def mark_replay_applied(self): return False diff --git a/tests/unit/test_asset_dependency_refresh.py b/tests/unit/test_asset_dependency_refresh.py index 3576f76..ee14e39 100644 --- a/tests/unit/test_asset_dependency_refresh.py +++ b/tests/unit/test_asset_dependency_refresh.py @@ -44,6 +44,7 @@ class _QueueReceiver: layered_replay_active = False sync_from = 1 origin = None + generation = 0 def __init__(self): self.messages = [] @@ -53,6 +54,12 @@ def drain_queue(self): messages, self.messages = self.messages, [] return messages + def reset_applied_progress(self): + pass + + def mark_applied_through(self, _generation, _sequence): + return True + def mark_replay_applied(self): return False diff --git a/tests/unit/test_blender_stage_author.py b/tests/unit/test_blender_stage_author.py index f17179d..2711022 100644 --- a/tests/unit/test_blender_stage_author.py +++ b/tests/unit/test_blender_stage_author.py @@ -1129,7 +1129,7 @@ def test_blender_emitter_releases_batch_only_after_send_succeeds(monkeypatch): emitter = MagicMock() emitter.prepare_events_for_send.return_value = events sender = MagicMock() - sender.sock = object() + sender.connected = True sender.send_events.side_effect = [False, True] monkeypatch.setattr(capture._state, "notice_emitter", emitter) monkeypatch.setattr(capture._state, "sender", sender) @@ -1137,7 +1137,7 @@ def test_blender_emitter_releases_batch_only_after_send_succeeds(monkeypatch): capture._try_send_dirty_events() emitter.mark_prepared_events_sent.assert_not_called() - sender.sock = object() + sender.connected = True capture._try_send_dirty_events() assert sender.send_events.call_count == 2 @@ -1317,7 +1317,7 @@ def test_blender_connect_reuses_sender_with_unacknowledged_outbox(monkeypatch): from integrations.blender import capture sender = MagicMock() - sender.sock = None + sender.connected = False sender.host = "127.0.0.1" sender.port = 7200 sender.department = "animation" @@ -1450,18 +1450,6 @@ def test_receiver_discards_retained_replay_state(monkeypatch): assert receiver_addon._pending_shader_baseline_paths == set() -def test_receiver_thread_cleanup_tolerates_unstarted_thread(): - from integrations.blender import receiver_addon - - receiver = MagicMock() - receiver.join.side_effect = RuntimeError("cannot join thread before it is started") - - receiver_addon._stop_receiver_thread(receiver) - - receiver.stop.assert_called_once_with() - receiver.join.assert_called_once_with(timeout=2.0) - - def test_receiver_sequence_persistence_tolerates_released_scene(): from integrations.blender import receiver_addon diff --git a/tests/unit/test_client_lifecycle.py b/tests/unit/test_client_lifecycle.py index 78fe2eb..3e64aca 100644 --- a/tests/unit/test_client_lifecycle.py +++ b/tests/unit/test_client_lifecycle.py @@ -1,10 +1,11 @@ """Host-facing lifecycle guarantees shared by the high-level clients.""" import threading +import uuid from types import SimpleNamespace import pytest -from pxr import Usd +from pxr import Sdf, Usd, UsdGeom from openusdconnect import ( ManagedClient, @@ -19,7 +20,59 @@ ) from openusdconnect import sender as sender_module from openusdconnect.client_observer import StageMetadata -from tests.helpers import RecordingObserver, force_handshake +from openusdconnect.protocol_constants import LayerMode +from tests.helpers import ( + RecordingObserver, + connect_client, + embedded_server, + recorded_hellos, + wait_until, +) + + +@pytest.fixture(scope="module") +def managed_server(): + with embedded_server() as runtime: + yield runtime + + +@pytest.fixture(scope="module") +def shared_server(tmp_path_factory): + root = tmp_path_factory.mktemp("shared") / "root.usda" + Sdf.Layer.CreateNew(str(root)).Save() + with embedded_server(base_usd_path=str(root), layer_mode=LayerMode.SHARED_STAGE) as runtime: + yield runtime + + +@pytest.fixture(scope="module") +def token_servers(tmp_path_factory): + """Servers by client kind that issue tokens and author an up axis.""" + root = tmp_path_factory.mktemp("tokens") / "root.usda" + stage = Usd.Stage.CreateNew(str(root)) + UsdGeom.SetStageUpAxis(stage, UsdGeom.Tokens.z) + stage.GetRootLayer().Save() + with ( + embedded_server(base_usd_path=str(root), require_token=True) as managed, + embedded_server( + base_usd_path=str(root), require_token=True, layer_mode=LayerMode.SHARED_STAGE, + ) as shared, + ): + yield { + ManagedClient: managed, + SharedStageClient: shared, + UsdReceiver: managed, + UsdPublisher: managed, + } + + +@pytest.fixture +def ports(managed_server, shared_server): + """Live server ports by client kind; publishers connect nowhere in these tests.""" + return { + ManagedClient: managed_server.server_address[1], + SharedStageClient: shared_server.server_address[1], + UsdPublisher: 1, + } def test_wait_until_ready_returns_false_only_when_startup_expires(): @@ -54,17 +107,6 @@ def test_blocked_states_raise_instead_of_timing_out(phase, auth_rejected, failur _client_lifecycle.raise_if_blocked(SimpleNamespace(), status) -def test_each_phase_outranks_the_phases_after_it(): - flags = [ - "closed", "recovery_required", "rejected", "parked", "replaying", "ready", "connecting", - ] - for index, flag in enumerate(flags): - state = {name: position >= index for position, name in enumerate(flags)} - assert _client_lifecycle.compute_phase(**state) is client_types.ClientPhase(flag) - offline = _client_lifecycle.compute_phase(**dict.fromkeys(flags, False)) - assert offline is client_types.ClientPhase.OFFLINE - - @pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient, UsdReceiver, UsdPublisher]) def test_every_client_reports_the_same_lifecycle_phases(kind, tmp_path): stage = Usd.Stage.CreateNew(str(tmp_path / "scene.usda")) @@ -93,7 +135,7 @@ def test_waits_raise_when_nothing_will_reconnect(): publisher.start() publisher.disconnect() receiver.start() - receiver.receiver.join(timeout=5) + wait_until(lambda: receiver.receiver.stopped) for client in (publisher, receiver): assert client.status.phase is client_types.ClientPhase.OFFLINE with pytest.raises(ConnectionError, match="offline"): @@ -105,60 +147,6 @@ def test_waits_raise_when_nothing_will_reconnect(): receiver.close() -def test_backlog_hold_counts_only_messages_queued_before_the_batch(): - hold = _client_lifecycle.BacklogHold() - hold.freeze(3) - hold.drained(2, queued=5) - assert hold.holding - hold.drained(2, queued=5) - assert not hold.holding - hold.freeze(4) - hold.drained(0, queued=0) - assert not hold.holding, "a discarded queue has nothing left ahead of the batch" - - -def test_queued_notifications_run_on_update_thread_and_bound_each_drain(): - notifications = _client_lifecycle.ClientCallbackQueue() - received = [] - - def observe(value): - received.append((value, threading.get_ident())) - if value == "first": - callback("next-tick") - - callback = notifications.wrap(observe) - worker = threading.Thread(target=lambda: callback("first")) - worker.start() - worker.join(timeout=1) - assert not worker.is_alive() - assert received == [] - notifications.drain() - assert received == [("first", threading.get_ident())] - notifications.drain() - assert received == [("first", threading.get_ident()), ("next-tick", threading.get_ident())] - callback("queued") - notifications.close() - callback("late") - notifications.drain() - assert [value for value, _thread in received] == ["first", "next-tick", "queued"] - - -def test_queued_notification_failure_propagates_and_keeps_later_notifications(): - notifications = _client_lifecycle.ClientCallbackQueue() - received = [] - - def fail(value): - raise RuntimeError("observer failed") - - notifications.wrap(fail)(None) - notifications.wrap(received.append)("next") - with pytest.raises(RuntimeError, match="observer failed"): - notifications.drain() - assert received == [] - notifications.drain() - assert received == ["next"] - - def test_public_status_types_keep_compatibility_identity(): from openusdconnect import ClientPhase, ClientStatus, SyncUpdate @@ -213,69 +201,57 @@ def slow_load(host, port): assert credential.current() == "issued-during-load" -def test_sender_takes_its_token_from_the_provider_on_every_attempt(monkeypatch): - tokens = iter(["first", "second"]) - presented = [] - - def refuse(*args, **kwargs): - presented.append(sender.token) - raise OSError("refused") - - monkeypatch.setattr(sender_module.socket, "create_connection", refuse) - sender = sender_module.EventSender( - "localhost", 1, client_id="token-attempts", token="stale", - token_provider=lambda: next(tokens), - ) - assert not sender.connect(timeout=0.5) - assert not sender.connect(timeout=0.5) - assert presented == ["first", "second"] - - @pytest.mark.parametrize("failure", [None, "persistence", "observer"]) def test_issued_token_is_persisted_before_notifying_and_used_by_both_roles( - failure, tmp_path, monkeypatch, + failure, tmp_path, monkeypatch, token_servers, ): calls = [] def record(name, token): - calls.append((name, token)) + calls.append((name, token, threading.get_ident())) if failure == name: raise RuntimeError(f"injected {name} failure") + monkeypatch.setattr(_client_utils, "load_token", lambda host, port: None) monkeypatch.setattr( _client_utils, "save_token", lambda host, port, token: record("persistence", token), ) - stage = Usd.Stage.CreateNew(str(tmp_path / "scene.usda")) - observer = RecordingObserver(on_call=lambda name, token: record("observer", token)) + hellos = recorded_hellos(monkeypatch) + observer = RecordingObserver( + on_call=lambda name, token: name == "token_issued" and record("observer", token), + ) client = ManagedClient( - stage, app_name="shared-credentials", token="configured", persist_token=True, - observer=observer, + Usd.Stage.CreateNew(str(tmp_path / "scene.usda")), app_name="shared-credentials", + client_id=uuid.uuid4().hex, port=token_servers[ManagedClient].server_address[1], + persist_token=True, observer=observer, ) try: - callback = client._sender._on_token_issued - if failure == "persistence": - with pytest.raises(RuntimeError, match="injected persistence failure"): - callback("replacement") + client.start() + wait_until(lambda: calls) + # The connection thread adopted and persisted the token; the host hears it in update(). + [(name, issued, thread)] = calls + assert name == "persistence" and thread != threading.get_ident() + if failure == "observer": + with pytest.raises(RuntimeError, match="injected observer failure"): + client.update() + # Later notifications wait for the next call. + assert [call[0] for call in observer.calls] == ["token_issued"] else: - callback("replacement") - # The host observer is queued; the token is adopted and persisted first. - assert calls == [("persistence", "replacement")] - for endpoint in (client._sender, client._receiver): - assert endpoint._token_provider() == "replacement" - - if failure != "persistence": - if failure == "observer": - with pytest.raises(RuntimeError, match="injected observer failure"): - client._callbacks.drain() - else: - client._callbacks.drain() - assert calls[-1] == ("observer", "replacement") + client.update() + assert calls[-1] == ("observer", issued, threading.get_ident()) + assert client.wait_until_ready(5) + client.update() + # Both roles' handshakes carried the same metadata. + assert [call[0] for call in observer.calls].count("stage_metadata") == 1 + assert [(hello["role"], hello.get("token")) for hello in hellos] == [ + ("receiver", None), ("emitter", issued), + ] finally: client.close() @pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient, UsdPublisher]) -def test_disconnected_update_does_no_token_io(kind, tmp_path, monkeypatch): +def test_disconnected_update_does_no_token_io(kind, tmp_path, monkeypatch, ports): reads = [] monkeypatch.setattr(_client_utils, "load_token", lambda host, port: reads.append(1)) requests = [] @@ -284,9 +260,12 @@ def test_disconnected_update_does_no_token_io(kind, tmp_path, monkeypatch): lambda self, timeout=2.0: requests.append(True) or True, ) stage = Usd.Stage.CreateNew(str(tmp_path / "scene.usda")) - client = kind(stage, app_name="token-io", port=1, persist_token=True) + client = kind(stage, app_name="token-io", port=ports[kind], persist_token=True) try: - force_handshake(client) + if kind is UsdPublisher: + client.start() + else: + connect_client(client) reads.clear() for _ in range(20): client.update() @@ -298,17 +277,20 @@ def test_disconnected_update_does_no_token_io(kind, tmp_path, monkeypatch): @pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient]) def test_update_schedules_handshake_without_waiting_or_touching_stage_in_worker( - kind, tmp_path, monkeypatch + kind, tmp_path, monkeypatch, ports ): stage = Usd.Stage.CreateNew(str(tmp_path / "scene.usda")) - client = kind(stage, app_name="background-connect-test", persist_token=False) + client = kind( + stage, app_name="background-connect-test", port=ports[kind], persist_token=False, + ) entered = threading.Event() release = threading.Event() finished = threading.Event() owner = threading.get_ident() threads = [] - def connect(*args, **kwargs): + def provide(): + # The handshake waits here, after the connection opened. threads.append(threading.get_ident()) entered.set() try: @@ -317,9 +299,9 @@ def connect(*args, **kwargs): finally: finished.set() - monkeypatch.setattr(sender_module.socket, "create_connection", connect) - force_handshake(client) try: + connect_client(client) + monkeypatch.setattr(client._sender, "_token_provider", provide) result = client.update() assert entered.wait(2) assert not finished.is_set() @@ -336,10 +318,8 @@ def connect(*args, **kwargs): @pytest.mark.parametrize("kind", [ManagedClient, UsdPublisher]) -def test_flush_shares_timeout_between_reconnect_and_acknowledgement(kind, monkeypatch): +def test_flush_shares_timeout_between_reconnect_and_acknowledgement(kind, monkeypatch, ports): clock = [10.0] - monkeypatch.setattr(_client_lifecycle.time, "monotonic", lambda: clock[0]) - client = kind(Usd.Stage.CreateInMemory(), app_name="flush-budget", persist_token=False) calls = [] def connect(timeout=None): @@ -351,15 +331,19 @@ def flush(timeout=None): calls.append(("flush", timeout)) return True - monkeypatch.setattr(client._sender, "connect", connect) - monkeypatch.setattr(client._sender, "flush", flush) - client._transform_coalescing = SimpleNamespace(buffering=True, force=lambda emitter: []) - if isinstance(client, ManagedClient): - client._receiver.connected = True - client._receiver._synchronized_event.set() - else: - monkeypatch.setattr(client, "_is_synchronized", lambda: True) + client = kind( + Usd.Stage.CreateInMemory(), app_name="flush-budget", port=ports[kind], + persist_token=False, + ) try: + if isinstance(client, ManagedClient): + connect_client(client) + else: + monkeypatch.setattr(client, "_is_synchronized", lambda: True) + monkeypatch.setattr(_client_lifecycle.time, "monotonic", lambda: clock[0]) + monkeypatch.setattr(client._sender, "connect", connect) + monkeypatch.setattr(client._sender, "flush", flush) + client._transform_coalescing = SimpleNamespace(buffering=True, force=lambda emitter: []) assert client.flush(timeout=1.0) assert calls == [("connect", 1.0), ("flush", 0.25)] finally: @@ -389,23 +373,25 @@ def test_can_author_combines_readiness_role_and_edit_target( @pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient, UsdReceiver, UsdPublisher]) -def test_close_delivers_notifications_queued_by_network_threads(kind, tmp_path): +def test_close_delivers_notifications_queued_by_network_threads(kind, tmp_path, token_servers): observer = RecordingObserver() stage = Usd.Stage.CreateNew(str(tmp_path / "scene.usda")) - client = kind(stage, app_name="notifications", persist_token=False, observer=observer) - endpoint = client._sender if kind is UsdPublisher else client._receiver + client = kind( + stage, app_name="notifications", client_id=uuid.uuid4().hex, + port=token_servers[kind].server_address[1], persist_token=False, observer=observer, + ) try: - worker = threading.Thread(target=lambda: ( - endpoint._on_token_issued("issued"), - endpoint._on_stage_metadata({"upAxis": "Y"}), - )) - worker.start() - worker.join(timeout=1) - assert observer.calls == [] + if kind is UsdPublisher: + assert client.connect(timeout=5) + else: + client.start() + assert client._receiver.wait_connected(5) + token = client._endpoints[0].token + assert token and observer.calls == [] finally: client.close() # A host that stores tokens itself must receive one issued just before close. - assert observer.calls == [ - ("token_issued", "issued", threading.get_ident()), - ("stage_metadata", StageMetadata(up_axis="Y"), threading.get_ident()), + assert [call for call in observer.calls if call[0] != "playback_state"] == [ + ("token_issued", token, threading.get_ident()), + ("stage_metadata", StageMetadata(up_axis="Z"), threading.get_ident()), ] diff --git a/tests/unit/test_client_observer.py b/tests/unit/test_client_observer.py index 5ce05ca..b2c62ff 100644 --- a/tests/unit/test_client_observer.py +++ b/tests/unit/test_client_observer.py @@ -11,10 +11,15 @@ PlaybackState, UsdReceiver, ) -from openusdconnect._observer_hooks import ObserverHooks, observer_hooks from openusdconnect.codec import encode_message -from openusdconnect.protocol_constants import K_ENSURE_PRIM, K_SET_REFERENCE, K_SET_VISIBILITY -from tests.helpers import RecordingObserver +from openusdconnect.protocol_constants import ( + K_ENSURE_PRIM, + K_SET_REFERENCE, + K_SET_VISIBILITY, + MSG_PLAYBACK_CLAIMED, + MSG_PLAYBACK_REJECTED, +) +from tests.helpers import RecordingObserver, embedded_server, wait_until def test_applied_batch_derives_sorted_unique_paths_once(): @@ -35,11 +40,18 @@ class Paths(ClientObserver): def on_applied(self, batch): pass - assert observer_hooks(None, lambda cb: cb) == ObserverHooks() - assert observer_hooks(ClientObserver(), lambda cb: cb) == ObserverHooks() - hooks = observer_hooks(Paths(), lambda cb: cb) - assert hooks.on_applied is not None - assert set(hooks.receiver_callbacks().values()) == {None} + def wiring(observer): + client = UsdReceiver( + Usd.Stage.CreateInMemory(), app_name="observer-wiring", persist_token=False, + observer=observer, + ) + dispatcher = client.dispatcher + client.close() + wired = client._notification_methods + return dispatcher.on_applied_events is not None, dispatcher.on_resync, wired + + assert wiring(None) == wiring(ClientObserver()) == (False, None, {}) + assert wiring(Paths()) == (True, None, {}) def _event_frames(*paths): @@ -116,17 +128,29 @@ def on_resync(self): def test_notification_payloads_are_typed(): observer = RecordingObserver() - callbacks = observer_hooks(observer, lambda callback: callback).receiver_callbacks() - callbacks["on_playback_state"]( - {"type": "playback_state", "playing": True, "time": 2.0, "rate": 1.0, - "leader_client_id": "a"} - ) - callbacks["on_playback_claimed"]({"type": "playback_claimed", "leader_client_id": "a"}) - callbacks["on_playback_rejected"]( - {"type": "playback_rejected", "reason": "busy", "current_leader_client_id": "b"} - ) - assert [value for _name, value, _thread in observer.calls] == [ - PlaybackState(True, 2.0, 1.0, "a"), + + def playback(): + return [value for name, value, _thread in observer.calls if name.startswith("playback")] + + with embedded_server() as server: + state = server.sync_server + client = UsdReceiver( + Usd.Stage.CreateInMemory(), app_name="typed-notifications", + port=server.server_address[1], persist_token=False, observer=observer, + ) + try: + client.start() + # The server sends its playback state after every accepted hello. + wait_until(lambda: client.update() is not None and playback()) + state.broadcast_message({"type": MSG_PLAYBACK_CLAIMED, "leader_client_id": "a"}) + state.broadcast_message({ + "type": MSG_PLAYBACK_REJECTED, "reason": "busy", "current_leader_client_id": "b", + }) + wait_until(lambda: client.update() is not None and len(playback()) == 3) + finally: + client.close() + assert playback() == [ + PlaybackState(**state.get_playback_state()), PlaybackClaim(True, "a"), PlaybackClaim(False, "b", "busy"), ] diff --git a/tests/unit/test_dispatcher.py b/tests/unit/test_dispatcher.py index feabbaa..a8306d0 100644 --- a/tests/unit/test_dispatcher.py +++ b/tests/unit/test_dispatcher.py @@ -34,6 +34,7 @@ class _QueuedReceiver: layered_replay_active = False sync_from = 1 origin = None + generation = 0 def __init__(self, messages): self.messages = list(messages) @@ -47,6 +48,12 @@ def drain_queue(self, max_messages=None): def request_replay_from(self, seq_start): self.replay_requests.append(seq_start) + def reset_applied_progress(self): + pass + + def mark_applied_through(self, _generation, _sequence): + return True + def mark_replay_applied(self): return False diff --git a/tests/unit/test_mcp_session.py b/tests/unit/test_mcp_session.py index e162e35..ca802b2 100644 --- a/tests/unit/test_mcp_session.py +++ b/tests/unit/test_mcp_session.py @@ -7,8 +7,9 @@ from integrations.mcp import session as session_mod from integrations.mcp.config import McpConfig from integrations.mcp.errors import ToolError -from openusdconnect import usd_client +from openusdconnect import _client_backend, usd_client from openusdconnect.checkpoints import MirrorCheckpoint +from openusdconnect.client_observer import PlaybackState from openusdconnect.client_types import SyncUpdate @@ -35,6 +36,13 @@ def _applied(count: int) -> SyncUpdate: return SyncUpdate(applied_events=count, submitted_events=0) +def _deliver_playback(session, playing, time, rate, leader_client_id): + """Deliver a playback state to the mirror's observer as update() does.""" + methods = session.receiver._notification_methods + on_playback_state, _convert = methods[_client_backend.PlaybackState] + on_playback_state(PlaybackState(playing, time, rate, leader_client_id)) + + def _patch_net(monkeypatch, started, stopped): class _FakeReceiver: synchronized = True @@ -46,25 +54,23 @@ class _FakeReceiver: server_instance = "test-server" replay_epoch = 0 stopped = False + rejection = None + generation = 0 def __init__(self, **kwargs): self.options = kwargs self.token = kwargs["token"] self.sync_from = kwargs["sync_from"] - self.joined = False def start(self): started.append(self) - def stop(self): - stopped.append(self) - - def is_alive(self): - return not self.joined + def snapshot(self): + return self - def join(self, timeout=None): - assert self in stopped - self.joined = True + def close(self, timeout=None): + stopped.append(self) + return True def drain_queue(self, max_messages=None): return [] @@ -73,7 +79,7 @@ def mark_replay_applied(self): return False monkeypatch.setattr(session_mod, "EventSender", _FakeSender) - monkeypatch.setattr(usd_client, "ReceiverThread", _FakeReceiver) + monkeypatch.setattr(usd_client, "EventReceiver", _FakeReceiver) monkeypatch.setattr(session_mod.token_client, "load_token", lambda host, port: None) @@ -94,7 +100,6 @@ def test_reconnect_stops_previous_receiver(monkeypatch): session.connect() assert first.receiver in stopped - assert first.receiver.joined assert session.receiver is not first # replaced by a fresh one assert session.receiver.receiver in started session.disconnect() @@ -122,14 +127,13 @@ def test_playback_status_reflects_broadcast(monkeypatch): assert session.playback_status()["observed"] is False # nothing broadcast yet - notify = session.receiver.receiver.options["on_playback_state"] - notify({"playing": True, "time": 12.0, "rate": 2.0, "leader_client_id": "mcp-x"}) + _deliver_playback(session, True, 12.0, 2.0, "mcp-x") st = session.playback_status() assert st["observed"] and st["playing"] is True assert st["time"] == 12.0 and st["rate"] == 2.0 assert st["has_leader"] is True and st["is_leader"] is True - notify({"playing": False, "time": 0.0, "rate": 1.0, "leader_client_id": "someone-else"}) + _deliver_playback(session, False, 0.0, 1.0, "someone-else") st2 = session.playback_status() assert st2["is_leader"] is False assert st2["leader_client_id"] == "someone-else" @@ -145,7 +149,6 @@ def test_disconnect_stops_receiver(monkeypatch): receiver = session.receiver session.disconnect() assert receiver.receiver in stopped - assert receiver.receiver.joined assert session.receiver is None assert session.sender is None assert session.mirror_stage is None @@ -390,9 +393,7 @@ def connect(sender): assert options["client_id"] == "mcp-identity" assert options["origin"] == f"{session._origin_base}-recv" assert options["layered_replay"] is True - options["on_playback_state"]( - {"playing": True, "time": 1.0, "rate": 1.0, "leader_client_id": ""} - ) + _deliver_playback(session, True, 1.0, 1.0, "") assert session.playback_status()["playing"] is True dispatcher = session.receiver.dispatcher dispatcher.last_seq = 3 @@ -436,13 +437,12 @@ def fail_start(self): with pytest.raises(RuntimeError, match="start failed"): session.connect() assert len(stopped) == 1 - assert stopped[0].joined assert session.sender is None assert session.receiver is None assert session.mirror_stage is None -def test_failed_send_closes_and_joins_mirror(monkeypatch): +def test_failed_send_closes_the_mirror(monkeypatch): started, stopped = [], [] _patch_net(monkeypatch, started, stopped) session = session_mod.ConnectionSession(McpConfig()) @@ -454,7 +454,6 @@ def test_failed_send_closes_and_joins_mirror(monkeypatch): assert error.value.code == "disconnected" assert not sender.is_connected assert stopped == started - assert stopped[0].joined def test_no_mirror_send_result_is_unchanged(monkeypatch): diff --git a/tests/unit/test_native_client.py b/tests/unit/test_native_client.py index c862db3..21a8a00 100644 --- a/tests/unit/test_native_client.py +++ b/tests/unit/test_native_client.py @@ -1,7 +1,6 @@ from __future__ import annotations import os -import struct import subprocess import sys from pathlib import Path @@ -36,307 +35,6 @@ def test_source_package_finds_native_extension_without_editable_hook(tmp_path): assert Path(result.stdout.strip()).resolve() == Path(native.__file__).resolve() -def test_frame_decoder_handles_fragmented_and_coalesced_input(): - decoder = native.FrameDecoder() - stream = native.encode_frame(b"alpha") + native.encode_frame(b"beta") - - assert decoder.feed(stream[:3]) == [] - assert decoder.buffered_bytes == 3 - assert decoder.feed(stream[3:8]) == [] - assert decoder.feed(stream[8:]) == [b"alpha", b"beta"] - assert decoder.buffered_bytes == 0 - - -def test_frame_decoder_rejects_invalid_size_at_header_boundary(): - decoder = native.FrameDecoder(8) - - with pytest.raises(native.FrameError): - decoder.feed(struct.pack(">I", 9)) - - -def test_python_boundary_preserves_frame_validation_exceptions(): - with pytest.raises(ValueError): - native.FrameDecoder(0) - with pytest.raises(native.FrameError): - native.encode_frame(b"") - with pytest.raises(native.FrameError): - native.encode_frame(b"oversized", max_frame_size=4) - - -def test_receiver_inbox_preserves_replay_live_boundary(): - inbox = native.ReceiverInbox(initial_sync_from=1, max_messages=8) - connection = inbox.begin_connection() - - assert connection.sync_from == 1 - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 1, b"event-1") - == native.AcceptResult.ACCEPTED - ) - assert ( - inbox.accept_replay_complete(connection.generation, head_seq=1, epoch=7) - == native.AcceptResult.ACCEPTED - ) - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 2, b"event-2") - == native.AcceptResult.ACCEPTED - ) - - assert inbox.mark_replay_applied() is False - assert inbox.drain(1) == [b"event-1"] - assert inbox.mark_replay_applied() is True - assert inbox.replay_head_sequence == 1 - assert inbox.replay_epoch == 7 - assert inbox.drain() == [b"event-2"] - - -def test_receiver_inbox_retains_python_bytes_without_copying(): - inbox = native.ReceiverInbox(initial_sync_from=1, max_messages=2) - connection = inbox.begin_connection() - payload = b"immutable-receiver-payload" - - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 1, payload) - == native.AcceptResult.ACCEPTED - ) - # The binding queues the immutable Python object itself; it does not copy its bytes. - drained = inbox.drain() - assert drained == [payload] - assert drained[0] is payload - assert inbox.drain() == [] - - -def test_receiver_inbox_rejects_stale_generation_without_mutation(): - inbox = native.ReceiverInbox(initial_sync_from=4, max_messages=8) - first = inbox.begin_connection() - second = inbox.begin_connection() - - assert ( - inbox.accept(first.generation, native.ReceiverMessageKind.EVENT, 4, b"stale") - == native.AcceptResult.STALE_GENERATION - ) - assert inbox.size == 0 - assert second.sync_from == 4 - - -def test_receiver_inbox_overflow_is_bounded_and_replayable(): - inbox = native.ReceiverInbox(initial_sync_from=1, max_messages=1) - connection = inbox.begin_connection() - - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 1, b"one") - == native.AcceptResult.ACCEPTED - ) - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 2, b"two") - == native.AcceptResult.QUEUE_FULL - ) - assert inbox.overflowed is True - assert inbox.drain() == [b"one"] - - inbox.request_replay_from(2) - replay = inbox.begin_connection() - assert replay.sync_from == 2 - assert inbox.overflowed is False - - -@pytest.mark.parametrize("require_contiguous", [False, True]) -@pytest.mark.parametrize("queued_prefix", [False, True]) -def test_receiver_reset_reconnects_from_one_without_discarding_queue( - require_contiguous, queued_prefix, -): - inbox = native.ReceiverInbox( - initial_sync_from=4, - max_messages=2 if queued_prefix else 1, - require_contiguous=require_contiguous, - ) - connection = inbox.begin_connection() - assert connection.sync_from == 4 - expected = [] - if queued_prefix: - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 4, b"old-4") - == native.AcceptResult.ACCEPTED - ) - expected.append(b"old-4") - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.RESYNC, 0, b"reset") - == native.AcceptResult.ACCEPTED - ) - expected.append(b"reset") - assert inbox.size == len(expected) - assert inbox.last_sequence == 0 - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 1, b"new-1") - == native.AcceptResult.QUEUE_FULL - ) - # Repeated disconnects before the first new event must retain both the - # reset cursor and every queued frame, even before the consumer drains. - for _ in range(2): - inbox.disconnect(connection.generation) - connection = inbox.begin_connection() - assert connection.sync_from == 1 - assert inbox.size == len(expected) - assert inbox.drain() == expected - inbox.clear_overflow() - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 1, b"new-1") - == native.AcceptResult.ACCEPTED - ) - inbox.disconnect(connection.generation) - assert inbox.begin_connection().sync_from == 2 - assert inbox.drain() == [b"new-1"] - - -def test_receiver_full_replay_cursor_survives_disconnect_before_any_frames(): - inbox = native.ReceiverInbox(initial_sync_from=4, max_messages=1) - assert inbox.begin_connection().sync_from == 4 - inbox.request_replay_from(1) - for _ in range(2): - connection = inbox.begin_connection() - assert connection.sync_from == 1 - inbox.disconnect(connection.generation) - - -def test_receiver_rejected_reset_preserves_snapshot_cursor_and_queue(): - inbox = native.ReceiverInbox(initial_sync_from=4, max_messages=1) - connection = inbox.begin_connection() - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 4, b"old-4") - == native.AcceptResult.ACCEPTED - ) - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.RESYNC, 0, b"reset") - == native.AcceptResult.QUEUE_FULL - ) - assert inbox.last_sequence == 4 - inbox.disconnect(connection.generation) - assert inbox.begin_connection().sync_from == 5 - assert inbox.drain() == [b"old-4"] - - -def test_receiver_session_can_enforce_contiguous_delivery(): - inbox = native.ReceiverInbox( - initial_sync_from=1, - max_messages=4, - require_contiguous=True, - ) - connection = inbox.begin_connection() - - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 2, b"gap") - == native.AcceptResult.SEQUENCE_GAP - ) - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 1, b"one") - == native.AcceptResult.ACCEPTED - ) - assert ( - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 1, b"duplicate") - == native.AcceptResult.DUPLICATE - ) - assert inbox.drain() == [b"one"] - assert inbox.mark_applied_through(connection.generation, 1) is True - assert inbox.last_applied_sequence == 1 - - -def test_python_boundary_validates_receiver_preconditions(): - with pytest.raises(ValueError): - native.ReceiverInbox(initial_sync_from=0, max_messages=1) - with pytest.raises(ValueError): - native.ReceiverInbox(initial_sync_from=1, max_messages=0) - - inbox = native.ReceiverInbox(initial_sync_from=1, max_messages=1) - connection = inbox.begin_connection() - with pytest.raises(ValueError): - inbox.accept(connection.generation, native.ReceiverMessageKind.EVENT, 0, b"event") - with pytest.raises(ValueError): - inbox.accept_replay_complete(connection.generation, -1, 0) - with pytest.raises(ValueError): - inbox.drain(0) - with pytest.raises(ValueError): - inbox.request_replay_from(0) - - -def test_producer_session_replays_only_the_unacknowledged_suffix(): - session = native.ProducerSession(capacity=4) - first = session.begin_connection() - assert session.accept_hello(first.generation, 0) == native.ProducerResult.ACCEPTED - assert ( - session.append(first.generation, 1, b"txn-1", 2, "layout") == native.ProducerResult.ACCEPTED - ) - assert ( - session.append(first.generation, 2, b"txn-2", 3, "animation") - == native.ProducerResult.ACCEPTED - ) - assert session.claim_next_unsent(first.generation) == (1, b"txn-1", 2, "layout") - assert session.claim_next_unsent(first.generation) == (2, b"txn-2", 3, "animation") - - assert session.disconnect(first.generation) == native.ProducerResult.ACCEPTED - second = session.begin_connection() - assert session.accept_hello(second.generation, 1) == native.ProducerResult.ACCEPTED - assert session.claim_next_unsent(second.generation) == (2, b"txn-2", 3, "animation") - assert session.claim_next_unsent(second.generation) is None - assert session.pending_transaction_count == 1 - assert session.pending_event_count == 3 - assert session.acknowledged_transaction_count == 1 - assert session.acknowledged_event_count == 2 - assert session.drain_acknowledged_event_count() == 2 - assert session.drain_acknowledged_event_count() == 0 - - -def test_producer_session_retains_python_bytes_without_copying(): - session = native.ProducerSession(capacity=1) - connection = session.begin_connection() - payload = b"immutable-producer-payload" - - assert session.accept_hello(connection.generation, 0) == native.ProducerResult.ACCEPTED - assert session.append(connection.generation, 1, payload, 1) == native.ProducerResult.ACCEPTED - claimed = session.claim_next_unsent(connection.generation) - assert claimed is not None - assert claimed[1] is payload - - -def test_producer_session_quarantines_highwater_contradictions(): - session = native.ProducerSession(capacity=2) - connection = session.begin_connection() - - assert session.accept_hello(connection.generation, 1) == native.ProducerResult.HIGHWATER_AHEAD - assert session.phase == native.ProducerPhase.RECOVERY_REQUIRED - assert session.recovery_required is True - assert session.begin_connection() is None - - -def test_producer_session_repairs_recoverable_transaction_at_same_id(): - session = native.ProducerSession(capacity=2) - first = session.begin_connection() - assert session.accept_hello(first.generation, 0) == native.ProducerResult.ACCEPTED - assert session.append(first.generation, 1, b"stale", 1) == native.ProducerResult.ACCEPTED - assert ( - session.reject( - first.generation, - 1, - native.ProducerRecoveryDisposition.RECOVERABLE_CONFLICT, - ) - == native.ProducerResult.ACCEPTED - ) - - assert session.repair_rejected(b"repaired", 2) == native.ProducerResult.ACCEPTED - second = session.begin_connection() - assert session.accept_hello(second.generation, 0) == native.ProducerResult.ACCEPTED - assert session.claim_next_unsent(second.generation) == (1, b"repaired", 2, "") - - -def test_python_boundary_validates_producer_preconditions(): - with pytest.raises(ValueError): - native.ProducerSession(capacity=0) - - session = native.ProducerSession(capacity=1) - connection = session.begin_connection() - assert session.accept_hello(connection.generation, 0) == native.ProducerResult.ACCEPTED - with pytest.raises(ValueError): - session.append(connection.generation, 1, b"txn", 0) - - def test_shared_core_and_unreal_module_do_not_enable_cpp_exceptions(): core = REPO_ROOT / "native" / "client_core" sources = [*core.rglob("*.h"), *core.rglob("*.cpp")] diff --git a/tests/unit/test_receiver.py b/tests/unit/test_receiver.py index e89afc8..d023d8f 100644 --- a/tests/unit/test_receiver.py +++ b/tests/unit/test_receiver.py @@ -1,1075 +1,409 @@ -"""Tests for ReceiverThread.""" +"""EventReceiver's wrapper contract, exercised against a live server.""" -import logging +import gc import socket +import threading import time +import uuid +from contextlib import nullcontext as does_not_raise import pytest -from pxr import Usd +from pxr import Usd, UsdGeom from openusdconnect import _client_backend -from openusdconnect.adapters import UsdStageAdapter -from openusdconnect.codec import HelloRejectionCode, encode_message, message_to_dict -from openusdconnect.dispatcher import EventDispatcher -from openusdconnect.framing import recv_framed, send_framed -from openusdconnect.receiver import ReceiverThread - - -def _make_server(): - """Create a listening socket on a random port, return (socket, port).""" - srv = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - srv.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) - srv.bind(("127.0.0.1", 0)) - srv.listen(1) - return srv, srv.getsockname()[1] - - -def _accept_and_hello(srv, timeout=2, hello_ok=None): - """Accept one connection, consume hello, send hello_ok. Return conn.""" - srv.settimeout(timeout) - conn, _ = srv.accept() - conn.settimeout(timeout) - recv_framed(conn) # consume hello - if hello_ok is None: - hello_ok = {"type": "hello_ok", "layered_replay": True} - send_framed(conn, encode_message(hello_ok)) - return conn - - -def _accept(srv, timeout=2): - """Accept one connection (no handshake).""" - srv.settimeout(timeout) - conn, _ = srv.accept() - conn.settimeout(timeout) - return conn - - -def _recv_hello(conn): - """Read and decode a hello message from connection.""" - buf = recv_framed(conn) - return message_to_dict(buf) - - -def _send_event(conn, seq, event=None): - """Send a broadcast event message.""" - if event is None: - event = {"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"} - msg = {"type": "event", "seq": seq, "event": event} - send_framed(conn, encode_message(msg)) - - -def _flood_events(conn, seqs): - """Send events, tolerating the receiver closing the socket mid-flood. - - Overflowing the bounded queue makes the receiver disconnect by design, so - pushing past that point legitimately races with an RST from the peer the - sender just stops. The test's real assertion is the receiver's reaction. - """ - for i in seqs: - try: - _send_event(conn, i) - except ConnectionError: - return - - -def _send_ping(conn): - """Send a ping message.""" - send_framed(conn, encode_message({"type": "ping"})) - - -def _send_replay_complete(conn, head_seq, epoch=1): - send_framed( - conn, - encode_message( - {"type": "replay_complete", "head_seq": head_seq, "epoch": epoch} - ), - ) +from openusdconnect._client_utils import ClientCredential +from openusdconnect.codec import HelloRejectionCode, message_to_dict +from openusdconnect.protocol_constants import ( + MSG_PLAYBACK_CLAIMED, + MSG_PLAYBACK_REJECTED, + MSG_PLAYBACK_STATE, + LayerMode, +) +from openusdconnect.receiver import EventReceiver +from tests.helpers import client_registered, embedded_server, ensure_prim_event, wait_until + +FAST = {"reconnect_base_delay": 0.01, "reconnect_max_delay": 0.04} +METADATA = {"timeCodesPerSecond": 24.0, "upAxis": "Z"} + + +@pytest.fixture(scope="module") +def server(tmp_path_factory): + """A server that requires tokens and whose base authors stage metadata.""" + base = tmp_path_factory.mktemp("receiver") / "base.usda" + stage = Usd.Stage.CreateNew(str(base)) + stage.SetTimeCodesPerSecond(METADATA["timeCodesPerSecond"]) + UsdGeom.SetStageUpAxis(stage, UsdGeom.Tokens.z) + stage.GetRootLayer().Save() + with embedded_server(base_usd_path=str(base), require_token=True) as runtime: + yield runtime + + +@pytest.fixture +def receivers(server): + """Build receivers of the server that are closed when the test ends.""" + created = [] + + def make(**options): + receiver = EventReceiver( + **{ + "host": "127.0.0.1", + "port": server.server_address[1], + "client_id": uuid.uuid4().hex, + **FAST, + **options, + } + ) + created.append(receiver) + return receiver + yield make + for receiver in created: + assert receiver.close(5) -def _poll_until(predicate, timeout=2, interval=0.02): - """Poll predicate() until truthy or timeout.""" - deadline = time.monotonic() + timeout - while time.monotonic() < deadline: - result = predicate() - if result: - return result - time.sleep(interval) - return predicate() - - -def _teardown(rt, conn, srv): - """Clean shutdown: close server-side socket first so recv unblocks.""" - conn.close() - rt.stop() - rt.join(timeout=1) - srv.close() - - -class TestReceiverThread: - def test_bounded_drain_preserves_suffix_and_replay_watermark(self): - rt = ReceiverThread(reconnect=False) - rt.connected = True - connection = rt._inbox.begin_connection() - for sequence, frame in enumerate((b"one", b"two", b"three"), start=1): - rt._inbox.accept( - connection.generation, - _client_backend.ReceiverMessageKind.EVENT, - sequence, - frame, - ) - rt._inbox.accept_replay_complete(connection.generation, 3, 7) - - assert list(rt.drain_queue(max_messages=2)) == [b"one", b"two"] - assert rt.queued_message_count == 1 - assert rt._inbox.size == 1 - assert not rt.mark_replay_applied() - - assert list(rt.drain_queue(max_messages=2)) == [b"three"] - assert rt.mark_replay_applied() - assert rt.synchronized - assert rt.replay_head_seq == 3 - assert rt.replay_epoch == 7 - - @pytest.mark.parametrize("limit", [0, -1, True, 1.5]) - def test_bounded_drain_rejects_invalid_limit(self, limit): - with pytest.raises(ValueError, match="max_messages"): - ReceiverThread(reconnect=False).drain_queue(max_messages=limit) - - def test_replay_request_advances_past_discarded_queue_serials(self): - rt = ReceiverThread(reconnect=False) - connection = rt._inbox.begin_connection() - rt._inbox.accept( - connection.generation, - _client_backend.ReceiverMessageKind.EVENT, - 1, - b"stale-one", - ) - rt._inbox.accept( - connection.generation, - _client_backend.ReceiverMessageKind.EVENT, - 2, - b"stale-two", - ) - rt.request_replay_from(4) - - assert rt._inbox.size == 0 - - def test_terminal_transport_failure_wakes_connection_waiter(self, monkeypatch): - def _fail_connect(*_args, **_kwargs): - raise OSError("injected connection failure") - - monkeypatch.setattr(socket, "create_connection", _fail_connect) - rt = ReceiverThread(host="127.0.0.1", port=1, reconnect=False) - rt.start() - try: - assert not rt.wait_connected(timeout=1) - rt.join(timeout=1) - assert not rt.is_alive() - assert isinstance(rt.connection_error, OSError) - finally: - rt.stop() - rt.join(timeout=1) - - def test_token_callback_failure_does_not_abort_handshake(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - on_token_issued=lambda _token: (_ for _ in ()).throw( - RuntimeError("injected callback failure") - ), - ) - rt.start() - conn = _accept(srv) - try: - _recv_hello(conn) - send_framed( - conn, - encode_message( - { - "type": "hello_ok", - "layered_replay": True, - "token": "issued-token", - } - ), - ) - assert rt.wait_connected(timeout=1) - assert rt.connected - assert rt.token == "issued-token" - finally: - _teardown(rt, conn, srv) - - def test_connects_and_sends_hello(self): - """ReceiverThread connects and sends a hello message.""" - srv, port = _make_server() - rt = ReceiverThread(host="127.0.0.1", port=port, sync_from=5, reconnect=False) - rt.start() - conn = _accept(srv) - try: - hello = _recv_hello(conn) - assert hello["type"] == "hello" - assert hello["role"] == "receiver" - assert hello["sync_from"] == 5 - assert hello["layered_replay"] is True - finally: - _teardown(rt, conn, srv) - - def test_socket_has_nodelay(self): - """Small control frames must not sit in Nagle's buffer.""" - srv, port = _make_server() - rt = ReceiverThread(host="127.0.0.1", port=port, reconnect=False) - rt.start() - conn = _accept_and_hello(srv) - try: - _poll_until(lambda: rt.connected) - assert rt.connected - assert rt.sock.getsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY) != 0 - finally: - _teardown(rt, conn, srv) - - def test_replay_ready_only_after_preceding_frames_are_drained_and_applied(self): - srv, port = _make_server() - rt = ReceiverThread(host="127.0.0.1", port=port, reconnect=False) - rt.start() - conn = _accept_and_hello(srv) - try: - _send_event(conn, 1) - _send_replay_complete(conn, 1, epoch=4) - _poll_until(lambda: rt.last_seq == 1) - - assert not rt.synchronized - queued = rt.drain_queue() - assert len(queued) == 1 - assert message_to_dict(queued[0])["type"] == "event" - assert _poll_until(rt.mark_replay_applied) - assert rt.synchronized - assert rt.replay_head_seq == 1 - assert rt.replay_epoch == 4 - finally: - _teardown(rt, conn, srv) - - def test_reconnect_clears_ready_until_the_new_replay_marker_is_applied(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=True, - reconnect_base_delay=0.02, - reconnect_max_delay=0.05, - ) - rt.start() - conn1 = _accept_and_hello(srv) - conn2 = None - try: - _send_replay_complete(conn1, 0, epoch=1) - assert _poll_until(rt.mark_replay_applied) - assert rt.synchronized - - conn1.close() - _poll_until(lambda: not rt.connected) - assert not rt.synchronized - - conn2 = _accept_and_hello(srv, timeout=5) - _poll_until(lambda: rt.connected) - assert not rt.synchronized - _send_replay_complete(conn2, 0, epoch=1) - assert _poll_until(rt.mark_replay_applied) - assert rt.synchronized - finally: - _teardown(rt, conn2 or conn1, srv) - - def test_negotiates_layered_replay(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - layered_replay=True, - ) - rt.start() - conn = _accept(srv) - try: - hello = _recv_hello(conn) - assert hello["layered_replay"] is True - send_framed( - conn, - encode_message( - {"type": "hello_ok", "layered_replay": True}, - ), - ) - assert _poll_until(lambda: rt.connected) - assert rt.layered_replay_active is True - finally: - _teardown(rt, conn, srv) - - def test_sends_department(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - department="layout", - ) - rt.start() - conn = _accept(srv) - try: - hello = _recv_hello(conn) - assert hello["department"] == "layout" - send_framed( - conn, - encode_message({"type": "hello_ok", "layered_replay": True}), - ) - assert _poll_until(lambda: rt.connected) - finally: - _teardown(rt, conn, srv) - - def test_unacknowledged_layered_replay_rejects_handshake(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - layered_replay=True, - ) - rt.start() - conn = _accept_and_hello(srv, hello_ok={"type": "hello_ok"}) - try: - rt.join(timeout=1) - assert not rt.is_alive() - assert not rt.connected - assert rt.hello_rejected - assert rt.layered_replay_active is False - finally: - _teardown(rt, conn, srv) - - def test_explicit_flat_replay_accepts_unlayered_handshake(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - layered_replay=False, - ) - rt.start() - conn = _accept_and_hello(srv, hello_ok={"type": "hello_ok"}) - try: - assert _poll_until(lambda: rt.connected) - assert not rt.layered_replay_active - assert not rt.hello_rejected - finally: - _teardown(rt, conn, srv) - - def test_hello_rejection_stops_reconnects(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=True, - reconnect_base_delay=0.01, - ) - rt.start() - conn = _accept(srv) - try: - _recv_hello(conn) - send_framed( - conn, - encode_message( - { - "type": "hello_rejected", - "code": HelloRejectionCode.LayeredReplayRequired, - "reason": "layered replay is required", - } - ), - ) - rt.join(timeout=1) - assert not rt.is_alive() - assert not rt.connected - assert not rt.auth_rejected - assert rt.hello_rejected - assert rt.rejection_code == HelloRejectionCode.LayeredReplayRequired - assert rt.rejection_reason == "layered replay is required" - finally: - _teardown(rt, conn, srv) - - def test_auth_rejection_stores_reason(self): - srv, port = _make_server() - rt = ReceiverThread(host="127.0.0.1", port=port, reconnect=True) - rt.start() - conn = _accept(srv) - try: - _recv_hello(conn) - send_framed( - conn, - encode_message({"type": "auth_rejected", "reason": "invalid token"}), - ) - rt.join(timeout=1) - assert not rt.is_alive() - assert rt.auth_rejected - assert not rt.hello_rejected - assert rt.rejection_reason == "invalid token" - finally: - _teardown(rt, conn, srv) - - def test_receives_and_drains(self): - """ReceiverThread queues incoming FB messages for drain_queue.""" - srv, port = _make_server() - rt = ReceiverThread(host="127.0.0.1", port=port, reconnect=False) - rt.start() - conn = _accept_and_hello(srv) - try: - _send_event(conn, 1) - _send_event(conn, 2) - - collected = [] - - def _drain_all(): - collected.extend(rt.drain_queue()) - return len(collected) >= 2 - - _poll_until(_drain_all) - assert len(collected) == 2 - assert rt.last_seq == 2 - assert len(rt.drain_queue()) == 0 - finally: - _teardown(rt, conn, srv) - - def test_stop_on_server_close(self): - """ReceiverThread stops cleanly when server closes connection.""" - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - socket_timeout=0.1, - ) - rt.start() - conn = _accept_and_hello(srv) - try: - _poll_until(lambda: rt.connected) - assert rt.connected - conn.close() - rt.join(timeout=1) - assert not rt.connected - finally: - if rt.is_alive(): - rt.stop() - rt.join(timeout=1) - srv.close() - - def test_drain_empty_before_connect(self): - """drain_queue returns empty deque before any data arrives.""" - srv, port = _make_server() - rt = ReceiverThread(host="127.0.0.1", port=port, reconnect=False) - assert len(rt.drain_queue()) == 0 - rt.start() - conn = _accept_and_hello(srv) - try: - assert len(rt.drain_queue()) == 0 - finally: - _teardown(rt, conn, srv) - - -class TestReconnection: - """ReceiverThread reconnects automatically after connection loss.""" - - def test_backoff_resets_after_successful_handshake(self, monkeypatch): - class StopAfterThreeWaits: - def __init__(self): - self.waits = [] - self.stopped = False - - def is_set(self): - return self.stopped - - def wait(self, timeout): - self.waits.append(timeout) - if len(self.waits) == 3: - self.stopped = True - return True - return False - - receiver = ReceiverThread( - reconnect=True, - reconnect_base_delay=1.0, - reconnect_max_delay=8.0, - ) - stop_event = StopAfterThreeWaits() - receiver._stop_event = stop_event - attempts = 0 - - def connect(): - nonlocal attempts - attempts += 1 - if attempts < 3: - raise OSError("injected connection failure") - receiver.connected = True - - monkeypatch.setattr(receiver, "_connect_and_recv", connect) - - receiver.run() - - assert stop_event.waits == [1.0, 2.0, 1.0] - - def test_reconnects_after_server_close(self): - """After server closes, receiver reconnects to a new server.""" - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=True, - reconnect_base_delay=0.05, - reconnect_max_delay=0.2, - ) - rt.start() - - # First connection - conn1 = _accept_and_hello(srv) - _poll_until(lambda: rt.connected) - assert rt.connected - - # Server drops connection - conn1.close() - _poll_until(lambda: not rt.connected) - - # Receiver should reconnect - conn2 = _accept_and_hello(srv) - _poll_until(lambda: rt.connected, timeout=2) - assert rt.connected - - _teardown(rt, conn2, srv) - - def test_reconnect_uses_last_seq(self): - """Reconnection sends sync_from based on last received seq.""" - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=True, - reconnect_base_delay=0.05, - reconnect_max_delay=0.2, - ) - rt.start() - - # First connection send some events - conn1 = _accept_and_hello(srv) - _send_event(conn1, 10) - _poll_until(lambda: rt.last_seq == 10) - - # Drop connection - conn1.close() - _poll_until(lambda: not rt.connected) - - # Reconnect should request sync_from=11 - conn2 = _accept(srv, timeout=2) - hello = _recv_hello(conn2) - assert hello["sync_from"] == 11 - - _teardown(rt, conn2, srv) - - def test_resync_rewinds_the_reconnect_cursor(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=True, - reconnect_base_delay=0.05, - reconnect_max_delay=0.2, - ) - rt.start() - conn1 = _accept_and_hello(srv) - conn2 = None - try: - _send_event(conn1, 10) - assert _poll_until(lambda: rt.last_seq == 10) - send_framed(conn1, encode_message({"type": "resync"})) - _send_event(conn1, 1) - assert _poll_until(lambda: rt.last_seq == 1) - - conn1.close() - assert _poll_until(lambda: not rt.connected) - conn2 = _accept(srv, timeout=2) - assert _recv_hello(conn2)["sync_from"] == 2 - finally: - if conn2 is not None: - _teardown(rt, conn2, srv) - else: - rt.stop() - rt.join(timeout=1) - srv.close() - - def test_in_place_resync_clears_ready_until_new_watermark_is_applied(self): - srv, port = _make_server() - rt = ReceiverThread(host="127.0.0.1", port=port, reconnect=False) - rt.start() - conn = _accept_and_hello(srv) - try: - _send_replay_complete(conn, 0, epoch=1) - assert _poll_until(rt.mark_replay_applied) - assert rt.synchronized - - send_framed(conn, encode_message({"type": "resync"})) - _send_event(conn, 1) - _send_replay_complete(conn, 1, epoch=2) - assert _poll_until(lambda: rt.last_seq == 1) - assert not rt.synchronized - - queued = rt.drain_queue() - assert [message_to_dict(raw)["type"] for raw in queued] == [ - "resync", - "event", - ] - assert _poll_until(rt.mark_replay_applied) - assert rt.synchronized - assert rt.replay_head_seq == 1 - assert rt.replay_epoch == 2 - finally: - _teardown(rt, conn, srv) - - def test_requested_replay_rewinds_sequence_and_clears_queue(self, caplog): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=True, - reconnect_base_delay=0.05, - reconnect_max_delay=0.2, - ) - rt.start() - conn1 = _accept_and_hello(srv) - conn2 = None - try: - _send_event(conn1, 1) - _send_event(conn1, 2) - assert _poll_until(lambda: rt.last_seq == 2) - - caplog.set_level(logging.WARNING, logger="openusdconnect.receiver") - rt.request_replay_from(2) - assert rt.last_seq == 1 - assert len(rt.drain_queue()) == 0 - - conn2 = _accept(srv, timeout=2) - hello = _recv_hello(conn2) - assert hello["sync_from"] == 2 - send_framed( - conn2, - encode_message({"type": "hello_ok", "layered_replay": True}), - ) - assert _poll_until(lambda: rt.connected) - - _send_event(conn2, 2) - assert _poll_until(lambda: rt.last_seq == 2) - assert "socket error during read" not in caplog.text - finally: - conn1.close() - if conn2 is not None: - _teardown(rt, conn2, srv) - else: - rt.stop() - rt.join(timeout=1) - srv.close() - - def test_requested_replay_closes_current_socket(self): - rt = ReceiverThread() - sock = socket.socket() - rt.sock = sock - try: - rt.request_replay_from(1) - assert sock.fileno() == -1 - assert rt.sock is None - finally: - sock.close() - - def test_close_socket_logs_cleanup_errors(self, caplog): - class BrokenSocket: - def shutdown(self, _how): - raise OSError("shutdown failed") - - def close(self): - raise OSError("close failed") - - rt = ReceiverThread() - sock = BrokenSocket() - rt.sock = sock - - with caplog.at_level(logging.DEBUG, logger="openusdconnect.receiver"): - rt._close_socket() - - assert rt.sock is None - assert "socket shutdown failed during close" in caplog.text - assert "socket close failed" in caplog.text - - def test_no_reconnect_when_disabled(self): - """With reconnect=False, thread exits after connection loss.""" - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - socket_timeout=0.1, - ) - rt.start() +def _commit(server, count): + """Commit new prims, which receivers that connect later replay, and return the head.""" + state = server.sync_server + state._commit_events([ensure_prim_event(f"/P{uuid.uuid4().hex}") for _ in range(count)]) + return state.store.get_max_seq() - conn = _accept_and_hello(srv) - _poll_until(lambda: rt.connected) - conn.close() - rt.join(timeout=1) - assert not rt.is_alive() - srv.close() +def _messages(frames): + return [message_to_dict(frame) for frame in frames] -class TestSocketTimeout: - """Socket timeout prevents hanging on unresponsive server.""" +@pytest.mark.parametrize( + ("max_queue", "outcome"), [(0, pytest.raises(ValueError)), (1, does_not_raise())] +) +def test_invalid_settings_raise_value_error(max_queue, outcome): + with outcome: + EventReceiver(max_queue=max_queue) - def test_timeout_does_not_kill_connection(self): - """Socket timeout triggers but connection stays alive if server responds later.""" - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - socket_timeout=0.05, - ) - rt.start() - conn = _accept_and_hello(srv) - try: - _poll_until(lambda: rt.connected) - - # Wait longer than socket timeout - time.sleep(0.1) - - # Connection should still be alive timeout just means no data - assert rt.connected - - # Send data after timeout should still be received - _send_event(conn, 1) - collected = [] - _poll_until(lambda: collected.extend(rt.drain_queue()) or len(collected) >= 1) - assert len(collected) == 1 - finally: - _teardown(rt, conn, srv) - - -class TestBoundedQueue: - """Queue overflow triggers reconnect instead of unbounded growth.""" - - def test_queue_overflow_triggers_reconnect(self): - """When queue is full, receiver disconnects and reconnects for replay.""" - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=True, - max_queue=3, - reconnect_base_delay=0.05, - reconnect_max_delay=0.2, - ) - rt.start() - - # First connection - conn1 = _accept_and_hello(srv) - - # Send 5 events into a queue with max depth 3 overflow disconnects - # the receiver mid-flood, so tolerate the RST from its closed socket. - _flood_events(conn1, range(1, 6)) - - # Wait for overflow to trigger disconnect - _poll_until(lambda: not rt.connected, timeout=5) - - # Queue should have the events it managed to buffer (up to 3) - msgs = rt.drain_queue() - assert len(msgs) <= 3 - - # Receiver should reconnect automatically - conn2 = _accept(srv, timeout=5) - hello = _recv_hello(conn2) - assert hello["type"] == "hello" - # Should request replay from where it left off - assert hello["sync_from"] > 0 - - _teardown(rt, conn2, srv) - conn1.close() - - def test_queue_overflow_no_reconnect_when_disabled(self): - """With reconnect=False, overflow stops the thread.""" - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - max_queue=3, - socket_timeout=0.1, - ) - rt.start() - conn = _accept_and_hello(srv) - try: - _flood_events(conn, range(1, 6)) - - rt.join(timeout=5) - assert not rt.is_alive() - finally: - conn.close() - srv.close() - - -class TestPingHandling: - """Server pings are handled transparently by the receiver.""" - - def test_ping_not_queued(self): - """Ping messages from server are silently dropped, not queued.""" - srv, port = _make_server() - rt = ReceiverThread(host="127.0.0.1", port=port, reconnect=False) - rt.start() - conn = _accept_and_hello(srv) - try: - _send_event(conn, 1) - _send_ping(conn) - _send_event(conn, 2) - - collected = [] - _poll_until(lambda: collected.extend(rt.drain_queue()) or len(collected) >= 2) - assert len(collected) == 2 - assert rt.last_seq == 2 - finally: - _teardown(rt, conn, srv) - - def test_ping_resets_timeout_counter(self): - """Receiving a ping prevents consecutive timeout disconnect.""" - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - socket_timeout=0.05, - ) - import openusdconnect.receiver as recv_mod - - original = recv_mod._MAX_CONSECUTIVE_TIMEOUTS - recv_mod._MAX_CONSECUTIVE_TIMEOUTS = 4 - try: - rt.start() - conn = _accept_and_hello(srv) - _poll_until(lambda: rt.connected) - # Each wait is below four timeouts, while their sum is above it. - time.sleep(0.12) - _send_ping(conn) - time.sleep(0.12) - # Should still be alive because ping reset the counter - assert rt.connected - finally: - recv_mod._MAX_CONSECUTIVE_TIMEOUTS = original - _teardown(rt, conn, srv) - - -class TestConsecutiveTimeouts: - """Receiver disconnects after too many consecutive recv timeouts.""" - - def test_max_consecutive_timeouts(self): - srv, port = _make_server() - rt = ReceiverThread( - host="127.0.0.1", - port=port, - reconnect=False, - socket_timeout=0.02, - ) - import openusdconnect.receiver as recv_mod - - original = recv_mod._MAX_CONSECUTIVE_TIMEOUTS - recv_mod._MAX_CONSECUTIVE_TIMEOUTS = 3 - try: - rt.start() - conn = _accept_and_hello(srv) - _poll_until(lambda: rt.connected) - # Don't send anything let timeouts accumulate - rt.join(timeout=2) - assert not rt.is_alive() - finally: - recv_mod._MAX_CONSECUTIVE_TIMEOUTS = original - conn.close() - srv.close() - - -def _accept_identity_hello(receiver, *, instance="server", supported=True, sync_from=1): - return receiver._handle_handshake_message( - encode_message( - { - "type": "hello_ok", - "server_instance": instance, - "layered_replay": True, - "replay_identity": supported, - } - ), - sync_from, + +def test_consumer_calls_reject_invalid_arguments(): + receiver = EventReceiver() + for limit in (0, True): + with pytest.raises(ValueError, match="max_messages"): + receiver.drain_queue(max_messages=limit) + with pytest.raises(ValueError, match="at least 1"): + receiver.request_replay_from(0) + + +def test_settings_read_back_and_state_starts_empty(): + receiver = EventReceiver( + port=7300, sync_from=5, reconnect=False, max_queue=7, socket_timeout=2.5 ) + assert (receiver.port, receiver.sync_from) == (7300, 5) + assert (receiver.max_queue, receiver.socket_timeout) == (7, 2.5) + assert receiver.layered_replay and receiver.layer_mode is LayerMode.MANAGED + assert not receiver.reconnect + receiver.reconnect = True + assert receiver.reconnect + + assert not receiver.connected and not receiver.synchronized + assert receiver.last_seq == 4 + assert receiver.queued_message_count == 0 and len(receiver.drain_queue()) == 0 + assert receiver.stage_metadata == {} + assert not receiver.auth_rejected and not receiver.hello_rejected + assert receiver.rejection_code == HelloRejectionCode.Unspecified + assert receiver.connection_error is None + + shared = EventReceiver(layer_mode="shared_stage", layered_replay=False) + assert shared.layer_mode is LayerMode.SHARED_STAGE and not shared.layered_replay + + +def test_close_stops_the_thread_and_reports_whether_it_exited(receivers): + assert EventReceiver().close(0), "a receiver that never started has no thread" + entered, release = threading.Event(), threading.Event() + + def hold(_token): + entered.set() + assert release.wait(5) + + receiver = receivers(on_token_issued=hold) + assert not receiver.running and not receiver.stopped + started = time.monotonic() + assert not receiver.wait_connected(5) + assert not receiver.wait_synchronized(5) + assert time.monotonic() - started < 1, "a receiver that never started was waited on" + + receiver.start() + assert receiver.running and not receiver.stopped + with pytest.raises(RuntimeError, match="started once"): + receiver.start() + # The callback holds the thread, so it cannot exit yet. + assert entered.wait(5) + assert receiver.connected + assert not receiver.close(timeout=0) + assert receiver.running + + release.set() + started = time.monotonic() + assert receiver.close() + assert time.monotonic() - started < 2, "close waited for the read timeout" + assert receiver.stopped and not receiver.running and not receiver.connected + assert receiver.close(0) + with pytest.raises(RuntimeError, match="started once"): + receiver.start() + + with receivers() as receiver: + receiver.start() + assert receiver.wait_connected(5) + assert receiver.stopped and not receiver.connected + + +def test_close_from_a_callback_returns_without_waiting_for_its_own_thread(receivers): + results = [] + + def close_on_token(_token): + started = time.monotonic() + results.append(receiver.close()) + results.append(time.monotonic() - started) + + receiver = receivers(on_token_issued=close_on_token) + receiver.start() + wait_until(lambda: len(results) == 2) + closed, waited = results + assert closed is False and waited < 1 + assert receiver.close(5) + assert receiver.stopped + + +def test_collecting_a_receiver_stops_its_connection(server): + receiver = EventReceiver("127.0.0.1", server.server_address[1], client_id=uuid.uuid4().hex) + client_id = receiver.client_id + receiver.start() + assert receiver.wait_connected(5) + wait_until(lambda: client_registered(server, client_id)) + del receiver + gc.collect() + wait_until(lambda: not client_registered(server, client_id)) + + +def test_close_before_start_ends_without_connecting(receivers, server): + receiver = receivers() + assert receiver.close() + receiver.start() + wait_until(lambda: receiver.stopped) + assert not receiver.running + assert not server.sync_server.token_store.has_token(receiver.client_id) + + +def test_handshake_state_reads_through_after_connecting(receivers, server): + state = server.sync_server + head = _commit(server, 2) + receiver = receivers(origin="origin", department="layout") + receiver.start() + assert receiver.wait_connected(5) + + assert receiver.layered_replay_active + assert receiver.layer_mode_active is LayerMode.MANAGED + assert receiver.stage_metadata == METADATA + with state.clients_lock: + clients = [ + (info.role, info.client_id, info.origin, info.department) + for info in state.clients.values() + ] + assert ("receiver", receiver.client_id, "origin", "layout") in clients + wait_until(lambda: receiver.last_seq == head) + assert not receiver.synchronized and receiver.server_instance == "" + + +def test_token_provider_supplies_each_attempt_and_the_issued_token_is_presented(receivers): + credential = ClientCredential("127.0.0.1", 0, None, persist=False) + presented = [] + + def provide(): + presented.append(credential.current()) + return presented[-1] + + receiver = receivers(token_provider=provide, on_token_issued=credential.issued) + receiver.start() + assert receiver.wait_connected(5) + issued = credential.token + assert issued and receiver.token == issued + receiver.request_replay_from(1) + assert receiver.wait_connected(5), "the server rejected the issued token" + assert not receiver.auth_rejected + assert presented == [None, issued] -def _receive_identity_message(receiver, generation, **message): - return receiver._handle_data_message(encode_message(message), generation) + +def test_token_provider_failure_is_the_connection_error(receivers, caplog): + def fail(): + raise RuntimeError("injected provider failure") + + receiver = receivers(token_provider=fail, reconnect=False) + receiver.start() + wait_until(lambda: receiver.stopped) + assert not receiver.connected + assert isinstance(receiver.connection_error, RuntimeError) + assert "EventReceiver: token provider failed" in caplog.text -def test_received_prefix_identity_is_not_published_until_applied(): - receiver = ReceiverThread() - first = receiver._inbox.begin_connection() - assert _accept_identity_hello(receiver, instance="old") - assert _receive_identity_message( - receiver, first.generation, type="replay_complete", head_seq=0, epoch=2 +def test_callbacks_receive_message_dicts_once_on_the_connection_thread(receivers, server, caplog): + state = server.sync_server + calls = [] + + def record(name): + def callback(value): + calls.append((name, value, threading.get_ident())) + if name == "token": + raise RuntimeError("injected callback failure") + + return callback + + receiver = receivers( + on_token_issued=record("token"), + on_stage_metadata=record("metadata"), + on_playback_state=record("state"), + on_playback_claimed=record("claimed"), + on_playback_rejected=record("rejected"), ) - assert receiver._received_replay_identity == ("old", 2) - assert receiver.server_instance == "" - assert receiver.mark_replay_applied() - assert receiver.server_instance == "old" - - receiver.connected = False - second = receiver._inbox.begin_connection() - assert _accept_identity_hello(receiver, instance="new") - assert receiver.server_instance == "old" - assert not receiver.synchronized - assert _receive_identity_message(receiver, second.generation, type="resync") - assert receiver._received_replay_identity is None - assert _receive_identity_message( - receiver, second.generation, type="replay_complete", head_seq=0, epoch=0 + receiver.start() + assert receiver.wait_connected(5) + # The server sends its playback state after every accepted hello. + wait_until(lambda: len(calls) == 3) + claimed = {"type": MSG_PLAYBACK_CLAIMED, "leader_client_id": "leader"} + rejected = { + "type": MSG_PLAYBACK_REJECTED, + "reason": "already led", + "current_leader_client_id": "leader", + } + state.broadcast_message(claimed) + state.broadcast_message(rejected) + wait_until(lambda: len(calls) == 5) + + assert [(name, value) for name, value, _thread in calls] == [ + ("token", receiver.token), + ("metadata", METADATA), + ("state", {"type": MSG_PLAYBACK_STATE, **state.get_playback_state()}), + ("claimed", claimed), + ("rejected", rejected), + ] + assert threading.get_ident() not in {thread for _name, _value, thread in calls} + assert "EventReceiver: on_token_issued callback failed" in caplog.text + assert receiver.connected + + +def test_a_given_queue_takes_the_notifications_and_snapshot_reads_the_status(receivers, server): + with pytest.raises(ValueError, match="on_playback_state"): + EventReceiver(notifications=_client_backend.NotificationQueue(), on_playback_state=print) + head = _commit(server, 1) + notifications = _client_backend.NotificationQueue() + issued = [] + receiver = receivers( + notifications=notifications, + on_token_issued=lambda token: issued.append((token, threading.get_ident())), ) - assert not receiver.mark_replay_applied() - receiver.drain_queue() - assert receiver.mark_replay_applied() - assert receiver.server_instance == "new" - assert receiver.replay_epoch == 0 - - -def test_interrupted_reset_does_not_reuse_old_prefix_identity(): - receiver = ReceiverThread() - first = receiver._inbox.begin_connection() - assert _accept_identity_hello(receiver) - assert _receive_identity_message( - receiver, first.generation, type="replay_complete", head_seq=0, epoch=3 + receiver.start() + wait_until(lambda: issued and receiver.last_seq == head) + [(token, thread)] = issued + assert token == receiver.token and thread != threading.get_ident() + kinds = [type(notification).__name__ for notification in notifications.drain()] + assert kinds[:3] == ["TokenIssued", "StageMetadata", "Connected"] + + snapshot = receiver.snapshot() + assert (snapshot.connected, snapshot.last_sequence, snapshot.layered_replay_active) == ( + receiver.connected, + receiver.last_seq, + receiver.layered_replay_active, ) - assert receiver.mark_replay_applied() - assert _receive_identity_message(receiver, first.generation, type="resync") - assert receiver._received_replay_identity is None - assert not receiver.synchronized - receiver.connected = False - receiver._inbox.begin_connection() - assert receiver._received_replay_identity is None - assert not receiver.mark_replay_applied() -def test_explicit_full_replay_queues_reset_before_colliding_events(): - receiver = ReceiverThread() - stage = Usd.Stage.CreateInMemory() - dispatcher = EventDispatcher(receiver=receiver, adapter=UsdStageAdapter(stage)) - first = receiver._inbox.begin_connection() - assert _accept_identity_hello(receiver) - stack = {"type": "layer_stack_state", "layers": [{"layer_key": "shared"}]} - assert _receive_identity_message(receiver, first.generation, **stack) - assert _receive_identity_message( - receiver, - first.generation, - type="event", - seq=1, - layer_key="shared", - event={"k": "ensure_prim", "prim": "/Old", "typeName": "Xform"}, - ) - dispatcher.drain_and_apply() - assert stage.GetPrimAtPath("/Old") - receiver.connected = False - second = receiver._inbox.begin_connection() - assert second.sync_from == 2 +def test_replay_drains_in_batches_and_is_ready_once_marked_applied(receivers, server): + state = server.sync_server + head = _commit(server, 3) + receiver = receivers() + receiver.start() + wait_until(lambda: receiver.last_seq == head) + waited = [] + waiter = threading.Thread(target=lambda: waited.append(receiver.wait_synchronized(5))) + waiter.start() + + generation = receiver.generation + frames = list(receiver.drain_queue(max_messages=2)) + assert len(frames) == 2 and receiver.queued_message_count == head - 1 + assert not receiver.mark_replay_applied() + frames.extend(receiver.drain_queue()) + messages = _messages(frames) + assert [message["type"] for message in messages] == ["layer_stack_state"] + ["event"] * head + assert [message["seq"] for message in messages[1:]] == list(range(1, head + 1)) + assert receiver.mark_applied_through(generation, head) + # The completion marker follows the last replayed event. + wait_until(receiver.mark_replay_applied) + + waiter.join(5) + assert waited == [True] + assert receiver.synchronized and receiver.wait_synchronized(0) + assert receiver.replay_head_seq == head + assert receiver.replay_epoch == state.get_replay_token()[0] + assert receiver.server_instance == state.server_instance + + +def test_replay_request_reconnects_and_replays_from_the_requested_sequence(receivers, server): + head = _commit(server, 2) + receiver = receivers(reconnect=False) + receiver.start() + wait_until(lambda: receiver.last_seq == head) + marker = receiver.freeze_marker() + assert not receiver.drained_through(marker) + + # Reconnecting after the request shows the setting reached the connection. + receiver.reconnect = True receiver.request_replay_from(1) - third = receiver._inbox.begin_connection() - assert third.sync_from == 1 - assert _accept_identity_hello(receiver, sync_from=1) - assert _receive_identity_message(receiver, third.generation, **stack) - assert _receive_identity_message( - receiver, - third.generation, - type="event", - seq=1, - layer_key="shared", - event={"k": "ensure_prim", "prim": "/Own", "typeName": "Xform"}, + assert receiver.drained_through(marker) + assert receiver.queued_message_count == 0 and receiver.last_seq == 0 + assert receiver.wait_connected(5) + wait_until(lambda: receiver.last_seq == head) + + generation = receiver.generation + messages = _messages(receiver.drain_queue()) + assert [message["type"] for message in messages] == ( + ["resync", "layer_stack_state"] + ["event"] * head ) - assert _receive_identity_message( - receiver, third.generation, type="replay_complete", head_seq=1, epoch=0 - ) - dispatcher.drain_and_apply() - assert stage.GetPrimAtPath("/Own") - assert not stage.GetPrimAtPath("/Old") - assert dispatcher.last_seq == 1 + receiver.reset_applied_progress() + assert receiver.mark_applied_through(generation, head) + wait_until(receiver.mark_replay_applied) assert receiver.synchronized - dispatcher.close() - - -@pytest.mark.parametrize("instance", ["server", "replacement"]) -@pytest.mark.parametrize("queue_full", [False, True]) -def test_changed_hello_identity_waits_for_accepted_reset(instance, queue_full): - receiver = ReceiverThread(max_queue=1) - first = receiver._inbox.begin_connection() - receiver._received_replay_identity = ("server", 0) - assert receiver._handle_data_message( - encode_message( - { - "type": "event", - "seq": 1, - "event": {"k": "ensure_prim", "prim": "/Old", "typeName": "Xform"}, - } - ), - first.generation, - ) - receiver.drain_queue() - receiver.connected = False - second = receiver._inbox.begin_connection() - receiver._prefix_validation_requested = True - assert second.sync_from == 2 - assert receiver._handle_handshake_message( - encode_message( - { - "type": "hello_ok", - "server_instance": instance, - "replay_identity": True, - "layered_replay": True, - "replay_epoch": 1, - } - ), - second.sync_from, - second.generation, - ) - assert receiver._received_replay_identity == ("server", 0) - assert receiver.last_seq == 1 - assert receiver.server_instance == "" - assert not receiver.synchronized - if queue_full: - assert receiver._handle_data_message( - encode_message({"type": "layer_stack_state", "layers": []}), - second.generation, + + +def test_refused_connection_is_the_connection_error(receivers): + with socket.socket() as probe: + probe.bind(("127.0.0.1", 0)) + port = probe.getsockname()[1] + receiver = receivers(port=port, reconnect=False) + receiver.start() + waited = [] + waiter = threading.Thread(target=lambda: waited.append(receiver.wait_connected(10))) + waiter.start() + wait_until(lambda: receiver.stopped, timeout=10) + waiter.join(5) + assert waited == [False] + assert isinstance(receiver.connection_error, ConnectionRefusedError) + + +@pytest.mark.parametrize("rejection", ["auth", "hello"]) +def test_rejection_is_reported_and_stops_reconnecting(receivers, server, rejection): + if rejection == "auth": + client_id = uuid.uuid4().hex + server.sync_server.token_store.issue(client_id) + receiver = receivers(client_id=client_id, token="not-issued") + expected = (True, False, HelloRejectionCode.Unspecified, "invalid or missing token") + else: + receiver = receivers(layered_replay=False, layer_mode=LayerMode.SHARED_STAGE) + expected = ( + False, + True, + HelloRejectionCode.LayerModeMismatch, + "server uses 'managed' layer mode, client requested 'shared_stage'", ) - accepted = receiver._handle_data_message( - encode_message({"type": "resync"}), - second.generation, - ) - assert accepted is not queue_full - assert receiver.last_seq == (1 if queue_full else 0) - assert receiver._received_replay_identity == (("server", 0) if queue_full else (instance, 1)) - assert receiver.server_instance == "" - assert not receiver.synchronized - - -def test_replay_request_during_hello_does_not_publish_stale_identity(): - receiver = ReceiverThread(on_token_issued=lambda _token: receiver.request_replay_from(2)) - connection = receiver._inbox.begin_connection() - assert not receiver._handle_handshake_message( - encode_message( - { - "type": "hello_ok", - "token": "issued", - "server_instance": "server", - "replay_identity": True, - "layered_replay": True, - "replay_epoch": 0, - } - ), - connection.sync_from, - connection.generation, - ) - assert receiver._received_replay_identity is None - assert not receiver.connected + receiver.start() + wait_until(lambda: receiver.stopped) + assert not receiver.running and not receiver.connected + assert ( + receiver.auth_rejected, + receiver.hello_rejected, + receiver.rejection_code, + receiver.rejection_reason, + ) == expected diff --git a/tests/unit/test_recovery.py b/tests/unit/test_recovery.py new file mode 100644 index 0000000..03bb37a --- /dev/null +++ b/tests/unit/test_recovery.py @@ -0,0 +1,23 @@ +"""Rejection codes map to stable names and recovery dispositions.""" + +import pytest + +from openusdconnect import RejectionDisposition, TransactionFailure + + +@pytest.mark.parametrize( + ("code", "name", "disposition"), + [ + (0, "unknown_0", RejectionDisposition.SESSION_FATAL), + (1, "invalid_identity", RejectionDisposition.SESSION_FATAL), + (2, "unexpected_id", RejectionDisposition.SESSION_FATAL), + (3, "stale_layer_graph", RejectionDisposition.RECOVERABLE_CONFLICT), + (4, "invalid_transaction", RejectionDisposition.INVALID_OPERATION), + (9, "unknown_9", RejectionDisposition.SESSION_FATAL), + ], +) +def test_rejection_code_name_and_disposition(code, name, disposition): + failure = TransactionFailure(txn_id=1, code=code, reason="") + + assert failure.code_name == name + assert failure.disposition is disposition diff --git a/tests/unit/test_sender.py b/tests/unit/test_sender.py index eda1473..08f4e7c 100644 --- a/tests/unit/test_sender.py +++ b/tests/unit/test_sender.py @@ -1,652 +1,622 @@ -"""Tests for EventSender.""" +"""EventSender's wrapper contract, exercised against a live server.""" +import gc import socket import threading import time +import uuid +import weakref +from contextlib import contextmanager +from contextlib import nullcontext as does_not_raise import pytest +from pxr import Sdf, Usd, UsdGeom from openusdconnect import _client_backend -from openusdconnect.codec import TransactionRejectionCode, encode_message, message_to_dict -from openusdconnect.framing import recv_framed, send_framed -from openusdconnect.protocol import make_transaction_result +from openusdconnect._client_utils import ClientCredential +from openusdconnect.checkpoints import MirrorCheckpoint +from openusdconnect.codec import message_to_dict +from openusdconnect.protocol import make_claim_playback, make_playback_control from openusdconnect.protocol_constants import LayerMode -from openusdconnect.recovery import RejectionDisposition, TransactionFailure +from openusdconnect.recovery import RejectionDisposition +from openusdconnect.sdf_spec_delta import serialize_spec_fields from openusdconnect.sender import EventSender, TransactionRejectedError - - -def _make_server(): - """Create a listening socket on a random port, return (socket, port).""" - srv = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - srv.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) - srv.bind(("127.0.0.1", 0)) - srv.listen(1) - return srv, srv.getsockname()[1] - - -def _accept_and_hello_ok(srv, timeout=2): - """Accept one connection, consume hello, send hello_ok. Return (conn, hello).""" - srv.settimeout(timeout) - conn, _ = srv.accept() - conn.settimeout(timeout) - hello = message_to_dict(recv_framed(conn)) - send_framed(conn, encode_message({"type": "hello_ok"})) - return conn, hello - - -class TestEventSenderConnect: - def test_constructor_rejects_non_emitter_role(self): - with pytest.raises(ValueError, match="must be 'emitter'"): - EventSender("127.0.0.1", 1, client_id="bad-role", role="receiver") - - def test_malformed_events_fail_before_encoding_or_outbox_ownership(self): - sender = EventSender("127.0.0.1", 1, client_id="validation-client") - with pytest.raises(ValueError, match="transform fields"): - sender.send_events([{"k": "set_xform_trs", "prim": "/World/X", "fields": ["bogus"]}]) - assert sender.pending_transaction_count == 0 - assert sender._next_txn_id == 1 - - @pytest.mark.parametrize("background_send", [False, True]) - def test_handshake_and_send_events(self, background_send): - srv, port = _make_server() - sender = EventSender( - "127.0.0.1", port, client_id="test-client", background_send=background_send - ) - conn = None - try: - import threading - - result = {} - - def _serve(): - result["conn"], result["hello"] = _accept_and_hello_ok(srv) - - t = threading.Thread(target=_serve) - t.start() - assert sender.connect() is True - t.join(timeout=2) - conn = result["conn"] - - assert result["hello"]["type"] == "hello" - assert result["hello"]["client_id"] == "test-client" - assert result["hello"]["producer_session_id"] == sender.session_id - - events = [{"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"}] - assert sender.send_events(events) is True - txn = message_to_dict(recv_framed(conn)) - assert txn["type"] == "txn" - assert txn["events"] == events - finally: - sender.disconnect() - if conn is not None: - conn.close() - srv.close() - - def test_concurrent_calls_preserve_transaction_id_wire_order(self, monkeypatch): - """Socket order must match IDs even when the second caller wins the race.""" - - class _SecondCallerFirstLock: - def __init__(self): - self._lock = threading.Lock() - self._state_lock = threading.Lock() - self._calls = 0 - self._second_released = threading.Event() - - def __enter__(self): - with self._state_lock: - self._calls += 1 - call = self._calls - if call == 1: - assert self._second_released.wait(timeout=2) - self._lock.acquire() - return self - - def __exit__(self, exc_type, exc_value, traceback): - self._lock.release() - with self._state_lock: - if self._calls >= 2 and not self._second_released.is_set(): - self._second_released.set() - - sender = EventSender("127.0.0.1", 1, client_id="concurrent-client") - connection = sender._session.begin_connection() - assert ( - sender._session.accept_hello(connection.generation, 0) - == _client_backend.ProducerResult.ACCEPTED - ) - sender._socket_generation = connection.generation - sender.sock = object() - sender._send_lock = _SecondCallerFirstLock() - wire_txn_ids = [] - monkeypatch.setattr( - "openusdconnect.sender.send_raw", - lambda _sock, payload: wire_txn_ids.append(message_to_dict(payload)["txn_id"]), - ) - - results = [] - threads = [ - threading.Thread( - target=lambda prim=prim: results.append( - sender.send_events([{"k": "ensure_prim", "prim": prim, "typeName": "Xform"}]) - ) - ) - for prim in ("/World/First", "/World/Second") - ] - for thread in threads: - thread.start() - for thread in threads: - thread.join(timeout=5) - - assert all(not thread.is_alive() for thread in threads) - assert results == [True, True] - assert wire_txn_ids == [1, 2] - assert [entry[0] for entry in sender._session.entries()] == [1, 2] - - def test_socket_has_nodelay(self): - """Interactive txn frames are small; Nagle must not delay them.""" - srv, port = _make_server() - sender = EventSender("127.0.0.1", port, client_id="test-client") - conn = None - try: - import threading - - result = {} - - def _serve(): - result["conn"], _ = _accept_and_hello_ok(srv) - - t = threading.Thread(target=_serve) - t.start() - assert sender.connect() is True - t.join(timeout=2) - conn = result["conn"] - - assert sender.sock.getsockopt(socket.IPPROTO_TCP, socket.TCP_NODELAY) != 0 - finally: - sender.disconnect() - if conn is not None: - conn.close() - srv.close() - - def test_connect_failure_returns_false(self): - srv, port = _make_server() - srv.close() # nothing listening on the port anymore - sender = EventSender("127.0.0.1", port, client_id="test-client", handshake_timeout=0.5) - assert sender.connect() is False - assert sender.sock is None - - def test_layer_mode_mismatch_is_a_terminal_hello_rejection(self): - srv, port = _make_server() - sender = EventSender("127.0.0.1", port, client_id="mode-mismatch") - - def _serve(): - conn = srv.accept()[0] - try: - recv_framed(conn) - send_framed( - conn, - encode_message( - { - "type": "hello_ok", - "layer_mode": LayerMode.SHARED_STAGE.value, - } - ), - ) - finally: - conn.close() - - thread = threading.Thread(target=_serve) - thread.start() - try: - assert sender.connect() is False - assert sender.hello_rejected is True - assert "shared_stage instead of managed" in sender.rejection_reason - finally: - sender.disconnect() - thread.join(timeout=2) - srv.close() - - def test_server_highwater_ahead_of_local_session_requires_recovery(self): - srv, port = _make_server() - sender = EventSender( +from openusdconnect.server import connection +from tests.helpers import ( + client_registered, + embedded_server, + ensure_prim_event, + recorded_hellos, + wait_until, +) + +METADATA = {"timeCodesPerSecond": 24.0, "upAxis": "Z"} + + +@pytest.fixture(scope="module") +def server(tmp_path_factory): + """A managed server that requires tokens and whose base authors stage metadata.""" + base = tmp_path_factory.mktemp("sender") / "base.usda" + stage = Usd.Stage.CreateNew(str(base)) + stage.SetTimeCodesPerSecond(METADATA["timeCodesPerSecond"]) + UsdGeom.SetStageUpAxis(stage, UsdGeom.Tokens.z) + stage.GetRootLayer().Save() + with embedded_server(base_usd_path=str(base), require_token=True) as runtime: + yield runtime + + +@pytest.fixture(scope="module") +def shared_server(tmp_path_factory): + root = tmp_path_factory.mktemp("sender-shared") / "root.usda" + Sdf.Layer.CreateNew(str(root)).Save() + with embedded_server(base_usd_path=str(root), layer_mode=LayerMode.SHARED_STAGE) as runtime: + yield runtime + + +@pytest.fixture(scope="module") +def limited_server(): + """Admits one transaction per connection at once, then five per second.""" + with embedded_server(txn_rate=5.0, txn_burst=1) as runtime: + yield runtime + + +@pytest.fixture +def senders(): + """Build senders with fresh identities that are closed when the test ends.""" + created = [] + + def make(port, **options): + sender = EventSender("127.0.0.1", port, **{"client_id": uuid.uuid4().hex, **options}) + created.append(sender) + return sender + + yield make + for sender in created: + assert sender.close(5) + + +@contextmanager +def silent_listener(): + """A port that completes TCP connections and never answers a Hello.""" + with socket.create_server(("127.0.0.1", 0)) as listener: + listener.settimeout(5) + yield listener + + +def _port(runtime) -> int: + return runtime.server_address[1] + + +def _committed_through(runtime, sender) -> int: + return runtime.sync_server.producer_committed_through(sender.client_id, sender.session_id) + + +def _spec_event(path): + source = Sdf.Layer.CreateAnonymous() + Sdf.CreatePrimInLayer(source, path).specifier = Sdf.SpecifierDef + return { + "k": "set_sdf_spec_fields", + "prim": path, + "spec_path": path, + "spec_kind": "prim", + "fields": ["specifier"], + "fragment": serialize_spec_fields( + source, + path, + "prim", + ["specifier"], + stabilize_asset_paths=False, + ), + "removed": False, + } + + +@pytest.mark.parametrize( + ("handshake_timeout", "outcome"), [(0, pytest.raises(ValueError)), (1, does_not_raise())] +) +def test_invalid_settings_raise_value_error(handshake_timeout, outcome): + with outcome: + EventSender("127.0.0.1", 7300, client_id="client", handshake_timeout=handshake_timeout) + + +def test_settings_read_back_and_state_starts_empty(): + sender = EventSender( + "127.0.0.1", 7300, client_id="client", handshake_timeout=2.5, max_pending_transactions=7 + ) + assert (sender.port, sender.handshake_timeout, sender.max_pending_transactions) == ( + 7300, + 2.5, + 7, + ) + assert sender.layer_mode is LayerMode.MANAGED + with pytest.raises(AttributeError): + sender.host = "elsewhere" + assert len(sender.session_id) == 32 + + assert not sender.connected and not sender.is_connected + assert sender.layer_mode_active is LayerMode.MANAGED + assert sender.stage_metadata == {} + assert not sender.auth_rejected and not sender.hello_rejected + assert (sender.pending_transaction_count, sender.pending_event_count) == (0, 0) + assert (sender.acknowledged_transaction_count, sender.acknowledged_event_count) == (0, 0) + assert sender.drain_acknowledged_event_count() == 0 + assert sender.acknowledged_checkpoint is None + assert sender.transaction_failure is None + assert not sender.recovery_required and sender.recovery_disposition is None + assert sender.recovery_incident is None and sender.recovery_artifact is None + assert sender.cancel_connect() + sender.disconnect() + # Nothing is pending, so there is nothing to wait for. + assert sender.flush(timeout=None) + + +def test_calls_that_need_a_connection_or_a_failure_refuse_without_one(): + sender = EventSender("127.0.0.1", 7300, client_id="client") + event = ensure_prim_event("/A") + assert not sender.send_events([]) + assert not sender.send_events([event]) + assert not sender.send_message(make_claim_playback("client")) + assert not sender.claim_playback(1.0) + assert not sender.send_playback_control("play") + assert not sender.connect(timeout=0) + with pytest.raises(ValueError, match="transform fields"): + sender.send_events([{"k": "set_xform_trs", "prim": "/World/X", "fields": ["bogus"]}]) + with pytest.raises(ValueError, match="must not be empty"): + sender.repair_rejected_transaction([]) + with pytest.raises(RuntimeError, match="no rejected transaction"): + sender.repair_rejected_transaction([event]) + with pytest.raises(ValueError, match="1-128 characters"): + sender.abandon_rejected_session(session_id="s" * 129) + with pytest.raises(RuntimeError, match="no rejected producer session"): + sender.abandon_rejected_session() + assert sender.pending_transaction_count == 0 + + +def test_handshake_presents_the_identity_and_reads_state_through( + senders, server, monkeypatch, caplog +): + hellos = recorded_hellos(monkeypatch) + calls = [] + + def record(name): + def callback(value): + calls.append((name, value, threading.get_ident())) + if name == "token": + raise RuntimeError("injected callback failure") + + return callback + + sender = senders( + _port(server), + origin="origin", + department="layout", + session_id="identity-session", + on_token_issued=record("token"), + on_stage_metadata=record("metadata"), + ) + assert sender.connect(timeout=5) + wait_until(lambda: len(calls) == 2) + assert [(name, value) for name, value, _thread in calls] == [ + ("token", sender.token), + ("metadata", METADATA), + ] + assert sender.token + assert threading.get_ident() not in {thread for _name, _value, thread in calls} + assert "EventSender: on_token_issued callback failed" in caplog.text + assert sender.connected and sender.stage_metadata == METADATA + assert sender.layer_mode_active is LayerMode.MANAGED + + assert [hello["role"] for hello in hellos] == ["emitter"] + hello = hellos[0] + assert (hello["client_id"], hello["origin"], hello["department"]) == ( + sender.client_id, + "origin", + "layout", + ) + assert hello["producer_session_id"] == "identity-session" and "token" not in hello + wait_until(lambda: client_registered(server, sender.client_id)) + + +def test_token_provider_supplies_each_attempt_and_the_issued_token_is_presented( + senders, server, monkeypatch, caplog +): + hellos = recorded_hellos(monkeypatch) + credential = ClientCredential("127.0.0.1", 0, None, persist=False) + sender = senders( + _port(server), + token="stale", + token_provider=credential.current, + on_token_issued=credential.issued, + ) + assert sender.connect(timeout=5) + wait_until(lambda: credential.token is not None) + assert sender.token == credential.token + sender.disconnect() + assert sender.connect(timeout=5), sender.rejection_reason + assert not sender.auth_rejected + assert [hello.get("token") for hello in hellos] == [None, credential.token] + + def fail(): + raise RuntimeError("injected provider failure") + + failing = senders(_port(server), token_provider=fail) + assert not failing.connect(timeout=5) + assert "EventSender: token provider failed" in caplog.text + assert len(hellos) == 2 + + +def test_a_given_queue_takes_the_notifications_and_snapshot_reads_the_status(senders, server): + with pytest.raises(ValueError, match="on_stage_metadata"): + EventSender( "127.0.0.1", - port, - client_id="highwater-client", - session_id="highwater-session", + _port(server), + client_id="refused", + notifications=_client_backend.NotificationQueue(), + on_stage_metadata=print, ) + notifications = _client_backend.NotificationQueue() + issued = [] + sender = senders( + _port(server), + notifications=notifications, + on_token_issued=lambda token: issued.append((token, threading.get_ident())), + ) + assert sender.connect(timeout=5) + wait_until(lambda: issued) + [(token, thread)] = issued + assert token == sender.token and thread != threading.get_ident() + kinds = [type(notification).__name__ for notification in notifications.drain()] + assert kinds == ["TokenIssued", "StageMetadata", "Connected"] + + assert sender.send_events([ensure_prim_event("/SnapshotPrim")]) + assert sender.flush(5) + snapshot = sender.snapshot() + assert ( + snapshot.connected, + snapshot.session_id, + snapshot.pending_events, + snapshot.acknowledged_events, + ) == ( + sender.connected, + sender.session_id, + sender.pending_event_count, + sender.acknowledged_event_count, + ) - def _serve(): - conn = srv.accept()[0] - recv_framed(conn) - send_framed( - conn, - encode_message({"type": "hello_ok", "committed_through": 1}), - ) - conn.close() - - thread = threading.Thread(target=_serve) - thread.start() - try: - assert sender.connect() is False - assert sender.recovery_required is True - assert sender.recovery_disposition is RejectionDisposition.SESSION_FATAL - assert "ahead of local transaction 0" in sender.transaction_error - finally: - sender.disconnect() - thread.join(timeout=2) - srv.close() - - def test_token_callback_failure_does_not_poison_completed_handshake(self): - srv, port = _make_server() - sender = EventSender( - "127.0.0.1", - port, - client_id="callback-client", - on_token_issued=lambda _token: (_ for _ in ()).throw( - RuntimeError("injected callback failure") - ), - ) - accepted = {} - - def _serve(): - conn = srv.accept()[0] - recv_framed(conn) - send_framed( - conn, - encode_message({"type": "hello_ok", "token": "issued-token"}), - ) - accepted["conn"] = conn - - thread = threading.Thread(target=_serve) - thread.start() - try: - assert sender.connect() - assert sender.connected - assert sender.token == "issued-token" - finally: - sender.disconnect() - thread.join(timeout=2) - if conn := accepted.get("conn"): - conn.close() - srv.close() - - def test_connect_timeout_never_extends_configured_handshake_timeout(self, monkeypatch): - observed = [] - - def _fail(_endpoint, *, timeout): - observed.append(timeout) - raise TimeoutError - - monkeypatch.setattr(socket, "create_connection", _fail) - sender = EventSender( - "127.0.0.1", - 7200, - client_id="timeout-client", - handshake_timeout=0.5, - ) - assert sender.connect(timeout=0.1) is False - assert sender.connect(timeout=2.0) is False - assert sender.connect(timeout=0.0) is False - assert len(observed) == 2 - assert 0 < observed[0] <= 0.1 - assert 0 < observed[1] <= 0.5 - - @pytest.mark.parametrize("background_send", [False, True]) - def test_reconnect_replays_identical_bytes_until_duplicate_ack(self, background_send): - srv, port = _make_server() - sender = EventSender( - "127.0.0.1", - port, - client_id="test-client", - session_id="stable-session", - background_send=background_send, - ) - observed = [] - first_closed = threading.Event() - - def _serve(): - first, _ = _accept_and_hello_ok(srv) - observed.append(recv_framed(first)) - first.close() # committed outcome was lost with the connection - first_closed.set() - - second, _ = _accept_and_hello_ok(srv) - observed.append(recv_framed(second)) - send_framed( - second, - encode_message( - make_transaction_result( - 1, - status="acknowledged", - ) - ), - ) - time.sleep(0.05) - second.close() - - thread = threading.Thread(target=_serve, daemon=True) - thread.start() - try: - assert sender.connect() - events = [{"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"}] - assert sender.send_events(events) - assert first_closed.wait(timeout=2) - deadline = time.monotonic() + 2 - while sender.connected and time.monotonic() < deadline: - time.sleep(0.01) - assert not sender.connected - assert sender.pending_transaction_count == 1 - - assert sender.connect() - assert sender.flush(timeout=2) - assert observed[0] == observed[1] - assert sender.pending_transaction_count == 0 - assert sender.acknowledged_event_count == 1 - assert sender.drain_acknowledged_event_count() == 1 - finally: - sender.disconnect() - thread.join(timeout=2) - srv.close() - - @pytest.mark.parametrize("background_send", [False, True]) - def test_bounded_outbox_and_rejection_are_terminal(self, background_send): - srv, port = _make_server() - sender = EventSender( - "127.0.0.1", - port, - client_id="test-client", - session_id="bounded-session", - max_pending_transactions=1, - background_send=background_send, - ) - - def _serve(): - conn, _ = _accept_and_hello_ok(srv) - recv_framed(conn) - send_framed( - conn, - encode_message( - make_transaction_result( - 1, - status="rejected", - expected_txn_id=1, - rejection_code="unexpected_id", - reason="injected rejection", - ) - ), +def test_concurrent_submissions_commit_in_transaction_order(senders, server): + sender = senders(_port(server)) + assert sender.connect(timeout=5) + barrier = threading.Barrier(4) + accepted = [] + + def submit(worker): + barrier.wait(timeout=5) + for item in range(25): + accepted.append(sender.send_events([ensure_prim_event(f"/Order{worker}_{item}")])) + + workers = [threading.Thread(target=submit, args=(index,)) for index in range(4)] + for worker in workers: + worker.start() + for worker in workers: + worker.join(timeout=10) + assert accepted == [True] * 100 + # The server refuses an ID out of order, so success proves the order. + assert sender.flush(timeout=5) + assert sender.transaction_failure is None + assert _committed_through(server, sender) == 100 + assert (sender.pending_transaction_count, sender.pending_event_count) == (0, 0) + assert (sender.acknowledged_transaction_count, sender.acknowledged_event_count) == (100, 100) + assert sender.drain_acknowledged_event_count() == 100 + assert sender.drain_acknowledged_event_count() == 0 + + +def test_a_transaction_above_the_frame_limit_is_refused(senders, server): + sender = senders(_port(server)) + assert sender.connect(timeout=5) + oversized = { + "k": "set_connectable_input", + "prim": "/World/Shader", + "info_id": "", + "inputs": {"large": "x" * (16 * 1024 * 1024)}, + } + assert not sender.send_events([oversized]) + assert sender.pending_transaction_count == 0 and sender.connected + + +def test_acknowledged_checkpoint_names_the_servers_replay_position(senders, server): + state = server.sync_server + sender = senders(_port(server)) + assert sender.connect(timeout=5) + assert sender.acknowledged_checkpoint is None + assert sender.send_events([ensure_prim_event(f"/Checkpoint{uuid.uuid4().hex}")]) + assert sender.flush(timeout=5) + epoch, head = state.get_replay_token() + assert sender.acknowledged_checkpoint == MirrorCheckpoint(state.server_instance, epoch, head) + # A Hello acknowledges no mirror position. + sender.disconnect() + assert sender.connect(timeout=5) + assert sender.acknowledged_checkpoint is None + + +def test_flush_waits_out_the_rate_limit_and_replays_the_outbox( + senders, limited_server, monkeypatch +): + hellos = recorded_hellos(monkeypatch) + sender = senders(_port(limited_server)) + assert sender.connect(timeout=5) + started = time.monotonic() + for name in ("First", "Second"): + assert sender.send_events([ensure_prim_event(f"/{name}{uuid.uuid4().hex}")]) + assert sender.flush(timeout=5) + # The server refused the second, so it was replayed on a second connection + # once the retry window passed. + assert time.monotonic() - started >= 0.1 + assert len(hellos) == 2 + assert _committed_through(limited_server, sender) == 2 + assert sender.transaction_failure is None + + +@pytest.mark.parametrize("cancel", ["cancel_connect", "disconnect"]) +def test_cancelled_handshake_leaves_the_session_healthy(senders, server, monkeypatch, cancel): + state = server.sync_server + entered, release = threading.Event(), threading.Event() + authenticate = state.authenticate + + def held(*args): + entered.set() + assert release.wait(5) + return authenticate(*args) + + # The token is issued up front, since the held answer would issue it unseen. + client_id = uuid.uuid4().hex + sender = senders(_port(server), client_id=client_id, token=state.token_store.issue(client_id)) + monkeypatch.setattr(state, "authenticate", held) + try: + assert sender.request_connect() + # The server has the Hello and holds its answer. + assert entered.wait(5) + assert not sender.request_connect() + result = getattr(sender, cancel)() + if cancel == "cancel_connect": + assert result is False + wait_until(sender.cancel_connect) + assert not sender.connected + finally: + release.set() + monkeypatch.setattr(state, "authenticate", authenticate) + + assert sender.connect(timeout=5) + assert sender.send_events([ensure_prim_event(f"/After{uuid.uuid4().hex}")]) + assert sender.flush(timeout=5) + assert sender.transaction_failure is None + assert _committed_through(server, sender) == 1 + + +def test_request_connect_makes_one_bounded_attempt(): + with silent_listener() as listener: + sender = EventSender(*listener.getsockname(), client_id="bounded") + started = time.monotonic() + assert sender.request_connect(timeout=0.05) + assert not sender.request_connect() + peer, _ = listener.accept() + with peer: + peer.settimeout(5) + assert message_to_dict(peer.recv(4096)[4:])["type"] == "hello" + # The attempt ends at its deadline and closes the connection. + assert peer.recv(1) == b"" + assert time.monotonic() - started < 1 + assert not sender.connected and not sender.auth_rejected and not sender.hello_rejected + # Cancelling ends the backoff that the failed request started. + wait_until(sender.cancel_connect) + assert sender.request_connect() + sender.disconnect() + + +def test_connect_never_waits_past_its_timeout_or_the_handshake_timeout(): + with silent_listener() as listener: + sender = EventSender(*listener.getsockname(), client_id="slow", handshake_timeout=0.2) + for timeout, budget in ((0.05, 0.05), (5, 0.2)): + started = time.monotonic() + assert not sender.connect(timeout=timeout) + assert budget <= time.monotonic() - started < budget + 1 + started = time.monotonic() + assert not sender.connect(timeout=0) + assert time.monotonic() - started < 0.5 + + +@pytest.mark.parametrize("rejection", ["auth", "hello", "empty"]) +def test_rejection_is_reported_until_an_explicit_connect(senders, server, monkeypatch, rejection): + if rejection == "auth": + client_id = uuid.uuid4().hex + issued = server.sync_server.token_store.issue(client_id) + sender = senders(_port(server), client_id=client_id, token="not-issued") + expected = (True, False, "invalid or missing token") + else: + if rejection == "empty": + reject = connection.ConnectionHandler._reject_hello + monkeypatch.setattr( + connection.ConnectionHandler, + "_reject_hello", + lambda handler, code, _reason: reject(handler, code, ""), ) - time.sleep(0.05) - conn.close() - - thread = threading.Thread(target=_serve, daemon=True) - thread.start() - event = {"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"} - try: - assert sender.connect() - assert sender.send_events([event]) - assert not sender.send_events([event]) - with pytest.raises(TransactionRejectedError, match="injected rejection"): - sender.flush(timeout=2) - assert sender.pending_transaction_count == 1 - assert sender.transaction_error - assert sender.transaction_failure.code_name == "unexpected_id" - assert sender.transaction_failure.expected_txn_id == 1 - assert sender.recovery_disposition is RejectionDisposition.SESSION_FATAL - finally: - sender.disconnect() - thread.join(timeout=2) - srv.close() - - @pytest.mark.parametrize("background_send", [False, True]) - def test_rejection_closes_transport_and_quarantines_later_transactions(self, background_send): - srv, port = _make_server() - sender = EventSender( - "127.0.0.1", - port, - client_id="test-client", - session_id="quarantine-session", - background_send=background_send, + sender = senders(_port(server), layer_mode=LayerMode.SHARED_STAGE) + reason = ( + "connection rejected" + if rejection == "empty" + else "server uses 'managed' layer mode, client requested 'shared_stage'" ) - received = threading.Event() - - def _serve(): - conn, _ = _accept_and_hello_ok(srv) - recv_framed(conn) - recv_framed(conn) - received.set() - send_framed( - conn, - encode_message( - make_transaction_result( - 1, - status="rejected", - rejection_code="invalid_transaction", - reason="invalid first transaction", - ) - ), - ) - time.sleep(0.1) - conn.close() - - thread = threading.Thread(target=_serve, daemon=True) - thread.start() - event = {"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"} - try: - assert sender.connect() - assert sender.send_events([event]) - assert sender.send_events([event]) - assert received.wait(timeout=2) - with pytest.raises(TransactionRejectedError, match="invalid first"): - sender.flush(timeout=2) - assert sender.recovery_required - assert sender.transaction_failure.code_name == "invalid_transaction" - assert sender.recovery_disposition is RejectionDisposition.INVALID_OPERATION - assert not sender.connected - assert sender.pending_transaction_count == 2 - incident = sender.recovery_incident - assert incident is not None - assert incident.incident_id == "quarantine-session:1" - assert incident.transaction_ids == (1, 2) - assert incident.transaction_count == 2 - assert incident.event_count == 2 - artifact = sender.recovery_artifact - assert artifact is not None - assert artifact.producer_session_id == "quarantine-session" - assert [transaction.txn_id for transaction in artifact.transactions] == [1, 2] - assert all(transaction.payload for transaction in artifact.transactions) - assert not sender.send_events([event]) - finally: - sender.disconnect() - thread.join(timeout=2) - srv.close() - - @pytest.mark.parametrize( - ("rejection_code", "expected"), - [ - ("stale_layer_graph", RejectionDisposition.RECOVERABLE_CONFLICT), - ("invalid_transaction", RejectionDisposition.INVALID_OPERATION), - ("invalid_identity", RejectionDisposition.SESSION_FATAL), - ], + expected = (False, True, reason) + assert not sender.connect(timeout=5) + assert (sender.auth_rejected, sender.hello_rejected, sender.rejection_reason) == expected + assert not sender.request_connect() + assert not sender.connected + if rejection == "auth": + sender.token = issued + assert sender.connect(timeout=5) + assert not sender.auth_rejected and sender.rejection_reason == "" + + +def test_hello_highwater_contradiction_requires_abandoning_the_session(senders, server): + first = senders(_port(server), session_id="shared-session") + assert first.connect(timeout=5) + assert first.send_events([ensure_prim_event(f"/Committed{uuid.uuid4().hex}")]) + assert first.flush(timeout=5) + first.disconnect() + + # A second outbox claims the same producer session. + second = senders( + _port(server), + client_id=first.client_id, + session_id="shared-session", + token=first.token, ) - def test_rejection_exposes_recovery_disposition(self, rejection_code, expected): - srv, port = _make_server() - sender = EventSender("127.0.0.1", port, client_id="classified-client") - - def _serve(): - conn, _ = _accept_and_hello_ok(srv) - recv_framed(conn) - send_framed( - conn, - encode_message( - make_transaction_result( - 1, - status="rejected", - rejection_code=rejection_code, - reason="classified rejection", - ) - ), - ) - time.sleep(0.05) - conn.close() - - thread = threading.Thread(target=_serve, daemon=True) - thread.start() - try: - assert sender.connect() - assert sender.send_events( - [{"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"}] - ) - with pytest.raises(TransactionRejectedError) as caught: - sender.flush(timeout=2) - assert caught.value.failure is sender.transaction_failure - assert sender.recovery_disposition is expected - finally: - sender.disconnect() - thread.join(timeout=2) - srv.close() - - @pytest.mark.parametrize("background_send", [False, True]) - def test_recoverable_rejection_reuses_boundary_before_later_transactions(self, background_send): - srv, port = _make_server() - sender = EventSender( - "127.0.0.1", - port, - client_id="retry-client", - session_id="retry-session", - background_send=background_send, - ) - observed = [] - - def _serve(): - first, _ = _accept_and_hello_ok(srv) - observed.append(message_to_dict(recv_framed(first))) - observed.append(message_to_dict(recv_framed(first))) - send_framed( - first, - encode_message( - make_transaction_result( - 1, - status="rejected", - rejection_code="stale_layer_graph", - reason="layer was remapped", - ) - ), - ) - first.close() - - second, _ = _accept_and_hello_ok(srv) - observed.append(message_to_dict(recv_framed(second))) - observed.append(message_to_dict(recv_framed(second))) - send_framed(second, encode_message(make_transaction_result(2))) - time.sleep(0.05) - second.close() - - thread = threading.Thread(target=_serve, daemon=True) - thread.start() - stale = {"k": "ensure_prim", "prim": "/World/Stale", "typeName": "Xform"} - later = {"k": "ensure_prim", "prim": "/World/Later", "typeName": "Xform"} - repaired = {"k": "ensure_prim", "prim": "/World/Repaired", "typeName": "Xform"} - try: - assert sender.connect() - assert sender.send_events([stale], layer_key="old-layer") - assert sender.send_events([later], layer_key="stable-layer") - with pytest.raises(TransactionRejectedError): - sender.flush(timeout=2) - - assert sender.repair_rejected_transaction([repaired], layer_key="new-layer") == 1 - assert sender.connect() - assert sender.flush(timeout=2) - - assert [txn["txn_id"] for txn in observed] == [1, 2, 1, 2] - assert observed[2]["layer_key"] == "new-layer" - assert observed[2]["events"] == [repaired] - assert observed[3] == observed[1] - finally: - sender.disconnect() - thread.join(timeout=2) - srv.close() - - def test_nonrecoverable_rejection_cannot_be_retried(self): - sender = EventSender("127.0.0.1", 1, client_id="fatal-client") - sender._failure = TransactionFailure( - txn_id=1, - code=TransactionRejectionCode.UnexpectedId, - reason="sequence gap", - ) - - with pytest.raises(RuntimeError, match="not recoverable"): - sender.repair_rejected_transaction( - [{"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"}] - ) - - @pytest.mark.parametrize("background_send", [False, True]) - def test_abandon_rejected_session_never_replays_its_suffix(self, background_send): - srv, port = _make_server() - sender = EventSender( - "127.0.0.1", - port, - client_id="abandon-client", - session_id="rejected-session", - background_send=background_send, - ) - observed = [] - - def _serve(): - first, first_hello = _accept_and_hello_ok(srv) - observed.append(first_hello) - observed.append(message_to_dict(recv_framed(first))) - observed.append(message_to_dict(recv_framed(first))) - send_framed( - first, - encode_message( - make_transaction_result( - 1, - status="rejected", - rejection_code="invalid_transaction", - reason="injected failure", - ) - ), - ) - first.close() - - second, second_hello = _accept_and_hello_ok(srv) - observed.append(second_hello) - rebuilt = message_to_dict(recv_framed(second)) - observed.append(rebuilt) - send_framed(second, encode_message(make_transaction_result(1))) - time.sleep(0.05) - second.close() - - thread = threading.Thread(target=_serve, daemon=True) - thread.start() - rejected = {"k": "ensure_prim", "prim": "/World/Rejected", "typeName": "Xform"} - suffix = {"k": "ensure_prim", "prim": "/World/Suffix", "typeName": "Xform"} - rebuilt = {"k": "ensure_prim", "prim": "/World/Rebuilt", "typeName": "Xform"} - try: - assert sender.connect() - assert sender.send_events([rejected]) - assert sender.send_events([suffix]) - with pytest.raises(TransactionRejectedError): - sender.flush(timeout=2) - - artifact = sender.recovery_artifact - assert artifact is not None - assert sender.abandon_rejected_session(session_id="replacement-session") is artifact - assert sender.session_id == "replacement-session" - assert sender.pending_transaction_count == 0 - assert not sender.recovery_required - assert sender.recovery_incident is None - - assert sender.connect() - assert sender.send_events([rebuilt]) - assert sender.flush(timeout=2) - - assert observed[0]["producer_session_id"] == "rejected-session" - assert [observed[1]["txn_id"], observed[2]["txn_id"]] == [1, 2] - assert observed[3]["producer_session_id"] == "replacement-session" - assert observed[4]["txn_id"] == 1 - assert observed[4]["events"] == [rebuilt] - finally: - sender.disconnect() - thread.join(timeout=2) - srv.close() + assert not second.connect(timeout=5) + failure = second.transaction_failure + assert (failure.txn_id, failure.code_name) == (1, "unexpected_id") + assert failure.reason == "server producer highwater 1 is ahead of local transaction 0" + assert second.rejection_reason == failure.reason + assert not second.auth_rejected and not second.hello_rejected + assert second.recovery_required and second.acknowledged_checkpoint is None + assert second.recovery_disposition is RejectionDisposition.SESSION_FATAL + artifact = second.recovery_artifact + assert artifact.producer_session_id == "shared-session" and artifact.transactions == () + assert second.recovery_incident.incident_id == "shared-session:1" + with pytest.raises(TransactionRejectedError) as caught: + second.flush(timeout=0) + assert caught.value.failure is second.transaction_failure + assert not second.connect(timeout=5) + with pytest.raises(RuntimeError, match="unexpected_id is session_fatal, not recoverable"): + second.repair_rejected_transaction([ensure_prim_event("/Repaired")]) + with pytest.raises(ValueError, match="must differ"): + second.abandon_rejected_session(session_id="shared-session") + + assert second.abandon_rejected_session(session_id="fresh-session") is artifact + assert second.session_id == "fresh-session" + assert not second.recovery_required and second.recovery_artifact is None + assert second.rejection_reason == "" + assert second.connect(timeout=5) + assert second.send_events([ensure_prim_event(f"/Fresh{uuid.uuid4().hex}")]) + assert second.flush(timeout=5) + assert _committed_through(server, second) == 1 + + +def test_recoverable_rejection_is_repaired_at_the_same_id(senders, shared_server): + state = shared_server.sync_server + root_key = state.shared_layer_graph.root_layer_key + sender = senders(_port(shared_server), layer_mode=LayerMode.SHARED_STAGE) + assert sender.connect(timeout=5) + assert sender.send_events([_spec_event("/Stale")], layer_key="unmapped") + assert sender.send_events([_spec_event("/Later")], layer_key=root_key) + with pytest.raises(TransactionRejectedError, match="stale_layer_graph") as caught: + sender.flush(timeout=5) + + failure = sender.transaction_failure + assert caught.value.failure is failure and failure.txn_id == 1 + assert sender.transaction_error == str(failure) + assert sender.recovery_disposition is RejectionDisposition.RECOVERABLE_CONFLICT + assert sender.recovery_required and not sender.connected + assert sender.pending_transaction_count == 2 + artifact = sender.recovery_artifact + assert artifact is sender.recovery_artifact + assert artifact.producer_session_id == sender.session_id + assert [transaction.txn_id for transaction in artifact.transactions] == [1, 2] + assert artifact.layer_keys == ("unmapped", root_key) + quarantined = [message_to_dict(transaction.payload) for transaction in artifact.transactions] + assert [(message["txn_id"], message["layer_key"]) for message in quarantined] == [ + (1, "unmapped"), + (2, root_key), + ] + incident = sender.recovery_incident + assert incident.incident_id == f"{sender.session_id}:1" + assert (incident.transaction_ids, incident.event_count) == ((1, 2), 2) + assert not sender.send_events([_spec_event("/Refused")], layer_key=root_key) + assert not sender.connect(timeout=5) + + assert sender.repair_rejected_transaction([_spec_event("/Repaired")], layer_key=root_key) == 1 + assert sender.transaction_failure is None and sender.recovery_artifact is None + assert sender.connect(timeout=5) + assert sender.flush(timeout=5) + assert _committed_through(shared_server, sender) == 2 + assert state.stage.GetPrimAtPath("/Repaired") and state.stage.GetPrimAtPath("/Later") + assert not state.stage.GetPrimAtPath("/Stale") + + +def test_playback_messages_follow_the_connection(senders, server): + state = server.sync_server + sender = senders(_port(server)) + assert not sender.claim_playback(2.0) + assert sender.connect(timeout=5) + assert sender.claim_playback(2.0) + wait_until(lambda: state.get_playback_state()["leader_client_id"] == sender.client_id) + assert sender.send_playback_control("play") + wait_until(lambda: state.get_playback_state()["playing"]) + assert sender.send_message(make_playback_control("pause")) + wait_until(lambda: not state.get_playback_state()["playing"]) + sender.disconnect() + assert not sender.send_playback_control("pause") + wait_until(lambda: state.get_playback_state()["leader_client_id"] != sender.client_id) + + +def test_close_stops_the_thread_permanently(senders, server): + assert senders(_port(server)).close(0), "a sender that never connected has no thread" + sender = senders(_port(server)) + assert sender.connect(timeout=5) + wait_until(lambda: client_registered(server, sender.client_id)) + assert sender.close() + assert not sender.connected + wait_until(lambda: not client_registered(server, sender.client_id)) + assert not sender.connect(timeout=5) + assert not sender.request_connect() + assert sender.close(0) + + with senders(_port(server)) as sender: + assert sender.connect(timeout=5) + wait_until(lambda: client_registered(server, sender.client_id)) + assert sender.close(0) and not sender.connected + wait_until(lambda: not client_registered(server, sender.client_id)) + + +def test_close_writes_the_queued_transactions(): + with embedded_server() as runtime: + sender = EventSender("127.0.0.1", _port(runtime), client_id=uuid.uuid4().hex) + assert sender.connect(timeout=5) + assert sender.send_events([ensure_prim_event(f"/Queued{index}") for index in range(3)]) + assert sender.close(5) + wait_until(lambda: runtime.sync_server.get_event_count() == 3) + + +def test_collecting_a_sender_stops_its_connection(server): + sender = EventSender("127.0.0.1", _port(server), client_id=uuid.uuid4().hex) + client_id, session_id = sender.client_id, sender.session_id + assert sender.connect(timeout=5) + wait_until(lambda: client_registered(server, client_id)) + assert sender.send_events([ensure_prim_event(f"/Collected{uuid.uuid4().hex}")]) + del sender + gc.collect() + wait_until(lambda: not client_registered(server, client_id)) + wait_until(lambda: server.sync_server.producer_committed_through(client_id, session_id) == 1) + + +def test_dropping_a_sender_inside_its_callback_writes_and_stops(server): + holder, requested, dropped = {}, threading.Event(), threading.Event() + + def on_token_issued(_token): + requested.wait(5) + sender = holder.pop("sender") + assert sender.send_events([ensure_prim_event(f"/Dropped{uuid.uuid4().hex}")]) + del sender + dropped.set() + + client_id = uuid.uuid4().hex + holder["sender"] = EventSender( + "127.0.0.1", _port(server), client_id=client_id, on_token_issued=on_token_issued + ) + session_id, collected = holder["sender"].session_id, weakref.ref(holder["sender"]) + assert holder["sender"].request_connect(timeout=5) + requested.set() + assert dropped.wait(5) + gc.collect() + wait_until(lambda: server.sync_server.producer_committed_through(client_id, session_id) == 1) + wait_until(lambda: not client_registered(server, client_id)) + assert collected() is None diff --git a/tests/unit/test_sender_background.py b/tests/unit/test_sender_background.py deleted file mode 100644 index ce8a0c0..0000000 --- a/tests/unit/test_sender_background.py +++ /dev/null @@ -1,261 +0,0 @@ -"""Deterministic backpressure and worker lifecycle checks for background writes.""" - -import socket -import threading - -from openusdconnect.codec import decode_envelope, encode_message, message_to_dict, resolve_payload -from openusdconnect.framing import recv_framed, send_framed -from openusdconnect.protocol import make_transaction_result -from openusdconnect.sender import EventSender -from openusdconnect.transport import send_raw - - -def _events(name): - return [{"k": "ensure_prim", "prim": f"/World/{name}", "typeName": "Xform"}] - - -class _Socket: - def __init__(self): - self.closed = threading.Event() - - def settimeout(self, _timeout): - pass - - def setsockopt(self, *_args): - pass - - def shutdown(self, _how): - self.closed.set() - - def close(self): - self.closed.set() - - -def _mock_connections(monkeypatch, sender, sockets, *, highwater=0): - sockets = iter(sockets) - monkeypatch.setattr(socket, "create_connection", lambda *_args, **_kwargs: next(sockets)) - monkeypatch.setattr("openusdconnect.sender.send_msg", lambda *_args: None) - monkeypatch.setattr( - "openusdconnect.sender.recv_framed", - lambda _sock: encode_message({"type": "hello_ok", "committed_through": highwater}), - ) - monkeypatch.setattr(sender, "_read_results", lambda sock, _generation: sock.closed.wait(5)) - - -def _acknowledge(sender, txn_id): - _, result = resolve_payload(decode_envelope(encode_message(make_transaction_result(txn_id)))) - sender._accept_result(result, sender._socket_generation) - - -def _join(*threads): - for thread in threads: - if thread is not None: - thread.join(timeout=2) - assert not thread.is_alive() - - -def test_blocked_writer_does_not_block_submission_capacity_check_or_disconnect(monkeypatch): - sender = EventSender( - "localhost", 1, client_id="blocked-writer", background_send=True, max_pending_transactions=2 - ) - sock = _Socket() - _mock_connections(monkeypatch, sender, [sock]) - entered = threading.Event() - - def blocked_write(current, _payload): - entered.set() - assert current.closed.wait(5) - raise OSError("socket interrupted by shutdown") - - monkeypatch.setattr("openusdconnect.sender.send_raw", blocked_write) - assert sender.connect() - writer, reader = sender._writer_thread, sender._reader_thread - finished = threading.Event() - accepted = [] - events = _events("First") - try: - assert sender.send_events(events) - assert entered.wait(2) - events[0]["prim"] = "/World/ChangedAfterSubmission" - - def submit_then_disconnect(): - accepted.append(sender.send_events(_events("Second"))) - accepted.append(sender.send_events(_events("Full"))) - sender.disconnect() - finished.set() - - caller = threading.Thread(target=submit_then_disconnect) - caller.start() - assert finished.wait(2), "submission or disconnect waited on the blocked writer" - _join(caller, writer, reader) - assert accepted == [True, False] - assert not sender.connected - assert sender.pending_transaction_count == 2 - assert sender.pending_event_count == 2 - assert sender._next_txn_id == 3 - assert message_to_dict(sender._session.entries()[0][1])["events"] == _events("First") - assert sender._writer_thread is None - finally: - sock.closed.set() - sender.disconnect() - _join(writer, reader) - - -def test_concurrent_submissions_have_one_ordered_writer_and_acknowledged_capacity(monkeypatch): - sender = EventSender( - "localhost", - 1, - client_id="ordered-writer", - background_send=True, - max_pending_transactions=16, - ) - sock = _Socket() - _mock_connections(monkeypatch, sender, [sock]) - wire = [] - writers = set() - received = threading.Event() - - def write(_sock, payload): - wire.append(message_to_dict(payload)) - writers.add(threading.current_thread()) - if len(wire) == 16: - received.set() - - monkeypatch.setattr("openusdconnect.sender.send_raw", write) - assert sender.connect() - writer, reader = sender._writer_thread, sender._reader_thread - barrier = threading.Barrier(16) - accepted = [] - - def submit(index): - barrier.wait(timeout=2) - accepted.append(sender.send_events(_events(f"Item{index}"))) - - callers = [threading.Thread(target=submit, args=(index,)) for index in range(16)] - try: - for caller in callers: - caller.start() - _join(*callers) - assert accepted == [True] * 16 - assert received.wait(2) - assert [txn["txn_id"] for txn in wire] == list(range(1, 17)) - assert {txn["events"][0]["prim"] for txn in wire} == { - f"/World/Item{index}" for index in range(16) - } - assert writers == {writer} - assert not sender.send_events(_events("Full")) - _acknowledge(sender, 16) - assert sender.flush(timeout=0) - assert sender.acknowledged_event_count == 16 - assert sender.send_events(_events("AfterAcknowledgement")) - finally: - sender.disconnect() - _join(writer, reader, *callers) - - -def test_old_writer_finishing_after_reconnect_cannot_close_the_new_connection(monkeypatch): - sender = EventSender("localhost", 1, client_id="stale-writer", background_send=True) - first, second = _Socket(), _Socket() - _mock_connections(monkeypatch, sender, [first, second]) - old_finishing, release_old, replayed = (threading.Event() for _ in range(3)) - entered = threading.Event() - wire = [] - close = sender._close - - def delayed_close(*, expected=None): - if expected is first: - old_finishing.set() - assert release_old.wait(5) - close(expected=expected) - - def write(sock, payload): - if sock is first: - entered.set() - assert sock.closed.wait(5) - raise OSError("old connection failed") - wire.append(payload) - replayed.set() - - monkeypatch.setattr(sender, "_close", delayed_close) - monkeypatch.setattr("openusdconnect.sender.send_raw", write) - assert sender.connect() - first_writer, first_reader = sender._writer_thread, sender._reader_thread - try: - assert sender.send_events(_events("Retained")) - assert entered.wait(2) - original = sender._session.entries()[0][1] - sender.disconnect() - assert old_finishing.wait(2) - # A decoded acknowledgement from the disconnected reader is also stale. - _acknowledge(sender, 1) - assert sender.transaction_failure is None - assert sender.pending_transaction_count == 1 - assert sender.connect() - assert replayed.wait(2) - release_old.set() - _join(first_writer, first_reader) - assert sender.sock is second - assert not second.closed.is_set() - assert wire == [original] - _acknowledge(sender, 1) - assert sender.flush(timeout=0) - finally: - release_old.set() - writer, reader = sender._writer_thread, sender._reader_thread - sender.disconnect() - _join(first_writer, first_reader, writer, reader) - - -def test_disconnect_interrupts_a_real_socket_with_a_full_send_buffer(monkeypatch): - entered, sent, stop_server = (threading.Event() for _ in range(3)) - with socket.socket() as server: - server.bind(("127.0.0.1", 0)) - server.listen() - server.settimeout(3) - - def serve(): - conn, _ = server.accept() - with conn: - conn.settimeout(3) - recv_framed(conn) - conn.setsockopt(socket.SOL_SOCKET, socket.SO_RCVBUF, 4096) - send_framed(conn, encode_message({"type": "hello_ok"})) - assert stop_server.wait(5) # Deliberately never consume transaction bytes. - - def write(sock, payload): - entered.set() - send_raw(sock, payload) - sent.set() - - monkeypatch.setattr("openusdconnect.sender.send_raw", write) - peer = threading.Thread(target=serve) - peer.start() - sender = EventSender( - *server.getsockname(), client_id="real-backpressure", background_send=True - ) - writer = reader = None - try: - assert sender.connect() - writer, reader = sender._writer_thread, sender._reader_thread - sender.sock.setsockopt(socket.SOL_SOCKET, socket.SO_SNDBUF, 4096) - assert sender.send_events( - [ - { - "k": "set_connectable_input", - "prim": "/World/Shader", - "info_id": "", - "inputs": {"large_string": "x" * (8 * 1024 * 1024)}, - } - ] - ) - assert entered.wait(2) - assert not sent.is_set() - sender.disconnect() - _join(writer, reader) - assert not sent.is_set() - assert sender.pending_transaction_count == 1 - assert not sender.connected - finally: - stop_server.set() - sender.disconnect() - _join(writer, reader, peer) diff --git a/tests/unit/test_sender_reconnect.py b/tests/unit/test_sender_reconnect.py deleted file mode 100644 index 3882a4b..0000000 --- a/tests/unit/test_sender_reconnect.py +++ /dev/null @@ -1,337 +0,0 @@ -"""Background reconnect lifecycle and real socket handshake tests.""" - -import socket -import threading -import time -from unittest.mock import MagicMock - -import pytest - -from openusdconnect import _client_backend -from openusdconnect.codec import encode_message -from openusdconnect.framing import recv_framed, send_framed -from openusdconnect.sender import EventSender - - -def _finish(sender): - with sender._condition: - assert sender._condition.wait_for( - lambda: sender._connect_thread is None, - timeout=3, - ) - - -@pytest.mark.parametrize("cancel", ["cancel_connect", "disconnect"]) -def test_cancel_blocked_creation_prevents_late_publication(monkeypatch, cancel): - sender = EventSender("localhost", 1, client_id="cancel") - entered, release = threading.Event(), threading.Event() - sock = MagicMock() - - def create(*args, **kwargs): - entered.set() - assert release.wait(3) - return sock - - monkeypatch.setattr(socket, "create_connection", create) - try: - assert sender.request_connect() - assert entered.wait(2) - assert not sender.request_connect() - started = time.monotonic() - result = getattr(sender, cancel)() - assert time.monotonic() - started < 0.5 - if cancel == "cancel_connect": - assert result is False - assert not sender.request_connect() - finally: - release.set() - _finish(sender) - assert sender.sock is None - assert sender.cancel_connect() - sock.close.assert_called_once() - sock.sendall.assert_not_called() - assert not sender.recovery_required - - -def test_connect_timeout_includes_lock_wait(): - sender = EventSender("localhost", 1, client_id="lock") - sender._connect_lock.acquire() - try: - started = time.monotonic() - assert not sender.connect(timeout=0.03) - assert time.monotonic() - started < 0.5 - finally: - sender._connect_lock.release() - - -def test_disconnect_invalidates_synchronous_handshake(monkeypatch): - sender = EventSender("localhost", 1, client_id="sync-cancel") - entered, release = threading.Event(), threading.Event() - sock = MagicMock() - results = [] - - def create(*args, **kwargs): - entered.set() - assert release.wait(3) - return sock - - monkeypatch.setattr(socket, "create_connection", create) - worker = threading.Thread(target=lambda: results.append(sender.connect())) - worker.start() - try: - assert entered.wait(2) - sender.disconnect() - finally: - release.set() - worker.join(3) - assert not worker.is_alive() - assert results == [False] - assert not sender.connected - sock.close.assert_called_once() - - -def test_background_success_and_reconnect_replays_outbox(): - with socket.socket() as server: - server.bind(("127.0.0.1", 0)) - server.listen() - server.settimeout(3) - sender = EventSender(*server.getsockname(), client_id="replay") - try: - assert sender.request_connect() - conn, _ = server.accept() - with conn: - conn.settimeout(2) - recv_framed(conn) - send_framed(conn, encode_message({"type": "hello_ok"})) - _finish(sender) - assert sender.connected - assert not sender.request_connect() - assert sender.send_events( - [ - {"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"}, - ] - ) - original = recv_framed(conn) - sender.disconnect() - assert sender.request_connect() - conn, _ = server.accept() - with conn: - conn.settimeout(2) - recv_framed(conn) - send_framed(conn, encode_message({"type": "hello_ok"})) - assert recv_framed(conn) == original - _finish(sender) - assert sender.connected - sender.disconnect() - finally: - sender.disconnect() - _finish(sender) - - -@pytest.mark.parametrize("error_type", [OSError, RuntimeError]) -def test_failed_replay_closes_connection_and_retains_outbox(monkeypatch, error_type): - sender = EventSender("localhost", 1, client_id="failed-replay") - connection = sender._session.begin_connection() - assert ( - sender._session.accept_hello(connection.generation, 0) - == _client_backend.ProducerResult.ACCEPTED - ) - payload = encode_message( - { - "type": "txn", - "txn_id": 1, - "events": [{"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"}], - } - ) - assert ( - sender._session.append(connection.generation, 1, payload, 1, "") - == _client_backend.ProducerResult.ACCEPTED - ) - sender._session.disconnect(connection.generation) - - failed_sock, retry_sock = MagicMock(), MagicMock() - monkeypatch.setattr( - socket, "create_connection", MagicMock(side_effect=[failed_sock, retry_sock]) - ) - monkeypatch.setattr( - "openusdconnect.sender.recv_framed", - MagicMock(return_value=encode_message({"type": "hello_ok"})), - ) - send = MagicMock(side_effect=error_type("injected replay failure")) - monkeypatch.setattr("openusdconnect.sender.send_raw", send) - monkeypatch.setattr(sender, "_read_results", MagicMock()) - - if error_type is OSError: - assert not sender.connect() - else: - with pytest.raises(RuntimeError, match="injected replay failure"): - sender.connect() - - assert not sender.connected - assert sender.pending_transaction_count == 1 - assert sender._reader_thread is None - failed_sock.close.assert_called_once() - assert sender._send_lock.acquire(blocking=False) - sender._send_lock.release() - - send.side_effect = None - try: - assert sender.connect() - assert [call.args[1] for call in send.call_args_list] == [payload, payload] - finally: - sender.disconnect() - if sender._reader_thread is not None: - sender._reader_thread.join(timeout=2) - - -def test_retry_backoff_and_default_budget(monkeypatch): - sender = EventSender("localhost", 1, client_id="backoff") - clock = [100.0] - monkeypatch.setattr("openusdconnect.sender.time.monotonic", lambda: clock[0]) - attempt = MagicMock(return_value=False) - monkeypatch.setattr(sender, "_connect_attempt", attempt) - for delay in (1, 2, 4, 8, 8): - assert sender.request_connect() - _finish(sender) - assert sender._connect_retry_at == clock[0] + delay - assert not sender.request_connect() - clock[0] += delay - assert attempt.call_args.args[0] == 2.0 - attempt.return_value = True - assert sender.request_connect() - _finish(sender) - assert sender._connect_retry_delay == 1.0 - - -@pytest.mark.parametrize( - "response,flag", - [ - ({"type": "auth_rejected", "reason": "denied"}, "auth_rejected"), - ({"type": "hello_rejected", "reason": "denied"}, "hello_rejected"), - ({"type": "hello_ok", "committed_through": 9}, "recovery_required"), - ], -) -def test_terminal_handshake_stops_background_retry(response, flag): - with socket.socket() as server: - server.bind(("127.0.0.1", 0)) - server.listen() - server.settimeout(3) - sender = EventSender(*server.getsockname(), client_id="rejected") - try: - assert sender.request_connect() - conn, _ = server.accept() - with conn: - conn.settimeout(2) - recv_framed(conn) - send_framed(conn, encode_message(response)) - _finish(sender) - assert getattr(sender, flag) - sender._connect_retry_at = 0 - assert not sender.request_connect() - assert not sender.connected - finally: - sender.disconnect() - - -def test_cancel_interrupts_real_blocked_handshake(): - with socket.socket() as server: - server.bind(("127.0.0.1", 0)) - server.listen() - server.settimeout(3) - sender = EventSender(*server.getsockname(), client_id="blocked") - try: - assert sender.request_connect() - conn, _ = server.accept() - with conn: - conn.settimeout(2) - recv_framed(conn) - sender.disconnect() - _finish(sender) - assert conn.recv(1) == b"" - assert not sender.connected - finally: - sender.disconnect() - - -def test_background_handshake_honors_short_timeout(): - with socket.socket() as server: - server.bind(("127.0.0.1", 0)) - server.listen() - server.settimeout(3) - sender = EventSender(*server.getsockname(), client_id="timeout") - try: - started = time.monotonic() - assert sender.request_connect(timeout=0.05) - conn, _ = server.accept() - with conn: - conn.settimeout(2) - recv_framed(conn) - _finish(sender) - assert time.monotonic() - started < 1.0 - assert not sender.connected - assert not sender.auth_rejected - assert not sender.hello_rejected - finally: - sender.disconnect() - - -@pytest.mark.parametrize("background_send", [False, True], ids=["synchronous", "background"]) -def test_cancel_after_hello_before_publication(monkeypatch, background_send): - with socket.socket() as server: - server.bind(("127.0.0.1", 0)) - server.listen() - server.settimeout(3) - entered, release = threading.Event(), threading.Event() - sender = EventSender( - *server.getsockname(), client_id="late", background_send=background_send, - ) - accept = sender._accept_handshake_response - - def blocked(*args): - entered.set() - assert release.wait(3) - return accept(*args) - - def handshake(conn): - conn.settimeout(2) - recv_framed(conn) - send_framed(conn, encode_message({"type": "hello_ok"})) - - monkeypatch.setattr(sender, "_accept_handshake_response", blocked) - try: - assert sender.request_connect() - conn, _ = server.accept() - with conn: - handshake(conn) - assert entered.wait(2) - sender.disconnect() - release.set() - _finish(sender) - assert not sender.connected - - # The canceled attempt must not quarantine the producer session. - monkeypatch.setattr(sender, "_accept_handshake_response", accept) - assert sender.request_connect() - conn, _ = server.accept() - with conn: - handshake(conn) - _finish(sender) - assert sender.send_events( - [{"k": "ensure_prim", "prim": "/After", "typeName": "Xform"}], - ), sender.transaction_error - assert recv_framed(conn) - finally: - release.set() - sender.disconnect() - _finish(sender) - - -def test_flush_passes_remaining_budget(monkeypatch): - sender = EventSender("localhost", 1, client_id="flush") - session = MagicMock() - session.empty = False - monkeypatch.setattr(sender, "_session", session) - attempt = MagicMock(return_value=False) - monkeypatch.setattr(sender, "connect", attempt) - assert not sender.flush(timeout=0.02) - assert 0 < attempt.call_args.kwargs["timeout"] <= 0.02 diff --git a/tests/unit/test_shared_stage_client.py b/tests/unit/test_shared_stage_client.py index dfcc7fd..f8206df 100644 --- a/tests/unit/test_shared_stage_client.py +++ b/tests/unit/test_shared_stage_client.py @@ -1,12 +1,15 @@ -"""USD-native shared-stage client lifecycle without network transport.""" +"""USD-native shared-stage client lifecycle and recovery.""" from __future__ import annotations +from types import SimpleNamespace + import pytest from pxr import Ar, Sdf, Usd from openusdconnect import ClientPhase, RecoveryError from openusdconnect.codec import ReceivedEvent, TransactionRejectionCode, encode_message +from openusdconnect.protocol_constants import LayerMode from openusdconnect.recovery import ( QuarantinedTransaction, RecoveryArtifact, @@ -15,7 +18,7 @@ ) from openusdconnect.sdf_spec_delta import serialize_spec_fields from openusdconnect.shared_stage_client import SharedStageClient -from tests.helpers import PeerTraffic +from tests.helpers import PeerTraffic, connect_client, embedded_server def _create_root(path) -> Usd.Stage: @@ -24,7 +27,18 @@ def _create_root(path) -> Usd.Stage: return Usd.Stage.Open(root) +def _sender_snapshot(sender) -> SimpleNamespace: + return SimpleNamespace( + connected=sender.connected, + rejection=None, + pending_events=sender.pending_event_count, + acknowledged_events=sender.acknowledged_event_count, + ) + + class _RecoverySender: + snapshot = _sender_snapshot + def __init__(self, artifact: RecoveryArtifact): self.connected = False self.auth_rejected = False @@ -51,8 +65,9 @@ def abandon_rejected_session(self, *, session_id=None): self.pending_transaction_count = 0 return artifact - def disconnect(self): + def close(self, timeout=None): self.connected = False + return True def connect(self, timeout=None): self.connect_timeouts.append(timeout) @@ -60,6 +75,65 @@ def connect(self, timeout=None): return True +class _ReceiverStub: + """A receiver whose handshake and replay the test completes; it queues nothing.""" + + stopped = False + auth_rejected = False + hello_rejected = False + rejection = None + rejection_reason = "" + reconnect = False + generation = 1 + + def __init__(self): + self.connected = False + self.synchronized = False + + def complete_replay(self): + self.connected = self.synchronized = True + + def snapshot(self): + return self + + def start(self): + pass + + def close(self, timeout=None): + return True + + def freeze_marker(self): + return 0 + + def drained_through(self, _marker): + return True + + def drain_queue(self, max_messages=None): + return [] + + def mark_replay_applied(self): + return self.synchronized + + def mark_applied_through(self, _generation, _sequence): + return True + + def reset_applied_progress(self): + pass + + def request_replay_from(self, _seq_start): + pass + + +def _start_with_receiver(client, *, replayed=True): + """Start *client* on a receiver stub, replayed unless a stubbed recovery replays it.""" + receiver = _ReceiverStub() + if replayed: + receiver.complete_replay() + client._receiver = receiver + client.start() + return receiver + + def _stale_artifact(layer_key: str) -> RecoveryArtifact: failure = TransactionFailure( txn_id=1, @@ -178,6 +252,7 @@ def test_status_exposes_shared_stage_partial_connection(tmp_path): original_sender = client._sender class _StatusSender: + snapshot = _sender_snapshot connected = False transaction_failure = None rejection_reason = "" @@ -191,9 +266,7 @@ class _StatusSender: sender = _StatusSender() client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver._synchronized_event.set() + _start_with_receiver(client) try: assert client.status.phase is ClientPhase.CONNECTING assert client.status.receiver_connected is True @@ -311,8 +384,7 @@ def test_unresolved_layer_events_apply_after_dependency_refresh(tmp_path): } assert not client._apply_record(ReceivedEvent(seq=2, event=event, layer_key=child_key)) assert client.status.deferred_events == 1 - client._receiver.connected = True - client._receiver._synchronized_event.set() + _start_with_receiver(client) assert client.status.synchronized assert client.status.deferred_events == 1 assert client.status.deferred_layer_keys == (child_key,) @@ -529,7 +601,7 @@ def test_shared_use_server_abandons_only_after_rejected_layer_detaches( Sdf.CreatePrimInLayer(child, "/Local/Rejected") sender = _RecoverySender(_stale_artifact("layer:child")) client._sender = sender - client._started = True + receiver = _start_with_receiver(client, replayed=False) def _detach(_timeout): client._graph.apply_sublayers( @@ -544,8 +616,7 @@ def _detach(_timeout): ) client._tracker.sync_graph(force=True) client._last_seq = 2 - client._receiver.connected = True - client._receiver._synchronized_event.set() + receiver.complete_replay() monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", _detach) try: @@ -841,9 +912,7 @@ def test_shared_external_recovery_completes_a_structured_reachable_assessment( _bind_child_graph(client) sender = _RecoverySender(_stale_artifact("layer:child")) client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver._synchronized_event.set() + _start_with_receiver(client) monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) try: assessment = client.refresh_recovery_assessment() @@ -877,9 +946,7 @@ def test_shared_external_recovery_rejects_an_assessment_from_another_incident( _bind_child_graph(client) sender = _RecoverySender(_stale_artifact("layer:child")) client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver._synchronized_event.set() + _start_with_receiver(client) monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) try: assessment = client.refresh_recovery_assessment() @@ -908,9 +975,7 @@ def test_shared_external_recovery_rejects_a_stale_graph_assessment( _bind_child_graph(client) sender = _RecoverySender(_stale_artifact("layer:child")) client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver._synchronized_event.set() + _start_with_receiver(client) monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) try: assessment = client.refresh_recovery_assessment() @@ -941,7 +1006,7 @@ def test_shared_rebind_recovery_preserves_work_and_replays_clean_stage( _bind_child_graph(client) sender = _RecoverySender(_stale_artifact("layer:child")) client._sender = sender - client._started = True + receiver = _start_with_receiver(client, replayed=False) fresh_child = Sdf.Layer.CreateNew(str(tmp_path / "fresh-child.usda")) fresh_child.Save() @@ -955,8 +1020,7 @@ def _refresh(_timeout): with client._tracker.suppressed(): _bind_child_graph(client) client._last_seq = 4 - client._receiver.connected = True - client._receiver._synchronized_event.set() + receiver.complete_replay() monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", _refresh) try: @@ -1007,7 +1071,7 @@ def test_shared_rebind_recovery_resumes_after_replacement_replay_timeout( original_sender = client._sender sender = _RecoverySender(_stale_artifact("layer:child")) client._sender = sender - client._started = True + receiver = _start_with_receiver(client, replayed=False) fresh_stage = _create_root(tmp_path / "fresh-root.usda") fresh_child = Sdf.Layer.CreateNew(str(tmp_path / "fresh-child.usda")) @@ -1017,8 +1081,7 @@ def test_shared_rebind_recovery_resumes_after_replacement_replay_timeout( def refresh(_timeout): checkpoints.append(client.stage) - client._receiver.connected = True - client._receiver._synchronized_event.set() + receiver.complete_replay() if len(checkpoints) == 2: with client._tracker.suppressed(): _bind_child_graph(client) @@ -1116,9 +1179,7 @@ def test_shared_rebind_recovery_preflights_the_clean_stage(tmp_path, monkeypatch ) sender = _RecoverySender(_stale_artifact("layer:root")) client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver._synchronized_event.set() + _start_with_receiver(client) monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) clean_stage = _create_root(tmp_path / "clean-root.usda") @@ -1168,9 +1229,7 @@ def test_shared_rebind_recovery_rejects_a_detached_source_reused_by_clean_stage( sender = _RecoverySender(_stale_artifact("layer:child")) client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver._synchronized_event.set() + _start_with_receiver(client) monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) clean_stage = _create_root(tmp_path / "clean-root.usda") @@ -1190,35 +1249,51 @@ def test_shared_rebind_recovery_rejects_a_detached_source_reused_by_clean_stage( client.close() +def _prim_spec_event(path): + source = Sdf.Layer.CreateAnonymous() + Sdf.CreatePrimInLayer(source, path).specifier = Sdf.SpecifierDef + return { + "k": "set_sdf_spec_fields", + "prim": path, + "spec_path": path, + "spec_kind": "prim", + "fields": ["specifier"], + "fragment": serialize_spec_fields( + source, path, "prim", ["specifier"], stabilize_asset_paths=False, + ), + "removed": False, + } + + def test_shared_budget_releases_local_edits_under_sustained_traffic(tmp_path, monkeypatch): + base = tmp_path / "server-root.usda" + Sdf.Layer.CreateNew(str(base)).Save() stage = _create_root(tmp_path / "root.usda") - client = SharedStageClient(stage, app_name="shared-budget", persist_token=False) - traffic = PeerTraffic(client._receiver, monkeypatch, queued=3) sent = [] - - try: - client._started = True - client._graph.apply_state({ - "type": "layer_graph_state", "seq": 1, "generation": "graph-1", - "revision": 1, "root_layer_key": "layer:root", - "layers": [{"layer_key": "layer:root", "revision": 1, "sublayers": []}], - }) - client._tracker.sync_graph(force=True) - client._receiver.connected = True - client._receiver._synchronized_event.set() - monkeypatch.setattr(client._sender, "sock", object()) - monkeypatch.setattr( - client._sender, "send_events", - lambda events, layer_key="": sent.append(events) or True, + with embedded_server(base_usd_path=str(base), layer_mode=LayerMode.SHARED_STAGE) as server: + state = server.sync_server + client = SharedStageClient( + stage, app_name="shared-budget", persist_token=False, port=server.server_address[1], ) - stage.DefinePrim("/Shared", "Xform") - - submitted = [] - for _ in range(2): - submitted.append(client.update(max_messages=2).submitted_events) - traffic.arrive(2) - assert submitted[0] == 0 and submitted[1] > 0 - assert sent - finally: - client._sender.sock = None - client.close() + traffic = PeerTraffic( + state, client._receiver, _prim_spec_event, + layer_key=state.shared_layer_graph.root_layer_key, + ) + try: + connect_client(client) + traffic.arrive(3) + assert client._sender.connect(timeout=5) + monkeypatch.setattr( + client._sender, "send_events", + lambda events, layer_key="": sent.append(events) or True, + ) + stage.DefinePrim("/Shared", "Xform") + + submitted = [] + for _ in range(2): + submitted.append(client.update(max_messages=2).submitted_events) + traffic.arrive(2) + assert submitted[0] == 0 and submitted[1] > 0 + assert sent + finally: + client.close() diff --git a/tests/unit/test_transaction_checkpoint.py b/tests/unit/test_transaction_checkpoint.py index 1ba0f07..099e968 100644 --- a/tests/unit/test_transaction_checkpoint.py +++ b/tests/unit/test_transaction_checkpoint.py @@ -1,14 +1,12 @@ -"""Optional wire checkpoints and producer acknowledgement ownership.""" +"""Optional wire checkpoints the server captures for durable acknowledgements.""" from unittest.mock import Mock import pytest -from openusdconnect import _client_backend -from openusdconnect.checkpoints import MirrorCheckpoint, TransactionCheckpoint -from openusdconnect.codec import decode_envelope, encode_message, message_to_dict, resolve_payload +from openusdconnect.checkpoints import TransactionCheckpoint +from openusdconnect.codec import encode_message, message_to_dict from openusdconnect.protocol import make_transaction_result -from openusdconnect.sender import EventSender from openusdconnect.server import UsdSyncServer from openusdconnect.server.transactions import TransactionRequest @@ -148,60 +146,6 @@ def test_optional_checkpoint_roundtrip(checkpoint): assert "checkpoint" not in decoded -def test_checkpoint_requires_current_ack_and_no_pending_transactions(monkeypatch): - sender = EventSender("127.0.0.1", 1, client_id="checkpoint") - connection = sender._session.begin_connection() - generation = connection.generation - assert sender._session.accept_hello(generation, 0) == _client_backend.ProducerResult.ACCEPTED - sender._socket_generation = generation - sender._server_instance = "server" - sender.sock = object() - monkeypatch.setattr("openusdconnect.sender.send_raw", lambda *args: None) - - def ack(txn_id, checkpoint=None): - envelope = decode_envelope( - encode_message(make_transaction_result(txn_id, checkpoint=checkpoint)) - ) - return resolve_payload(envelope)[1] - - event = {"k": "ensure_prim", "prim": "/Own", "typeName": "Xform"} - assert sender.send_events([event]) - assert sender.acknowledged_checkpoint is None - sender._accept_result(ack(1, TransactionCheckpoint(2, 8)), generation) - assert sender.acknowledged_checkpoint == MirrorCheckpoint("server", 2, 8) - sender._accept_result(ack(1, TransactionCheckpoint(9, 999)), generation + 1) - assert sender.acknowledged_checkpoint == MirrorCheckpoint("server", 2, 8) - assert sender.send_events([event]) - assert sender.acknowledged_checkpoint is None - sender._accept_result(ack(2), generation) - assert sender.flush(timeout=0) - assert sender.acknowledged_checkpoint is None - sender.sock = None - - -def test_hello_highwater_recovery_does_not_confirm_mirror(monkeypatch): - sender = EventSender("127.0.0.1", 1, client_id="checkpoint") - generation = sender._session.begin_connection().generation - sender._session.accept_hello(generation, 0) - sender._socket_generation = generation - sender.sock = Mock() - monkeypatch.setattr("openusdconnect.sender.send_raw", lambda *args: None) - try: - assert sender.send_events([{"k": "ensure_prim", "prim": "/Own", "typeName": "Xform"}]) - sender._acknowledged_checkpoint = MirrorCheckpoint("old-instance", 0, 10) - generation = sender._session.begin_connection().generation - envelope = decode_envelope(encode_message({ - "type": "hello_ok", "server_instance": "new-instance", "committed_through": 1, - })) - assert sender._accept_handshake_response( - sender.sock, envelope, envelope.PayloadType(), generation, - ) - assert sender._session.empty - assert sender.acknowledged_checkpoint is None - finally: - sender.sock = None - - def test_duplicate_after_purge_has_no_original_visibility_proof(tmp_path, monkeypatch): state = UsdSyncServer(log_path=str(tmp_path / "checkpoint.db")) try: diff --git a/tests/unit/test_unreal_protocol_consistency.py b/tests/unit/test_unreal_protocol_consistency.py index 6d9827e..11fb89d 100644 --- a/tests/unit/test_unreal_protocol_consistency.py +++ b/tests/unit/test_unreal_protocol_consistency.py @@ -28,103 +28,78 @@ def test_native_wire_versions_match_python_core(): assert protocol and int(protocol.group(1)) == PROTOCOL_VERSION -def test_native_unreal_reliability_architecture_stays_explicit(): +def test_native_unreal_plugin_runs_on_the_client_engine(): root = Path(__file__).resolve().parents[2] plugin = root / "integrations" / "unreal" / "OpenUSDConnect" / "Source" - emitter = (plugin / "OpenUSDConnect" / "Private" / "EmitClient.h").read_text(encoding="utf-8") - emitter_source = (plugin / "OpenUSDConnect" / "Private" / "EmitClient.cpp").read_text( - encoding="utf-8" - ) - framing = (plugin / "OpenUSDConnect" / "Private" / "USDWireFraming.h").read_text( - encoding="utf-8" - ) - receiver_source = (plugin / "OpenUSDConnect" / "Private" / "SyncClient.cpp").read_text( - encoding="utf-8" - ) - transaction_builder = (plugin / "OpenUSDConnect" / "Private" / "TxnBuilder.cpp").read_text( - encoding="utf-8" - ) - core = ( - root - / "native" - / "client_core" - / "include" - / "openusdconnect" - / "client" - / "producer_session.h" - ) - protocol = ( - root - / "native" - / "client_core" - / "include" - / "openusdconnect" - / "client" - / "protocol_codec.h" - ) - protocol_source = protocol.read_text(encoding="utf-8") - receiver = (plugin / "OpenUSDConnect" / "Private" / "SyncClient.h").read_text(encoding="utf-8") - subsystem = (plugin / "OpenUSDConnect" / "Private" / "USDConnectSubsystem.cpp").read_text( - encoding="utf-8" - ) + private = plugin / "OpenUSDConnect" / "Private" + client = root / "native" / "client_core" / "include" / "openusdconnect" / "client" + runner = (private / "EndpointRunner.h").read_text(encoding="utf-8") + subsystem = (private / "USDConnectSubsystem.cpp").read_text(encoding="utf-8") subsystem_header = (plugin / "OpenUSDConnect" / "Public" / "USDConnectSubsystem.h").read_text( encoding="utf-8" ) + transaction_builder = (private / "TxnBuilder.cpp").read_text(encoding="utf-8") + protocol_source = (client / "protocol_codec.h").read_text(encoding="utf-8") applier = (plugin / "OpenUSDConnectPXR" / "Public" / "USDEventApplier.h").read_text( encoding="utf-8" ) + plugin_sources = "\n".join( + path.read_text(encoding="utf-8") + for module in ("OpenUSDConnect", "OpenUSDConnectPXR") + for path in (plugin / module).rglob("*") + if path.suffix in {".h", ".cpp"} + ) + + # The connection protocol lives in the client core's endpoints; the plugin + # only moves bytes, keeps no replay or outbox state, and builds no handshake. + assert "template \nclass FEndpointRunner final : public FRunnable" in runner + assert "Target.TakeActions()" in runner + assert "Target.OnReadTimeout()" in runner + for retired in ( + "ReceiverReplayIdentity", + "OrderedReceiverSession", + "OrderedProducerSession", + "BuildHelloFrame", + "HandshakeResponseView", + "AsyncTask", + ): + assert retired not in plugin_sources, retired + for retired in ("EmitClient.h", "SyncClient.h", "USDWireFraming.h"): + assert not (private / retired).exists() + + assert "TSharedPtr Receiver" in subsystem_header + assert "TSharedPtr Producer" in subsystem_header + for queue in ("ReceiverNotifications", "ProducerNotifications"): + assert f"TSharedPtr {queue}" in subsystem_header + # The receive path follows the engine's drain contract. + for call in ( + "Receiver->Generation()", + "Receiver->DrainFrames(1)", + "Receiver->MarkAppliedThrough(Generation, Seq)", + "Receiver->ResetAppliedProgress()", + "Receiver->MarkReplayApplied()", + "Receiver->RequestReplayFrom(", + "FUSDEventApplier::ApplyValidatedFrame(", + "Queue.Drain()", + ): + assert call in subsystem, call + # A transaction ID is paired with the frame that encodes it under one lock. + submit = subsystem[subsystem.index("bool UUSDConnectSubsystem::SubmitTransaction") :] + lock = submit.index("FScopeLock Lock(&SubmitCS)") + assert lock < submit.index("Producer->NextTransactionId()") < submit.index("Producer->Append(") + assert "bOwnEcho" not in subsystem - assert "class FProducerEndpointState" in emitter - assert "OrderedProducerSession" in emitter - assert "TSharedPtr" in emitter - assert "std::vector" not in emitter - assert "std::shared_ptr" not in emitter - assert "PendingTxns" not in emitter - assert "Session.AcknowledgeThrough" in emitter_source - assert "Session.ClaimNextUnsent" in emitter_source - assert "MakeShared(MoveTemp(Frame))" in emitter_source - assert "FinishEnvelopeFrame(Builder, RootOffset)" in framing - assert "Builder.Release()" in framing - assert "FinishEnvelopeBuffer(builder, envelope)" in protocol_source - assert "FinishSizePrefixedEnvelopeBuffer(builder, envelope)" not in protocol_source - assert "builder.PushBytes(header, kFrameHeaderSize)" in protocol_source - assert "WriteFrameHeader(payload_size, header, max_frame_size)" in protocol_source - assert "EncodeFrameInto" not in framing - assert core.is_file() - assert protocol.is_file() - assert "BuildHelloFrame(Builder, Parameters)" in framing - assert "ReplayPrefixClaim ReplayPrefix" in framing - assert "ReplayPrefixClaim ReplayPrefix" in protocol_source assert "FinishTransactionFrame(" in transaction_builder assert "BuildXformTrsEvent(" in transaction_builder assert "BuildVisibilityEvent(" in transaction_builder assert "BuildConnectableInputValue(" in transaction_builder - assert "HandshakeResponseView" in emitter_source - assert "ControlMessageView" in emitter_source + assert "FinishEnvelopeBuffer(builder, envelope)" in protocol_source + assert "FinishSizePrefixedEnvelopeBuffer(builder, envelope)" not in protocol_source + assert "builder.PushBytes(header, kFrameHeaderSize)" in protocol_source + assert "WriteFrameHeader(payload_size, header, max_frame_size)" in protocol_source assert "std::vector" not in protocol_source - assert "TSharedPtr ProducerState" in subsystem_header - assert "NextProducerTxnId" not in subsystem_header - assert "struct FValidatedReceiverFrame" in receiver - assert "OrderedReceiverSession" in receiver - assert "FReceiverSession ReceiverSession" in receiver - assert "ReceiverReplayIdentity ReplayIdentityState" in receiver - assert "ReplayIdentityState.BeginConnection()" in receiver_source - assert "ReplayIdentityState.AcceptHello(" in receiver_source - assert "ReplayIdentityState.AcceptResync()" in receiver_source - assert "ReplayIdentityState.AcceptReplayComplete(" in receiver_source - assert "ReplayIdentityState.MarkReplayApplied()" in receiver_source - assert "FQueuedReceiverFrame" not in subsystem_header - assert "OnReceiverReplayGenerationChanged" in subsystem - assert "RequestReceiverReplay(" in subsystem - assert "SyncClient->TryPopFrame(Frame)" in subsystem - assert "DrainFrames" not in subsystem - assert "bOwnEcho" not in subsystem - assert "FUSDEventApplier::ApplyValidatedFrame(Frame.Bytes" in subsystem - assert "FUSDEventApplier::FrameUsesChangeBlock(Frame" not in subsystem - assert "static bool ApplyFrame" in applier + assert "static bool EventUsesChangeBlock" in applier assert "static bool ApplyValidatedFrame" in applier - assert "WorkEvent->Trigger()" in emitter_source - assert "WorkEvent->Wait(WaitMilliseconds)" in emitter_source def test_native_unreal_department_receiver_fails_closed(): diff --git a/tests/unit/test_unreal_test_harness.py b/tests/unit/test_unreal_test_harness.py index 45a281f..3daffe1 100644 --- a/tests/unit/test_unreal_test_harness.py +++ b/tests/unit/test_unreal_test_harness.py @@ -298,44 +298,64 @@ def test_plugin_fingerprint_ignores_non_build_documentation(tmp_path): assert _plugin_fingerprint(plugin) != fingerprint -def test_unreal_source_staging_vendors_canonical_client_core(tmp_path): +def _fake_client_core(root: Path) -> Path: + for relative in ( + "include/openusdconnect/client/producer_session.h", + "src/frame_codec.cpp", + "src/engine/receiver_endpoint.cpp", + "src/driver/threaded_producer_driver.cpp", + "src/platform/bsd_socket.cpp", + "CMakeLists.txt", + ): + path = root / relative + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(f"// canonical {relative}", encoding="utf-8") + return root + + +def test_unreal_source_staging_vendors_only_the_compiled_client_core(tmp_path): plugin = tmp_path / "plugin" - (plugin / "Source").mkdir(parents=True) - (plugin / "OpenUSDConnect.uplugin").write_text("{}\n", encoding="utf-8") - stale = ( - plugin - / "Source" - / "ThirdParty" - / "OpenUSDConnectClientCore" - / "include" - / "openusdconnect" - / "client" - / "producer_session.h" - ) + module = plugin / "Source" / "OpenUSDConnectClientCore" + build_rules = module / "OpenUSDConnectClientCore.Build.cs" + stale = module / "src" / "engine" / "removed_endpoint.cpp" stale.parent.mkdir(parents=True) - stale.write_text("// stale staged core\n", encoding="utf-8") - core = tmp_path / "client_core" - header = core / "include" / "openusdconnect" / "client" / "producer_session.h" - header.parent.mkdir(parents=True) - header.write_text("// canonical core\n", encoding="utf-8") - - staged = _stage_plugin_source( - plugin, - tmp_path / "staged", - client_core_source=core, + stale.write_text("// stale staged core", encoding="utf-8") + build_rules.write_text("// module rules", encoding="utf-8") + (plugin / "OpenUSDConnect.uplugin").write_text("{}", encoding="utf-8") + core = _fake_client_core(tmp_path / "client_core") + + staged = _stage_plugin_source(plugin, tmp_path / "staged", client_core_source=core) + + staged_module = staged / "Source" / "OpenUSDConnectClientCore" + assert sorted( + path.relative_to(staged_module).as_posix() + for path in staged_module.rglob("*") + if path.is_file() + ) == [ + "OpenUSDConnectClientCore.Build.cs", + "include/openusdconnect/client/producer_session.h", + "src/engine/receiver_endpoint.cpp", + "src/frame_codec.cpp", + ] + header = staged_module / "include" / "openusdconnect" / "client" / "producer_session.h" + assert header.read_text(encoding="utf-8") == ( + "// canonical include/openusdconnect/client/producer_session.h" ) - staged_header = ( - staged - / "Source" - / "ThirdParty" - / "OpenUSDConnectClientCore" - / "include" - / "openusdconnect" - / "client" - / "producer_session.h" - ) - assert staged_header.read_text(encoding="utf-8") == "// canonical core\n" + +def test_plugin_fingerprint_covers_only_the_staged_client_core(tmp_path): + plugin = tmp_path / "plugin" + plugin.mkdir() + (plugin / "OpenUSDConnect.uplugin").write_text("{}", encoding="utf-8") + core = _fake_client_core(tmp_path / "client_core") + fingerprint = _plugin_fingerprint(plugin, core) + + (core / "src" / "driver" / "threaded_producer_driver.cpp").write_text("// edited") + (core / "CMakeLists.txt").write_text("# edited") + assert _plugin_fingerprint(plugin, core) == fingerprint + + (core / "src" / "engine" / "receiver_endpoint.cpp").write_text("// edited") + assert _plugin_fingerprint(plugin, core) != fingerprint def test_engine_fingerprint_distinguishes_builds_and_installations(tmp_path): diff --git a/tests/unit/test_usd_client.py b/tests/unit/test_usd_client.py index 6bbd114..b716f59 100644 --- a/tests/unit/test_usd_client.py +++ b/tests/unit/test_usd_client.py @@ -2,6 +2,8 @@ from __future__ import annotations +from types import SimpleNamespace + import pytest from pxr import Sdf, Usd, UsdGeom @@ -22,7 +24,25 @@ make_recovery_incident, ) from openusdconnect.usd_client import UsdPublisher, UsdReceiver -from tests.helpers import PeerTraffic, RecordingObserver, force_handshake +from tests.helpers import ( + PeerTraffic, + RecordingObserver, + connect_client, + embedded_server, + ensure_prim_event, + wait_until, +) + + +@pytest.fixture(scope="module") +def server(): + with embedded_server() as runtime: + yield runtime + + +@pytest.fixture +def port(server): + return server.server_address[1] class _SenderStub: @@ -48,6 +68,14 @@ def __init__(self, results: list[bool]): self.abandoned_session_ids: list[str | None] = [] self.connect_requests = 0 + def snapshot(self): + return SimpleNamespace( + connected=self.connected, + rejection=None, + pending_events=self.pending_event_count, + acknowledged_events=self.acknowledged_event_count, + ) + def send_events(self, events: list[dict]) -> bool: self.batches.append(events) accepted = next(self.results) @@ -74,6 +102,10 @@ def disconnect(self) -> None: self.connected = False self.disconnect_count += 1 + def close(self, timeout=None) -> bool: + self.connected = False + return True + def flush(self, timeout=None) -> bool: return True @@ -199,7 +231,7 @@ def test_receiver_rebinds_an_explicit_usd_stage_adapter(): receiver.close() -def test_receiver_surfaces_and_acknowledges_native_scene_rebuild(): +def test_receiver_surfaces_and_acknowledges_native_scene_rebuild(port): class _ProjectionState: native_scene_rebuild_required = True @@ -219,13 +251,12 @@ def close(self): adapter=MockAdapter(), persist_token=False, reconnect=False, + port=port, ) - state = _ProjectionState() - receiver._dispatcher._projection_state = state - receiver._started = True - receiver._receiver.connected = True - receiver._receiver._synchronized_event.set() try: + connect_client(receiver) + state = _ProjectionState() + receiver._dispatcher._projection_state = state assert receiver.status.phase is ClientPhase.RECOVERY_REQUIRED assert "must be rebuilt" in receiver.status.reason @@ -237,21 +268,25 @@ def close(self): receiver.close() -def test_receiver_status_distinguishes_connecting_replay_and_ready(): +def test_receiver_status_distinguishes_connecting_replay_and_ready(server): receiver = UsdReceiver( Usd.Stage.CreateInMemory(), app_name="status-receiver", persist_token=False, reconnect=False, + port=server.server_address[1], ) try: assert receiver.status.phase is ClientPhase.OFFLINE - receiver._started = True - assert receiver.status.phase is ClientPhase.CONNECTING - receiver._receiver.connected = True + # While a transaction holds the server's barrier, it cannot answer the hello. + with server.sync_server.txn_barrier.shared(): + receiver.start() + assert receiver.status.phase is ClientPhase.CONNECTING + assert receiver.receiver.wait_connected(5) assert receiver.status.phase is ClientPhase.REPLAYING - receiver._receiver._synchronized_event.set() - assert receiver.status.phase is ClientPhase.READY + wait_until( + lambda: receiver.update() is not None and receiver.status.phase is ClientPhase.READY + ) assert receiver.status.receiver_connected is True assert receiver.status.sender_connected is None finally: @@ -297,20 +332,21 @@ def test_publisher_context_start_is_nonblocking_and_update_connects_in_backgroun assert publisher.status.phase is ClientPhase.CLOSED -def test_managed_status_exposes_partial_connection_and_event_counts(): +def test_managed_status_exposes_partial_connection_and_event_counts(port): client = ManagedClient( Usd.Stage.CreateInMemory(), app_name="status-managed", persist_token=False, reconnect=False, + port=port, ) sender = _SenderStub([]) sender.connected = False sender.pending_event_count = 4 sender.acknowledged_event_count = 7 client._sender = sender - force_handshake(client, synchronized=True) try: + connect_client(client) status = client.status assert status.phase is ClientPhase.CONNECTING assert status.connected is False @@ -327,15 +363,15 @@ def test_managed_status_exposes_partial_connection_and_event_counts(): @pytest.fixture -def ready_managed_client(monkeypatch): +def ready_managed_client(monkeypatch, port): client = ManagedClient( - Usd.Stage.CreateInMemory(), app_name="managed-api", persist_token=False, + Usd.Stage.CreateInMemory(), app_name="managed-api", persist_token=False, port=port, ) sender = _SenderStub([True] * 10) client._sender = sender - force_handshake(client, synchronized=True) - monkeypatch.setattr(client._dispatcher, "drain_and_apply", lambda max_messages=None: 0) try: + connect_client(client) + monkeypatch.setattr(client._dispatcher, "drain_and_apply", lambda max_messages=None: 0) yield client, sender finally: client.close() @@ -375,13 +411,15 @@ def test_managed_metadata_only_changes_count_as_unsent_work(ready_managed_client def test_managed_snapshot_waits_for_replay_without_losing_newer_edits( - ready_managed_client, monkeypatch, + ready_managed_client, server, monkeypatch, ): client, sender = ready_managed_client starts = [] client._started = False monkeypatch.setattr(client._receiver, "start", lambda: starts.append(True)) - client._receiver._synchronized_event.clear() + # A purge resets the connected receiver in place. + server.sync_server.purge() + wait_until(lambda: not client.receiver.synchronized) value = UsdGeom.Sphere.Define(client.stage, "/Local").GetRadiusAttr() value.Set(1) @@ -392,7 +430,11 @@ def test_managed_snapshot_waits_for_replay_without_losing_newer_edits( with pytest.raises(RuntimeError, match="earlier publisher batch"): client.publish_current_edit_target() value.Set(2) - client._receiver._synchronized_event.set() + # The stubbed dispatcher applies nothing, so take the replay directly. + wait_until( + lambda: client.receiver.drain_queue() is not None + and client.receiver.mark_replay_applied() + ) assert client.update().submitted_events > 0 assert client.status.has_unsent_changes assert client.update().submitted_events > 0 @@ -449,28 +491,29 @@ def test_managed_parked_client_is_not_ready(ready_managed_client): def test_managed_queued_callback_can_close_before_stage_work(monkeypatch): - client = ManagedClient( - Usd.Stage.CreateInMemory(), app_name="close-from-callback", persist_token=False, - observer=RecordingObserver(on_call=lambda _name, _value: client.close()), - ) - client._started = True - monkeypatch.setattr( - client.dispatcher, "drain_and_apply", lambda: pytest.fail("closed client applied work"), - ) - try: - client.receiver._on_playback_state( - {"playing": False, "time": 0.0, "rate": 1.0, "leader_client_id": ""} + with embedded_server(require_token=True) as server: + client = ManagedClient( + Usd.Stage.CreateInMemory(), app_name="close-from-callback", + port=server.server_address[1], persist_token=False, + observer=RecordingObserver(on_call=lambda _name, _value: client.close()), ) - assert client.update() == SyncUpdate(applied_events=0, submitted_events=0) - assert client.status.phase is ClientPhase.CLOSED - with pytest.raises(RuntimeError, match="ManagedClient is closed"): - client.update() - finally: - client.close() + monkeypatch.setattr( + client.dispatcher, "drain_and_apply", lambda: pytest.fail("closed client applied work"), + ) + try: + client.start() + # The issued token is queued once the handshake has completed. + assert client.receiver.wait_connected(5) + assert client.update() == SyncUpdate(applied_events=0, submitted_events=0) + assert client.status.phase is ClientPhase.CLOSED + with pytest.raises(RuntimeError, match="ManagedClient is closed"): + client.update() + finally: + client.close() @pytest.mark.parametrize("reconnects", [True, False]) -def test_managed_use_server_preserves_and_clears_owned_authoring_layer(reconnects): +def test_managed_use_server_preserves_and_clears_owned_authoring_layer(reconnects, port): stage = Usd.Stage.CreateInMemory() stage.SetEditTarget(Usd.EditTarget(stage.GetSessionLayer())) client = ManagedClient( @@ -478,6 +521,7 @@ def test_managed_use_server_preserves_and_clears_owned_authoring_layer(reconnect app_name="managed-recovery", persist_token=False, reconnect=False, + port=port, ) authoring = client.authoring_layer assert authoring is not None @@ -507,8 +551,8 @@ def test_managed_use_server_preserves_and_clears_owned_authoring_layer(reconnect sender.connect_result = reconnects client._sender = sender client._replay_to_fresh_checkpoint = lambda timeout: None - force_handshake(client, synchronized=True) try: + connect_client(client) assert client.recovery_artifact is artifact result = client.recover_use_server(session_id="replacement-session") @@ -527,18 +571,19 @@ def test_managed_use_server_preserves_and_clears_owned_authoring_layer(reconnect client.close() -def test_managed_client_rejects_an_edit_target_switch_before_publishing(): +def test_managed_client_rejects_an_edit_target_switch_before_publishing(port): stage = Usd.Stage.CreateInMemory() client = ManagedClient( stage, app_name="managed-edit-target", persist_token=False, reconnect=False, + port=port, ) sender = _SenderStub([]) client._sender = sender - force_handshake(client, synchronized=True) try: + connect_client(client) stage.SetEditTarget(Usd.EditTarget(stage.GetRootLayer())) stage.DefinePrim("/World/WrongLayer", "Xform") @@ -549,13 +594,14 @@ def test_managed_client_rejects_an_edit_target_switch_before_publishing(): client.close() -def test_managed_use_server_refuses_an_edit_target_switch(): +def test_managed_use_server_refuses_an_edit_target_switch(port): stage = Usd.Stage.CreateInMemory() client = ManagedClient( stage, app_name="managed-custom-recovery", persist_token=False, reconnect=False, + port=port, ) authoring = client.authoring_layer assert authoring is not None @@ -567,8 +613,8 @@ def test_managed_use_server_refuses_an_edit_target_switch(): sender.recovery_required = True client._sender = sender client._replay_to_fresh_checkpoint = lambda timeout: None - force_handshake(client, synchronized=True) try: + connect_client(client) with pytest.raises(RecoveryError, match="edit target changed") as error: client.recover_use_server() assert error.value.code == "edit_target_changed" @@ -577,7 +623,7 @@ def test_managed_use_server_refuses_an_edit_target_switch(): client.close() -def test_managed_use_server_restores_local_layer_when_session_abandonment_fails(): +def test_managed_use_server_restores_local_layer_when_session_abandonment_fails(port): stage = Usd.Stage.CreateInMemory() stage.SetEditTarget(Usd.EditTarget(stage.GetSessionLayer())) client = ManagedClient( @@ -585,6 +631,7 @@ def test_managed_use_server_restores_local_layer_when_session_abandonment_fails( app_name="managed-recovery-rollback", persist_token=False, reconnect=False, + port=port, ) authoring = client.authoring_layer assert authoring is not None @@ -598,8 +645,8 @@ def _fail_abandonment(*, session_id=None): sender.abandon_rejected_session = _fail_abandonment client._sender = sender client._replay_to_fresh_checkpoint = lambda timeout: None - force_handshake(client, synchronized=True) try: + connect_client(client) with pytest.raises(RuntimeError, match="injected abandonment failure"): client.recover_use_server() assert authoring.GetPrimAtPath("/World/Local") @@ -812,7 +859,7 @@ def test_publisher_rejects_invalid_transform_coalesce_window(value): def test_managed_client_gates_new_edits_until_replay_is_applied_but_not_on_acks( - monkeypatch, + monkeypatch, server, ): stage = Usd.Stage.CreateInMemory() client = ManagedClient( @@ -820,25 +867,28 @@ def test_managed_client_gates_new_edits_until_replay_is_applied_but_not_on_acks( app_name="readiness-gate", persist_token=False, reconnect=False, + port=server.server_address[1], ) sender = _SenderStub([True, True]) client._sender = sender - force_handshake(client) monkeypatch.setattr(client, "_connect_sender", lambda: None) try: prim = stage.DefinePrim("/World/Thing", "Xform") value = prim.CreateAttribute("value", Sdf.ValueTypeNames.Int) value.Set(1) - replaying = client.update() + # While a transaction holds the server's barrier, it cannot answer the hello. + with server.sync_server.txn_barrier.shared(): + client.start() + replaying = client.update() assert replaying.submitted_events == 0 assert sender.batches == [] assert not client.status.synchronized - client._receiver._synchronized_event.set() - first = client.update() + updates = [] + wait_until(lambda: updates.append(client.update()) or client.status.synchronized) + first = updates[-1] assert first.submitted_events > 0 - assert client.status.synchronized assert sender.pending_event_count == first.submitted_events value.Set(2) @@ -851,7 +901,7 @@ def test_managed_client_gates_new_edits_until_replay_is_applied_but_not_on_acks( client.close() -def test_managed_client_uses_the_same_pre_submission_transform_window(monkeypatch): +def test_managed_client_uses_the_same_pre_submission_transform_window(monkeypatch, port): clock = [0.0] monkeypatch.setattr(coalescing_module, "monotonic", lambda: clock[0]) stage = Usd.Stage.CreateInMemory() @@ -860,6 +910,7 @@ def test_managed_client_uses_the_same_pre_submission_transform_window(monkeypatc app_name="managed-coalescing", persist_token=False, reconnect=False, + port=port, transform_coalesce_seconds=0.1, ) prim = UsdGeom.Xform.Define(stage, "/World/Thing").GetPrim() @@ -869,9 +920,9 @@ def test_managed_client_uses_the_same_pre_submission_transform_window(monkeypatc xformable.AddScaleOp(UsdGeom.XformOp.PrecisionDouble) sender = _SenderStub([True, True]) client._sender = sender - force_handshake(client, synchronized=True) monkeypatch.setattr(client, "_connect_sender", lambda: None) try: + connect_client(client) translate.Set((1, 0, 0)) assert client.update().submitted_events > 0 translate.Set((2, 0, 0)) @@ -1081,23 +1132,27 @@ def test_app_name_is_required(): UsdReceiver(stage, app_name=" ", persist_token=False) -def test_managed_budget_releases_local_edits_under_sustained_traffic(monkeypatch): - client = ManagedClient( - Usd.Stage.CreateInMemory(), app_name="managed-budget", persist_token=False, - ) - client._sender = _SenderStub([True] * 10) - force_handshake(client, synchronized=True) - traffic = PeerTraffic(client.receiver, monkeypatch, queued=3) - try: - client.stage.DefinePrim("/Local", "Xform") - submitted = [] - for _ in range(2): - submitted.append(client.update(max_messages=2).submitted_events) - traffic.arrive(2) - # Held behind the three queued messages, then released although the - # peers keep every later drain at its budget. - assert submitted[0] == 0 and submitted[1] > 0 - assert not client.status.has_unsent_changes - finally: - client.close() +def test_managed_budget_releases_local_edits_under_sustained_traffic(): + # A server of its own: later clients of the shared one would replay the peer's records. + with embedded_server() as server: + client = ManagedClient( + Usd.Stage.CreateInMemory(), app_name="managed-budget", persist_token=False, + port=server.server_address[1], + ) + client._sender = _SenderStub([True] * 10) + traffic = PeerTraffic(server.sync_server, client.receiver, ensure_prim_event) + try: + connect_client(client) + traffic.arrive(3) + client.stage.DefinePrim("/Local", "Xform") + submitted = [] + for _ in range(2): + submitted.append(client.update(max_messages=2).submitted_events) + traffic.arrive(2) + # Held behind the three queued records, then released although the + # peer keeps every later drain at its budget. + assert submitted[0] == 0 and submitted[1] > 0 + assert not client.status.has_unsent_changes + finally: + client.close() diff --git a/verification/tla/ProducerConnection.cfg b/verification/tla/ProducerConnection.cfg new file mode 100644 index 0000000..a0128ef --- /dev/null +++ b/verification/tla/ProducerConnection.cfg @@ -0,0 +1,24 @@ +CONSTANTS + MaxTxn = 2 + MaxFaults = 3 + ServerMayDiverge = FALSE + +SPECIFICATION Spec + +INVARIANTS + TypeOK + OneHostSocket + NoForgottenSocket + ReadyExactlyWhileConnected + FailureExactlyWhileRecovering + AttemptsStartHealthy + CancelledHandshakeNeverPublishes + HighwaterFailsOnlyAfterDivergence + AcknowledgesOnlySubmitted + AcknowledgedIsDurable + ReplayNeverSkips + OrderedWire + +PROPERTIES + AcknowledgementNeverRegresses + EventuallyComplete diff --git a/verification/tla/ProducerConnection.tla b/verification/tla/ProducerConnection.tla new file mode 100644 index 0000000..be37a18 --- /dev/null +++ b/verification/tla/ProducerConnection.tla @@ -0,0 +1,410 @@ +------------------------- MODULE ProducerConnection ------------------------- +EXTENDS Integers, Sequences, TLC + +(*************************************************************************** +Bounded model of the producer endpoint's connection attempts, the host loop +that applies its actions, and one server. + +The host's application thread starts attempts, submits, cancels, and +disconnects at any time. Its I/O loop takes the queued actions as a batch, +blocks while a connect is pending, and reads only once a batch is applied, +so it can report a socket's end before it applies a close that the endpoint +queued for that socket. Applying a close always reports, even without a +socket. Endpoint transitions are atomic, as under the endpoint's lock. +Backoff and rate limits only delay attempts and are not modeled. + +The server is honest, except that when ServerMayDiverge holds its durable +progress for this session may change once while no socket is open: it loses +its newest transaction, or another producer with the same session id commits +past this client's outbox. Rejections and repair are covered by +TransactionRecovery.tla. +***************************************************************************) + +CONSTANTS MaxTxn, MaxFaults, ServerMayDiverge + +States == {"idle", "connecting", "handshaking", "connected", "closing"} +Phases == {"disconnected", "awaiting", "ready", "recovery"} +Failures == {"none", "ahead", "regressed", "invalid"} +TxnIds == 1..MaxTxn + +Hello == [type |-> "hello"] +Quit == [type |-> "quit"] +Txn(id) == [type |-> "txn", id |-> id] +ClientFrames == {Hello, Quit} \cup [type : {"txn"}, id : TxnIds] +ServerFrames == + [type : {"hellook"}, high : 0..MaxTxn] \cup [type : {"ack"}, id : 0..MaxTxn] + +ConnectAct == [kind |-> "connect"] +CloseAct == [kind |-> "close"] +Send(frame) == [kind |-> "send", frame |-> frame] +Actions == {ConnectAct, CloseAct} \cup [kind : {"send"}, frame : ClientFrames] + +VARIABLES + state, \* endpoint connection state + phase, \* producer session phase + generation, \* session generation counter + connGen, \* session generation of the attempt or connection in flight + failure, + nextId, + acked, + cancelled, \* generations whose handshake the host abandoned + queue, \* actions the endpoint queued that the host has not taken + batch, \* actions the host took and applies in order + socket, + doubleConnect, \* the host was asked to open a second socket + toServer, \* frames on the open socket from the client + toClient, \* frames on the open socket from the server + serverHigh, + diverged, + gap, \* the server received a transaction past its next one + faults + +endpointVars == <> +hostVars == <> +wireVars == <> +serverVars == <> +vars == <> + +RECURSIVE Replay(_, _) +Replay(first, last) == + IF first > last THEN <<>> ELSE <> \o Replay(first + 1, last) + +IsConnectAct(action) == action.kind = "connect" +NotConnect(action) == ~IsConnectAct(action) +NotForTheSocket(action) == action.kind \notin {"close", "send"} + +\* ProducerSession::Disconnect: a failure outlives its connection. +EndSession == phase' = IF phase = "recovery" THEN "recovery" ELSE "disconnected" + +Init == + /\ state = "idle" + /\ phase = "disconnected" + /\ generation = 0 + /\ connGen = 0 + /\ failure = "none" + /\ nextId = 1 + /\ acked = 0 + /\ cancelled = {} + /\ queue = <<>> + /\ batch = <<>> + /\ socket = "none" + /\ doubleConnect = FALSE + /\ toServer = <<>> + /\ toClient = <<>> + /\ serverHigh = 0 + /\ diverged = FALSE + /\ gap = FALSE + /\ faults = 0 + +--------------------------------------------------------------------------- +\* Endpoint reactions to host reports. + +\* OnDisconnected: the socket is gone, so its queued frames and close are void. +OnDisconnected == + IF state \in {"connecting", "handshaking", "connected", "closing"} + THEN /\ state' = "idle" + /\ IF state \in {"handshaking", "connected"} THEN EndSession ELSE UNCHANGED phase + /\ queue' = SelectSeq(queue, NotForTheSocket) + ELSE UNCHANGED <> + +OnConnected == + IF state = "connecting" + THEN /\ generation' = generation + 1 + /\ connGen' = generation + 1 + /\ phase' = "awaiting" + /\ state' = "handshaking" + /\ queue' = Append(queue, Send(Hello)) + ELSE UNCHANGED <> + +Fail(kind) == + /\ failure' = kind + /\ phase' = "recovery" + /\ state' = "closing" + /\ queue' = Append(queue, CloseAct) + /\ UNCHANGED acked + +\* A Hello for an ended generation would quarantine a healthy session. +AcceptHello(high) == + IF connGen # generation \/ phase # "awaiting" THEN Fail("invalid") + ELSE IF high > nextId - 1 THEN Fail("ahead") + ELSE IF high < acked THEN Fail("regressed") + ELSE /\ acked' = high + /\ phase' = "ready" + /\ state' = "connected" + /\ queue' = queue \o Replay(high + 1, nextId - 1) + /\ UNCHANGED failure + +Acknowledge(id) == + IF id > nextId - 1 THEN Fail("ahead") + ELSE IF id < acked THEN Fail("regressed") + ELSE /\ acked' = id + /\ UNCHANGED <> + +OnFrame(frame) == + IF state = "handshaking" /\ frame.type = "hellook" THEN AcceptHello(frame.high) + ELSE IF state = "connected" /\ frame.type = "ack" THEN Acknowledge(frame.id) + ELSE UNCHANGED <> + +--------------------------------------------------------------------------- +\* The host's application thread. + +\* RequestConnect or Connect. +StartAttempt == + /\ state = "idle" + /\ failure = "none" + /\ state' = "connecting" + /\ queue' = Append(queue, ConnectAct) + /\ UNCHANGED <> + /\ UNCHANGED <> + +\* CancelConnect, Disconnect, Stop, or the handshake deadline. An attempt the +\* host has not taken is withdrawn; any other is closed. +AbandonAttempt == + /\ faults < MaxFaults + /\ state \in {"connecting", "handshaking"} + /\ faults' = faults + 1 + /\ IF state = "connecting" /\ \E index \in 1..Len(queue) : IsConnectAct(queue[index]) + THEN /\ state' = "idle" + /\ queue' = SelectSeq(queue, NotConnect) + ELSE /\ state' = "closing" + /\ queue' = Append(queue, CloseAct) + /\ IF state = "handshaking" + THEN /\ EndSession + /\ cancelled' = cancelled \cup {connGen} + ELSE UNCHANGED <> + /\ UNCHANGED <> + /\ UNCHANGED <> + +DisconnectPublished == + /\ faults < MaxFaults + /\ state = "connected" + /\ faults' = faults + 1 + /\ state' = "closing" + /\ EndSession + /\ queue' = queue \o <> + /\ UNCHANGED <> + /\ UNCHANGED <> + +Submit == + /\ state = "connected" + /\ nextId <= MaxTxn + /\ nextId' = nextId + 1 + /\ queue' = Append(queue, Send(Txn(nextId))) + /\ UNCHANGED <> + /\ UNCHANGED <> + +--------------------------------------------------------------------------- +\* The host's I/O loop. + +Take == + /\ batch = <<>> + /\ queue # <<>> + /\ socket # "connecting" + /\ batch' = queue + /\ queue' = <<>> + /\ UNCHANGED <> + +ApplyConnect == + /\ batch # <<>> + /\ socket # "connecting" + /\ IsConnectAct(Head(batch)) + /\ batch' = Tail(batch) + /\ doubleConnect' = (doubleConnect \/ socket # "none") + /\ socket' = "connecting" + /\ UNCHANGED <> + +ApplySend == + /\ batch # <<>> + /\ socket # "connecting" + /\ Head(batch).kind = "send" + /\ batch' = Tail(batch) + /\ toServer' = IF socket = "open" THEN Append(toServer, Head(batch).frame) ELSE toServer + /\ UNCHANGED <> + +ApplyClose == + /\ batch # <<>> + /\ socket # "connecting" + /\ Head(batch).kind = "close" + /\ batch' = Tail(batch) + /\ socket' = "none" + /\ toServer' = <<>> + /\ toClient' = <<>> + /\ OnDisconnected + /\ UNCHANGED <> + /\ UNCHANGED <> + +ConnectSucceeds == + /\ socket = "connecting" + /\ socket' = "open" + /\ OnConnected + /\ UNCHANGED <> + /\ UNCHANGED <> + +\* The host interrupts the pending connect of an abandoned attempt. +ConnectInterrupted == + /\ socket = "connecting" + /\ state = "closing" + /\ socket' = "none" + /\ OnDisconnected + /\ UNCHANGED <> + /\ UNCHANGED <> + +ConnectFails == + /\ faults < MaxFaults + /\ socket = "connecting" + /\ faults' = faults + 1 + /\ socket' = "none" + /\ OnDisconnected + /\ UNCHANGED <> + /\ UNCHANGED <> + +Deliver == + /\ batch = <<>> + /\ socket = "open" + /\ toClient # <<>> + /\ toClient' = Tail(toClient) + /\ OnFrame(Head(toClient)) + /\ UNCHANGED <> + /\ UNCHANGED <> + +PeerCloses == + /\ faults < MaxFaults + /\ batch = <<>> + /\ socket = "open" + /\ faults' = faults + 1 + /\ socket' = "none" + /\ toServer' = <<>> + /\ toClient' = <<>> + /\ OnDisconnected + /\ UNCHANGED <> + /\ UNCHANGED <> + +--------------------------------------------------------------------------- +\* The server. + +ServerStep == + /\ socket = "open" + /\ toServer # <<>> + /\ toServer' = Tail(toServer) + /\ LET frame == Head(toServer) + IN CASE frame.type = "hello" -> + /\ toClient' = Append(toClient, [type |-> "hellook", high |-> serverHigh]) + /\ UNCHANGED <> + [] frame.type = "txn" /\ frame.id = serverHigh + 1 -> + /\ serverHigh' = frame.id + /\ toClient' = Append(toClient, [type |-> "ack", id |-> frame.id]) + /\ UNCHANGED gap + [] frame.type = "txn" /\ frame.id <= serverHigh -> + /\ toClient' = Append(toClient, [type |-> "ack", id |-> serverHigh]) + /\ UNCHANGED <> + [] frame.type = "txn" -> + /\ gap' = TRUE + /\ UNCHANGED <> + [] OTHER -> + UNCHANGED <> + /\ UNCHANGED <> + +ServerLosesProgress == + /\ ServerMayDiverge + /\ ~diverged + /\ socket = "none" + /\ serverHigh > 0 + /\ serverHigh' = serverHigh - 1 + /\ diverged' = TRUE + /\ UNCHANGED <> + +ServerRunsAhead == + /\ ServerMayDiverge + /\ ~diverged + /\ socket = "none" + /\ nextId <= MaxTxn + /\ serverHigh' = nextId + /\ diverged' = TRUE + /\ UNCHANGED <> + +--------------------------------------------------------------------------- + +Next == + \/ StartAttempt + \/ AbandonAttempt + \/ DisconnectPublished + \/ Submit + \/ Take + \/ ApplyConnect + \/ ApplySend + \/ ApplyClose + \/ ConnectSucceeds + \/ ConnectInterrupted + \/ ConnectFails + \/ Deliver + \/ PeerCloses + \/ ServerStep + \/ ServerLosesProgress + \/ ServerRunsAhead + +Spec == + /\ Init + /\ [][Next]_vars + /\ WF_vars(StartAttempt) + /\ WF_vars(Submit) + /\ WF_vars(Take) + /\ WF_vars(ApplyConnect) + /\ WF_vars(ApplySend) + /\ WF_vars(ApplyClose) + /\ WF_vars(ConnectSucceeds) + /\ WF_vars(ConnectInterrupted) + /\ WF_vars(Deliver) + /\ WF_vars(ServerStep) + +TypeOK == + /\ state \in States + /\ phase \in Phases + /\ generation \in Nat + /\ connGen \in 0..generation + /\ failure \in Failures + /\ nextId \in 1..(MaxTxn + 1) + /\ acked \in 0..MaxTxn + /\ cancelled \subseteq 1..generation + /\ queue \in Seq(Actions) + /\ batch \in Seq(Actions) + /\ socket \in {"none", "connecting", "open"} + /\ doubleConnect \in BOOLEAN + /\ toServer \in Seq(ClientFrames) + /\ toClient \in Seq(ServerFrames) + /\ serverHigh \in 0..MaxTxn + /\ diverged \in BOOLEAN + /\ gap \in BOOLEAN + /\ faults \in 0..MaxFaults + +OneHostSocket == ~doubleConnect + +\* Every report the endpoint acts on belongs to the socket it tracks. +NoForgottenSocket == state = "idle" => socket = "none" + +ReadyExactlyWhileConnected == (phase = "ready") <=> (state = "connected") + +FailureExactlyWhileRecovering == (failure # "none") <=> (phase = "recovery") + +AttemptsStartHealthy == state \in {"connecting", "handshaking"} => failure = "none" + +CancelledHandshakeNeverPublishes == state = "connected" => connGen \notin cancelled + +HighwaterFailsOnlyAfterDivergence == + failure \in IF diverged THEN {"none", "ahead", "regressed"} ELSE {"none"} + +AcknowledgesOnlySubmitted == acked < nextId + +AcknowledgedIsDurable == ~diverged => acked <= serverHigh + +ReplayNeverSkips == ~gap + +OrderedWire == + \A first, second \in 1..Len(toServer) : + (first < second /\ toServer[first].type = "txn" /\ toServer[second].type = "txn") + => toServer[first].id < toServer[second].id + +AcknowledgementNeverRegresses == [][acked' >= acked]_acked + +EventuallyComplete == <>(acked = MaxTxn) + +============================================================================= diff --git a/verification/tla/ProducerConnectionDivergence.cfg b/verification/tla/ProducerConnectionDivergence.cfg new file mode 100644 index 0000000..450eda4 --- /dev/null +++ b/verification/tla/ProducerConnectionDivergence.cfg @@ -0,0 +1,22 @@ +CONSTANTS + MaxTxn = 2 + MaxFaults = 3 + ServerMayDiverge = TRUE + +SPECIFICATION Spec + +INVARIANTS + TypeOK + OneHostSocket + NoForgottenSocket + ReadyExactlyWhileConnected + FailureExactlyWhileRecovering + AttemptsStartHealthy + CancelledHandshakeNeverPublishes + HighwaterFailsOnlyAfterDivergence + AcknowledgesOnlySubmitted + AcknowledgedIsDurable + ReplayNeverSkips + OrderedWire + +PROPERTY AcknowledgementNeverRegresses diff --git a/verification/tla/README.md b/verification/tla/README.md index 18c0877..6357b54 100644 --- a/verification/tla/README.md +++ b/verification/tla/README.md @@ -44,6 +44,28 @@ cannot advance after abandonment, the recovery artifact stays complete, and new-session transactions commit exactly once in order. Weak fairness also checks that recovery reaches the ready state and the new session completes. +### `ProducerConnection.tla` + +The producer endpoint's connection attempts against one server, with the host +loop that applies its actions. The application thread starts, cancels, and +disconnects attempts at any time; the I/O loop takes actions in batches, +blocks while connecting, and can report a socket's end before applying a +close the endpoint queued for it. The Hello carries the server's committed +highwater, which is checked before the outbox replays. + +The model checks that the host never opens a second socket or keeps one the +endpoint has forgotten, that a cancelled handshake never publishes or +quarantines the session, that the session is ready exactly while connected, +that replay never skips a transaction, and that acknowledgement never +regresses or covers an unsubmitted transaction. It covers the stale-close +hazard: an attempt the host has not taken is withdrawn instead of closed, and +a reported end voids the frames and close still queued for that socket. + +`ProducerConnection.cfg` uses an honest server and also checks that every +transaction is eventually acknowledged; `ProducerConnectionDivergence.cfg` +lets the server's progress for the session regress or run ahead once, and +checks that only that divergence fails the Hello highwater check. + ### `ReceiverSynchronization.tla` Replay and live frames flowing through a bounded receiver queue into the @@ -54,6 +76,50 @@ application failure followed by replay. The primary and tight-queue configurations verify that synchronization is published only after the advertised replay head has applied successfully. +### `ReceiverReplayIdentity.tla` + +The receiver's Hello and replay-identity flow across sequence domains: +compaction, purge, or snapshot replacement on a live connection, server +restarts, server-side sequence gaps, queue overflow, a consumer apply failure, +and a reconnect between draining a frame and reporting it applied. The server +resumes a Hello whose prefix claim matches its domain and otherwise resets. + +The model checks that the stage only ever holds a contiguous prefix of one +domain, that a replay marked applied is reflected in the stage, that the +receiver queues its own reset only ahead of a replay from one, and that the +receiver converges once the network stabilizes. It covers the live-reset +hazard: the applied cursor may still count the old domain, so a replay it +positions claims the applied identity and the server resets instead of +resuming. Injected gaps are always followed by a frame that reveals them. + +A replay request keeps the received identity unless a reset is pending, +queued or drained but not yet reported applied; only then does the claim +fall back to the applied identity. `ResetPendingTracksResets` checks the +inbox's flag against the frames it stands for. + +Two invariants tie the identities to the prefix they describe. +`ReceivedNamesHeldPrefix`: the received identity names the domain of the +newest held prefix (the frames after the last pending reset, or the stage +and pending frames), except after a replay request with a reset pending +whose consumer started a newer domain from an empty stage since its last +mark; the claim then names that older applied replay and the server resets. +`AppliedNamesStage`: the applied identity names the stage until the stage is +next empty. + +`NoSpuriousResetWhenKnown`: once the network is stable and the server's +domain has not changed since the receiver last connected, a receiver that +knew the domain of everything it kept when its connection ended is resumed. + +`NoSpuriousReset` asks the same of any receiver holding a prefix of the +current domain and is not checked. Two causes still reset it: a host-loaded +snapshot prefix is never claimed, and a live reset whose ReplayComplete never +arrived leaves its frames unproven because Resync carries no epoch. The +Python receiver shares both. + +`ReceiverReplayIdentity.cfg` starts a fresh receiver and allows two domain +changes; `ReceiverReplayIdentitySnapshot.cfg` starts from a host-loaded +snapshot cursor with a one-frame queue. + ### `TransactionCoordinator.tla` Two producer sessions in one managed transaction group. It explores direct @@ -110,15 +176,19 @@ The following results describe the model and configuration files in the commit that contains this snapshot. Regenerate the table after changing a `.tla` or `.cfg` file, or when adopting a different TLC version. -TLC2 2026.08.11.125311 results from 2026-08-16: +TLC2 2026.08.11.125311 results from 2026-10-04: | Model and scenario | Generated | Distinct | Depth | Result | |---|---:|---:|---:|---| | Transaction recovery: reject transaction 1 | 1,669 | 634 | 25 | No error | | Transaction recovery: reject transaction 3 | 929 | 372 | 25 | No error | | Recovery session rollover: reject transaction 2 | 28 | 24 | 14 | No error | +| Producer connection: honest server, eventual acknowledgement | 5,341 | 3,387 | 42 | No error | +| Producer connection: server progress diverges once | 10,694 | 6,732 | 47 | No error | | Receiver: three-frame queue, live apply failure | 15,041 | 3,792 | 27 | No error | | Receiver: one-frame queue, replay apply failure | 3,723 | 1,024 | 25 | No error | +| Replay identity: fresh receiver, two domain changes | 5,857,627 | 1,170,182 | 53 | No error | +| Replay identity: snapshot cursor, one-frame queue | 282,901 | 79,024 | 47 | No error | | Coordinator: valid group or infrastructure fallback | 237 | 153 | 11 | No error | | Coordinator: invalid middle transaction fallback | 106 | 64 | 11 | No error | | Shared-layer parent revision, stable identity, and recovery | 3,262 | 2,288 | 13 | No error | diff --git a/verification/tla/ReceiverReplayIdentity.cfg b/verification/tla/ReceiverReplayIdentity.cfg new file mode 100644 index 0000000..c7d4096 --- /dev/null +++ b/verification/tla/ReceiverReplayIdentity.cfg @@ -0,0 +1,21 @@ +CONSTANTS + MaxSeq = 2 + MaxDomain = 2 + QueueBound = 2 + InitialHead = 2 + InitialCursor = 1 + MaxFailures = 1 + +SPECIFICATION Spec + +INVARIANTS + TypeOK + StageIsOneDomainPrefix + ReadyMeansReplayApplied + OwnResetPrecedesReplayFromOne + ReceivedNamesHeldPrefix + AppliedNamesStage + ResetPendingTracksResets + NoSpuriousResetWhenKnown + +PROPERTY EventuallyConverged diff --git a/verification/tla/ReceiverReplayIdentity.tla b/verification/tla/ReceiverReplayIdentity.tla new file mode 100644 index 0000000..5a2823f --- /dev/null +++ b/verification/tla/ReceiverReplayIdentity.tla @@ -0,0 +1,546 @@ +----------------------- MODULE ReceiverReplayIdentity ----------------------- +EXTENDS Integers, Sequences, TLC + +(*************************************************************************** +Bounded model of the receiver's Hello and replay-identity flow. + +The server's history lives in one sequence domain (server instance and replay +epoch) until compaction, purge, snapshot replacement, or a restart replaces +it. Every Hello carries the receiver's cursor and, from the second Hello on, +a claim naming the domain of the prefix the receiver holds. The server +resumes a matching claim and otherwise sends Resync and replays from one; it +never resets a replay that starts at one. A live reset sends Resync, the new +domain's history, and ReplayComplete on the open connection. + +The receiver keeps its queue across reconnects, so it claims the received +identity: the domain of the newest prefix it holds. A replay request discards +the queue and keeps that claim unless a reset is pending, queued or drained +but not yet reported applied; the claim then falls back to the applied +identity, the domain of the last replay the consumer marked applied. A +request for sequence one makes the receiver queue its own reset. The +consumer drains and applies one frame at a time, may fail once, and a +reconnect may separate draining from reporting the applied cursor. Until +the network stabilizes it may drop connections and frames. + +The stage must always hold a contiguous prefix of exactly one domain's +history. In particular, a replay positioned by an applied cursor that still +counts the old domain after a live reset must never resume the new domain. +The invariants also tie both identities to the prefix they describe, and a +receiver that knew the domain of everything it kept is resumed rather than +reset. Without that knowledge a reconnect to an unchanged domain still +resets: the claim carries no proof for a host-loaded snapshot or for a live +reset whose ReplayComplete never arrived. +***************************************************************************) + +CONSTANTS MaxSeq, MaxDomain, QueueBound, InitialHead, InitialCursor, MaxFailures + +None == -1 +NoClaim == -2 +Domains == 0..MaxDomain +Identities == Domains \cup {None} + +Msg(type, dom, seq) == [type |-> type, dom |-> dom, seq |-> seq] +Frame(kind, dom, seq) == [kind |-> kind, dom |-> dom, seq |-> seq] +NoFrame == Frame("none", None, 0) +NoMarker == [head |-> -1, dom |-> None, backlog |-> 0] +NoReady == [head |-> -1, dom |-> None] + +EventsFrom(dom, first, last) == + [index \in 1..(last - first + 1) |-> Msg("event", dom, first + index - 1)] + +Max(a, b) == IF a > b THEN a ELSE b + +VARIABLES + \* Server. + dom, + head, + \* Connection: frames the server sent on the open socket, in order. + conn, + chan, + connSync, + \* Receiver inbox. + queue, + lastRecv, + lastApplied, + requested, + marker, + ready, + resetPending, + \* Receiver replay identity. + helloSent, + claimIncluded, + claimed, + handshake, + received, + pendingId, + appliedId, + proven, + resetRequired, + \* Consumer and its stage; stageDom is None while the stage is empty. + stageDom, + stageLen, + stageOk, + inflight, + inflightStale, + failures, + networkStable, + \* History only, for stating properties: they never constrain a transition. + domainChangedSinceConnect, + stageEmptiedSinceMark, + keptPrefixKnown + +serverVars == <> +connVars == <> +inboxVars == <> +identityVars == <> +consumerVars == <> +historyVars == <> +vars == <> + +Init == + /\ InitialHead \in 0..MaxSeq + /\ InitialCursor \in 1..(InitialHead + 1) + /\ QueueBound >= 1 + /\ dom = 0 + /\ head = InitialHead + /\ conn = "down" + /\ chan = <<>> + /\ connSync = 0 + /\ queue = <<>> + /\ lastRecv = InitialCursor - 1 + /\ lastApplied = InitialCursor - 1 + /\ requested = 0 + /\ marker = NoMarker + /\ ready = NoReady + /\ resetPending = FALSE + /\ helloSent = FALSE + /\ claimIncluded = FALSE + /\ claimed = None + /\ handshake = None + /\ received = None + /\ pendingId = None + /\ appliedId = None + /\ proven = FALSE + /\ resetRequired = FALSE + \* A cursor above one means the host loaded this server's snapshot. + /\ stageDom = IF InitialCursor > 1 THEN 0 ELSE None + /\ stageLen = InitialCursor - 1 + /\ stageOk = TRUE + /\ inflight = NoFrame + /\ inflightStale = FALSE + /\ failures = 0 + /\ networkStable = FALSE + /\ domainChangedSinceConnect = FALSE + /\ stageEmptiedSinceMark = FALSE + /\ keptPrefixKnown = FALSE + +(* Frames the consumer drained or will drain, in order. *) +PendingFrames == (IF inflight = NoFrame THEN <<>> ELSE <>) \o queue + +IsReset(frames, index) == frames[index].kind = "reset" + +QueuedReset == \E index \in 1..Len(queue) : IsReset(queue, index) + +ResetAmongPending == \E index \in 1..Len(PendingFrames) : IsReset(PendingFrames, index) + +(* The next Hello's cursor and claim. *) +HelloSyncFrom == IF requested > 0 THEN requested ELSE lastRecv + 1 +HelloClaim == IF helloSent THEN received ELSE NoClaim + +ServerResets(sync, claim) == + \/ sync > head + 1 + \/ sync > 1 /\ claim # NoClaim /\ claim # dom + +(* The server captures its replay atomically when it accepts the Hello. *) +ServerReplay(sync, claim) == + LET reset == ServerResets(sync, claim) + first == IF reset THEN 1 ELSE sync + IN <> + \o (IF reset THEN <> ELSE <<>>) + \o EventsFrom(dom, first, head) + \o <> + +Connect == + /\ conn = "down" + /\ chan' = ServerReplay(HelloSyncFrom, HelloClaim) + /\ connSync' = HelloSyncFrom + /\ conn' = "handshake" + /\ requested' = 0 + /\ marker' = NoMarker + /\ ready' = NoReady + /\ helloSent' = TRUE + /\ claimIncluded' = helloSent + /\ claimed' = received + /\ pendingId' = None + /\ inflightStale' = TRUE + /\ domainChangedSinceConnect' = FALSE + /\ keptPrefixKnown' = FALSE + /\ UNCHANGED <> + +AcceptReset == + /\ queue' = Append(queue, Frame("reset", None, 0)) + /\ lastRecv' = 0 + /\ marker' = NoMarker + /\ ready' = NoReady + /\ resetPending' = TRUE + /\ handshake' = None + /\ proven' = TRUE + /\ pendingId' = None + /\ resetRequired' = FALSE + +ReceiveHello == + /\ conn = "handshake" + /\ chan # <<>> + /\ Head(chan).type = "hello" + /\ LET d == Head(chan).dom + isProven == connSync = 1 \/ (claimIncluded /\ claimed = d) + IN IF resetRequired + THEN /\ AcceptReset + /\ received' = d + ELSE /\ received' = IF isProven THEN d ELSE received + /\ handshake' = d + /\ proven' = isProven + /\ UNCHANGED <> + /\ conn' = "up" + /\ chan' = Tail(chan) + /\ UNCHANGED <> + +(* Inbox overflow closes the connection; the queue survives the reconnect. *) +Overflow == + /\ conn' = "down" + /\ chan' = <<>> + /\ marker' = NoMarker + /\ ready' = NoReady + /\ keptPrefixKnown' = (received = dom) + /\ UNCHANGED <> + +(* A pending reset separates the consumer's prefix from the received frames, + so only the applied replay can name it. *) +RequestReplay(seq) == + /\ requested' = seq + /\ lastRecv' = seq - 1 + /\ lastApplied' = seq - 1 + /\ queue' = <<>> + /\ marker' = NoMarker + /\ ready' = NoReady + /\ resetPending' = FALSE + /\ received' = IF resetPending THEN appliedId ELSE received + /\ handshake' = None + /\ pendingId' = None + /\ resetRequired' = (seq = 1) + /\ inflightStale' = TRUE + /\ conn' = "down" + /\ chan' = <<>> + /\ keptPrefixKnown' = (received = dom /\ ~ResetAmongPending) + +ReceiveResync == + /\ conn = "up" + /\ chan # <<>> + /\ Head(chan).type = "resync" + /\ IF Len(queue) = QueueBound + THEN Overflow + ELSE /\ AcceptReset + /\ received' = handshake + /\ chan' = Tail(chan) + /\ UNCHANGED <> + /\ UNCHANGED <> + +ReceiveEvent == + /\ conn = "up" + /\ chan # <<>> + /\ Head(chan).type = "event" + /\ LET m == Head(chan) IN + \/ /\ m.seq <= lastRecv + /\ chan' = Tail(chan) + /\ UNCHANGED <> + \/ /\ m.seq > lastRecv + 1 + /\ RequestReplay(lastApplied + 1) + \/ /\ m.seq = lastRecv + 1 + /\ Len(queue) = QueueBound + /\ Overflow + /\ UNCHANGED <> + \/ /\ m.seq = lastRecv + 1 + /\ Len(queue) < QueueBound + /\ queue' = Append(queue, Frame("event", m.dom, m.seq)) + /\ lastRecv' = m.seq + /\ chan' = Tail(chan) + /\ UNCHANGED <> + /\ UNCHANGED <> + +(* A marker beyond the received records reveals a gap like an early frame. *) +ReceiveComplete == + /\ conn = "up" + /\ chan # <<>> + /\ Head(chan).type = "complete" + /\ LET m == Head(chan) + identity == IF proven THEN m.dom ELSE None + IN IF m.seq > lastRecv + THEN RequestReplay(lastApplied + 1) + ELSE /\ marker' = [head |-> m.seq, dom |-> m.dom, backlog |-> Len(queue)] + /\ received' = identity + /\ pendingId' = identity + /\ handshake' = None + /\ chan' = Tail(chan) + /\ UNCHANGED <> + /\ UNCHANGED <> + +Disconnect == + /\ ~networkStable + /\ conn # "down" + /\ conn' = "down" + /\ chan' = <<>> + /\ marker' = NoMarker + /\ ready' = NoReady + /\ keptPrefixKnown' = (received = dom) + /\ UNCHANGED <> + +(* A server-side sequence gap, which only a later frame can reveal. *) +DropEvent == + /\ ~networkStable + /\ conn = "up" + /\ Len(chan) > 1 + /\ Head(chan).type = "event" + /\ chan' = Tail(chan) + /\ UNCHANGED <> + +AppendLive == + /\ head < MaxSeq + /\ head' = head + 1 + /\ chan' = IF conn = "down" THEN chan ELSE Append(chan, Msg("event", dom, head + 1)) + /\ UNCHANGED <> + +(* A snapshot cursor is only meaningful against the server it came from. *) +DomainMayChange == dom < MaxDomain /\ (helloSent \/ InitialCursor = 1) + +LiveReset == + /\ DomainMayChange + /\ dom' = dom + 1 + /\ \E newHead \in 0..MaxSeq: + /\ head' = newHead + /\ chan' = IF conn = "down" THEN chan + ELSE chan \o <> \o EventsFrom(dom + 1, 1, newHead) + \o <> + /\ domainChangedSinceConnect' = TRUE + /\ UNCHANGED <> + +Restart == + /\ DomainMayChange + /\ dom' = dom + 1 + /\ head' \in 0..MaxSeq + /\ conn' = "down" + /\ chan' = <<>> + /\ marker' = NoMarker + /\ ready' = NoReady + /\ domainChangedSinceConnect' = TRUE + /\ keptPrefixKnown' = FALSE + /\ UNCHANGED <> + +(* Reading the generation and draining are one step; reporting is later. *) +ConsumerDrain == + /\ inflight = NoFrame + /\ queue # <<>> + /\ inflight' = Head(queue) + /\ queue' = Tail(queue) + /\ inflightStale' = FALSE + /\ marker' = IF marker.backlog > 0 THEN [marker EXCEPT !.backlog = @ - 1] ELSE marker + /\ UNCHANGED <> + +(* Applies the frame, then reports the cursor with the drained generation; an + applied reset reports every reset drained so far. *) +ConsumerApply == + /\ inflight # NoFrame + /\ LET f == inflight + reset == f.kind = "reset" + extends == IF stageDom = None THEN f.seq = 1 + ELSE f.dom = stageDom /\ f.seq <= stageLen + 1 + newLen == IF reset THEN 0 ELSE Max(stageLen, f.seq) + cursor == IF reset THEN 0 ELSE lastApplied + reported == ~inflightStale /\ newLen >= cursor /\ newLen <= lastRecv + IN /\ stageDom' = IF reset THEN None ELSE f.dom + /\ stageLen' = newLen + /\ stageOk' = (stageOk /\ (reset \/ extends)) + /\ lastApplied' = IF reported THEN newLen ELSE cursor + /\ ready' = IF reset THEN NoReady ELSE ready + /\ resetPending' = IF reset THEN QueuedReset ELSE resetPending + /\ stageEmptiedSinceMark' = (stageEmptiedSinceMark \/ reset) + /\ inflight' = NoFrame + /\ UNCHANGED <> + +(* The batch failed, so the consumer replays from its own applied cursor. *) +ConsumerFail == + /\ inflight # NoFrame + /\ failures < MaxFailures + /\ failures' = failures + 1 + /\ inflight' = NoFrame + /\ RequestReplay(IF inflight.kind = "reset" THEN 1 ELSE stageLen + 1) + /\ UNCHANGED <> + +(* Every drained frame applied, so the replay head counts as applied. *) +ConsumerMark == + /\ inflight = NoFrame + /\ marker # NoMarker + /\ marker.backlog = 0 + /\ lastApplied' = Max(lastApplied, marker.head) + /\ ready' = [head |-> marker.head, dom |-> marker.dom] + /\ marker' = NoMarker + /\ appliedId' = pendingId + /\ pendingId' = None + /\ stageEmptiedSinceMark' = (stageDom = None) + /\ UNCHANGED <> + +StabilizeNetwork == + /\ ~networkStable + /\ networkStable' = TRUE + /\ UNCHANGED <> + +Next == + \/ Connect + \/ ReceiveHello + \/ ReceiveResync + \/ ReceiveEvent + \/ ReceiveComplete + \/ Disconnect + \/ DropEvent + \/ AppendLive + \/ LiveReset + \/ Restart + \/ ConsumerDrain + \/ ConsumerApply + \/ ConsumerFail + \/ ConsumerMark + \/ StabilizeNetwork + +Spec == + /\ Init + /\ [][Next]_vars + /\ WF_vars(StabilizeNetwork) + /\ WF_vars(Connect) + /\ WF_vars(ReceiveHello) + /\ WF_vars(ReceiveResync) + /\ WF_vars(ReceiveEvent) + /\ WF_vars(ReceiveComplete) + /\ WF_vars(ConsumerDrain) + /\ WF_vars(ConsumerApply) + /\ WF_vars(ConsumerMark) + +TypeOK == + /\ dom \in Domains + /\ head \in 0..MaxSeq + /\ conn \in {"down", "handshake", "up"} + /\ connSync \in 0..(MaxSeq + 1) + /\ Len(queue) <= QueueBound + /\ lastRecv \in 0..MaxSeq + /\ lastApplied \in 0..MaxSeq + /\ requested \in 0..(MaxSeq + 1) + /\ {resetPending, helloSent, claimIncluded, proven, resetRequired} \subseteq BOOLEAN + /\ {claimed, handshake, received, pendingId, appliedId} \subseteq Identities + /\ stageDom \in Identities + /\ stageLen \in 0..MaxSeq + /\ stageOk \in BOOLEAN + /\ failures \in 0..MaxFailures + /\ {domainChangedSinceConnect, stageEmptiedSinceMark, keptPrefixKnown} \subseteq BOOLEAN + +LastPendingReset == + LET resets == {index \in 1..Len(PendingFrames) : IsReset(PendingFrames, index)} + IN IF resets = {} THEN 0 ELSE CHOOSE index \in resets : \A other \in resets : other <= index + +NewestPending == SubSeq(PendingFrames, LastPendingReset + 1, Len(PendingFrames)) + +(* Domains of the newest prefix the receiver holds once everything pending + applies: the frames after the last pending reset, or else the stage and the + pending frames. A well-formed prefix has at most one. *) +HeldDomains == + {NewestPending[index].dom : index \in 1..Len(NewestPending)} + \cup (IF LastPendingReset = 0 /\ stageDom # None THEN {stageDom} ELSE {}) + +(* The inbox flag covers every queued reset and every reset drained in this + generation, and is only set while some reset is still queued or drained. *) +ResetPendingTracksResets == + /\ QueuedReset => resetPending + /\ (inflight.kind = "reset" /\ ~inflightStale) => resetPending + /\ resetPending => ResetAmongPending + +(* The received identity names the held prefix, except after a replay request + with a reset pending whose consumer started a newer domain from an empty + stage since its last mark: the claim then names that older applied replay, + so the server resets. Domain numbers grow with every replacement. *) +ReceivedNamesHeldPrefix == + received # None => + \/ HeldDomains \subseteq {received} + \/ /\ received = appliedId + /\ stageEmptiedSinceMark + /\ \A held \in HeldDomains : held > received + +(* The applied identity names the stage until the stage is next empty. *) +AppliedNamesStage == appliedId # None => stageEmptiedSinceMark \/ stageDom = appliedId + +(* Once stable, a receiver that knew the domain of everything it kept when its + connection last ended is resumed by an unchanged server. *) +NoSpuriousResetWhenKnown == + conn = "down" /\ networkStable /\ ~domainChangedSinceConnect /\ keptPrefixKnown + => ~ServerResets(HelloSyncFrom, HelloClaim) + +(* Not checked: a host-loaded snapshot prefix and a live reset whose + ReplayComplete never arrived both reset an unchanged server. *) +NoSpuriousReset == + conn = "down" /\ networkStable /\ ~domainChangedSinceConnect /\ HeldDomains = {dom} + => ~ServerResets(HelloSyncFrom, HelloClaim) + +StageIsOneDomainPrefix == stageOk + +ReadyMeansReplayApplied == + ready # NoReady => + /\ conn = "up" + /\ stageLen >= ready.head + /\ (ready.head = 0 \/ stageDom = ready.dom) + +(* Justifies queueing the receiver's own reset without an overflow check. *) +OwnResetPrecedesReplayFromOne == + resetRequired => + /\ queue = <<>> + /\ (conn = "handshake" => connSync = 1) + +Converged == + /\ conn = "up" + /\ chan = <<>> + /\ queue = <<>> + /\ inflight = NoFrame + /\ ready # NoReady + /\ stageOk + /\ stageLen = head + /\ (head = 0 \/ stageDom = dom) + +EventuallyConverged == <>[]Converged + +============================================================================= diff --git a/verification/tla/ReceiverReplayIdentitySnapshot.cfg b/verification/tla/ReceiverReplayIdentitySnapshot.cfg new file mode 100644 index 0000000..0f7c96c --- /dev/null +++ b/verification/tla/ReceiverReplayIdentitySnapshot.cfg @@ -0,0 +1,21 @@ +CONSTANTS + MaxSeq = 2 + MaxDomain = 1 + QueueBound = 1 + InitialHead = 2 + InitialCursor = 2 + MaxFailures = 1 + +SPECIFICATION Spec + +INVARIANTS + TypeOK + StageIsOneDomainPrefix + ReadyMeansReplayApplied + OwnResetPrecedesReplayFromOne + ReceivedNamesHeldPrefix + AppliedNamesStage + ResetPendingTracksResets + NoSpuriousResetWhenKnown + +PROPERTY EventuallyConverged