diff --git a/CHANGELOG.md b/CHANGELOG.md new file mode 100644 index 0000000..b6e32bb --- /dev/null +++ b/CHANGELOG.md @@ -0,0 +1,182 @@ +# Changelog + +Notable changes to OpenUSDConnect are recorded here. The format follows +[Keep a Changelog](https://keepachangelog.com/en/1.1.0/), and the project uses +[Semantic Versioning](https://semver.org/spec/v2.0.0.html). Before 1.0, a minor +release may contain breaking changes; each such release starts with a +**Migrating** section. + +The release version covers the Python package and the Blender add-on +(`openusdconnect/_version.py`). The Unreal plugin carries its own version in +`OpenUSDConnect.uplugin`. Each entry states wire compatibility separately: the +protocol version (`PROTOCOL_VERSION`) and wire schema version +(`SCHEMA_VERSION`) decide whether clients and servers from different releases +can connect. + +## [0.5.0] - Unreleased + +The high-level clients share one lifecycle, report all state through +`client.status`, and deliver host notifications through one `ClientObserver`. +Code written against 0.4 needs the changes in **Migrating from 0.4**. + +**Wire compatibility:** protocol 13 and wire schema 10, unchanged. New +handshake and transaction-result fields are optional: 0.4.0 peers still +connect but get no replay-identity validation or checkpoints. + +### Migrating from 0.4 + +Callbacks become one `observer=ClientObserver`: + +- `on_applied`, `on_applied_events`, and `on_imported` become + `on_applied(batch)`; read `batch.prim_paths`, `batch.events`, and + `batch.imported_paths`. +- `on_playback_claimed` and `on_playback_rejected` become + `on_playback_claim(result)`. +- Other `on_*` arguments become the observer method of the same name, which + receives a typed value instead of a dict: `StageMetadata`, `PlaybackState`, + or `PlaybackClaim` (a rejection's `current_leader_client_id` is its + `leader_client_id`). +- For `UsdReceiver.applying_seq`, use `batch.seq` inside `on_applied`. + +State moves to `client.status`: + +- Properties keep their names there, except the event counts: + `status.pending_events`, `prepared_events`, `deferred_events`, and + `acknowledged_events_total` replace the `*_event_count` properties. +- `transaction_failure` and `recovery_incident` become `status.failure` and + `status.recovery`; use `str(status.failure)` and + `status.failure.disposition` for `transaction_error` and + `recovery_disposition`. +- `recovery_required`, `native_scene_rebuild_required`, and + `connection_rejected` are phases: check `status.phase` (and + `status.auth_rejected` to tell rejections apart). +- `UsdReceiver.layered_replay_active` is removed. A connected receiver always + has layered replay; the server's handshake is rejected otherwise. + +Return values and defaults: + +- `UsdReceiver.update()` and `UsdPublisher.update()` return `SyncUpdate` + instead of `int`; read `applied_events` or `submitted_events`. +- `client.stage_metadata` returns `StageMetadata` instead of `dict`. +- `connect()` and `flush()` wait 10 seconds instead of indefinitely; pass + `timeout=None` for no limit. + +Behavior: + +- Token, metadata, and playback notifications run during `update()` or + `close()` on the calling thread instead of on network threads. +- Stage edits made in `on_resync` are no longer published, matching + `on_applied`. +- `UsdPublisher.update()` raises before `start()`. While disconnected it + reconnects in the background; `disconnect()` pauses that until `connect()`. +- `publish_current_edit_target()` captures its snapshot while disconnected and + sends it from a later `update()`, instead of returning 0 without capturing. +- `ManagedClient.rebind_stage()` refuses unacknowledged work (finish with + `submit_and_wait()`) and unsent edits unless `discard_unsent=True` drops + them. +- `ManagedClient.flush()` raises when the server rejects the connection, + 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 +callable arguments and properties. + +### Added + +- `ClientObserver` and its typed payloads: `AppliedBatch`, `StageMetadata`, + `PlaybackState`, `PlaybackClaim`. +- `wait_until_ready()` and `submit_and_wait()`. They return `False` only on + timeout and raise for states that more updates cannot fix. +- `update(max_messages=)` on `ManagedClient` and `SharedStageClient` spreads a + reconnect backlog over frames; local edits wait only for the messages queued + before them. +- `ClientStatus` fields `auth_rejected`, `has_unsent_changes`, + `deferred_events`, `deferred_layer_keys`, `edit_target_is_published`, and + `recovery_stage_pending`; the `can_author` property; `ClientPhase.PARKED`. +- `SharedStageClient.resume_recovery()` (error code + `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 + for each connection attempt. +- `EventDispatcher.drained_message_count`, `ReceiverThread.stopped`, and + `NoticeEmitter.has_local_changes`. +- Receiver replay identity and optional post-commit transaction checkpoints. + +### Changed + +- `EventDispatcher` starts its cursor at `receiver.sync_from - 1`, so + integrations no longer seed `last_seq` for continuation. + +### 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. +- 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 + sender reconnected. +- Retries against an unreachable server logged a traceback each time. +- `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. + +## [0.4.0] - 2026-09-18 + +**Wire compatibility:** protocol 13 (from 12), wire schema 10. Clients and +servers must both be 0.4.0 or later. + +### Added + +- `ServerRuntime`, `ServerConfig`, and `start_server()` embed the server in an + application without signal handlers; `run_server()` remains the blocking + command-line runner. +- `ClientPhase`, `ClientStatus`, and `SyncUpdate` in + `openusdconnect.client_types`, `client.status`, and `client.client_id`. +- `EventSender.request_connect()` and `cancel_connect()` for nonblocking + handshakes; bidirectional clients reconnect their sender from `update()`. +- `openusdconnect.shader_mapping` for shader mapper interfaces and + `openusdconnect.usd_authoring` helpers (`set_connectable_input_value`, + `resolve_shader_port_type`). +- Distributable server packages and configurable OpenUSD runtimes. +- Targeted time-sample erasure with native Sdf tracking. + +### Changed + +- The native Sdf notice bridge ABI is version 2; rebuild it with + `openusdconnect-build-sdf-notice-bridge`. + +### Fixed + +- Time-sample deletions keep their order through application and log + compaction. + +## [0.3.0] - 2026-08-30 + +**Wire compatibility:** protocol 12, wire schema 10. + +### Added + +- Shared native C++ client core (framing, producer outbox, receiver inbox) + used by the Python client through nanobind and by the Unreal plugin. +- `UsdReceiver(adapter=...)` projects composed changes into non-USD + native scenes. +- Cross-platform OpenUSD runtime configuration and managed OpenUSD source + builds. + +### Changed + +- `uv sync` builds the native client extension, which requires CMake and a + 64-bit C++17 toolchain. The Blender add-on vendors a build for Blender's + Python. + +## [0.2.0] - 2026-08-18 + +First versioned release. Protocol 12, wire schema 10. The release version is +defined once in `openusdconnect/_version.py` and mirrored by the Blender +add-on. diff --git a/README.md b/README.md index fc9c740..d8bb59a 100644 --- a/README.md +++ b/README.md @@ -167,7 +167,7 @@ Use `ManagedClient` for a bidirectional application that owns a ```python from pxr import Usd -from openusdconnect import ClientPhase, ManagedClient +from openusdconnect import ManagedClient stage = Usd.Stage.Open("scene.usda") @@ -177,12 +177,12 @@ with ManagedClient(stage, app_name="my-editor") as client: while application_is_running(): client.update() - if client.status.phase is ClientPhase.READY: + if client.status.can_author: edit_scene(stage) ``` -Call `update()` on the stage-owning thread. Use `flush(timeout)` at save, -publish, or orderly-shutdown boundaries when acknowledgement matters. Receive- +Call `update()` on the stage-owning thread. Use `submit_and_wait(timeout)` at +save, publish, or orderly-shutdown boundaries when acknowledgement matters. Receive- only tools can use `UsdReceiver`; send-only tools can use `UsdPublisher`. The [USD-native API guide](docs/usd-native-integration.md) covers ownership, replay, reconnection, recovery, and shared-stage clients. diff --git a/docs/README.md b/docs/README.md index b51353b..a6f4227 100644 --- a/docs/README.md +++ b/docs/README.md @@ -56,6 +56,8 @@ linked from the shorter workflow guides. - [Profiling](profiling.md): server and Blender sampling workflows. - [Development commands](cli-reference.md#development-commands): packaging, test launchers, benchmarks, and diagnostics. +- [Changelog](../CHANGELOG.md): release notes, breaking changes with migration + notes, and wire compatibility per release. ## Troubleshoot diff --git a/docs/client-recovery.md b/docs/client-recovery.md index 18bf86d..6a2b06d 100644 --- a/docs/client-recovery.md +++ b/docs/client-recovery.md @@ -7,9 +7,11 @@ producer session and transaction ID. A deterministic producer rejection requires an explicit policy. The rejected transaction and its ordered suffix are quarantined because later IDs cannot safely pass the gap. -`ManagedClient` and `SharedStageClient` report this through +`ManagedClient`, `SharedStageClient`, and `UsdPublisher` report this through `client.status.phase == ClientPhase.RECOVERY_REQUIRED`, -`client.status.recovery`, and `client.recovery_artifact`. +`client.status.recovery`, and `client.recovery_artifact`. `UsdPublisher` +recovers only by repair (`repair_and_resume(events)`), because it holds no +authoritative state to fall back to. Ordinary `update()` calls report the condition without raising. Explicit recovery commands may raise `RecoveryError`, `TimeoutError`, or @@ -90,14 +92,24 @@ This operation requires an equivalent clean stage whose loaded ```python clean_stage = open_clean_equivalent_stage() -assessment = client.recover_use_server(clean_stage=clean_stage, timeout=5) +previous_stage = client.stage +try: + assessment = client.recover_use_server(clean_stage=clean_stage, timeout=5) +finally: + # Recovery may bind the replacement before replay completes. + if client.stage is not previous_stage: + replace_stage_in_host(client.stage) for index, snapshot in enumerate(assessment.rejected_snapshots): snapshot.Export(f"rejected-work-{index}.usda") - -replace_stage_in_host(client.stage) ``` +If replay times out after the replacement is bound, +`client.status.recovery_stage_pending` +is `True`: keep authoring disabled and continue with +`client.resume_recovery(timeout=5)`, which keeps the original rejected +snapshots. A failure before replacement is retried with `recover_use_server()`. + Opening the same asset path again in the same process is usually not enough. OpenUSD's layer registry may return the same loaded `Sdf.Layer` objects. Filesystem integrations should open an isolated copy of the clean root and @@ -170,14 +182,18 @@ except (TimeoutError, ConnectionError): Stable codes include `no_incident`, `wrong_recovery_kind`, `stale_assessment`, `stage_not_synchronized`, `invalid_clean_stage`, `shared_loaded_layers`, `invalid_repair_target`, `local_changes_pending`, -`transactions_pending`, `stage_unavailable`, and `edit_target_changed`. +`transactions_pending`, `no_pending_recovery_stage`, `stage_unavailable`, and +`edit_target_changed`. ## UI guidance -Drive editing state from `client.status.phase`: +Enable authoring only when `client.status.can_author` is true; use +`client.status.phase` for the message: -- `READY`: enable authoring -- `CONNECTING` or `REPLAYING`: keep calling `update()`, but disable authoring +- `CONNECTING` or `REPLAYING`: keep calling `update()` +- `OFFLINE`: nothing will reconnect; call `connect()` after `disconnect()`, or + create a new client if its receiver used `reconnect=False` +- `PARKED`: bind a stage with `rebind_stage()` - `RECOVERY_REQUIRED`: disable authoring and present Use Server, repair, or application-specific merge choices - `REJECTED`: show the authentication or layer-mode reason diff --git a/docs/usd-native-integration.md b/docs/usd-native-integration.md index f802284..5b2d954 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 readers and automatic -reconnect attempts run on background threads; event submission still performs -a socket write on the calling thread. +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. ## Choose an API @@ -16,7 +16,10 @@ a socket write on the calling thread. Use `ManagedClient` by default. Choose a directional client for send-only or receive-only work, or `SharedStageClient` when the application must edit its -existing authored-layer graph. +existing authored-layer graph. These bidirectional clients observe a USD stage; +they do not capture edits to a host's native objects. Native-scene integrations +use an adapter-backed receiver and their own outbound capture bridge, as +described in [Adapter destination contract](#adapter-destination-contract). Managed clients must open equivalent base content and resolve referenced assets compatibly. The server synchronizes collaboration opinions and their ordered @@ -26,49 +29,46 @@ Unreal host integrations. ## Lifecycle and status -All high-level clients use the same lifecycle: - -1. Construction validates the stage and initializes the role-specific stage - state. -2. `start()` returns immediately. It starts the background receiver for - `UsdReceiver`, `ManagedClient`, and `SharedStageClient`; `UsdPublisher` - merely enters its nonblocking lifecycle. Entering a context manager calls - `start()`. -3. `connect(timeout)` waits for the applicable handshakes. For receiving - clients, it does not apply queued replay. -4. `update()` applies incoming work and, for bidirectional clients, submits - local work without waiting for a durable acknowledgement. If the sender - disconnects, it schedules a background handshake after the receiver has - connected. Repeated calls use a single attempt with retry backoff. Auth, - protocol, and recovery rejections stop automatic retries. -5. `flush(timeout)` waits for already submitted work. Call `update()` first if - the stage may still contain unsent edits. -6. `close()` stops networking. It does not implicitly turn every pending edit - into a blocking flush. - -`client.status` is an immutable `ClientStatus`; `phase` is one of -`OFFLINE`, `CONNECTING`, `REPLAYING`, `READY`, `RECOVERY_REQUIRED`, `REJECTED`, -or `CLOSED`. Bidirectional applications should enable editing only in `READY`. -The directional connection fields distinguish partial connectivity from a -role that is not present. - -`ClientPhase`, `ClientStatus`, and `SyncUpdate` are available from -`openusdconnect.client_types` and the package root. `client.client_id` exposes -the connection identity without accessing an underlying transport object. - -Keep lifecycle and stage operations on the host's owning thread. Apply callbacks -(`on_imported`, `on_resync`, `on_applied`, `on_applied_events`) run during -`update()` on that thread. Token, metadata, and playback callbacks run on the -thread handling the handshake or incoming message, which may be a worker. -Queue UI and USD work back to the owning thread from those callbacks. -Inside a receiver's apply callback, `receiver.applying_seq` is the candidate -batch tail; `last_seq` advances only after the complete apply succeeds. - -An adapter-backed `UsdReceiver` also enters `RECOVERY_REQUIRED` when resolver -recomposition makes incremental projection unsafe. Rebuild the native scene, -then call `acknowledge_native_scene_rebuilt()`. - -`ManagedClient.update()` and `SharedStageClient.update()` return `SyncUpdate`: +All high-level clients share one lifecycle: + +1. `start()` returns immediately; entering a context manager calls it. + `UsdPublisher` opens its socket on the first `update()`. +2. `wait_until_ready(timeout)` pumps `update()` until the client is `READY`. + `connect(timeout)` only completes the handshakes. +3. `update(max_messages=None)` applies incoming work and submits local work. + While a sender is disconnected it schedules a background handshake with + backoff; rejections stop the retries. +4. `submit_and_wait(timeout)` publishes pending edits and waits until they are + durable. `flush(timeout)` waits only for work already submitted (plus a + coalesced transform). +5. `close()` stops networking without flushing, then delivers notifications + that were still queued. + +Blocking calls (`connect`, `flush`, `wait_until_ready`, `submit_and_wait`) +default to a 10 second timeout and return `False` only when it expires; the +work stays queued. `flush` also returns `False` at once while a coalesced +transform cannot be submitted yet (replaying or parked). States that more +updates cannot fix raise: + +| State | Exception | +| --- | --- | +| Authentication rejected | `PermissionError` | +| Handshake rejected | `ConnectionError` | +| Offline and not reconnecting (after `disconnect()`, or `reconnect=False`) | `ConnectionError` | +| Transaction rejected | `TransactionRejectedError` | +| Closed, parked, or native-scene rebuild required | `RuntimeError` | + +`client.status` is an immutable `ClientStatus` and the one place to read client +state; the clients themselves expose data and operations. Its `phase` is `OFFLINE`, +`CONNECTING`, `REPLAYING`, `READY`, `RECOVERY_REQUIRED`, `REJECTED`, `PARKED` +(no bound stage), or `CLOSED`. `status.can_author` tells a UI whether edits to +the current edit target will be published now. The status also reports unsent +(`has_unsent_changes`) and unacknowledged (`pending_events`) work, and +`auth_rejected` separates an authentication rejection from a protocol one. Per-role +connection fields are `None` for a role the client lacks. `ClientPhase`, +`ClientStatus`, and `SyncUpdate` are importable from the package root. + +`update()` returns a `SyncUpdate` for the work done by that call: - `applied_events`: authoritative events applied during this call - `submitted_events`: local events accepted by the sender outbox @@ -76,6 +76,70 @@ then call `acknowledge_native_scene_rebuilt()`. - `pending_events`: currently submitted but unacknowledged events - `recovery`: a deterministic rejection that requires application action +### Observing the client + +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: + +```python +class HostObserver(ClientObserver): + def on_applied(self, batch): # AppliedBatch: seq, events, prim_paths + refresh_host_ui(batch.prim_paths) + + def on_playback_state(self, state): # PlaybackState + set_host_time(state.time) + +client = ManagedClient(stage, app_name="my-editor", observer=HostObserver()) +``` + +`on_applied` and `on_resync` are part of delivery: raising from one in +`update()` rolls the batch back and replays it, so they must be safe to retry. +Stage edits made in them are not published, and a `close()` called from one +takes effect once the batch has been applied. `on_stage_metadata`, +`on_playback_state`, `on_playback_claim`, and `on_token_issued` only observe: +raising propagates out of `update()` and later notifications wait for the next +call. Methods that do not apply to a client never fire; `UsdPublisher` reports +only tokens and stage metadata, and `SharedStageClient` has no delivery +methods. + +An adapter-backed `UsdReceiver` enters `RECOVERY_REQUIRED` when resolver +recomposition makes incremental projection unsafe. Rebuild the native scene, +then call `acknowledge_native_scene_rebuilt()`. + +### Host loop + +GUI hosts drive the client from a timer instead of waiting: + +```python +client = ManagedClient(stage, app_name="my-editor").start() + +def on_timer(): + if client.status.edit_target_is_published: # ManagedClient publishes only its authoring layer + client.update(max_messages=256) + set_editing_enabled(client.status.can_author) +``` + +Pass `max_messages` in interactive hosts. Without it, `update()` applies the +whole queued backlog in one call: about 20 µs per transform event, so a +reconnect with 10,000 queued events stalls one frame for about 200 ms. With a +budget the backlog spreads over frames at the same total cost, and local edits +are held until the backlog queued before them has been applied. +Receiving pauses while a `ManagedClient` edit target is foreign, because its +`update()` refuses to publish another layer's opinions. `SharedStageClient` +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. + ## Receive into a stage ```python @@ -86,19 +150,16 @@ from openusdconnect import ClientPhase, UsdReceiver stage = Usd.Stage.Open("shot.usda") with UsdReceiver(stage, app_name="my-viewer") as receiver: - if not receiver.connect(timeout=5): - raise ConnectionError("OpenUSDConnect server is unavailable") + if not receiver.wait_until_ready(timeout=5): + raise TimeoutError("OpenUSDConnect replay did not finish in time") while application_is_running(): receiver.update() show_loading(receiver.status.phase is not ClientPhase.READY) ``` -Interactive receive-only applications may bound one tick's work with -`receiver.update(max_messages=500)`. Ordered replay remains pending until all -messages preceding the server's synchronization watermark have been applied. -Bidirectional clients intentionally drain their complete queued prefix before -publishing local edits, so this budget applies only to `UsdReceiver`. +Ordered replay remains pending until all messages preceding the server's +synchronization watermark have been applied, however `max_messages` splits it. `UsdReceiver` always requests managed layered replay from sequence 1. It owns anonymous collaboration layers at the strong end of the stage's session-layer @@ -111,12 +172,8 @@ history over it would duplicate opinions. Snapshot continuation is a separate flat integration path used by the live-open host plugins. Use `rebind_stage(new_stage)` when a host replaces its stage. Passing `None` -parks stage application while the network queue continues to receive data. - -Application callbacks such as `on_applied` and `on_applied_events` run inside -`update()` on the calling thread. Transport callbacks, including token, -metadata, and playback notifications, may run on a background connection -thread and must be marshalled before touching a UI. +parks stage application (phase `PARKED`) while the network queue continues to +receive data. ### Receive into an application-owned scene @@ -136,7 +193,7 @@ with UsdReceiver( mirror_stage, app_name="my-host", adapter=adapter, - on_resync=adapter.reset, + observer=MyHostObserver(adapter), # on_resync resets the adapter ) as client: while application_is_running(): client.update() # call from the host's scene/UI thread @@ -162,14 +219,13 @@ stage = Usd.Stage.Open("shot.usda") stage.SetEditTarget(Usd.EditTarget(stage.GetSessionLayer())) with UsdPublisher(stage, app_name="layout") as publisher: - if not publisher.connect(timeout=5): - raise ConnectionError("OpenUSDConnect server is unavailable") + if not publisher.wait_until_ready(timeout=5): + raise TimeoutError("OpenUSDConnect server is unavailable") sphere = UsdGeom.Sphere.Define(stage, "/World/Sphere") UsdGeom.Xformable(sphere).AddTranslateOp().Set(Gf.Vec3d(1, 2, 3)) - publisher.update() - if not publisher.flush(timeout=5): + if not publisher.submit_and_wait(timeout=5): raise TimeoutError("changes were not durably acknowledged") ``` @@ -183,13 +239,12 @@ failure after an ambiguous write retains the exact transaction and resends it with the same producer session and transaction ID after reconnection. The server either commits it once or reports the existing durable high-water mark. -`UsdPublisher.update()` does not initiate a reconnect. While disconnected it -returns zero and leaves noticed edits dirty; call `connect()` and then -`update()` to submit them. +`disconnect()` pauses automatic reconnection until the next `connect()`. Use `publish_current_edit_target()` when attaching to a layer that was already authored before the publisher existed. It publishes authored opinions, not a -flattened composed stage. Retry any retained batch with `update()` first. +flattened composed stage. While disconnected the snapshot stays queued for a +later `update()`. Retry any retained batch with `update()` first. For high-frequency default-time transforms, set `transform_coalesce_seconds` to a small host-appropriate window. Only repeated @@ -203,7 +258,7 @@ other event kinds, and distinct animation samples remain ordering barriers. ```python from pxr import Gf, Usd, UsdGeom -from openusdconnect import ClientPhase, ManagedClient +from openusdconnect import ManagedClient stage = Usd.Stage.Open("shot.usda") @@ -213,14 +268,14 @@ with ManagedClient( department="layout", transform_coalesce_seconds=0.02, ) as client: - if not client.connect(timeout=5): - raise ConnectionError("OpenUSDConnect server is unavailable") + if not client.wait_until_ready(timeout=5): + raise TimeoutError("OpenUSDConnect replay did not finish in time") translate = None while application_is_running(): client.update() - if client.status.phase is ClientPhase.READY: + if client.status.can_author: if translate is None: sphere = UsdGeom.Sphere.Define(stage, "/World/Sphere") translate = UsdGeom.Xformable(sphere).AddTranslateOp() @@ -234,6 +289,16 @@ 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. +`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. +`rebind_stage()` refuses unsent or unacknowledged work: call +`submit_and_wait()` first, or pass `discard_unsent=True` to drop unsent edits. +`rebind_stage(None)` parks the client while networking stays active. + +`close()` detaches the collaboration layers but leaves the authoring layer on +the stage, so the composed scene can change. To keep what the user sees, +flatten first: `Usd.Stage.Open(client.stage.Flatten())`. + Use separate `UsdPublisher` and `UsdReceiver` stages when the host intentionally authors persistent layers or changes edit targets. Attaching those two low-level roles directly to the same stage duplicates opinion ownership; @@ -286,17 +351,17 @@ uv run openusdconnect-server --base shot.usda --layer-mode shared_stage ```python from pxr import Usd -from openusdconnect import ClientPhase, SharedStageClient +from openusdconnect import SharedStageClient stage = Usd.Stage.Open("shot.usda") with SharedStageClient(stage, app_name="layer-editor") as client: - if not client.connect(timeout=5): - raise ConnectionError("OpenUSDConnect server is unavailable") + if not client.wait_until_ready(timeout=5): + raise TimeoutError("OpenUSDConnect replay did not finish in time") while application_is_running(): - result = client.update() - set_editing_enabled(client.status.phase is ClientPhase.READY) + client.update() + set_editing_enabled(client.status.can_author) ``` Every process opens its own equivalent root document under its normal @@ -317,6 +382,11 @@ sublayers may resolve later; call `refresh_layer_graph()` after resolver or asset availability changes. Use `is_layer_reachable(layer)` before authoring into a newly attached layer. +`READY` does not wait for unresolved layers: their records are counted in +`status.deferred_events` / `deferred_layer_keys` and applied by +`refresh_layer_graph()`. A remote topology edit that removes the current edit +target selects the root layer instead. + The portable Python tracker keeps full in-memory layer snapshots. Native hosts can build an optional bridge against the exact OpenUSD installation they load: @@ -358,8 +428,8 @@ 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 -client exposes this through `client.native_scene_rebuild_required` and -`client.status`. Rebuild the native destination and call +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. @@ -371,7 +441,8 @@ belong to one integration should use the same `app_name` or explicit TOFU tokens are loaded and saved by default. Set `persist_token=False` for ephemeral tools or tests, pass `token=` when the host owns credential storage, -and use `on_token_issued` to integrate with a host-specific store. +and override `ClientObserver.on_token_issued` to integrate with a host-specific +store. ## Low-level APIs diff --git a/examples/shared_stage_client/README.md b/examples/shared_stage_client/README.md index ce13713..e2918a4 100644 --- a/examples/shared_stage_client/README.md +++ b/examples/shared_stage_client/README.md @@ -43,3 +43,5 @@ persistence remains an application decision. The author only changes the stage while `client.status.phase` is `READY`. During connection or replay it continues pumping `update()` without authoring; on a recoverable rejection it stops authoring and reports the server reason. +Normal shutdown uses `submit_and_wait()`, so a successful exit means its edits +are durable. diff --git a/examples/shared_stage_client/demo.py b/examples/shared_stage_client/demo.py index dd1c20b..e1cf646 100644 --- a/examples/shared_stage_client/demo.py +++ b/examples/shared_stage_client/demo.py @@ -12,8 +12,8 @@ if str(_REPO_ROOT) not in sys.path: sys.path.insert(0, str(_REPO_ROOT)) -from openusdconnect import ClientPhase, SharedStageClient # noqa: E402, I001 -from pxr import Gf, Usd, UsdGeom # noqa: E402 +from openusdconnect import ClientPhase, SharedStageClient, TransactionRejectedError # noqa: E402, I001 +from pxr import Gf, Sdf, Usd, UsdGeom # noqa: E402 DEFAULT_STAGE = Path(__file__).with_name("scene.usda") SPHERE_PATH = "/World/SharedSphere" @@ -54,6 +54,17 @@ def main() -> int: return 1 content = _content_layer(stage) + try: + return _run(args, stage, content) + except (PermissionError, ConnectionError) as exc: + print(f"server rejected the client: {exc}", file=sys.stderr) + return 1 + except TransactionRejectedError as exc: + print(str(exc), file=sys.stderr) + return 2 + + +def _run(args: argparse.Namespace, stage: Usd.Stage, content: Sdf.Layer) -> int: with SharedStageClient( stage, app_name=args.app_name, @@ -62,22 +73,10 @@ def main() -> int: persist_token=False, delegate_bridge_path=args.sdf_notice_bridge, ) as client: - if not client.connect(timeout=5): - print("server is unavailable", file=sys.stderr) + if not client.wait_until_ready(timeout=5): + print(f"client did not become ready: {client.status.phase.value}", file=sys.stderr) return 1 - deadline = time.monotonic() + 5.0 - while client.status.phase is not ClientPhase.READY: - client.update() - status = client.status - if status.phase in (ClientPhase.RECOVERY_REQUIRED, ClientPhase.REJECTED): - print(status.reason or status.phase.value, file=sys.stderr) - return 2 - if time.monotonic() >= deadline: - print(f"client did not become ready: {status.phase.value}", file=sys.stderr) - return 1 - time.sleep(0.01) - if not client.is_layer_reachable(content): print("content layer is outside the synchronized graph", file=sys.stderr) return 1 @@ -118,6 +117,9 @@ def main() -> int: next_tick += interval if (sleep_for := next_tick - time.monotonic()) > 0: time.sleep(sleep_for) + if not client.submit_and_wait(timeout=5): + print("local edits were not durably acknowledged", file=sys.stderr) + return 1 return 0 diff --git a/examples/usd_native_client/README.md b/examples/usd_native_client/README.md index eefc455..b9006fb 100644 --- a/examples/usd_native_client/README.md +++ b/examples/usd_native_client/README.md @@ -2,8 +2,8 @@ This example runs two independent USD-native clients: -- `demo.py` publishes a moving sphere while receiving authoritative layered - replay into a separate mirror stage. +- `demo.py` uses `ManagedClient` to publish a moving sphere and receive + authoritative layered replay on the same application stage. - `peer.py` pre-authors a cube in another stage and publishes its current edit target from a separate process. @@ -24,9 +24,8 @@ reports `local_valid=True` and `peer_valid=True`. The launcher then stops its temporary server and peer process and removes its temporary event log. Pressing `Ctrl+C` performs the same cleanup during an unbounded run. -The ownership rule is visible in `demo.py`: its publisher observes -the author stage's session layer, while `UsdReceiver` owns different session -layers on the mirror stage. All three stages share the same read-only base -layer. +The ownership rule is visible in `demo.py`: `ManagedClient` selects its own +transient authoring layer below the authoritative collaboration layers. Both +processes open equivalent read-only base content. See the [USD-native integration contract](../../docs/usd-native-integration.md) for the corresponding host integration rules. diff --git a/examples/usd_native_client/demo.py b/examples/usd_native_client/demo.py index 2a800f0..401dd7c 100644 --- a/examples/usd_native_client/demo.py +++ b/examples/usd_native_client/demo.py @@ -12,7 +12,7 @@ if str(_REPO_ROOT) not in sys.path: sys.path.insert(0, str(_REPO_ROOT)) -from openusdconnect import ManagedClient # noqa: E402, I001 +from openusdconnect import ClientPhase, ManagedClient, TransactionRejectedError # noqa: E402, I001 from pxr import Gf, Sdf, Usd, UsdGeom # noqa: E402 BASE_USD = Path(__file__).with_name("empty.usda") @@ -54,8 +54,8 @@ def run(args: argparse.Namespace, *, expect_peer: bool = False) -> int: try: with client: - if not client.connect(timeout=5): - print("could not connect", file=sys.stderr) + if not client.wait_until_ready(timeout=5): + print(f"client did not become ready: {client.status.phase.value}", file=sys.stderr) return 1 sphere = UsdGeom.Sphere.Define(stage, LOCAL_SPHERE_PATH) @@ -63,9 +63,7 @@ def run(args: argparse.Namespace, *, expect_peer: bool = False) -> int: sphere.CreateDisplayColorAttr([Gf.Vec3f(0.08, 0.45, 1.0)]) translate = UsdGeom.Xformable(sphere).AddTranslateOp() translate.Set(Gf.Vec3d(0.0, 1.25, 0.0)) - if client.update().submitted_events == 0: - print("initial sphere batch was not sent", file=sys.stderr) - return 1 + client.update() expected_paths = [LOCAL_SPHERE_PATH] if expect_peer: @@ -92,8 +90,13 @@ def run(args: argparse.Namespace, *, expect_peer: bool = False) -> int: print("publishing LocalSphere and receiving PeerCube; press Ctrl+C to stop") while args.seconds <= 0 or time.monotonic() - started < args.seconds: elapsed = time.monotonic() - started - translate.Set(Gf.Vec3d(math.sin(elapsed) * 2.5, 1.25, 0.0)) update = client.update() + status = client.status + if status.phase in (ClientPhase.RECOVERY_REQUIRED, ClientPhase.REJECTED): + print(status.reason or status.phase.value, file=sys.stderr) + return 2 + if status.phase is ClientPhase.READY: + translate.Set(Gf.Vec3d(math.sin(elapsed) * 2.5, 1.25, 0.0)) now = time.monotonic() if now >= next_report: @@ -110,7 +113,16 @@ def run(args: argparse.Namespace, *, expect_peer: bool = False) -> int: next_tick += interval if (sleep_for := next_tick - time.monotonic()) > 0: time.sleep(sleep_for) + if not client.submit_and_wait(timeout=5): + print("local edits were not durably acknowledged", file=sys.stderr) + return 1 return 0 + except (PermissionError, ConnectionError) as exc: + print(f"server rejected the client: {exc}", file=sys.stderr) + return 1 + except TransactionRejectedError as exc: + print(str(exc), file=sys.stderr) + return 2 except KeyboardInterrupt: return 0 diff --git a/examples/usd_native_client/peer.py b/examples/usd_native_client/peer.py index 4b8d174..64522ed 100644 --- a/examples/usd_native_client/peer.py +++ b/examples/usd_native_client/peer.py @@ -11,7 +11,7 @@ if str(_REPO_ROOT) not in sys.path: sys.path.insert(0, str(_REPO_ROOT)) -from openusdconnect import UsdPublisher # noqa: E402, I001 +from openusdconnect import TransactionRejectedError, UsdPublisher # noqa: E402, I001 from pxr import Gf, Sdf, Usd, UsdGeom # noqa: E402 BASE_USD = Path(__file__).with_name("empty.usda") @@ -52,10 +52,17 @@ def main() -> int: department="lookdev", persist_token=False, ) as publisher: - if not publisher.connect(timeout=5): - raise ConnectionError("OpenUSDConnect server is unavailable") + if not publisher.wait_until_ready(timeout=5): + print("OpenUSDConnect server is unavailable", file=sys.stderr) + return 1 sent = publisher.publish_current_edit_target() - except ConnectionError as exc: + if not publisher.submit_and_wait(timeout=5): + print("peer cube was not durably acknowledged", file=sys.stderr) + return 1 + except TransactionRejectedError as exc: + print(str(exc), file=sys.stderr) + return 2 + except (PermissionError, ConnectionError) as exc: print(f"peer could not connect: {exc}", file=sys.stderr) return 1 diff --git a/integrations/blender/__init__.py b/integrations/blender/__init__.py index 5656187..d5785c7 100644 --- a/integrations/blender/__init__.py +++ b/integrations/blender/__init__.py @@ -7,7 +7,7 @@ bl_info = { "name": "USD Connect", "author": "OpenUSDConnect", - "version": (0, 4, 0), + "version": (0, 5, 0), "blender": (4, 4, 0), "location": "View3D > Sidebar > USD Connect", "description": "Real-time USD sync: capture and receive transform edits over the network", diff --git a/integrations/mcp/session.py b/integrations/mcp/session.py index fa12e50..df0ece9 100644 --- a/integrations/mcp/session.py +++ b/integrations/mcp/session.py @@ -14,6 +14,7 @@ import uuid from openusdconnect import token_client +from openusdconnect.client_observer import AppliedBatch, ClientObserver, PlaybackState from openusdconnect.event_apply import apply_events from openusdconnect.protocol_constants import K_SET_STAGE_METADATA, STAGE_METADATA_KEYS from openusdconnect.sender import EventSender, TransactionRejectedError @@ -24,6 +25,20 @@ from .introspection import select_changes +class _MirrorObserver(ClientObserver): + """Feed changes_since() and playback_status() from the mirror receiver.""" + + def __init__(self, session: ConnectionSession): + self._session = session + + def on_applied(self, batch: AppliedBatch) -> None: + # A whole drain shares its final seq; enough for "changed since N". + for path in batch.prim_paths: + self._session._dirty[path] = batch.seq + + def on_playback_state(self, state: PlaybackState) -> None: + self._session._playback_state = state + class ConnectionSession: """Owns the network client and the mirror stage for one MCP process.""" @@ -34,11 +49,11 @@ def __init__(self, config: McpConfig): self.mirror_stage = None self.auth_rejected = False self._origin_base = f"mcp-{uuid.uuid4().hex[:8]}" - # prim_path -> sequence it last changed at, fed by the dispatcher's - # on_applied hook; powers changes_since() diff queries. + # prim_path -> sequence it last changed at, fed by _MirrorObserver; + # powers changes_since() diff queries. self._dirty: dict[str, int] = {} - # Latest PlaybackState the server broadcast, set on the receiver thread. - self._playback_state: dict | None = None + # Latest PlaybackState the server broadcast, delivered during pump(). + self._playback_state: PlaybackState | None = None @property def connected(self) -> bool: @@ -123,30 +138,10 @@ def _start_mirror(self) -> None: origin=f"{self._origin_base}-recv", token=recv_token, persist_token=False, - on_playback_state=self._on_playback_state, - on_applied=self._on_applied, + observer=_MirrorObserver(self), ) self.receiver.start() - def _on_applied(self, prim_paths: list) -> None: - """Stamp each applied prim with the current sequence so changes_since can - report it. Coarse at drain granularity (a whole drain shares its final - seq), which is fine for 'what changed since N' polling.""" - seq = self.receiver.applying_seq if self.receiver else 0 - for path in prim_paths: - self._dirty[path] = seq - - def _on_playback_state(self, msg: dict) -> None: - """Store the latest shared-playhead snapshot. Runs on the receiver - thread, so assign a fresh dict (an atomic reference swap) rather than - mutating in place.""" - self._playback_state = { - "playing": msg.get("playing"), - "time": msg.get("time"), - "rate": msg.get("rate"), - "leader_client_id": msg.get("leader_client_id") or "", - } - def _seed_metadata(self, metadata: dict | None) -> None: payload = {k: v for k, v in (metadata or {}).items() if k in STAGE_METADATA_KEYS} if payload: @@ -174,7 +169,7 @@ def send(self, events: list[dict]) -> dict: """Send one txn and drain the mirror until it reflects the write.""" if not self.connected: raise ToolError("not connected, call usd_connect first", code="not_connected") - if self.receiver is not None and not self.receiver.synchronized: + if self.receiver is not None and not self.receiver.status.synchronized: if not self._drain_initial_replay(): raise ToolError( "the mirror is still applying the initial replay", @@ -201,7 +196,7 @@ def _drain_after_write(self) -> bool: # A nonblocking poll avoids a reconnect handshake extending the read budget. try: acknowledged = self.sender.flush(timeout=0) - applied = self.receiver.update() + applied = self.receiver.update().applied_events if not acknowledged: # The acknowledgement may arrive while queued events are applied. acknowledged = self.sender.flush(timeout=0) @@ -213,7 +208,7 @@ def _drain_after_write(self) -> bool: return False if checkpoint is not None: if ( - self.receiver.synchronized + self.receiver.status.synchronized and self.receiver.server_instance == checkpoint.server_instance and self.receiver.replay_epoch == checkpoint.epoch and self.receiver.last_seq >= checkpoint.head_seq @@ -230,16 +225,16 @@ def _drain_initial_replay(self) -> bool: if self.receiver is None: return False deadline = time.monotonic() + self.config.read_after_write_timeout_s - while not self.receiver.synchronized and time.monotonic() < deadline: + while not self.receiver.status.synchronized and time.monotonic() < deadline: if self.pump() == 0: time.sleep(0.005) - return self.receiver.synchronized + return self.receiver.status.synchronized def pump(self) -> int: """Non-blocking drain so introspection reflects recent foreign edits.""" if self.receiver is None: return 0 - return self.receiver.update() + return self.receiver.update().applied_events def require_mirror(self): """Return the mirror stage or raise if introspection is unavailable.""" @@ -298,16 +293,17 @@ def playback_status(self) -> dict: mirror's receiver; disabled under --no-mirror).""" if not self.connected: raise ToolError("not connected, call usd_connect first", code="not_connected") + self.pump() state = self._playback_state if state is None: return {"ok": True, "observed": False} - leader = state.get("leader_client_id") or "" + leader = state.leader_client_id return { "ok": True, "observed": True, - "playing": bool(state.get("playing")), - "time": state.get("time"), - "rate": state.get("rate"), + "playing": state.playing, + "time": state.time, + "rate": state.rate, "leader_client_id": leader, "has_leader": bool(leader), "is_leader": bool(leader) and leader == self.config.client_id, @@ -326,7 +322,7 @@ def _mirror_prim_count(self) -> int: def status(self) -> dict: if self.receiver is not None: - if not self.receiver.synchronized: + if not self.receiver.status.synchronized: self.pump() return { "ok": True, @@ -336,7 +332,7 @@ def status(self) -> dict: "client_id": self.config.client_id, "department": self.config.department, "mirror_enabled": self.config.mirror_enabled, - "mirror_synchronized": bool(self.receiver and self.receiver.synchronized), + "mirror_synchronized": bool(self.receiver and self.receiver.status.synchronized), "mirror_prim_count": self._mirror_prim_count(), "last_seq": self.receiver.last_seq if self.receiver else 0, "auth_rejected": self.auth_rejected, diff --git a/integrations/usdview/connection.py b/integrations/usdview/connection.py index 2b982ed..a7b6813 100644 --- a/integrations/usdview/connection.py +++ b/integrations/usdview/connection.py @@ -19,6 +19,7 @@ sys.path.insert(0, _PROJECT_ROOT) from openusdconnect.cli_common import parse_bool # noqa: E402 +from openusdconnect.client_observer import AppliedBatch, ClientObserver # noqa: E402 from openusdconnect.defaults import DEFAULT_HOST, DEFAULT_SYNC_PORT # noqa: E402 from openusdconnect.usd_client import UsdReceiver # noqa: E402 @@ -105,8 +106,7 @@ def start( host=host, port=port, token=token, - on_applied=_on_applied if _translate_openpbr else None, - on_applied_events=_on_applied_events, + observer=_ViewerObserver(), ) except ValueError as exc: LOG.error("Cannot start receiver: %s", exc) @@ -157,7 +157,7 @@ def _tick() -> None: """Drain the receive queue each frame. Post-apply fixups (interface-input forwarding, OpenPBR translation) are - driven by the receiver's ``on_applied`` callback, scoped to the prims + driven by ``_ViewerObserver.on_applied``, scoped to the prims each batch actually changed. """ if _receiver is not None: @@ -199,26 +199,18 @@ def _forward_interface_edits(stage, events: list[dict]) -> None: consumer.Set(value) -def _on_applied_events(events: list[dict]) -> None: - """Post-apply conditioning scoped to what each event edited.""" - if _usdview_api is None: - return - stage = _usdview_api.dataModel.stage - if stage is None: - return - _forward_interface_edits(stage, events) - +class _ViewerObserver(ClientObserver): + """Condition each applied batch for usdview's renderer.""" -def _on_applied(prim_paths: list[str]) -> None: - """Translate OpenPBR materials owning the just-applied prims.""" - if _usdview_api is None: - return - stage = _usdview_api.dataModel.stage - if stage is None: - return - from integrations.openpbr_translate import translate_openpbr_for_paths + def on_applied(self, batch: AppliedBatch) -> None: + stage = _usdview_api.dataModel.stage if _usdview_api is not None else None + if stage is None: + return + _forward_interface_edits(stage, batch.events) + if _translate_openpbr: + from integrations.openpbr_translate import translate_openpbr_for_paths - translate_openpbr_for_paths(stage, prim_paths) + translate_openpbr_for_paths(stage, batch.prim_paths) def refresh_asset_dependency(asset_path: str | None = None) -> dict: @@ -271,7 +263,7 @@ def status() -> dict: "running": True, "host": f"{_receiver.receiver.host}:{_receiver.receiver.port}", "client_id": _receiver.receiver.client_id or "", - "receiver_connected": _receiver.connected, + "receiver_connected": _receiver.status.connected, "last_seq": _receiver.last_seq, "pending_asset_dependencies": list(_receiver.pending_asset_dependencies), "translate_openpbr": _translate_openpbr, diff --git a/openusdconnect/__init__.py b/openusdconnect/__init__.py index bcdf2c2..e1a3f56 100644 --- a/openusdconnect/__init__.py +++ b/openusdconnect/__init__.py @@ -24,6 +24,13 @@ UsdStageAdapter, ) from .checkpoints import MirrorCheckpoint, TransactionCheckpoint +from .client_observer import ( + AppliedBatch, + ClientObserver, + PlaybackClaim, + PlaybackState, + StageMetadata, +) from .client_types import ClientPhase, ClientStatus, SyncUpdate from .codec import ( DecodeResult, @@ -68,6 +75,8 @@ from .usd_client import UsdPublisher, UsdReceiver __all__ = [ + "AppliedBatch", + "ClientObserver", "DCCAdapter", "ClientPhase", "ClientStatus", @@ -82,6 +91,8 @@ "MirrorCheckpoint", "MockAdapter", "NoticeEmitter", + "PlaybackClaim", + "PlaybackState", "PluginEnvironmentError", "PluginEnvironmentResult", "ReceiverThread", @@ -97,6 +108,7 @@ "SharedRecoveryAssessment", "SharedRecoveryLayer", "SharedStageClient", + "StageMetadata", "SyncUpdate", "TransactionRejectedError", "TransactionFailure", diff --git a/openusdconnect/_client_base.py b/openusdconnect/_client_base.py new file mode 100644 index 0000000..336c78f --- /dev/null +++ b/openusdconnect/_client_base.py @@ -0,0 +1,441 @@ +"""Lifecycle shells the high-level clients build on. + +``ClientBase`` owns the lifecycle, status, observer hooks, 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. +""" + +from __future__ import annotations + +from collections.abc import Callable, Sequence + +from pxr import Usd + +from ._client_lifecycle import ( + DEFAULT_WAIT_TIMEOUT_S, + ClientCallbackQueue, + _pause_before_poll, + 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_types import ClientStatus, SyncUpdate +from .coalescing import TransformCoalescingWindow +from .emitter import NoticeEmitter, PrimChannel +from .receiver import ReceiverThread +from .recovery import RecoveryArtifact, RecoveryError, RejectionDisposition, TransactionFailure +from .sender import EventSender + + +class ClientBase: + """Lifecycle, status, observer hooks, 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. + """ + + _sender: EventSender | None = None + _receiver: ReceiverThread | None = None + # Reconnection deliberately suspended by the host (UsdPublisher.disconnect). + _paused = False + + def __init__( + self, + *, + host: str, + port: int, + token: str | None, + persist_token: bool, + observer: ClientObserver | None, + ): + self._host = host + self._port = port + self._started = False + 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, + ) + + @property + def client_id(self) -> str: + """Stable identity shared by every connection role.""" + return self._endpoints[0].client_id + + @property + def stage_metadata(self) -> StageMetadata: + return stage_metadata_from_message(self._endpoints[0].stage_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 + 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() + phase = compute_phase( + closed=self._closed, + recovery_required=failure is not None or bool(rebuild_reason), + rejected=rejected, + parked=self._is_parked(), + replaying=receiver is not None and receiver.connected and not receiver.synchronized, + ready=connected and synchronized, + connecting=( + self._started + and not self._paused + and not (receiver is not None and receiver.stopped) + ), + ) + if failure is not None: + reason = str(failure) + else: + reasons = [e.rejection_reason for e in (sender, receiver) if e is not None] + 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, + 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, + failure=failure, + recovery=None if sender is None else sender.recovery_incident, + reason=reason, + auth_rejected=auth_rejected, + has_unsent_changes=not self._closed and self._has_unsent_changes(), + **self._role_status(), + ) + + def start(self): + """Start background networking without blocking and return this client.""" + self._require_open() + if not self._started: + if self._receiver is not None: + self._receiver.start() + self._started = True + return self + + def connect(self, timeout: float | None = DEFAULT_WAIT_TIMEOUT_S) -> bool: + """Start and complete the handshakes within ``timeout``. + + Queued replay still needs :meth:`update` on the stage-owning thread. + """ + self.start() + if self._receiver is not None: + deadline = deadline_after(timeout) + if not self._receiver.wait_connected(timeout): + raise_if_rejected(self._receiver, type(self).__name__) + return False + timeout = remaining_time(deadline) + return self._sender is None or self._connect_sender(timeout) + + def update(self, *, max_messages: int | None = None) -> SyncUpdate: + """Apply received work and submit local work; call it every frame.""" + raise NotImplementedError + + def wait_until_ready(self, timeout: float | None = DEFAULT_WAIT_TIMEOUT_S) -> bool: + """Pump updates until ``READY``; ``False`` only on timeout.""" + return wait_until_ready(self, timeout) + + def close(self) -> None: + """Stop networking, deliver queued notifications, and release resources. + + Called from an observer while events are applied, it takes effect when + that apply returns. + """ + if self._closed: + return + if self._apply_depth: + self._close_requested = True + return + self._closed = True + try: + if self._sender is not None: + self._sender.disconnect() + if self._receiver is not None: + stop_receiver(self._receiver) + # A token issued by the last handshake must still reach the host. + self._callbacks.close() + finally: + self._release() + + def __enter__(self): + return self.start() + + def __exit__(self, exc_type, exc, traceback) -> bool: + self.close() + return False + + @property + 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 + + def _is_parked(self) -> bool: + """Status hook: no stage is bound.""" + return False + + def _rebuild_reason(self) -> str: + """Status hook: why the host must rebuild before continuing, else empty.""" + return "" + + def _prepared_events(self) -> int: + return 0 + + def _has_unsent_changes(self) -> bool: + return False + + def _role_status(self) -> dict: + """Status hook: client-specific ``ClientStatus`` fields.""" + return {} + + def _release(self) -> None: + """Release subclass resources after networking has stopped.""" + + def _require_open(self) -> None: + if self._closed: + raise RuntimeError(f"{type(self).__name__} is closed") + + def _require_started(self) -> None: + self._require_open() + if not self._started: + raise RuntimeError(f"{type(self).__name__} has not been started") + + def _dispatch(self, operation: Callable, /, *args, **kwargs): + """Run a dispatcher operation that may call observer delivery methods.""" + self._apply_depth += 1 + try: + return operation(*args, **kwargs) + finally: + self._apply_depth -= 1 + if self._close_requested and not self._apply_depth: + self.close() + + def _begin_update(self) -> bool: + """Deliver queued notifications; ``False`` when one of them closed the client.""" + self._require_started() + self._callbacks.drain() + return not self._closed + + def _connect_sender(self, timeout: float | None = None) -> bool: + if self._sender.connected: + return True + if not self._sender.connect(timeout=timeout): + raise_if_rejected(self._sender, type(self).__name__) + return False + return True + + def _progress(self, applied: int = 0, submitted: int = 0) -> SyncUpdate: + sender = self._sender + if sender is None: + return SyncUpdate(applied_events=applied, submitted_events=submitted) + return SyncUpdate( + applied_events=applied, + submitted_events=submitted, + acknowledged_events_delta=sender.drain_acknowledged_event_count(), + pending_events=sender.pending_event_count, + recovery=sender.recovery_incident, + ) + + +class PublishingClientBase(ClientBase): + """Adds the sender role: recovery, playback, and durable completion.""" + + @property + def recovery_artifact(self) -> RecoveryArtifact | None: + """Exact quarantined transactions for integration-owned recovery.""" + return self._sender.recovery_artifact + + def flush(self, timeout: float | None = DEFAULT_WAIT_TIMEOUT_S) -> bool: + """Wait until submitted work is durable; ``False`` on timeout.""" + self._require_open() + return self._sender.flush(timeout) + + def submit_and_wait(self, timeout: float | None = DEFAULT_WAIT_TIMEOUT_S) -> bool: + """Publish noticed edits and wait until durable; ``False`` only on timeout.""" + return submit_and_wait(self, timeout) + + def claim_playback(self, time: float | None = None) -> bool: + """Request the shared-playback leader role.""" + self._require_open() + return self._sender.claim_playback(time=time) + + def send_playback_control( + self, + action: str, + *, + time: float | None = None, + rate: float | None = None, + ) -> bool: + """Drive the shared playhead (leader only).""" + self._require_open() + return self._sender.send_playback_control(action, time=time, rate=rate) + + def _require_recoverable_failure(self) -> TransactionFailure: + self._require_open() + failure = self._sender.transaction_failure + if failure is None: + raise RecoveryError("no_incident", "there is no recovery incident to resolve") + if failure.disposition is not RejectionDisposition.RECOVERABLE_CONFLICT: + raise RecoveryError( + "wrong_recovery_kind", + f"{failure.code_name} is {failure.disposition.value}, not recoverable", + ) + return failure + + def _reconnect_repaired(self, txn_id: int) -> None: + if not self._connect_sender(): + raise ConnectionError( + f"transaction {txn_id} repaired but reconnect to " + f"{self._host}:{self._port} failed; it remains queued" + ) + + def _replay_to_fresh_checkpoint(self, timeout: float | None) -> None: + """Apply replay through a new server watermark before resolving recovery.""" + deadline = deadline_after(timeout) + receiver = self._receiver + reconnect = receiver.reconnect + receiver.reconnect = True + try: + receiver.request_replay_from(self.last_seq + 1) + while True: + self._apply_queued() + self._require_open() + if receiver.synchronized: + return + if not _pause_before_poll(deadline): + raise TimeoutError("authoritative recovery replay timed out") + finally: + receiver.reconnect = reconnect + + def _apply_queued(self, max_messages: int | None = None) -> int: + """Apply queued receiver messages; clients with a receiver implement it.""" + raise NotImplementedError + + def _resume_sender_after_recovery(self, timeout: float | None) -> None: + """Best-effort producer reconnect after recovery has committed.""" + try: + self._connect_sender(timeout=timeout) + except (PermissionError, ConnectionError): + # Recovery already completed and must not look rolled back. Status + # exposes rejection/offline state; update() retries ordinary loss. + pass + + +class EmitterClientBase(PublishingClientBase): + """Adds stage-edit capture through a NoticeEmitter with transform coalescing.""" + + def _init_emitter( + self, + stage: Usd.Stage, + *, + attr_filter: Callable[[str], bool] | None, + replicated_api_schemas: set[str] | None, + extra_channels: Sequence[PrimChannel] | None, + transform_coalesce_seconds: float, + ) -> None: + self._transform_coalescing = TransformCoalescingWindow(transform_coalesce_seconds) + self._emitter = NoticeEmitter( + stage, + attr_filter=attr_filter, + replicated_api_schemas=replicated_api_schemas, + extra_channels=extra_channels, + ) + + @property + def sender(self) -> EventSender: + """The underlying :class:`EventSender`; a diagnostic handle.""" + return self._sender + + def repair_and_resume(self, events: list[dict]) -> int: + """Replace a recoverable transaction and resume its ordered outbox. + + The application must first reconcile its stage with authoritative + state and rebuild *events* for that state. The repaired transaction is + assigned the original rejected ID; later quarantined transactions keep + their existing IDs and replay after it. + """ + self._require_recoverable_failure() + txn_id = self._sender.repair_rejected_transaction(events) + self._reconnect_repaired(txn_id) + return txn_id + + @property + def emitter(self) -> NoticeEmitter: + """The underlying :class:`NoticeEmitter`; a diagnostic handle.""" + return self._emitter + + def flush(self, timeout: float | None = DEFAULT_WAIT_TIMEOUT_S) -> bool: + """Submit a coalesced transform, then wait until submitted work is durable.""" + self._require_open() + deadline = deadline_after(timeout) + if self._transform_coalescing.buffering: + if not self._ready_to_publish(remaining_time(deadline)): + return False + events = self._transform_coalescing.force(self._emitter) + if events and not self._send(events): + return False + return self._sender.flush(remaining_time(deadline)) + + def publish_current_edit_target(self) -> int: + """Queue every opinion in the edit target; returns events this call submitted.""" + self._require_open() + if not self._can_capture_edit_target(): + return 0 + if self._emitter.prepared_event_count: + raise RuntimeError( + "an earlier publisher batch is still prepared; call update() " + "before publishing the current edit target" + ) + self.start() + self._emitter.prepare_snapshot_events_for_send() + return self.update().submitted_events + + def _ready_to_publish(self, timeout: float | None) -> bool: + """Connect the sender within *timeout*; ``False`` while publishing must wait.""" + return self._connect_sender(timeout) and self._is_synchronized() + + def _can_capture_edit_target(self) -> bool: + return True + + def _send(self, events: list[dict]) -> int: + if not events: + return 0 + if self._sender.send_events(events): + self._emitter.mark_prepared_events_sent(events) + self._transform_coalescing.mark_submitted() + return len(events) + return 0 + + def _prepare_outgoing_events(self) -> list[dict]: + return self._transform_coalescing.prepare(self._emitter) + + def _prepared_events(self) -> int: + return self._emitter.prepared_event_count + + def _has_unsent_changes(self) -> bool: + return self._emitter.has_local_changes + + def _release(self) -> None: + self._emitter.cleanup() diff --git a/openusdconnect/_client_lifecycle.py b/openusdconnect/_client_lifecycle.py index 691816c..5780657 100644 --- a/openusdconnect/_client_lifecycle.py +++ b/openusdconnect/_client_lifecycle.py @@ -1,21 +1,25 @@ -"""Shared transport lifecycle operations; no stage or layer policy.""" +"""Helpers shared by the high-level clients.""" from __future__ import annotations import logging +import queue import threading import time from collections.abc import Callable from typing import TYPE_CHECKING -from ._client_utils import resolve_client_token +from .client_types import ClientPhase, ClientStatus +from .sender import TransactionRejectedError if TYPE_CHECKING: from .receiver import ReceiverThread - from .sender import EventSender LOG = logging.getLogger(__name__) +DEFAULT_WAIT_TIMEOUT_S = 10.0 +_POLL_INTERVAL_S = 0.01 + def deadline_after(timeout: float | None) -> float | None: return None if timeout is None else time.monotonic() + max(timeout, 0.0) @@ -25,6 +29,168 @@ def remaining_time(deadline: float | None) -> float | None: return None if deadline is None else max(0.0, deadline - time.monotonic()) +def _pause_before_poll(deadline: float | None) -> bool: + remaining = remaining_time(deadline) + if remaining is not None and remaining <= 0: + return False + time.sleep(_POLL_INTERVAL_S if remaining is None else min(_POLL_INTERVAL_S, remaining)) + return True + + +def compute_phase( + *, + closed: bool, + recovery_required: bool, + rejected: bool, + parked: bool, + replaying: bool, + ready: bool, + 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 + + +def raise_if_blocked(client, status: ClientStatus) -> None: + """Raise for a state that further updates cannot resolve.""" + if status.failure is not None: + raise TransactionRejectedError(status.failure) + name = type(client).__name__ + if status.phase is ClientPhase.CLOSED: + raise RuntimeError(f"{name} is closed") + if status.phase is ClientPhase.REJECTED: + if status.auth_rejected: + raise PermissionError(status.reason or f"{name} authentication rejected") + raise ConnectionError(status.reason or f"{name} connection rejected") + if status.phase is ClientPhase.OFFLINE: + raise ConnectionError(status.reason or f"{name} is offline and not reconnecting") + if status.phase is ClientPhase.PARKED: + raise RuntimeError(f"{name} has no bound stage; call rebind_stage() first") + if status.phase is ClientPhase.RECOVERY_REQUIRED: + raise RuntimeError(status.reason or f"{name} requires recovery") + + +def wait_until_ready(client, timeout: float | None) -> bool: + """Pump updates on the calling thread until ready; ``False`` only on timeout.""" + client.start() + deadline = deadline_after(timeout) + while True: + client.update() + status = client.status + if status.phase is ClientPhase.READY: + return True + raise_if_blocked(client, status) + if not _pause_before_poll(deadline): + return False + + +def submit_and_wait(client, timeout: float | None) -> bool: + """Submit noticed edits and wait until durable; ``False`` only on timeout.""" + client.start() + deadline = deadline_after(timeout) + while True: + client.update() + status = client.status + raise_if_blocked(client, status) + # flush(0) also releases transform coalescing and must not block the pump. + if status.phase is ClientPhase.READY and client.flush(timeout=0): + status = client.status + if not status.has_unsent_changes and not status.pending_events: + return True + if not _pause_before_poll(deadline): + 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") @@ -32,38 +198,6 @@ def raise_if_rejected(endpoint, role: str) -> None: raise ConnectionError(endpoint.rejection_reason or f"{role} connection rejected") -def prepare_sender_token( - sender: EventSender, - receiver: ReceiverThread | None, - *, - host: str, - port: int, - persist_token: bool, -) -> None: - """Fill missing sender credentials; issued tokens are shared by callbacks.""" - if sender.token is not None: - return - token = receiver.token if receiver is not None else None - if token is None: - token = resolve_client_token(host, port, None, persist_token) - # A handshake can supply a token while stored credentials are being read. - if sender.token is None: - sender.token = token - - -def share_client_token( - token: str, - sender: EventSender, - receiver: ReceiverThread, - callback: Callable[[str], None] | None, -) -> None: - """Update both connections before persistence or application callbacks can fail.""" - sender.token = token - receiver.token = token - if callback is not None: - callback(token) - - def stop_receiver(receiver: ReceiverThread) -> None: receiver.stop() if receiver.is_alive() and receiver is not threading.current_thread(): diff --git a/openusdconnect/_client_utils.py b/openusdconnect/_client_utils.py index 33dc294..a15cd0c 100644 --- a/openusdconnect/_client_utils.py +++ b/openusdconnect/_client_utils.py @@ -2,6 +2,7 @@ from __future__ import annotations +import threading import uuid from collections.abc import Callable @@ -49,40 +50,55 @@ def resolve_client_token( return load_token(host, port) -def client_token_callback( - host: str, - port: int, - persist: bool, -) -> Callable[[str], None] | None: - if not persist: - return None - return lambda token: save_token(host, port, token) - - -def client_token_handlers( - host: str, - port: int, - persist: bool, - on_token_issued: Callable[[str], None] | None, -) -> Callable[[str], None] | None: - """Return a callback that chains *on_token_issued* with disk persistence.""" - persist_cb = client_token_callback(host, port, persist) - if on_token_issued is None or persist_cb is None: - return on_token_issued or persist_cb - - def _both(token: str) -> None: - persist_cb(token) - on_token_issued(token) - - return _both +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, + ): + 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) + + def current(self) -> str | None: + """The token for a connection attempt, loading a stored one if none is known.""" + with self._lock: + if self.token is None and self._persist: + self.token = load_token(self._host, self._port) + return self.token + + def issued(self, token: str) -> None: + """Adopt a server-issued token, persist it, then notify the host.""" + 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.""" + return { + "token": self.token, + "token_provider": self.current, + "on_token_issued": self.issued, + } __all__ = [ "ClientPhase", "ClientStatus", + "ClientCredential", "client_origin", - "client_token_callback", - "client_token_handlers", "require_app_name", "resolve_client_token", "SyncUpdate", diff --git a/openusdconnect/_observer_hooks.py b/openusdconnect/_observer_hooks.py new file mode 100644 index 0000000..8f59b4c --- /dev/null +++ b/openusdconnect/_observer_hooks.py @@ -0,0 +1,121 @@ +"""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/_version.py b/openusdconnect/_version.py index f8462e6..ab4462e 100644 --- a/openusdconnect/_version.py +++ b/openusdconnect/_version.py @@ -1,3 +1,3 @@ """Canonical OpenUSDConnect release version.""" -__version__ = "0.4.0" +__version__ = "0.5.0" diff --git a/openusdconnect/client_observer.py b/openusdconnect/client_observer.py new file mode 100644 index 0000000..1aade23 --- /dev/null +++ b/openusdconnect/client_observer.py @@ -0,0 +1,110 @@ +"""Host notifications for the high-level clients.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass + +from .events import Event +from .protocol_constants import IMPORT_KINDS + + +class AppliedBatch: + """Authoritative events applied by one delivery. + + ``seq`` is the sequence at the end of the drain that applied the batch; + every batch of one drain shares it. With an adapter-backed receiver, + ``events`` are the projected adapter events. + """ + + __slots__ = ("seq", "events", "_prim_paths", "_imported_paths") + + def __init__(self, seq: int, events: Sequence[Event]): + self.seq = seq + self.events = events + self._prim_paths: tuple[str, ...] | None = None + self._imported_paths: tuple[str, ...] | None = None + + @property + def prim_paths(self) -> tuple[str, ...]: + """Sorted unique prim paths the batch touched.""" + if self._prim_paths is None: + self._prim_paths = tuple(sorted({e["prim"] for e in self.events if e.get("prim")})) + return self._prim_paths + + @property + def imported_paths(self) -> tuple[str, ...]: + """Prims whose references or payloads may have brought in new content.""" + if self._imported_paths is None: + self._imported_paths = tuple( + e["prim"] for e in self.events if e.get("k") in IMPORT_KINDS and e.get("prim") + ) + return self._imported_paths + + +@dataclass(frozen=True, slots=True) +class StageMetadata: + """Stage-level settings; ``None`` where the server has no authored opinion.""" + + time_codes_per_second: float | None = None + frames_per_second: float | None = None + start_time_code: float | None = None + end_time_code: float | None = None + meters_per_unit: float | None = None + up_axis: str | None = None + + +@dataclass(frozen=True, slots=True) +class PlaybackState: + playing: bool + time: float + rate: float + leader_client_id: str + + +@dataclass(frozen=True, slots=True) +class PlaybackClaim: + """Reply to ``claim_playback()``.""" + + granted: bool + leader_client_id: str + reason: str = "" + + +class ClientObserver: + """Override the notifications a host needs; no method runs on a network thread. + + Delivery methods run while the client applies authoritative state: raising + from one in ``update()`` rolls the batch back and replays it, and stage + edits made in them are not published. Notification methods only observe + and arrive in ``update()`` or ``close()``: raising propagates, and later + notifications wait for the next call. A client calls only the methods it + supports and a subclass overrides. + """ + + def on_applied(self, batch: AppliedBatch) -> None: + """Delivery (ManagedClient, UsdReceiver): authoritative events were applied.""" + + 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.""" + + def on_playback_state(self, state: PlaybackState) -> None: + """Notification (receiving clients): the shared playhead changed.""" + + def on_playback_claim(self, result: PlaybackClaim) -> None: + """Notification (receiving clients): the server answered ``claim_playback()``.""" + + def on_token_issued(self, token: str) -> None: + """Notification: the server issued a credential for host-owned storage.""" + + +__all__ = [ + "AppliedBatch", + "ClientObserver", + "PlaybackClaim", + "PlaybackState", + "StageMetadata", +] diff --git a/openusdconnect/client_types.py b/openusdconnect/client_types.py index d0a2574..32f7240 100644 --- a/openusdconnect/client_types.py +++ b/openusdconnect/client_types.py @@ -18,16 +18,19 @@ class ClientPhase(StrEnum): RECOVERY_REQUIRED = "recovery_required" REJECTED = "rejected" CLOSED = "closed" + PARKED = "parked" @dataclass(frozen=True, slots=True) class ClientStatus: - """Immutable client state suitable for application and UI polling. - - A directional connection is ``None`` when that role is not present. - ``acknowledged_events_total`` is cumulative for the lifetime of the - client instance, including producer-session recovery. Connection and - replay readiness do not imply that every submitted edit is durable. + """Immutable client state for UI polling, taken on the stage-owning thread. + + ``connected`` means every role the client has is connected; a + directional field is ``None`` when that role is absent. + ``acknowledged_events_total`` is cumulative for the lifetime of the client + instance. Connection and replay readiness do not imply that every + submitted edit is durable. ``edit_target_is_published`` is ``None`` for + clients that publish from any edit target or do not publish. """ phase: ClientPhase @@ -41,11 +44,26 @@ class ClientStatus: failure: TransactionFailure | None = None recovery: RecoveryIncident | None = None reason: str = "" + auth_rejected: bool = False + has_unsent_changes: bool = False + deferred_events: int = 0 + deferred_layer_keys: tuple[str, ...] = () + edit_target_is_published: bool | None = None + recovery_stage_pending: bool = False + + @property + def can_author(self) -> bool: + """Whether edits to the current edit target will be published now.""" + return ( + self.phase is ClientPhase.READY + and self.sender_connected is not None + and self.edit_target_is_published is not False + ) @dataclass(frozen=True, slots=True) class SyncUpdate: - """Work completed by one bidirectional client update call. + """Work completed by one client update call; read ``client.status`` for state. ``acknowledged_events_delta`` is consumed by this update and therefore is not a cumulative counter. diff --git a/openusdconnect/dispatcher.py b/openusdconnect/dispatcher.py index bf28ed1..edb8b3d 100644 --- a/openusdconnect/dispatcher.py +++ b/openusdconnect/dispatcher.py @@ -407,7 +407,9 @@ def __init__( self.on_resync = on_resync self.on_applied = on_applied self.on_applied_events = on_applied_events - self._last_seq = 0 + # A receiver resuming after sequence N already holds 1..N. + self._last_seq = receiver.sync_from - 1 + self._drained_message_count = 0 self._applying_seq: int | None = None self._asset_stage = None self._asset_events: dict[tuple[str, str, str], _TrackedAssetEvent] = {} @@ -420,6 +422,11 @@ def __init__( def last_seq(self) -> int: return self._last_seq + @property + def drained_message_count(self) -> int: + """Messages the last :meth:`drain_and_apply` took from the receiver queue.""" + return self._drained_message_count + @last_seq.setter def last_seq(self, value: int) -> None: self._last_seq = value @@ -453,6 +460,7 @@ def drain_and_apply(self, *, max_messages: int | None = None) -> int: if max_messages is None else self.receiver.drain_queue(max_messages=max_messages) ) + self._drained_message_count = len(bufs) if not bufs: self.receiver.mark_replay_applied() return 0 @@ -479,7 +487,8 @@ def drain_and_apply(self, *, max_messages: int | None = None) -> int: if result.resync_requested: self._clear_asset_dependencies() if self.on_resync is not None: - self.on_resync() + with self._suppressed(): + self.on_resync() if self._layer_router is not None: applied = self._apply_layered( result.received_records, @@ -515,8 +524,7 @@ def bind_layered_stage(self, stage: Usd.Stage) -> None: """Move layered and shared receiver state to a replacement stage.""" if self._layer_router is None: return - suppress_ctx = self.emitter.suppressed() if self.emitter else nullcontext() - with suppress_ctx: + with self._suppressed(): self._layer_router.bind(stage) self._stage_session_state.bind(stage) if self._projection_state is not None: @@ -549,8 +557,7 @@ def close(self) -> None: self._release_layered_state() def _release_layered_state(self) -> None: - suppress_ctx = self.emitter.suppressed() if self.emitter else nullcontext() - with suppress_ctx: + with self._suppressed(): self._stage_session_state.close() if self._layer_router is not None: self._layer_router.close() @@ -716,7 +723,7 @@ def _apply_run(run, edit_target=None): self.emitter.invalidate_for_event(event) self._observe_asset_dependencies(run, edit_target=edit_target) - suppress_ctx = self.emitter.suppressed() if self.emitter else nullcontext() + suppress_ctx = self._suppressed() projection_ctx = projection if projection is not None else nullcontext() native_adapter_events: list[dict] = [] with projection_ctx, suppress_ctx: @@ -755,8 +762,7 @@ def _apply(self, events: list[dict]) -> int: """ from .event_apply import apply_events, atomic_apply - suppress_ctx = self.emitter.suppressed() if self.emitter else nullcontext() - with suppress_ctx: + with self._suppressed(): # Arc skip decisions must inspect composition before this batch. stage_skip = self._compute_stage_skip(events) adapter_skip = self._compute_adapter_skip(events, stage_skip) @@ -1012,8 +1018,7 @@ def refresh_resolver_context(self) -> bool: stage = self._asset_dependency_stage() if stage is None: return False - suppress_ctx = self.emitter.suppressed() if self.emitter else nullcontext() - with suppress_ctx: + with self._suppressed(): self._refresh_resolver_context_suppressed(stage) return True @@ -1039,8 +1044,7 @@ def refresh_asset_dependency( "pending": [], } - suppress_ctx = self.emitter.suppressed() if self.emitter else nullcontext() - with suppress_ctx: + with self._suppressed(): self._refresh_resolver_context_suppressed(stage) self._discard_stale_asset_events(stage) dependencies = { @@ -1196,11 +1200,12 @@ def _refresh_asset_dependency_suppressed( self.emitter.invalidate_for_event(event) tracked_event.dependencies = dependencies - affected = self._run_post_apply_callbacks( - adapter_events, - notify_empty_events=True, - unique_paths=True, - ) + with self._suppressed(): + affected = self._run_post_apply_callbacks( + adapter_events, + notify_empty_events=True, + unique_paths=True, + ) return { "status": "refreshed", @@ -1209,6 +1214,10 @@ def _refresh_asset_dependency_suppressed( "pending": list(self.pending_asset_dependencies), } + def _suppressed(self): + """Keep stage edits made by callbacks from being published.""" + return self.emitter.suppressed() if self.emitter else nullcontext() + def _compute_stage_skip(self, events: list[dict]) -> set[int]: """Find arc events whose composed state already matches the stage. diff --git a/openusdconnect/emitter.py b/openusdconnect/emitter.py index 6e6162d..13d452f 100644 --- a/openusdconnect/emitter.py +++ b/openusdconnect/emitter.py @@ -1939,7 +1939,6 @@ def __init__( self._renamed_prims: list[tuple[str, str]] = [] # (old_path, new_path) self._suppress_depth: int = 0 self._suppressed_edit_target: Usd.EditTarget | None = None - self.listener = Tf.Notice.Register(Usd.Notice.ObjectsChanged, self._on_changed, stage) self._prim_cache: dict[str, dict] = {} # Unfiltered info-only attr names. Channels use this for read gating; # the gprim attr scan applies _attr_filter later. @@ -2003,6 +2002,8 @@ def __init__( # User-provided attr_filter wins; otherwise derive it from the active # channel set. self._attr_filter = attr_filter or _make_attr_filter(self._channels) + # Last, so no notice reaches a partially constructed emitter. + self.listener = Tf.Notice.Register(Usd.Notice.ObjectsChanged, self._on_changed, stage) def _local_prim_spec(self, prim_path: str): return _edit_target_prim_spec(self.stage, prim_path) @@ -2311,6 +2312,8 @@ def cleanup(self): self._sdf_spec_fields.clear() self._local_property_spec_fields.clear() self._local_prim_states.clear() + self._stage_metadata_dirty = False + self._stage_metadata_cache.clear() self._prepared_events = None self._suppress_depth = 0 self._suppressed_edit_target = None @@ -2328,6 +2331,7 @@ def rebind_stage(self, stage: Usd.Stage) -> None: raise RuntimeError("cannot rebind an emitter while a prepared batch is pending") self.cleanup() self.stage = stage + self._stage_metadata_cache = read_stage_metadata(stage) self.listener = Tf.Notice.Register(Usd.Notice.ObjectsChanged, self._on_changed, stage) def seed_prim_cache(self, stage: Usd.Stage, prim_path: str): @@ -2424,17 +2428,24 @@ def suppressed(self): """Return a reentrant notice-suppression context manager.""" return _SuppressScope(self) + def _pending_notice_state(self) -> tuple: + """Containers of noticed changes not yet built into events.""" + return ( + self.dirty, + self._pending_deactivations, + self._removed_local_definition_prims, + self._renamed_prims, + self._dirty_attrs, + self._sample_dirty_attrs, + self._notice_resynced_prims, + self._dirty_sdf_specs, + self._dirty_sdf_subtrees, + ) + def clear_all(self): """Discard pending prim notices, retaining diff caches and any prepared batch.""" - self.dirty.clear() - self._pending_deactivations.clear() - self._removed_local_definition_prims.clear() - self._renamed_prims.clear() - self._dirty_attrs.clear() - self._sample_dirty_attrs.clear() - self._notice_resynced_prims.clear() - self._dirty_sdf_specs.clear() - self._dirty_sdf_subtrees.clear() + for pending in self._pending_notice_state(): + pending.clear() self._clear_pending_context() def _clear_pending_context(self) -> None: @@ -4079,6 +4090,9 @@ def _build_events_for_dirty_current_target( self._notice_resynced_prims.discard(prim_path) continue events.extend(self._build_dirty_prim_events(prim_path, prim, eps_trs)) + # A resynced path that was never dirty (a removed local definition) + # only drove the subtree walk above. + self._notice_resynced_prims.clear() if self._full_sdf_spec_scan or self._dirty_sdf_specs or self._dirty_sdf_subtrees: events.extend( @@ -4172,6 +4186,16 @@ def prepare_coalesced_transform_events_for_send( merge_latest_transform_events(self._prepared_events, incoming) return self._prepared_events + @property + def has_local_changes(self) -> bool: + """Whether notices or a prepared batch still need evaluation or submission.""" + return bool( + self._prepared_events + or self._full_sdf_spec_scan + or self._stage_metadata_dirty + or any(self._pending_notice_state()) + ) + @property def prepared_event_count(self) -> int: """Number of events retained for a later transport attempt.""" diff --git a/openusdconnect/managed_client.py b/openusdconnect/managed_client.py index 84e698c..faa78af 100644 --- a/openusdconnect/managed_client.py +++ b/openusdconnect/managed_client.py @@ -1,49 +1,33 @@ """Single bidirectional client for server-owned collaboration layers. ``ManagedClient`` composes the emitter, sender, receiver, and dispatcher so a -USD-native application authors and observes one stage. It is the managed-mode -counterpart of ``SharedStageClient`` with the same lifecycle shape -(``start`` / ``connect`` / ``update`` / ``close``). +USD-native application authors and observes one stage. """ from __future__ import annotations -import time from collections.abc import Callable, Sequence from dataclasses import dataclass from pxr import Sdf, Usd +from ._client_base import EmitterClientBase from ._client_lifecycle import ( + DEFAULT_WAIT_TIMEOUT_S, + BacklogHold, deadline_after, - prepare_sender_token, - raise_if_rejected, remaining_time, - share_client_token, - stop_receiver, -) -from ._client_utils import ( - client_origin, - client_token_handlers, - require_app_name, - resolve_client_token, - validate_layered_source, ) +from ._client_utils import client_origin, require_app_name, validate_layered_source from .adapters import UsdStageAdapter from .client_id import make_stable_client_id -from .client_types import ClientPhase, ClientStatus, SyncUpdate -from .coalescing import TransformCoalescingWindow +from .client_observer import ClientObserver +from .client_types import SyncUpdate from .defaults import DEFAULT_HOST, DEFAULT_SYNC_PORT from .dispatcher import AssetDependencyRefreshResult, EventDispatcher -from .emitter import NoticeEmitter, PrimChannel +from .emitter import PrimChannel from .receiver import ReceiverThread -from .recovery import ( - RecoveryArtifact, - RecoveryError, - RecoveryIncident, - RejectionDisposition, - TransactionFailure, -) +from .recovery import RecoveryArtifact, RecoveryError from .sender import EventSender @@ -55,7 +39,7 @@ class ManagedRecoveryResult: preserved_authoring_layer: Sdf.Layer -class ManagedClient: +class ManagedClient(EmitterClientBase): """Bidirectional sync through one client-owned transient authoring layer. The stage edit target is moved to :attr:`authoring_layer` at construction @@ -76,195 +60,83 @@ def __init__( token: str | None = None, persist_token: bool = True, reconnect: bool = True, - on_imported: Callable[[list[str]], None] | None = None, - on_resync: Callable[[], None] | None = None, - on_applied: Callable[[list[str]], None] | None = None, - on_applied_events: Callable[[list[dict]], 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, - on_token_issued: Callable[[str], None] | None = None, + observer: ClientObserver | None = None, attr_filter: Callable[[str], bool] | None = None, 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): raise TypeError("ManagedClient requires a Usd.Stage") - adapter = UsdStageAdapter(stage) validate_layered_source(stage) - stable_client_id = client_id or make_stable_client_id(app_name) - connection_origin = origin or client_origin(app_name, "sync") - resolved_token = resolve_client_token(host, port, token, persist_token) - token_callback = client_token_handlers(host, port, persist_token, on_token_issued) - - def _on_token_issued(token: str) -> None: - share_client_token(token, self._sender, self._receiver, token_callback) - - self._stage = stage - self._host = host - self._port = port - self._persist_token = persist_token + super().__init__( + host=host, port=port, token=token, persist_token=persist_token, observer=observer, + ) + self._stage: Usd.Stage | None = stage self._app_name = app_name - self._authoring_layer = self._create_authoring_layer(stage, app_name) + self._authoring_layer: Sdf.Layer | None = None self._last_recovery_result: ManagedRecoveryResult | None = None - self._transform_coalescing = TransformCoalescingWindow(transform_coalesce_seconds) - - self._emitter = NoticeEmitter( + self._backlog = BacklogHold() + self._init_emitter( stage, attr_filter=attr_filter, replicated_api_schemas=replicated_api_schemas, extra_channels=extra_channels, + transform_coalesce_seconds=transform_coalesce_seconds, ) + identity = { + "client_id": client_id or make_stable_client_id(app_name), + "origin": origin or client_origin(app_name, "sync"), + } + credential = self._credential.endpoint_kwargs() self._sender = EventSender( - host, - port, - client_id=stable_client_id, - origin=connection_origin, - department=department, - token=resolved_token, - on_token_issued=_on_token_issued, + host, port, department=department, background_send=background_send, + **identity, **credential, ) self._receiver = ReceiverThread( - host=host, - port=port, - sync_from=1, - reconnect=reconnect, - client_id=stable_client_id, - origin=connection_origin, - token=resolved_token, - on_token_issued=_on_token_issued, - on_stage_metadata=on_stage_metadata, - on_playback_state=on_playback_state, - on_playback_claimed=on_playback_claimed, - on_playback_rejected=on_playback_rejected, - layered_replay=True, + host=host, port=port, sync_from=1, reconnect=reconnect, layered_replay=True, + **identity, **credential, **self._hooks.receiver_callbacks(), ) self._dispatcher = EventDispatcher( receiver=self._receiver, - adapter=adapter, + adapter=UsdStageAdapter(stage), emitter=self._emitter, - on_imported=on_imported, - on_resync=on_resync, - on_applied=on_applied, - on_applied_events=on_applied_events, + on_resync=self._hooks.on_resync, ) - self._started = False - self._closed = False + self._dispatcher.on_applied_events = self._hooks.applied_events_for(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) @property def stage(self) -> Usd.Stage | None: """Application-owned stage, or ``None`` while parked.""" return self._stage - @property - def client_id(self) -> str: - """Stable identity used by both connection roles.""" - return self._receiver.client_id - - @property - def status(self) -> ClientStatus: - """Current transport, replay, durability, and recovery state.""" - failure = self._sender.transaction_failure - reason = str(failure) if failure is not None else ( - self._sender.rejection_reason or self._receiver.rejection_reason - ) - if self._closed: - phase = ClientPhase.CLOSED - elif failure is not None: - phase = ClientPhase.RECOVERY_REQUIRED - elif self.auth_rejected or self.connection_rejected: - phase = ClientPhase.REJECTED - elif self._receiver.connected and not self._receiver.synchronized: - phase = ClientPhase.REPLAYING - elif self.connected and self.synchronized: - phase = ClientPhase.READY - elif self._started: - phase = ClientPhase.CONNECTING - else: - phase = ClientPhase.OFFLINE - return ClientStatus( - phase=phase, - connected=self.connected, - synchronized=self.synchronized, - receiver_connected=self._receiver.connected, - sender_connected=self._sender.connected, - prepared_events=self.prepared_event_count, - pending_events=self.pending_event_count, - acknowledged_events_total=self._sender.acknowledged_event_count, - failure=failure, - recovery=self._sender.recovery_incident, - reason=reason, - ) - - @property - def sender(self): - """The underlying :class:`EventSender`. Read access is safe; mutating - configuration on this object is at your own risk.""" - return self._sender - - @property - def receiver(self): - """The underlying :class:`ReceiverThread`.""" - return self._receiver - - @property - def dispatcher(self): - """The underlying :class:`EventDispatcher`.""" - return self._dispatcher - - @property - def emitter(self): - """The underlying :class:`NoticeEmitter`.""" - return self._emitter - @property def authoring_layer(self) -> Sdf.Layer | None: """Client-owned transient layer for all local managed-mode edits.""" return self._authoring_layer @property - def connected(self) -> bool: - return not self._closed and self._receiver.connected and self._sender.connected - - @property - def synchronized(self) -> bool: - """Whether the local stage applied replay through the server watermark.""" - return ( - not self._closed and self._receiver.synchronized and not self._sender.recovery_required - ) - - @property - def recovery_required(self) -> bool: - """Whether deterministic rejection requires local-state reconciliation.""" - return self._sender.recovery_required - - @property - def pending_event_count(self) -> int: - """Number of submitted events not yet durably acknowledged.""" - return self._sender.pending_event_count - - @property - def transaction_error(self) -> str: - """Terminal producer rejection, or an empty string.""" - return self._sender.transaction_error + def receiver(self) -> ReceiverThread: + """The underlying :class:`ReceiverThread`; a diagnostic handle.""" + return self._receiver @property - def transaction_failure(self) -> TransactionFailure | None: - """Structured rejection including its recovery disposition, if any.""" - return self._sender.transaction_failure + def dispatcher(self) -> EventDispatcher: + """The underlying :class:`EventDispatcher`; a diagnostic handle.""" + return self._dispatcher @property - def recovery_incident(self) -> RecoveryIncident | None: - """Structured recovery summary for polling and host UI.""" - return self._sender.recovery_incident + def last_seq(self) -> int: + return self._dispatcher.last_seq @property - def recovery_artifact(self) -> RecoveryArtifact | None: - """Exact quarantined transactions for integration-owned recovery.""" - return self._sender.recovery_artifact + def pending_asset_dependencies(self) -> tuple[str, ...]: + return self._dispatcher.pending_asset_dependencies @property def last_recovery_result(self) -> ManagedRecoveryResult | None: @@ -275,11 +147,85 @@ def dismiss_recovery_result(self) -> None: """Release references held for the most recently resolved incident.""" self._last_recovery_result = None + def update(self, *, max_messages: int | None = None) -> SyncUpdate: + """Freeze local edits, apply the commit stream, then publish them. + + ``max_messages`` bounds one call's receive work; local edits are held + until the backlog queued before them has been applied. + """ + if not self._begin_update() or self._stage is None: + return self._progress() + + # A queued authoritative record may touch the same field as a newer + # local opinion. Freeze the local delta before dispatcher invalidation + # advances emitter baselines, then send that exact batch after the + # authoritative prefix has applied. SharedStageClient follows the same + # prepare/apply/restore ordering at the Sdf-layer level. + self._validate_authoring_target() + 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) + 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: + sent = self._send(outgoing) + return self._progress(received, sent) + + def rebind_stage(self, stage: Usd.Stage | None, *, discard_unsent: bool = False) -> None: + """Move sending and receiving to a new stage with a fresh authoring layer. + + ``None`` parks the client. Refuses while work is unacknowledged, or + unsent unless ``discard_unsent=True`` drops it. + """ + self._require_open() + adapter = None + if stage is not None: + adapter = UsdStageAdapter(stage) + validate_layered_source(stage) + if self._sender.pending_event_count: + raise RuntimeError("cannot rebind with submitted work pending; call submit_and_wait()") + if self._has_unsent_changes() and not discard_unsent: + raise RuntimeError( + "cannot rebind with unsent changes; call submit_and_wait() " + "or pass discard_unsent=True" + ) + if discard_unsent: + self._emitter.discard_prepared_events() + self._transform_coalescing.mark_submitted() + if stage is None: + self._dispatcher.unbind_stage() + self._emitter.cleanup() + self._stage = None + self._authoring_layer = None + return + self._dispatcher.adapter = adapter + self._dispatcher.bind_layered_stage(stage) + self._authoring_layer = self._create_authoring_layer(stage, self._app_name) + self._emitter.rebind_stage(stage) + self._stage = stage + + def refresh_asset_dependency( + self, + asset_path: str | None = None, + ) -> AssetDependencyRefreshResult: + """Retry dependencies under the stage's current resolver context.""" + self._require_open() + return self._dispatch(self._dispatcher.refresh_asset_dependency, asset_path) + def recover_use_server( self, *, session_id: str | None = None, - timeout: float | None = 10.0, + timeout: float | None = DEFAULT_WAIT_TIMEOUT_S, ) -> ManagedRecoveryResult: """Discard local optimistic opinions and start a fresh producer session. @@ -287,30 +233,20 @@ def recover_use_server( layer is cleared. The producer reconnect is attempted within the same timeout budget; if it cannot complete, the normal update loop retries. """ - if self._closed: - raise RuntimeError("ManagedClient is closed") - if not self._started: - raise RuntimeError("ManagedClient has not been started") + self._require_started() stage = self._stage authoring = self._authoring_layer if stage is None or authoring is None: - raise RecoveryError( - "stage_unavailable", - "ManagedClient has no bound stage to recover", - ) + raise RecoveryError("stage_unavailable", "ManagedClient has no bound stage to recover") if not self._sender.recovery_required: - raise RecoveryError( - "no_incident", - "there is no recovery incident to resolve", - ) + raise RecoveryError("no_incident", "there is no recovery incident to resolve") if stage.GetEditTarget().GetLayer() is not authoring: raise RecoveryError( - "edit_target_changed", - "the active edit target changed during recovery", + "edit_target_changed", "the active edit target changed during recovery", ) deadline = deadline_after(timeout) - self._refresh_recovery_checkpoint(timeout) + self._replay_to_fresh_checkpoint(timeout) preserved = Sdf.Layer.CreateAnonymous("openusdconnect-recovery-authoring") preserved.TransferContent(authoring) @@ -333,286 +269,48 @@ def recover_use_server( finally: self._emitter.rebind_stage(stage) self._transform_coalescing.mark_submitted() - remaining = remaining_time(deadline) - self._resume_sender_after_recovery(remaining) + self._resume_sender_after_recovery(remaining_time(deadline)) return result - def _resume_sender_after_recovery(self, timeout: float | None) -> None: - """Best-effort producer reconnect after state recovery has committed.""" - try: - self._connect_sender(timeout=timeout) - except (PermissionError, ConnectionError): - # Recovery already completed and must not look rolled back. Status - # exposes rejection/offline state; update() retries ordinary loss. - pass - - def _refresh_recovery_checkpoint(self, timeout: float | None) -> None: - """Replay through a new server head before resolving optimistic state.""" - deadline = deadline_after(timeout) - reconnect = self._receiver.reconnect - self._receiver.reconnect = True - try: - self._receiver.request_replay_from(self._dispatcher.last_seq + 1) - while True: - self._dispatcher.drain_and_apply() - if self._receiver.synchronized: - return - if deadline is not None and time.monotonic() >= deadline: - raise TimeoutError("authoritative recovery replay timed out") - time.sleep(0.01) - finally: - self._receiver.reconnect = reconnect - - @property - def recovery_disposition(self) -> RejectionDisposition | None: - """Recovery policy category for the current rejection, if any.""" - return self._sender.recovery_disposition - - def repair_and_resume(self, events: list[dict]) -> int: - """Replace a recoverable transaction and resume its ordered outbox. - - The application must first reconcile its stage with authoritative - state and rebuild *events* for that state. The repaired transaction is - assigned the original rejected ID; later quarantined transactions keep - their existing IDs and replay after it. - """ - if self._closed: - raise RuntimeError("ManagedClient is closed") - self._require_recoverable_failure() - txn_id = self._sender.repair_rejected_transaction(events) - if not self._connect_sender(): - raise ConnectionError( - f"transaction {txn_id} repaired but reconnect to " - f"{self._host}:{self._port} failed; it remains queued" - ) - return txn_id - - def _require_recoverable_failure(self) -> TransactionFailure: - failure = self._sender.transaction_failure - if failure is None: - raise RecoveryError("no_incident", "there is no recovery incident to resolve") - if failure.disposition is not RejectionDisposition.RECOVERABLE_CONFLICT: - raise RecoveryError( - "wrong_recovery_kind", - f"{failure.code_name} is {failure.disposition.value}, not recoverable", - ) - return failure - - def flush(self, timeout: float | None = None) -> bool: - """Submit any coalesced transform, then wait for durable acknowledgement.""" - if self._closed: - raise RuntimeError("ManagedClient is closed") - deadline = deadline_after(timeout) - if self._transform_coalescing.buffering: - try: - if not self._connect_sender(timeout=remaining_time(deadline)): - return False - except (PermissionError, ConnectionError): - return False - if not self.synchronized: - return False - events = self._transform_coalescing.force(self._emitter) - if events and not self._send(events): - return False - return self._sender.flush(remaining_time(deadline)) - - @property - def auth_rejected(self) -> bool: - return self._receiver.auth_rejected or self._sender.auth_rejected - - @property - def connection_rejected(self) -> bool: - return self._receiver.hello_rejected or self._sender.hello_rejected - - @property - def last_seq(self) -> int: - return self._dispatcher.last_seq - - @property - def stage_metadata(self) -> dict: - return dict(self._receiver.stage_metadata) - - @property - def prepared_event_count(self) -> int: - """Number of events retained after an unsuccessful transport write.""" - return self._emitter.prepared_event_count - - @property - def pending_asset_dependencies(self) -> tuple[str, ...]: - return self._dispatcher.pending_asset_dependencies - - def start(self) -> ManagedClient: - """Start the background socket reader and return this client.""" - if self._closed: - raise RuntimeError("ManagedClient is closed") - if not self._started: - if self._receiver.token is None and self._persist_token: - self._receiver.token = resolve_client_token(self._host, self._port, None, True) - self._receiver.start() - self._started = True - return self - - def connect(self, timeout: float | None = None) -> bool: - """Start and complete both handshakes within ``timeout``. - - Queued replay still requires :meth:`update` on the stage-owning thread. - """ - self.start() - deadline = deadline_after(timeout) - connected = self._receiver.wait_connected(timeout) - if connected: - self._require_layered_replay() - remaining = remaining_time(deadline) - return self._connect_sender(timeout=remaining) - elif self._receiver.auth_rejected: - raise PermissionError("authentication rejected") - elif self._receiver.hello_rejected: - raise ConnectionError(self._receiver.rejection_reason or "connection rejected") - return False - - def _require_layered_replay(self) -> None: - if self._receiver.connected and not self._receiver.layered_replay_active: - self.close() - raise RuntimeError("server did not negotiate required layered replay") - - def _connect_sender(self, timeout: float | None = None) -> bool: - if self._sender.connected: - return True - self._prepare_sender_token() - if not self._sender.connect(timeout=timeout): - raise_if_rejected(self._sender, "sender") - return False - return True - - def _prepare_sender_token(self) -> None: - prepare_sender_token( - self._sender, self._receiver, - host=self._host, port=self._port, persist_token=self._persist_token, + def _is_synchronized(self) -> bool: + return ( + self._stage is not None + and self._receiver.synchronized + and not self._sender.recovery_required ) - def _send(self, events: list[dict]) -> int: - if not events: - return 0 - if self._sender.send_events(events): - self._emitter.mark_prepared_events_sent(events) - self._transform_coalescing.mark_submitted() - return len(events) - return 0 + def _is_parked(self) -> bool: + return self._stage is None - def _prepare_outgoing_events(self) -> list[dict]: - return self._transform_coalescing.prepare(self._emitter) + def _has_unsent_changes(self) -> bool: + return self._stage is not None and super()._has_unsent_changes() - def claim_playback(self, time: float | None = None) -> bool: - """Request the shared-playback leader role.""" - if self._closed: - raise RuntimeError("ManagedClient is closed") - return self._sender.claim_playback(time=time) + def _role_status(self) -> dict: + return { + "edit_target_is_published": ( + self._stage is not None + and self._stage.GetEditTarget().GetLayer() is self._authoring_layer + ), + } - def send_playback_control( - self, - action: str, - *, - time: float | None = None, - rate: float | None = None, - ) -> bool: - """Drive the shared playhead (leader only).""" - if self._closed: - raise RuntimeError("ManagedClient is closed") - return self._sender.send_playback_control(action, time=time, rate=rate) + def _apply_queued(self, max_messages: int | None = None) -> int: + return self._dispatch(self._dispatcher.drain_and_apply, max_messages=max_messages) - def update(self) -> SyncUpdate: - """Freeze local edits, apply the commit stream, then publish them.""" - if self._closed: - raise RuntimeError("ManagedClient is closed") - if not self._started: - raise RuntimeError("ManagedClient has not been started") + def _can_capture_edit_target(self) -> bool: if self._stage is None: - return SyncUpdate( - applied_events=0, - submitted_events=0, - acknowledged_events_delta=self._sender.drain_acknowledged_event_count(), - pending_events=self._sender.pending_event_count, - recovery=self._sender.recovery_incident, - ) - self._require_layered_replay() - - # A queued authoritative record may touch the same field as a newer - # local opinion. Freeze the local delta before dispatcher invalidation - # advances emitter baselines, then send that exact batch after the - # authoritative prefix has applied. SharedStageClient follows the same - # prepare/apply/restore ordering at the Sdf-layer level. + return False self._validate_authoring_target() - outgoing = self._prepare_outgoing_events() - received = self._dispatcher.drain_and_apply() - - sent = 0 - if self._receiver.connected and not self._sender.connected: - self._prepare_sender_token() - self._sender.request_connect() - if self._sender.connected and self.synchronized: - sent = self._send(outgoing) - - return SyncUpdate( - applied_events=received, - submitted_events=sent, - acknowledged_events_delta=self._sender.drain_acknowledged_event_count(), - pending_events=self._sender.pending_event_count, - recovery=self._sender.recovery_incident, - ) - - def publish_current_edit_target(self) -> int: - """Publish all opinions currently authored in the active edit target. + return True - An earlier retained batch must be retried with :meth:`update` first. - This keeps one call from ambiguously mixing two transport transactions. - """ - if self._closed: - raise RuntimeError("ManagedClient is closed") - if self._stage is None: - return 0 - self._validate_authoring_target() - if not self._sender.connected: - try: - self._connect_sender() - except (PermissionError, ConnectionError): - return 0 - if self._emitter.prepared_event_count: + def _validate_authoring_target(self) -> None: + if self._stage.GetEditTarget().GetLayer() is not self._authoring_layer: raise RuntimeError( - "an earlier publisher batch is still prepared; call update() " - "before publishing the current edit target" + "ManagedClient publishes only from client.authoring_layer; " + "restore that edit target before update()" ) - return self._send(self._emitter.prepare_snapshot_events_for_send()) - - def rebind_stage(self, stage: Usd.Stage | None) -> None: - """Move sending and receiving to a new stage and select a fresh authoring layer. - - Pass ``None`` to park: the receiver stays connected and the queue - continues to fill, but ``update()`` returns zero until a new stage - is bound. - """ - if self._closed: - raise RuntimeError("ManagedClient is closed") - if self._emitter.prepared_event_count: - raise RuntimeError("cannot rebind while a prepared publisher batch is pending") - if stage is None: - self._dispatcher.unbind_stage() - self._emitter.cleanup() - self._stage = None - self._authoring_layer = None - return - adapter = UsdStageAdapter(stage) - validate_layered_source(stage) - self._dispatcher.adapter = adapter - self._dispatcher.bind_layered_stage(stage) - self._authoring_layer = self._create_authoring_layer(stage, self._app_name) - self._emitter.rebind_stage(stage) - self._stage = stage @staticmethod - def _create_authoring_layer( - stage: Usd.Stage, - label: str, - ) -> Sdf.Layer: + def _create_authoring_layer(stage: Usd.Stage, label: str) -> Sdf.Layer: """Create and select the one transient layer owned by this client.""" session = stage.GetSessionLayer() authoring = Sdf.Layer.CreateAnonymous(f"openusdconnect-{label}-authoring") @@ -621,42 +319,9 @@ def _create_authoring_layer( stage.SetEditTarget(Usd.EditTarget(authoring)) return authoring - def _validate_authoring_target(self) -> None: - stage = self._stage - if stage is None: - return - layer = stage.GetEditTarget().GetLayer() - if layer is not self._authoring_layer: - raise RuntimeError( - "ManagedClient publishes only from client.authoring_layer; " - "restore that edit target before update()" - ) - - def refresh_asset_dependency( - self, - asset_path: str | None = None, - ) -> AssetDependencyRefreshResult: - """Retry dependencies under the stage's current resolver context.""" - if self._closed: - raise RuntimeError("ManagedClient is closed") - return self._dispatcher.refresh_asset_dependency(asset_path) - - def close(self) -> None: - """Stop networking and release receiver-owned collaboration layers.""" - if self._closed: - return - self._sender.disconnect() - stop_receiver(self._receiver) + def _release(self) -> None: self._dispatcher.close() - self._emitter.cleanup() - self._closed = True - - def __enter__(self) -> ManagedClient: - return self.start() - - def __exit__(self, exc_type, exc, traceback) -> bool: - self.close() - return False + super()._release() -__all__ = ["ManagedClient"] +__all__ = ["ManagedClient", "ManagedRecoveryResult"] diff --git a/openusdconnect/receiver.py b/openusdconnect/receiver.py index 3201277..49303b0 100644 --- a/openusdconnect/receiver.py +++ b/openusdconnect/receiver.py @@ -71,6 +71,7 @@ def __init__( origin: str | None = None, department: str | None = None, token: str | None = None, + token_provider: Callable[[], str | 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, @@ -94,6 +95,7 @@ def __init__( self.layered_replay_active = False self.layer_mode = LayerMode(layer_mode) self.layer_mode_active = LayerMode.MANAGED + self._token_provider = token_provider self._on_token_issued = on_token_issued self._on_stage_metadata = on_stage_metadata self._on_playback_state = on_playback_state @@ -150,6 +152,11 @@ def synchronized(self) -> bool: def last_seq(self) -> int: return self._inbox.last_sequence + @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() + @property def queued_message_count(self) -> int: """Number of received messages waiting for the owning thread to drain.""" @@ -262,6 +269,8 @@ def _connect_and_recv(self) -> None: prefix_identity = self._received_replay_identity self._synchronized_event.clear() + if self._token_provider is not None: + self.token = self._token_provider() hello = make_hello( "receiver", sync_from=sync_from, diff --git a/openusdconnect/sender.py b/openusdconnect/sender.py index 5dd01c3..7e46bc5 100644 --- a/openusdconnect/sender.py +++ b/openusdconnect/sender.py @@ -62,6 +62,8 @@ class EventSender: 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. """ def __init__( @@ -77,9 +79,11 @@ def __init__( handshake_timeout: float = _HANDSHAKE_TIMEOUT_S, on_token_issued: Callable[[str], None] | None = None, on_stage_metadata: Callable[[dict], None] | None = None, + token_provider: Callable[[], str | None] | None = None, layer_mode: LayerMode | str = LayerMode.MANAGED, session_id: str | None = None, max_pending_transactions: int = _MAX_PENDING_TRANSACTIONS, + background_send: bool = False, ): if role != "emitter": raise ValueError("EventSender role must be 'emitter'") @@ -99,8 +103,10 @@ def __init__( 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._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 @@ -108,6 +114,7 @@ def __init__( 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 @@ -118,6 +125,7 @@ def __init__( 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 @@ -318,12 +326,13 @@ def _connect_locked(self, deadline: float, epoch: int) -> bool: connect_timeout = deadline - time.monotonic() if connect_timeout <= 0.0: return False + if self._token_provider is not None: + self.token = self._token_provider() connection = self._session.begin_connection() if connection is None: return False generation = connection.generation - self._socket_generation = generation self.auth_rejected = False self.hello_rejected = False @@ -370,21 +379,27 @@ def _connect_locked(self, deadline: float, epoch: int) -> bool: 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 - 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 + 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: - LOG.exception("EventSender: handshake failed") + 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): @@ -406,9 +421,24 @@ def _connect_locked(self, deadline: float, epoch: int) -> bool: name=f"openusdconnect-ack-{self.client_id}", daemon=True, ) - with self._condition: - self._reader_thread = reader - reader.start() + 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, @@ -498,6 +528,10 @@ def _notify_handshake_callback(callback, value, *, name: str) -> None: 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 @@ -523,11 +557,11 @@ def send_events(self, events: list, *, layer_key: str = "") -> bool: 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: - # Transaction identity and wire order are one operation. Without - # this outer lock, concurrent callers can allocate IDs 1 then 2 - # but acquire the socket lock and transmit them as 2 then 1. - with self._send_lock: + # 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 @@ -547,6 +581,9 @@ def send_events(self, events: list, *, layer_key: str = "") -> bool: ) 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) @@ -558,6 +595,37 @@ def send_events(self, events: list, *, layer_key: str = "") -> bool: self._close(expected=sock) 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. @@ -731,6 +799,8 @@ 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) @@ -837,11 +907,12 @@ def _highwater_failure_reason(self, result, transaction_id: int) -> str: def _close(self, *, expected: socket.socket | None = None) -> None: with self._condition: sock = self.sock - if expected is not None and sock is not expected: + if sock is None or (expected is not None and sock is not expected): return self.sock = None - generation = self._socket_generation - self._session.disconnect(generation) + # 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) diff --git a/openusdconnect/shared_stage_client.py b/openusdconnect/shared_stage_client.py index 00041e7..e9835fd 100644 --- a/openusdconnect/shared_stage_client.py +++ b/openusdconnect/shared_stage_client.py @@ -3,29 +3,23 @@ from __future__ import annotations import logging -import time -from collections.abc import Callable, Iterable +from collections.abc import Iterable from dataclasses import dataclass from pathlib import Path from pxr import Sdf, Usd +from ._client_base import PublishingClientBase from ._client_lifecycle import ( + DEFAULT_WAIT_TIMEOUT_S, + BacklogHold, deadline_after, - prepare_sender_token, - raise_if_rejected, remaining_time, - share_client_token, - stop_receiver, -) -from ._client_utils import ( - client_origin, - client_token_handlers, - require_app_name, - resolve_client_token, ) +from ._client_utils import client_origin, require_app_name from .client_id import make_stable_client_id -from .client_types import ClientPhase, ClientStatus, SyncUpdate +from .client_observer import ClientObserver +from .client_types import SyncUpdate from .codec import ReceivedEvent, decode_messages from .defaults import DEFAULT_HOST, DEFAULT_SYNC_PORT from .event_apply import apply_events, atomic_apply, atomic_apply_prim_paths @@ -37,15 +31,10 @@ LayerMode, ) from .receiver import ReceiverThread -from .recovery import ( - RecoveryArtifact, - RecoveryError, - RejectionDisposition, -) +from .recovery import RecoveryArtifact, RecoveryError from .sdf_layer_tracker import SdfLayerChangeTracker from .sender import EventSender from .shared_layer_graph import SharedLayerGraph -from .token_client import load_token LOG = logging.getLogger(__name__) @@ -119,7 +108,7 @@ def rejected_snapshots(self) -> tuple[Sdf.Layer, ...]: ) -class SharedStageClient: +class SharedStageClient(PublishingClientBase): """Synchronize authored opinions in a stage's root-layer graph. Every process opens its own equivalent stage and resolver context. Opaque @@ -139,11 +128,8 @@ def __init__( token: str | None = None, persist_token: bool = True, reconnect: bool = True, - 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, - on_token_issued: Callable[[str], None] | None = None, + background_send: bool = False, + observer: ClientObserver | None = None, delegate_bridge_path: str | Path | None = None, ): if not isinstance(stage, Usd.Stage): @@ -151,18 +137,10 @@ def __init__( app_name = require_app_name(app_name) if Sdf.Layer.IsAnonymousLayerIdentifier(stage.GetRootLayer().identifier): raise ValueError("shared-stage synchronization requires a portable root layer") - stable_client_id = client_id or make_stable_client_id(app_name) - connection_origin = origin or client_origin(app_name, "shared") - resolved_token = resolve_client_token(host, port, token, persist_token) - token_callback = client_token_handlers(host, port, persist_token, on_token_issued) - - def _on_token_issued(token: str) -> None: - share_client_token(token, self._sender, self._receiver, token_callback) - + super().__init__( + host=host, port=port, token=token, persist_token=persist_token, observer=observer, + ) self._stage = stage - self._host = host - self._port = port - self._persist_token = persist_token self._graph = SharedLayerGraph(stage) self._graph._validate_local_graph() @@ -170,36 +148,25 @@ def _on_token_issued(token: str) -> None: self._delegate_bridge_path = delegate_bridge_path or _find_bridge() self._tracker = self._make_tracker(stage, self._graph) + identity = { + "client_id": client_id or make_stable_client_id(app_name), + "origin": origin or client_origin(app_name, "shared"), + } + credential = self._credential.endpoint_kwargs() self._receiver = ReceiverThread( - host=host, - port=port, - sync_from=1, - reconnect=reconnect, - client_id=stable_client_id, - origin=connection_origin, - token=resolved_token, - on_token_issued=_on_token_issued, - on_stage_metadata=on_stage_metadata, - on_playback_state=on_playback_state, - on_playback_claimed=on_playback_claimed, - on_playback_rejected=on_playback_rejected, - layered_replay=False, - layer_mode=LayerMode.SHARED_STAGE, + host=host, port=port, sync_from=1, reconnect=reconnect, + layered_replay=False, layer_mode=LayerMode.SHARED_STAGE, + **identity, **credential, **self._hooks.receiver_callbacks(), ) self._sender = EventSender( - host, - port, - client_id=stable_client_id, - origin=connection_origin, - token=resolved_token, - on_token_issued=_on_token_issued, - layer_mode=LayerMode.SHARED_STAGE, + host, port, layer_mode=LayerMode.SHARED_STAGE, background_send=background_send, + **identity, **credential, ) self._last_seq = 0 - self._pending_records: list[ReceivedEvent] = [] + self._backlog = BacklogHold() + self._set_deferred([]) self._last_recovery_assessment: SharedRecoveryAssessment | None = None - self._started = False - self._closed = False + self._recovery_rebind_artifact: RecoveryArtifact | None = None def _make_tracker(self, stage: Usd.Stage, graph: SharedLayerGraph): bridge_path = self._delegate_bridge_path @@ -218,74 +185,22 @@ def stage(self) -> Usd.Stage: return self._stage @property - def client_id(self) -> str: - """Stable identity used by both connection roles.""" - return self._receiver.client_id - - @property - def status(self) -> ClientStatus: - """Current transport, replay, durability, and recovery state.""" - failure = self._sender.transaction_failure - reason = str(failure) if failure is not None else ( - self._sender.rejection_reason or self._receiver.rejection_reason - ) - if self._closed: - phase = ClientPhase.CLOSED - elif failure is not None: - phase = ClientPhase.RECOVERY_REQUIRED - elif ( - self._receiver.auth_rejected - or self._sender.auth_rejected - or self._receiver.hello_rejected - or self._sender.hello_rejected - ): - phase = ClientPhase.REJECTED - elif self._receiver.connected and not self._receiver.synchronized: - phase = ClientPhase.REPLAYING - elif self.connected and self.synchronized: - phase = ClientPhase.READY - elif self._started: - phase = ClientPhase.CONNECTING - else: - phase = ClientPhase.OFFLINE - return ClientStatus( - phase=phase, - connected=self.connected, - synchronized=self.synchronized, - receiver_connected=self._receiver.connected, - sender_connected=self._sender.connected, - prepared_events=self.prepared_event_count, - pending_events=self.pending_event_count, - acknowledged_events_total=self._sender.acknowledged_event_count, - failure=failure, - recovery=self._sender.recovery_incident, - reason=reason, - ) + def last_seq(self) -> int: + return self._last_seq - @property - def connected(self) -> bool: - return not self._closed and self._receiver.connected and self._sender.connected + def is_layer_reachable(self, layer: Sdf.Layer) -> bool: + """Return whether *layer* is reachable in the synchronized root graph.""" + layer_key = self._graph.key_for(layer) + return bool(layer_key and layer_key in self._graph.reachable_layer_keys()) @property - def synchronized(self) -> bool: - """Whether the local layer graph applied the server replay watermark.""" + def _recovery_stage_pending(self) -> bool: + """Whether a replacement stage is bound but recovery is incomplete.""" return ( - not self._closed - and self._graph.ready - and self._receiver.synchronized - and not self._sender.recovery_required + self._recovery_rebind_artifact is not None + and self._recovery_rebind_artifact is self._sender.recovery_artifact ) - @property - def pending_event_count(self) -> int: - """Number of submitted events not yet durably acknowledged.""" - return self._sender.pending_event_count - - @property - def recovery_artifact(self) -> RecoveryArtifact | None: - """Exact quarantined transactions for integration-owned recovery.""" - return self._sender.recovery_artifact - def repair_and_resume(self, events: list[dict], *, layer: Sdf.Layer) -> int: """Replace a recoverable layer transaction and resume its outbox. @@ -293,8 +208,7 @@ def repair_and_resume(self, events: list[dict], *, layer: Sdf.Layer) -> int: rebuild *events* against *layer*. The layer must still be mapped by the current graph; no semantic merge or layer redirection is inferred. """ - if self._closed: - raise RuntimeError("SharedStageClient is closed") + self._require_open() layer_key = self._graph.key_for(layer) if not layer_key or layer_key not in self._graph.reachable_layer_keys(): raise RecoveryError( @@ -304,17 +218,14 @@ def repair_and_resume(self, events: list[dict], *, layer: Sdf.Layer) -> int: self._require_recoverable_artifact() txn_id = self._sender.repair_rejected_transaction(events, layer_key=layer_key) self._last_recovery_assessment = None - if not self._connect_sender(): - raise ConnectionError( - f"transaction {txn_id} repaired but reconnect to " - f"{self._host}:{self._port} failed; it remains queued" - ) + self._recovery_rebind_artifact = None + self._reconnect_repaired(txn_id) return txn_id def refresh_recovery_assessment( self, *, - timeout: float | None = 10.0, + timeout: float | None = DEFAULT_WAIT_TIMEOUT_S, ) -> SharedRecoveryAssessment: """Replay to a fresh checkpoint and classify every quarantined layer.""" artifact = self._require_recoverable_artifact() @@ -342,7 +253,7 @@ def refresh_recovery_assessment( rejected_snapshot.TransferContent(layer) captured.append((layer_key, layer, rejected_snapshot)) - self._refresh_recovery_checkpoint(timeout) + self._replay_to_fresh_checkpoint(timeout) assessment = self._build_recovery_assessment(artifact, captured) self._last_recovery_assessment = assessment return assessment @@ -352,7 +263,7 @@ def recover_use_server( *, clean_stage: Usd.Stage, session_id: str | None = None, - timeout: float | None = 10.0, + timeout: float | None = DEFAULT_WAIT_TIMEOUT_S, ) -> SharedRecoveryAssessment: """Select server state by replaying onto a clean equivalent stage. @@ -363,14 +274,55 @@ def recover_use_server( preserved in the returned assessment before the new stage is replayed. Producer reconnect is attempted within the same timeout budget; if it cannot complete, the normal update loop retries. + + If replay fails after the replacement is bound, it stays bound + (``recovery_stage_pending``): keep the host on ``client.stage`` and call + :meth:`resume_recovery`. """ - self._validate_clean_recovery_stage(clean_stage) deadline = deadline_after(timeout) + self._validate_clean_recovery_stage(clean_stage) assessment = self.refresh_recovery_assessment(timeout=timeout) self._validate_clean_recovery_stage(clean_stage, assessment=assessment) self._rebind_stage_for_recovery(clean_stage) + self._recovery_rebind_artifact = assessment.recovery_artifact + return self._finish_use_server(assessment, session_id=session_id, deadline=deadline) + + def resume_recovery( + self, + *, + session_id: str | None = None, + timeout: float | None = DEFAULT_WAIT_TIMEOUT_S, + ) -> SharedRecoveryAssessment: + """Continue a Use Server recovery whose replacement replay did not finish. + + Rejected snapshots captured by the original attempt are preserved. + """ + self._require_recoverable_artifact() + if not self._recovery_stage_pending: + raise RecoveryError( + "no_pending_recovery_stage", + "no replacement stage is waiting for recovery to complete", + ) + if self._tracker.has_local_changes: + raise RecoveryError( + "local_changes_pending", + "cannot resume replacement-stage recovery while unsent edits remain", + ) + return self._finish_use_server( + self._last_recovery_assessment, + session_id=session_id, + deadline=deadline_after(timeout), + ) + + def _finish_use_server( + self, + assessment: SharedRecoveryAssessment, + *, + session_id: str | None, + deadline: float | None, + ) -> SharedRecoveryAssessment: remaining = remaining_time(deadline) - self._refresh_recovery_checkpoint(remaining) + self._replay_to_fresh_checkpoint(remaining) assessment = self._build_recovery_assessment( assessment.recovery_artifact, ( @@ -397,9 +349,14 @@ def _validate_clean_recovery_stage( if not isinstance(clean_stage, Usd.Stage): raise TypeError("SharedStageClient requires a Usd.Stage") if clean_stage is self._stage: + hint = ( + "; call resume_recovery() to continue the pending replacement" + if self._recovery_stage_pending + else "" + ) raise RecoveryError( "invalid_clean_stage", - "Use Server recovery requires a different clean stage", + f"Use Server recovery requires a different clean stage{hint}", ) if Sdf.Layer.IsAnonymousLayerIdentifier(clean_stage.GetRootLayer().identifier): raise RecoveryError( @@ -456,24 +413,9 @@ def complete_recovery( return self._complete_recovery(assessment, session_id=session_id) def _require_recoverable_artifact(self) -> RecoveryArtifact: - if self._closed: - raise RuntimeError("SharedStageClient is closed") - if not self._started: - raise RuntimeError("SharedStageClient has not been started") - artifact = self._sender.recovery_artifact - failure = self._sender.transaction_failure - if artifact is None or failure is None: - raise RecoveryError( - "no_incident", - "there is no recovery incident to resolve", - ) - if failure.disposition is not RejectionDisposition.RECOVERABLE_CONFLICT: - raise RecoveryError( - "wrong_recovery_kind", - f"{failure.code_name} is {failure.disposition.value}, not a " - "recoverable shared-stage conflict", - ) - return artifact + self._require_started() + self._require_recoverable_failure() + return self._sender.recovery_artifact def _validate_recovery_assessment( self, @@ -531,146 +473,45 @@ def _complete_recovery( self._validate_recovery_assessment(assessment) self._sender.abandon_rejected_session(session_id=session_id) self._last_recovery_assessment = None + self._recovery_rebind_artifact = None self._tracker.sync_graph(force=True) return assessment - def _resume_sender_after_recovery(self, timeout: float | None) -> None: - """Best-effort producer reconnect after state recovery has committed.""" - try: - self._connect_sender(timeout=timeout) - except (PermissionError, ConnectionError): - # Recovery already completed and must not look rolled back. Status - # exposes rejection/offline state; update() retries ordinary loss. - pass - - def _refresh_recovery_checkpoint(self, timeout: float | None) -> None: - """Apply through a fresh shared-stage replay watermark.""" - deadline = deadline_after(timeout) - reconnect = self._receiver.reconnect - self._receiver.reconnect = True - try: - self._receiver.request_replay_from(self._last_seq + 1) - while True: - self._tracker.prepare_local_changes() - try: - self._apply_incoming() - finally: - self._tracker.restore_prepared() - if self._receiver.synchronized: - return - if deadline is not None and time.monotonic() >= deadline: - raise TimeoutError("authoritative shared-stage recovery replay timed out") - time.sleep(0.01) - finally: - self._receiver.reconnect = reconnect - - def flush(self, timeout: float | None = None) -> bool: - """Wait for every submitted layer edit to be durably committed.""" - if self._closed: - raise RuntimeError("SharedStageClient is closed") - return self._sender.flush(timeout) - - @property - def stage_metadata(self) -> dict: - return dict(self._receiver.stage_metadata) - - @property - def last_seq(self) -> int: - return self._last_seq - - @property - def prepared_event_count(self) -> int: - return self._tracker.prepared_event_count - - @property - def deferred_event_count(self) -> int: - return len(self._pending_records) - - @property - def deferred_layer_keys(self) -> tuple[str, ...]: - return tuple(dict.fromkeys(record.layer_key or "" for record in self._pending_records)) - - def is_layer_reachable(self, layer: Sdf.Layer) -> bool: - """Return whether *layer* is reachable in the synchronized root graph.""" - layer_key = self._graph.key_for(layer) - return bool(layer_key and layer_key in self._graph.reachable_layer_keys()) - - def start(self) -> SharedStageClient: - """Start the background receiver; connect the sender after handshake.""" - if self._closed: - raise RuntimeError("SharedStageClient is closed") - if not self._started: - if self._receiver.token is None and self._persist_token: - self._receiver.token = load_token(self._host, self._port) - self._receiver.start() - self._started = True - return self - - def connect(self, timeout: float | None = None) -> bool: - """Start and complete both shared-stage handshakes within ``timeout``.""" - self.start() - deadline = deadline_after(timeout) - if not self._receiver.wait_connected(timeout): - if self._receiver.auth_rejected: - raise PermissionError("shared-stage receiver authentication rejected") - if self._receiver.hello_rejected: - raise ConnectionError( - self._receiver.rejection_reason or "shared-stage receiver rejected" - ) - return False - if self._receiver.layer_mode_active is not LayerMode.SHARED_STAGE: - raise RuntimeError("server did not negotiate shared-stage mode") - remaining = remaining_time(deadline) - return self._connect_sender(timeout=remaining) - - def _connect_sender(self, timeout: float | None = None) -> bool: - if self._sender.connected: - return True - self._prepare_sender_token() - if not self._sender.connect(timeout=timeout): - raise_if_rejected(self._sender, "shared-stage sender") - return False - return True - - def _prepare_sender_token(self) -> None: - prepare_sender_token( - self._sender, self._receiver, - host=self._host, port=self._port, persist_token=self._persist_token, - ) - - def update(self) -> SyncUpdate: - """Apply queued authoritative records, then publish local layer edits.""" - if self._closed: - raise RuntimeError("SharedStageClient is closed") - if not self._started: - raise RuntimeError("SharedStageClient has not been started") + def update(self, *, max_messages: int | None = None) -> SyncUpdate: + """Apply queued authoritative records, then publish local layer edits. - self._tracker.prepare_local_changes() - try: - received = self._apply_incoming() - finally: - self._tracker.restore_prepared() + ``max_messages`` bounds one call's receive work; local edits are held + until the backlog queued before them has been applied. + """ + if not self._begin_update(): + return self._progress() + received = self._apply_queued(max_messages) sent = 0 if self._graph.ready and self._receiver.connected and not self._sender.connected: - self._prepare_sender_token() self._sender.request_connect() - if self._sender.connected and self._graph.ready and self.synchronized: + if self._sender.connected and self._is_synchronized() and not self._backlog.holding: while routed := self._tracker.next_routed_batch(): batch, layer_key, events = routed if not self._sender.send_events(events, layer_key=layer_key): break sent += len(events) self._tracker.mark_prepared_sent(batch) - return SyncUpdate( - applied_events=received, - submitted_events=sent, - acknowledged_events_delta=self._sender.drain_acknowledged_event_count(), - pending_events=self._sender.pending_event_count, - recovery=self._sender.recovery_incident, - ) + return self._progress(received, sent) + + def _apply_queued(self, max_messages: int | None = None) -> int: + """Apply queued records while local edits are frozen out of the layers.""" + 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) + try: + return self._apply_incoming(max_messages) + finally: + self._tracker.restore_prepared() - def _apply_incoming(self) -> int: - buffers = self._receiver.drain_queue() + def _apply_incoming(self, max_messages: int | None = None) -> int: + 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 @@ -683,7 +524,7 @@ def _apply_incoming(self) -> int: require_contiguous=True, ) if result.resync_requested: - self._pending_records.clear() + self._set_deferred([]) applied_seq = 0 if result.resync_requested else self._last_seq applied = 0 try: @@ -749,7 +590,7 @@ def _apply_record(self, record: ReceivedEvent) -> bool: layer = self._graph.layer_for(layer_key) if layer is None: - self._pending_records.append(record) + self._defer(record) return False self._apply_layer_events(layer, [event]) return True @@ -773,6 +614,19 @@ def _restore_edit_target(self, preferred: Usd.EditTarget) -> None: else: self._stage.SetEditTarget(Usd.EditTarget(self._stage.GetRootLayer())) + def _defer(self, record: ReceivedEvent) -> None: + """Hold a record until its layer is mapped by the local graph.""" + self._pending_records.append(record) + key = record.layer_key or "" + if key not in self._deferred_layer_keys: + self._deferred_layer_keys += (key,) + + def _set_deferred(self, records: list[ReceivedEvent]) -> None: + self._pending_records = records + self._deferred_layer_keys = tuple( + dict.fromkeys(record.layer_key or "" for record in records) + ) + def _apply_pending(self) -> int: if not self._pending_records: return 0 @@ -792,13 +646,12 @@ def _apply_pending(self) -> int: events = [record.event for record in records] self._apply_layer_events(layer, events) applied += len(records) - self._pending_records = retained + self._set_deferred(retained) return applied def refresh_layer_graph(self) -> tuple[str, ...]: """Retry unresolved graph edges under this stage's resolver context.""" - if self._closed: - raise RuntimeError("SharedStageClient is closed") + self._require_open() with self._tracker.suppressed(): mapped = self._graph.refresh_dependencies() self._tracker.sync_graph(force=True) @@ -818,8 +671,7 @@ def _rebind_stage_for_recovery(self, stage: Usd.Stage) -> None: sender outbox is allowed because its exact bytes and layer snapshots are retained by the active recovery incident and assessment. """ - if self._closed: - raise RuntimeError("SharedStageClient is closed") + self._require_open() if not isinstance(stage, Usd.Stage): raise TypeError("SharedStageClient requires a Usd.Stage") if Sdf.Layer.IsAnonymousLayerIdentifier(stage.GetRootLayer().identifier): @@ -848,25 +700,38 @@ def _rebind_stage_for_recovery(self, stage: Usd.Stage) -> None: self._stage = stage self._graph = graph self._tracker = tracker - self._pending_records.clear() + self._set_deferred([]) self._last_seq = 0 old_tracker.close() - def close(self) -> None: - if self._closed: - return - self._sender.disconnect() - stop_receiver(self._receiver) - self._tracker.close() - self._last_recovery_assessment = None - self._closed = True + def _is_synchronized(self) -> bool: + return ( + self._graph.ready + and self._receiver.synchronized + and not self._sender.recovery_required + ) - def __enter__(self) -> SharedStageClient: - return self.start() + def _prepared_events(self) -> int: + return self._tracker.prepared_event_count + + def _has_unsent_changes(self) -> bool: + return self._tracker.has_local_changes - def __exit__(self, exc_type, exc, traceback) -> bool: - self.close() - return False + def _role_status(self) -> dict: + target = self._stage.GetEditTarget().GetLayer() + return { + "deferred_events": len(self._pending_records), + "deferred_layer_keys": self._deferred_layer_keys, + "edit_target_is_published": ( + target in self._stage.GetLayerStack(includeSessionLayers=False) + ), + "recovery_stage_pending": self._recovery_stage_pending, + } + + def _release(self) -> None: + self._tracker.close() + self._last_recovery_assessment = None + self._recovery_rebind_artifact = None __all__ = [ diff --git a/openusdconnect/usd_client.py b/openusdconnect/usd_client.py index 7208835..d4be293 100644 --- a/openusdconnect/usd_client.py +++ b/openusdconnect/usd_client.py @@ -1,4 +1,4 @@ -"""Layer-preserving client lifecycle for USD-native Python applications. +"""Directional clients for USD-native Python applications. The classes in this module compose the low-level sender, receiver, emitter, and dispatcher without changing their event or stage semantics. Applications @@ -12,40 +12,20 @@ from pxr import Usd -from ._client_lifecycle import ( - deadline_after, - prepare_sender_token, - raise_if_rejected, - remaining_time, - stop_receiver, -) -from ._client_utils import ( - client_origin, - client_token_handlers, - require_app_name, - resolve_client_token, - validate_layered_source, -) +from ._client_base import ClientBase, EmitterClientBase +from ._client_utils import client_origin, require_app_name, validate_layered_source from .adapters import DCCAdapter, UsdStageAdapter from .client_id import make_stable_client_id -from .client_types import ClientPhase, ClientStatus -from .coalescing import TransformCoalescingWindow +from .client_observer import ClientObserver +from .client_types import SyncUpdate from .defaults import DEFAULT_HOST, DEFAULT_SYNC_PORT from .dispatcher import AssetDependencyRefreshResult, EventDispatcher -from .emitter import NoticeEmitter, PrimChannel +from .emitter import PrimChannel from .receiver import ReceiverThread -from .recovery import ( - RecoveryArtifact, - RecoveryError, - RecoveryIncident, - RejectionDisposition, - TransactionFailure, -) from .sender import EventSender -from .token_client import load_token -class UsdReceiver: +class UsdReceiver(ClientBase): """Receive authoritative layered replay through a high-level lifecycle. ``start`` launches only the socket reader. ``update`` drains and applies @@ -68,27 +48,18 @@ def __init__( persist_token: bool = True, reconnect: bool = True, adapter: DCCAdapter | None = None, - on_imported: Callable[[list[str]], None] | None = None, - on_resync: Callable[[], None] | None = None, - on_applied: Callable[[list[str]], None] | None = None, - on_applied_events: Callable[[list[dict]], 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, - on_token_issued: Callable[[str], None] | None = None, + observer: ClientObserver | None = None, ): app_name = require_app_name(app_name) if not isinstance(stage, Usd.Stage): raise TypeError("UsdReceiver requires a Usd.Stage composition source") validate_layered_source(stage) + super().__init__( + host=host, port=port, token=token, persist_token=persist_token, observer=observer, + ) + self._stage: Usd.Stage | None = stage self._owns_stage_adapter = adapter is None - destination_adapter = adapter or UsdStageAdapter(stage) - resolved_token = resolve_client_token(host, port, token, persist_token) - self._stage = stage - self._host = host - self._port = port - self._persist_token = persist_token + destination = adapter or UsdStageAdapter(stage) self._receiver = ReceiverThread( host=host, port=port, @@ -96,25 +67,17 @@ def __init__( reconnect=reconnect, client_id=client_id or make_stable_client_id(app_name), origin=origin or client_origin(app_name, "recv"), - token=resolved_token, - on_token_issued=client_token_handlers(host, port, persist_token, on_token_issued), - on_stage_metadata=on_stage_metadata, - on_playback_state=on_playback_state, - on_playback_claimed=on_playback_claimed, - on_playback_rejected=on_playback_rejected, layered_replay=True, + **self._credential.endpoint_kwargs(), + **self._hooks.receiver_callbacks(), ) self._dispatcher = EventDispatcher( receiver=self._receiver, - adapter=destination_adapter, - mirror_stage=(None if destination_adapter.targets_stage() is stage else stage), - on_imported=on_imported, - on_resync=on_resync, - on_applied=on_applied, - on_applied_events=on_applied_events, + adapter=destination, + mirror_stage=None if destination.targets_stage() is stage else stage, + on_resync=self._hooks.on_resync, ) - self._started = False - self._closed = False + self._dispatcher.on_applied_events = self._hooks.applied_events_for(self._dispatcher) @property def stage(self) -> Usd.Stage | None: @@ -126,66 +89,15 @@ def stage(self) -> Usd.Stage | None: return self._stage @property - def status(self) -> ClientStatus: - """Current receiver transport and replay state.""" - if self._closed: - phase = ClientPhase.CLOSED - elif self.auth_rejected or self.connection_rejected: - phase = ClientPhase.REJECTED - elif self.native_scene_rebuild_required: - phase = ClientPhase.RECOVERY_REQUIRED - elif self.synchronized: - phase = ClientPhase.READY - elif self.connected: - phase = ClientPhase.REPLAYING - elif self._started: - phase = ClientPhase.CONNECTING - else: - phase = ClientPhase.OFFLINE - return ClientStatus( - phase=phase, - connected=self.connected, - synchronized=self.synchronized, - receiver_connected=self._receiver.connected, - reason=( - "the adapter-owned scene must be rebuilt after resolver recomposition" - if self.native_scene_rebuild_required - else self._receiver.rejection_reason - ), - ) - - @property - def receiver(self): - """The underlying :class:`ReceiverThread`.""" + def receiver(self) -> ReceiverThread: + """The underlying :class:`ReceiverThread`; a diagnostic handle.""" return self._receiver @property - def client_id(self) -> str: - """Stable identity used by this receiver.""" - return self._receiver.client_id - - @property - def applying_seq(self) -> int: - """Sequence of the current delivery, available inside apply callbacks.""" - return self._dispatcher.applying_seq - - @property - def dispatcher(self): - """The underlying :class:`EventDispatcher`.""" + def dispatcher(self) -> EventDispatcher: + """The underlying :class:`EventDispatcher`; a diagnostic handle.""" return self._dispatcher - @property - def connected(self) -> bool: - return not self._closed and self._receiver.connected - - @property - def synchronized(self) -> bool: - return not self._closed and self._receiver.synchronized - - @property - def layered_replay_active(self) -> bool: - return not self._closed and self._receiver.layered_replay_active - @property def last_seq(self) -> int: return self._dispatcher.last_seq @@ -200,83 +112,30 @@ def replay_epoch(self) -> int: """Epoch of the replay that has been fully applied.""" return self._receiver.replay_epoch - @property - def auth_rejected(self) -> bool: - return self._receiver.auth_rejected - - @property - def connection_rejected(self) -> bool: - return self._receiver.hello_rejected - - @property - def stage_metadata(self) -> dict: - return dict(self._receiver.stage_metadata) - @property def pending_asset_dependencies(self) -> tuple[str, ...]: return self._dispatcher.pending_asset_dependencies - @property - def native_scene_rebuild_required(self) -> bool: - """Whether an external adapter destination needs a complete rebuild.""" - return self._dispatcher.native_scene_rebuild_required - - def start(self) -> UsdReceiver: - """Start the background socket reader and return this receiver.""" - if self._closed: - raise RuntimeError("UsdReceiver is closed") - if not self._started: - if self._receiver.token is None and self._persist_token: - self._receiver.token = load_token(self._host, self._port) - self._receiver.start() - self._started = True - return self - - def connect(self, timeout: float | None = None) -> bool: - """Start and complete the receiver handshake within ``timeout``. - - Queued replay still requires :meth:`update` on the stage-owning thread. - """ - self.start() - connected = self._receiver.wait_connected(timeout) - if connected: - self._require_layered_replay() - elif self._receiver.auth_rejected: - raise PermissionError("receiver authentication rejected") - elif self._receiver.hello_rejected: - raise ConnectionError(self._receiver.rejection_reason or "receiver connection rejected") - return connected - - def _require_layered_replay(self) -> None: - if self._receiver.connected and not self._receiver.layered_replay_active: - self.close() - raise RuntimeError("server did not negotiate required layered replay") - - def update(self, *, max_messages: int | None = None) -> int: + def update(self, *, max_messages: int | None = None) -> SyncUpdate: """Apply queued messages on the calling thread. ``max_messages`` bounds one call's receive work for interactive applications. Replay becomes ready only after every message preceding the server's synchronization watermark has been applied. """ - if self._closed: - raise RuntimeError("UsdReceiver is closed") - if not self._started: - raise RuntimeError("UsdReceiver has not been started") - if self._stage is None: - return 0 - self._require_layered_replay() - return self._dispatcher.drain_and_apply(max_messages=max_messages) + if not self._begin_update() or self._stage is None: + return self._progress() + return self._progress( + self._dispatch(self._dispatcher.drain_and_apply, max_messages=max_messages) + ) def rebind_stage(self, stage: Usd.Stage | None) -> None: """Move receive-side composition and managed layers to a new stage. - Pass ``None`` to park: the receiver stays connected and the queue - continues to fill, but ``update()`` returns zero until a new stage - is bound. A caller-provided external adapter remains attached. + ``None`` parks the receiver: it stays connected and the queue keeps + filling until a stage is bound. A caller-provided adapter stays attached. """ - if self._closed: - raise RuntimeError("UsdReceiver is closed") + self._require_open() if stage is None: self._stage = None self._dispatcher.unbind_stage() @@ -299,33 +158,30 @@ def refresh_asset_dependency( asset_path: str | None = None, ) -> AssetDependencyRefreshResult: """Retry dependencies under the stage's current resolver context.""" - if self._closed: - raise RuntimeError("UsdReceiver is closed") - return self._dispatcher.refresh_asset_dependency(asset_path) + self._require_open() + return self._dispatch(self._dispatcher.refresh_asset_dependency, asset_path) def acknowledge_native_scene_rebuilt(self) -> None: """Resume projection after rebuilding an external adapter destination.""" - if self._closed: - raise RuntimeError("UsdReceiver is closed") + self._require_open() self._dispatcher.acknowledge_native_scene_rebuilt() - def close(self) -> None: - """Stop networking and release receiver-owned collaboration layers.""" - if self._closed: - return - stop_receiver(self._receiver) - self._dispatcher.close() - self._closed = True + def _is_synchronized(self) -> bool: + return self._stage is not None and self._receiver.synchronized - def __enter__(self) -> UsdReceiver: - return self.start() + def _is_parked(self) -> bool: + return self._stage is None - def __exit__(self, exc_type, exc, traceback) -> bool: - self.close() - return False + def _rebuild_reason(self) -> str: + if self._dispatcher.native_scene_rebuild_required: + return "the adapter-owned scene must be rebuilt after resolver recomposition" + return "" + def _release(self) -> None: + self._dispatcher.close() -class UsdPublisher: + +class UsdPublisher(EmitterClientBase): """Publish current-edit-target opinions authored on a USD stage.""" def __init__( @@ -340,26 +196,26 @@ def __init__( department: str | None = None, token: str | None = None, persist_token: bool = True, - on_token_issued: Callable[[str], None] | None = None, - on_stage_metadata: Callable[[dict], None] | None = None, + observer: ClientObserver | None = None, attr_filter: Callable[[str], bool] | None = None, 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): raise TypeError("UsdPublisher requires a Usd.Stage") + super().__init__( + host=host, port=port, token=token, persist_token=persist_token, observer=observer, + ) self._stage = stage - self._host = host - self._port = port - self._persist_token = persist_token - self._transform_coalescing = TransformCoalescingWindow(transform_coalesce_seconds) - self._emitter = NoticeEmitter( + self._init_emitter( stage, attr_filter=attr_filter, replicated_api_schemas=replicated_api_schemas, extra_channels=extra_channels, + transform_coalesce_seconds=transform_coalesce_seconds, ) self._sender = EventSender( host, @@ -367,240 +223,43 @@ def __init__( client_id=client_id or make_stable_client_id(app_name), origin=origin or client_origin(app_name, "emit"), department=department, - token=resolve_client_token(host, port, token, persist_token), - on_token_issued=client_token_handlers(host, port, persist_token, on_token_issued), - on_stage_metadata=on_stage_metadata, + on_stage_metadata=self._hooks.on_stage_metadata, + background_send=background_send, + **self._credential.endpoint_kwargs(), ) - self._closed = False - self._started = False - self._connecting = False @property def stage(self) -> Usd.Stage: """Application-owned stage observed for authored changes.""" return self._stage - @property - def status(self) -> ClientStatus: - """Current publisher transport, durability, and recovery state.""" - failure = self._sender.transaction_failure - reason = str(failure) if failure is not None else self._sender.rejection_reason - if self._closed: - phase = ClientPhase.CLOSED - elif failure is not None: - phase = ClientPhase.RECOVERY_REQUIRED - elif self._sender.auth_rejected or self._sender.hello_rejected: - phase = ClientPhase.REJECTED - elif self.connected: - phase = ClientPhase.READY - elif self._connecting: - phase = ClientPhase.CONNECTING - else: - phase = ClientPhase.OFFLINE - return ClientStatus( - phase=phase, - connected=self.connected, - synchronized=self.synchronized, - sender_connected=self._sender.connected, - prepared_events=self.prepared_event_count, - pending_events=self.pending_event_count, - acknowledged_events_total=self.acknowledged_event_count, - failure=failure, - recovery=self._sender.recovery_incident, - reason=reason, - ) - - @property - def sender(self): - """The underlying :class:`EventSender`.""" - return self._sender - - @property - def client_id(self) -> str: - """Stable identity used by this publisher.""" - return self._sender.client_id - - @property - def emitter(self): - """The underlying :class:`NoticeEmitter`.""" - return self._emitter - - @property - def connected(self) -> bool: - return not self._closed and self._sender.connected - - @property - def synchronized(self) -> bool: - """Send-only clients are synchronized whenever their transport is connected.""" - return self.connected - - @property - def auth_rejected(self) -> bool: - return self._sender.auth_rejected - - @property - def stage_metadata(self) -> dict: - return dict(self._sender.stage_metadata) - - @property - def prepared_event_count(self) -> int: - """Number of events not yet accepted by the sender outbox.""" - return self._emitter.prepared_event_count - - @property - def pending_event_count(self) -> int: - """Submitted events not yet durably acknowledged by the server.""" - return self._sender.pending_event_count - - @property - def acknowledged_event_count(self) -> int: - """Cumulative events durably acknowledged by the server.""" - return self._sender.acknowledged_event_count - - @property - def transaction_error(self) -> str: - """Terminal producer rejection, or an empty string.""" - return self._sender.transaction_error - - @property - def transaction_failure(self) -> TransactionFailure | None: - """Structured rejection including its recovery disposition, if any.""" - return self._sender.transaction_failure - - @property - def recovery_incident(self) -> RecoveryIncident | None: - """Structured recovery summary for polling and host UI.""" - return self._sender.recovery_incident - - @property - def recovery_artifact(self) -> RecoveryArtifact | None: - """Exact quarantined transactions for integration-owned recovery.""" - return self._sender.recovery_artifact - - @property - def recovery_disposition(self) -> RejectionDisposition | None: - """Recovery policy category for the current rejection, if any.""" - return self._sender.recovery_disposition - - def repair_and_resume(self, events: list[dict]) -> int: - """Replace a recoverable transaction and resume its ordered outbox. - - The application must first reconcile its stage with authoritative - state and rebuild *events* for that state. The repaired transaction is - assigned the original rejected ID; later quarantined transactions keep - their existing IDs and replay after it. - """ - if self._closed: - raise RuntimeError("UsdPublisher is closed") - failure = self._sender.transaction_failure - if failure is None: - raise RecoveryError("no_incident", "there is no recovery incident to resolve") - if failure.disposition is not RejectionDisposition.RECOVERABLE_CONFLICT: - raise RecoveryError( - "wrong_recovery_kind", - f"{failure.code_name} is {failure.disposition.value}, not recoverable", - ) - txn_id = self._sender.repair_rejected_transaction(events) - if not self.connect(): - raise ConnectionError( - f"transaction {txn_id} repaired but reconnect to " - f"{self._host}:{self._port} failed; it remains queued" - ) - return txn_id - - def flush(self, timeout: float | None = None) -> bool: - """Submit any coalesced transform, then wait for durable acknowledgement.""" - if self._closed: - raise RuntimeError("UsdPublisher is closed") - deadline = deadline_after(timeout) - if self._transform_coalescing.buffering: - if not self.connected and not self.connect(timeout=remaining_time(deadline)): - return False - events = self._transform_coalescing.force(self._emitter) - if events and not self._send(events): - return False - return self._sender.flush(remaining_time(deadline)) - - def start(self) -> UsdPublisher: - """Enter the nonblocking lifecycle without opening a socket.""" - if self._closed: - raise RuntimeError("UsdPublisher is closed") - self._started = True - return self - - def connect(self, timeout: float | None = None) -> bool: - """Start and complete the publisher handshake within ``timeout``.""" - self.start() - prepare_sender_token( - self._sender, None, - host=self._host, port=self._port, persist_token=self._persist_token, - ) - self._connecting = True - try: - connected = self._sender.connect(timeout=timeout) - finally: - self._connecting = False - if connected: - return True - raise_if_rejected(self._sender, "publisher") - return False - def disconnect(self) -> None: - """Close the socket while retaining dirty and prepared emitter state.""" + """Close the socket and pause reconnection until :meth:`connect`.""" if not self._closed: + self._paused = True self._sender.disconnect() - def _send(self, events: list[dict]) -> int: - if not events: - return 0 - if self._sender.send_events(events): - self._emitter.mark_prepared_events_sent(events) - self._transform_coalescing.mark_submitted() - return len(events) - return 0 - - def _prepare_outgoing_events(self) -> list[dict]: - return self._transform_coalescing.prepare(self._emitter) - - def update(self) -> int: - """Build and send one retryable batch of authored stage changes.""" - if self._closed: - raise RuntimeError("UsdPublisher is closed") - if not self.connected: - return 0 - return self._send(self._prepare_outgoing_events()) - - def publish_current_edit_target(self) -> int: - """Publish all opinions currently authored in the active edit target. - - An earlier retained batch must be retried with :meth:`update` first. - This keeps one call from ambiguously mixing two transport transactions. - """ - if self._closed: - raise RuntimeError("UsdPublisher is closed") - if not self.connected: - return 0 - if self._emitter.prepared_event_count: - raise RuntimeError( - "an earlier publisher batch is still prepared; call update() " - "before publishing the current edit target" - ) - return self._send(self._emitter.prepare_snapshot_events_for_send()) - - def close(self) -> None: - """Disconnect and release the stage notice listener.""" - if self._closed: - return - self._sender.disconnect() - self._emitter.cleanup() - self._closed = True + def update(self, *, max_messages: int | None = None) -> SyncUpdate: + """Submit one retryable batch, or schedule a reconnect while disconnected. - def __enter__(self) -> UsdPublisher: - return self.start() - - def __exit__(self, exc_type, exc, traceback) -> bool: - self.close() - return False + ``max_messages`` is accepted for a uniform host loop; a publisher + receives nothing. + """ + if not self._begin_update(): + return self._progress() + sent = 0 + if self._sender.connected: + sent = self._send(self._prepare_outgoing_events()) + elif not self._paused: + 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) __all__ = ["UsdPublisher", "UsdReceiver"] diff --git a/scripts/check_versions.py b/scripts/check_versions.py index 2643427..c4332a8 100644 --- a/scripts/check_versions.py +++ b/scripts/check_versions.py @@ -111,6 +111,9 @@ def collect_errors() -> list[str]: if not SEMVER.fullmatch(release): errors.append(f"OpenUSDConnect release must use X.Y.Z SemVer, got {release!r}") + newest_entry = re.search(r"^## \[([^\]]+)\]", _read("CHANGELOG.md"), re.MULTILINE) + if newest_entry is None or newest_entry.group(1) != release: + errors.append(f"CHANGELOG.md must start with a '## [{release}]' section") if pyproject["project"].get("dynamic") != ["version"]: errors.append("pyproject project.version must be dynamic") if pyproject.get("tool", {}).get("hatch", {}).get("version", {}).get("path") != ( diff --git a/tests/helpers.py b/tests/helpers.py index 91064ea..079dc59 100644 --- a/tests/helpers.py +++ b/tests/helpers.py @@ -6,8 +6,11 @@ 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, @@ -211,3 +214,69 @@ def mcp_session_with_receiver(port): # These tests drive one connection attempt directly to control reconnect timing. session.receiver._started = True return session + + +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)), + ) + + 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 + + +class RecordingObserver(ClientObserver): + """Records (method, value, thread id) for every observer call.""" + + def __init__(self, on_call=None): + self.calls = [] + self._on_call = on_call + + def _record(self, name, value): + self.calls.append((name, value, threading.get_ident())) + if self._on_call is not None: + self._on_call(name, value) + + def on_applied(self, batch): + self._record("applied", batch) + + def on_resync(self): + self._record("resync", None) + + def on_stage_metadata(self, metadata): + self._record("stage_metadata", metadata) + + def on_playback_state(self, state): + self._record("playback_state", state) + + def on_playback_claim(self, result): + self._record("playback_claim", result) + + def on_token_issued(self, token): + self._record("token_issued", token) diff --git a/tests/integration/scripts/blender_emitter_reconnect_script.py b/tests/integration/scripts/blender_emitter_reconnect_script.py index 4f0cd8d..7c6e132 100644 --- a/tests/integration/scripts/blender_emitter_reconnect_script.py +++ b/tests/integration/scripts/blender_emitter_reconnect_script.py @@ -57,6 +57,12 @@ def _find_cube(): def _tick(): global _phase, _edit_time try: + if _phase == "setup": + # Importing the add-on first would load its native module, which + # Windows then refuses to overwrite during the reinstall. + bpy.ops.preferences.addon_install(filepath=ARGS.addon, overwrite=True) + bpy.ops.preferences.addon_enable(module="usd_connect") + from usd_connect import capture if time.monotonic() > _deadline: @@ -71,8 +77,6 @@ def _tick(): return None if _phase == "setup": - bpy.ops.preferences.addon_install(filepath=ARGS.addon, overwrite=True) - bpy.ops.preferences.addon_enable(module="usd_connect") scene = bpy.context.scene scene.usd_connect_live_auto_start_emitter = False scene.usd_connect_live_auto_start_receiver = False diff --git a/tests/integration/scripts/shared_stage_custom_resolver_client.py b/tests/integration/scripts/shared_stage_custom_resolver_client.py index 2d9d417..2105fd4 100644 --- a/tests/integration/scripts/shared_stage_custom_resolver_client.py +++ b/tests/integration/scripts/shared_stage_custom_resolver_client.py @@ -65,7 +65,7 @@ def main() -> int: raise RuntimeError( "first custom-resolver edit was not committed: " f"update={update!r}, mapped={client.is_layer_reachable(content)}, " - f"prepared={client.prepared_event_count}, value={value.Get()!r}, " + f"prepared={client.status.prepared_events}, value={value.Get()!r}, " f"content_default={content_spec.default if content_spec else None!r}, " f"dirty={content.dirty}, editable={content.permissionToEdit}, " f"tracker={type(client._tracker).__name__}, " @@ -80,7 +80,7 @@ def main() -> int: raise RuntimeError( "second custom-resolver edit was not committed: " f"update={update!r}, mapped={client.is_layer_reachable(content)}, " - f"prepared={client.prepared_event_count}, value={value.Get()!r}" + f"prepared={client.status.prepared_events}, value={value.Get()!r}" ) print( diff --git a/tests/integration/test_client_completion.py b/tests/integration/test_client_completion.py new file mode 100644 index 0000000..0c2c336 --- /dev/null +++ b/tests/integration/test_client_completion.py @@ -0,0 +1,113 @@ +"""Readiness and local-publication completion through the public client APIs.""" + +import time + +import pytest +from pxr import Gf, Usd, UsdGeom + +from openusdconnect import ( + ClientPhase, + ManagedClient, + ServerConfig, + ServerRuntime, + SharedStageClient, + UsdPublisher, + UsdReceiver, +) +from openusdconnect.protocol_constants import LayerMode + + +def _stage(path): + stage = Usd.Stage.CreateNew(str(path)) + prim = stage.DefinePrim("/World", "Xform") + UsdGeom.Xformable(prim).AddTranslateOp().Set(Gf.Vec3d(0)) + stage.GetRootLayer().Save() + return stage + + +@pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient]) +def test_submit_and_wait_publishes_unprepared_edits_and_waits_for_commit(tmp_path, kind): + # Distinct files prevent the in-process server and client from sharing a + # mutable Sdf.Layer through USD's layer registry. + server_base = tmp_path / "server.usda" + _stage(server_base) + stage = _stage(tmp_path / "client.usda") + managed = kind is ManagedClient + config = ServerConfig( + host="127.0.0.1", + port=0, + base_usd_path=str(server_base), + log_path=str(tmp_path / "events.db"), + layer_mode=LayerMode.MANAGED if managed else LayerMode.SHARED_STAGE, + ) + with ServerRuntime(config) as server: + options = {"transform_coalesce_seconds": 30.0} if managed else {} + with kind( + stage, + app_name="completion-test", + port=server.server_address[1], + persist_token=False, + **options, + ) as client: + assert client.wait_until_ready(timeout=5) + assert client.status.phase is ClientPhase.READY + assert client.status.edit_target_is_published + + stage.GetAttributeAtPath("/World.xformOp:translate").Set(Gf.Vec3d(2, 3, 4)) + assert client.status.has_unsent_changes + assert client.status.prepared_events == 0 + assert client.status.pending_events == 0 + assert client.flush(timeout=0) # ACK-only flush has no submitted work yet. + + assert client.submit_and_wait(timeout=5) + assert not client.status.has_unsent_changes + assert client.status.pending_events == 0 + assert client.status.acknowledged_events_total > 0 + assert server.sync_server.stage.GetAttributeAtPath( + "/World.xformOp:translate" + ).Get() == Gf.Vec3d(2, 3, 4) + + +def test_directional_clients_share_the_blocking_helpers(tmp_path): + server_base = tmp_path / "server.usda" + _stage(server_base) + author = _stage(tmp_path / "author.usda") + author.SetEditTarget(Usd.EditTarget(author.GetSessionLayer())) + viewer = _stage(tmp_path / "viewer.usda") + config = ServerConfig( + host="127.0.0.1", + port=0, + base_usd_path=str(server_base), + log_path=str(tmp_path / "events.db"), + ) + with ServerRuntime(config) as server: + port = server.server_address[1] + with ( + UsdPublisher( + author, app_name="completion-author", port=port, persist_token=False, + ) as publisher, + UsdReceiver( + viewer, app_name="completion-viewer", port=port, persist_token=False, + ) as receiver, + ): + assert publisher.wait_until_ready(timeout=5) + assert receiver.wait_until_ready(timeout=5) + assert publisher.status.phase is ClientPhase.READY + assert receiver.status.phase is ClientPhase.READY + + author.GetAttributeAtPath("/World.xformOp:translate").Set(Gf.Vec3d(2, 3, 4)) + assert publisher.status.has_unsent_changes + assert publisher.submit_and_wait(timeout=5) + assert not publisher.status.has_unsent_changes + + received = viewer.GetAttributeAtPath("/World.xformOp:translate") + deadline = time.monotonic() + 5 + while received.Get() != Gf.Vec3d(2, 3, 4) and time.monotonic() < deadline: + assert receiver.update().submitted_events == 0 + time.sleep(0.01) + assert received.Get() == Gf.Vec3d(2, 3, 4) + + receiver.rebind_stage(None) + assert receiver.status.phase is ClientPhase.PARKED + with pytest.raises(RuntimeError, match="no bound stage"): + receiver.wait_until_ready(timeout=0) diff --git a/tests/integration/test_managed_client.py b/tests/integration/test_managed_client.py index 3383c04..3473d1e 100644 --- a/tests/integration/test_managed_client.py +++ b/tests/integration/test_managed_client.py @@ -65,6 +65,9 @@ def _client_stage(tmp_path, name="client"): def _translation(prim: Usd.Prim): + """Local translation, or ``None`` before the prim has arrived.""" + if not prim: + return None m = UsdGeom.Xformable(prim).GetLocalTransformation(Usd.TimeCode.Default()) return (m[3][0], m[3][1], m[3][2]) @@ -108,11 +111,12 @@ def test_managed_client_tight_loop_round_trips_without_crash(live_server, tmp_pa try: assert _drain_until( client, - lambda: _translation(sync_server.stage.GetPrimAtPath("/World/Test"))[0] == 199.0, + lambda: _translation(sync_server.stage.GetPrimAtPath("/World/Test")) + == (199.0, 0.0, 0.0), ) assert _drain_until( client, - lambda: _translation(stage.GetPrimAtPath("/World/Test"))[0] == 199.0, + lambda: _translation(stage.GetPrimAtPath("/World/Test")) == (199.0, 0.0, 0.0), ) assert _translation(sync_server.stage.GetPrimAtPath("/World/Test")) == (199.0, 0.0, 0.0) assert _translation(stage.GetPrimAtPath("/World/Test")) == (199.0, 0.0, 0.0) @@ -154,7 +158,9 @@ def counting_send(events): for i in range(50): tr.Set(Gf.Vec3d(float(i), 0, 0)) client.update() - _drain_until(client, lambda: _translation(stage.GetPrimAtPath("/World/Test"))[0] == 49.0) + _drain_until( + client, lambda: _translation(stage.GetPrimAtPath("/World/Test")) == (49.0, 0.0, 0.0), + ) client.close() # /World and /World/Test are both locally defined by the first @@ -248,13 +254,13 @@ def reject_marked_transaction(events, *args, **kwargs): try: client.start() assert client.connect(timeout=5) - assert _drain_until(client, lambda: client.synchronized) + assert _drain_until(client, lambda: client.status.synchronized) rejected_session = client.sender.session_id stage.DefinePrim("/World/Rejected", "Xform") assert client.update().submitted_events > 0 - assert _drain_until(client, lambda: client.recovery_required) - incident = client.recovery_incident + assert _drain_until(client, lambda: client.status.phase is ClientPhase.RECOVERY_REQUIRED) + incident = client.status.recovery assert incident is not None assert incident.producer_session_id == rejected_session assert incident.event_count > 0 @@ -265,8 +271,7 @@ def reject_marked_transaction(events, *args, **kwargs): assert recovered.preserved_authoring_layer.GetPrimAtPath("/World/Rejected") assert not stage.GetPrimAtPath("/World/Rejected") assert client.sender.session_id == "managed-replacement-session" - assert not client.recovery_required - assert client.connected + assert client.status.connected assert client.status.phase is ClientPhase.READY assert client.receiver.reconnect is False assert client.receiver.replay_head_seq == sync_server.store.get_max_seq() @@ -405,7 +410,7 @@ def test_managed_client_shares_reissued_tokens(tmp_path, background, first_recon try: assert client.connect(timeout=5) sender_readers.append(client.sender._reader_thread) - assert _drain_until(client, lambda: client.synchronized) + assert _drain_until(client, lambda: client.status.synchronized) old_token = client.sender.token assert old_token == client.receiver.token assert runtime.sync_server.revoke_token(client.client_id) @@ -427,9 +432,9 @@ def test_managed_client_shares_reissued_tokens(tmp_path, background, first_recon client.sender.disconnect() if background: assert _drain_until( - client, lambda: client.connected or client.sender.auth_rejected, + client, lambda: client.status.connected or client.sender.auth_rejected, ) - assert client.connected + assert client.status.connected else: assert client.connect(timeout=3) sender_readers.append(client.sender._reader_thread) @@ -437,15 +442,14 @@ def test_managed_client_shares_reissued_tokens(tmp_path, background, first_recon assert not client.sender.auth_rejected assert sender_tokens[0] != old_token assert sender_tokens[1] == sender_tokens[0] - assert client.receiver.token == sender_tokens[0] # The other connection must also authenticate with the replacement. client.receiver.request_replay_from(1) assert not client.receiver.synchronized assert _drain_until( - client, lambda: client.synchronized or client.receiver.auth_rejected, + client, lambda: client.status.synchronized or client.receiver.auth_rejected, ) - assert client.synchronized + assert client.status.synchronized assert not client.receiver.auth_rejected assert client.sender.token == client.receiver.token == sender_tokens[0] finally: @@ -484,10 +488,10 @@ def test_managed_client_hands_ephemeral_tofu_token_to_sender(tmp_path): try: client.start() assert client.connect(timeout=5) - assert client.connected + assert client.status.connected assert client.receiver.token assert client.sender.token == client.receiver.token - assert not client.auth_rejected + assert not client.status.auth_rejected finally: client.close() tcp_server.shutdown() diff --git a/tests/integration/test_receiver_replay_identity.py b/tests/integration/test_receiver_replay_identity.py index 6fb2e8c..32160cf 100644 --- a/tests/integration/test_receiver_replay_identity.py +++ b/tests/integration/test_receiver_replay_identity.py @@ -21,7 +21,7 @@ def _drain_ready(session): def ready(): session.receiver.update() - return session.receiver.synchronized + return session.receiver.status.synchronized wait_until(ready) diff --git a/tests/integration/test_shared_stage_client.py b/tests/integration/test_shared_stage_client.py index 689a59b..1251ee3 100644 --- a/tests/integration/test_shared_stage_client.py +++ b/tests/integration/test_shared_stage_client.py @@ -92,7 +92,7 @@ def test_shared_client_shares_reissued_tokens(tmp_path, background, first_reconn try: assert client.connect(timeout=5) sender_readers.append(client._sender._reader_thread) - assert _pump_until([client], lambda: client.synchronized) + assert _pump_until([client], lambda: client.status.synchronized) old_token = client._sender.token assert old_token == client._receiver.token assert runtime.sync_server.revoke_token(client.client_id) @@ -114,9 +114,9 @@ def test_shared_client_shares_reissued_tokens(tmp_path, background, first_reconn client._sender.disconnect() if background: assert _pump_until( - [client], lambda: client.connected or client._sender.auth_rejected, + [client], lambda: client.status.connected or client._sender.auth_rejected, ) - assert client.connected + assert client.status.connected else: assert client.connect(timeout=3) sender_readers.append(client._sender._reader_thread) @@ -124,15 +124,14 @@ def test_shared_client_shares_reissued_tokens(tmp_path, background, first_reconn assert not client._sender.auth_rejected assert sender_tokens[0] != old_token assert sender_tokens[1] == sender_tokens[0] - assert client._receiver.token == sender_tokens[0] # The other connection must also authenticate with the replacement. client._receiver.request_replay_from(1) assert not client._receiver.synchronized assert _pump_until( - [client], lambda: client.synchronized or client._receiver.auth_rejected, + [client], lambda: client.status.synchronized or client._receiver.auth_rejected, ) - assert client.synchronized + assert client.status.synchronized assert not client._receiver.auth_rejected assert client._sender.token == client._receiver.token == sender_tokens[0] finally: @@ -689,8 +688,8 @@ def test_clients_reconnect_and_converge_after_server_restart(tmp_path): assert _pump_until( [first, second], lambda: ( - first.connected - and second.connected + first.status.connected + and second.status.connected and _value(first_stage) == 5 and _value(second_stage) == 5 and _value(restarted_stage) == 5 @@ -918,7 +917,7 @@ def test_shared_client_use_server_recovers_after_layer_detach_race(tmp_path, reb assert client._sender.connected is rebind # update schedules a background handshake; readiness arrives on a later tick. - assert _pump_until([client], lambda: client.connected) + assert _pump_until([client], lambda: client.status.connected) finally: client.close() tcp_server.shutdown() diff --git a/tests/integration/test_usd_client.py b/tests/integration/test_usd_client.py index b2131e3..8bc190c 100644 --- a/tests/integration/test_usd_client.py +++ b/tests/integration/test_usd_client.py @@ -89,11 +89,10 @@ def test_layered_receiver_preserves_and_clears_an_override(live_server): assert publisher.connect() source.GetAttributeAtPath(_VALUE_PATH).Set(17) UsdGeom.SetStageUpAxis(source, UsdGeom.Tokens.z) - assert publisher.update() > 0 + assert publisher.update().submitted_events > 0 receiver.start() assert receiver.connect(timeout=2) - assert receiver.layered_replay_active assert _pump_until(receiver, lambda: _value(target) == 17) router = receiver._dispatcher.layer_router @@ -112,7 +111,7 @@ def test_layered_receiver_preserves_and_clears_an_override(live_server): assert UsdGeom.GetStageUpAxis(target) == UsdGeom.Tokens.z source.GetAttributeAtPath(_VALUE_PATH).Clear() - assert publisher.update() > 0 + assert publisher.update().submitted_events > 0 assert _pump_until(receiver, lambda: _value(target) == 5) assert not managed.GetAttributeAtPath(_VALUE_PATH).HasInfo("default") diff --git a/tests/integration/test_vfs_webdav.py b/tests/integration/test_vfs_webdav.py index 9e2800d..f2ca16e 100644 --- a/tests/integration/test_vfs_webdav.py +++ b/tests/integration/test_vfs_webdav.py @@ -797,7 +797,7 @@ def test_translated_put_restores_live_client_readiness_and_publishing( def pump_until_ready(): client.update() - return client.synchronized + return client.status.synchronized try: client.start() @@ -814,7 +814,7 @@ def pump_until_ready(): ) assert 200 <= status < 300 - assert _wait_until(lambda: client.connected and not client.synchronized) + assert _wait_until(lambda: client.status.connected and not client.status.synchronized) assert _wait_until( lambda: pump_until_ready() and bool(stage.GetPrimAtPath("/Root/FromVfsPut")) diff --git a/tests/unit/layered_replay_test_support.py b/tests/unit/layered_replay_test_support.py index 34e8e63..e2584ba 100644 --- a/tests/unit/layered_replay_test_support.py +++ b/tests/unit/layered_replay_test_support.py @@ -35,6 +35,7 @@ def _state( class _LayeredQueue: layered_replay = True layered_replay_active = True + sync_from = 1 origin = None def __init__(self, messages): diff --git a/tests/unit/test_asset_dependency_refresh.py b/tests/unit/test_asset_dependency_refresh.py index 6ec9cc1..3576f76 100644 --- a/tests/unit/test_asset_dependency_refresh.py +++ b/tests/unit/test_asset_dependency_refresh.py @@ -27,6 +27,7 @@ class _NullReceiver: layered_replay_active = False + sync_from = 1 origin = None def drain_queue(self): @@ -41,6 +42,7 @@ def request_replay_from(self, _seq_start): class _QueueReceiver: layered_replay_active = False + sync_from = 1 origin = None def __init__(self): diff --git a/tests/unit/test_client_lifecycle.py b/tests/unit/test_client_lifecycle.py index 3500f9e..78fe2eb 100644 --- a/tests/unit/test_client_lifecycle.py +++ b/tests/unit/test_client_lifecycle.py @@ -9,6 +9,8 @@ from openusdconnect import ( ManagedClient, SharedStageClient, + TransactionFailure, + TransactionRejectedError, UsdPublisher, UsdReceiver, _client_lifecycle, @@ -16,6 +18,145 @@ client_types, ) from openusdconnect import sender as sender_module +from openusdconnect.client_observer import StageMetadata +from tests.helpers import RecordingObserver, force_handshake + + +def test_wait_until_ready_returns_false_only_when_startup_expires(): + client = SimpleNamespace( + start=lambda: None, + update=lambda: None, + status=SimpleNamespace(phase=client_types.ClientPhase.REPLAYING, failure=None), + ) + assert not _client_lifecycle.wait_until_ready(client, timeout=0) + + +@pytest.mark.parametrize( + ("phase", "auth_rejected", "failure", "expected"), + [ + ("rejected", True, None, PermissionError), + ("rejected", False, None, ConnectionError), + ( + "recovery_required", False, TransactionFailure(1, 0, "rejected"), + TransactionRejectedError, + ), + ("recovery_required", False, None, RuntimeError), + ("parked", False, None, RuntimeError), + ("closed", False, None, RuntimeError), + ], +) +def test_blocked_states_raise_instead_of_timing_out(phase, auth_rejected, failure, expected): + status = SimpleNamespace( + phase=client_types.ClientPhase(phase), failure=failure, reason="", + auth_rejected=auth_rejected, + ) + with pytest.raises(expected): + _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")) + client = kind(stage, app_name="phases", port=1, persist_token=False) + try: + assert client.status.phase is client_types.ClientPhase.OFFLINE + assert not client.status.connected + client.start() + assert client.status.phase is client_types.ClientPhase.CONNECTING + finally: + client.close() + assert client.status.phase is client_types.ClientPhase.CLOSED + with pytest.raises(RuntimeError, match=f"{kind.__name__} is closed"): + client.update() + + +def test_waits_raise_when_nothing_will_reconnect(): + publisher = UsdPublisher( + Usd.Stage.CreateInMemory(), app_name="paused", port=1, persist_token=False, + ) + receiver = UsdReceiver( + Usd.Stage.CreateInMemory(), app_name="one-shot", port=1, persist_token=False, + reconnect=False, + ) + try: + publisher.start() + publisher.disconnect() + receiver.start() + receiver.receiver.join(timeout=5) + for client in (publisher, receiver): + assert client.status.phase is client_types.ClientPhase.OFFLINE + with pytest.raises(ConnectionError, match="offline"): + client.wait_until_ready(timeout=5) + with pytest.raises(ConnectionError, match="offline"): + publisher.submit_and_wait(timeout=5) + finally: + publisher.close() + 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(): @@ -30,48 +171,74 @@ def test_public_status_types_keep_compatibility_identity(): assert value is getattr(_client_utils, name) -@pytest.mark.parametrize( - ("sender_token", "receiver", "persist", "expected"), - [ - ("current", SimpleNamespace(token="stale"), False, "current"), - ("current", SimpleNamespace(token="stale"), True, "current"), - (None, SimpleNamespace(token="issued"), True, "issued"), - ("configured", SimpleNamespace(token=None), True, "configured"), - ("configured", None, True, "configured"), - (None, SimpleNamespace(token=None), True, "stored"), - (None, None, True, "stored"), - (None, SimpleNamespace(token=None), False, None), - (None, None, False, None), - ], -) -def test_sender_token_preparation_only_fills_missing_credentials( - monkeypatch, sender_token, receiver, persist, expected, -): +def test_credential_reads_storage_only_while_no_token_is_known(monkeypatch): + stored = [None] reads = [] def load_token(host, port): reads.append((host, port)) - return "stored" + return stored[0] monkeypatch.setattr(_client_utils, "load_token", load_token) - sender = SimpleNamespace(token=sender_token) - _client_lifecycle.prepare_sender_token( - sender, receiver, host="test-host", port=7200, persist_token=persist, + credential = _client_utils.ClientCredential("test-host", 7200, None, True) + assert credential.current() is None + stored[0] = "stored" + assert credential.current() == "stored" + reads.clear() + assert credential.current() == "stored" + assert _client_utils.ClientCredential("test-host", 7200, None, False).current() is None + assert reads == [] + + +def test_credential_keeps_a_token_issued_while_storage_loads(monkeypatch): + monkeypatch.setattr(_client_utils, "save_token", lambda host, port, token: None) + monkeypatch.setattr(_client_utils, "load_token", lambda host, port: None) + credential = _client_utils.ClientCredential("localhost", 1, None, True) + loading, release = threading.Event(), threading.Event() + + def slow_load(host, port): + loading.set() + assert release.wait(5) + return "old-stored-token" + + monkeypatch.setattr(_client_utils, "load_token", slow_load) + reader = threading.Thread(target=credential.current) + reader.start() + assert loading.wait(5) + issuer = threading.Thread(target=credential.issued, args=("issued-during-load",)) + issuer.start() + release.set() + reader.join(5) + issuer.join(5) + 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 sender.token == expected - assert reads == ([("test-host", 7200)] if expected == "stored" else []) + assert not sender.connect(timeout=0.5) + assert not sender.connect(timeout=0.5) + assert presented == ["first", "second"] -@pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient]) -@pytest.mark.parametrize("issuer", ["_sender", "_receiver"]) @pytest.mark.parametrize("failure", [None, "persistence", "observer"]) -def test_issued_token_updates_both_connections_before_callbacks( - kind, issuer, failure, tmp_path, monkeypatch, +def test_issued_token_is_persisted_before_notifying_and_used_by_both_roles( + failure, tmp_path, monkeypatch, ): calls = [] def record(name, token): - calls.append((name, token, client._sender.token, client._receiver.token)) + calls.append((name, token)) if failure == name: raise RuntimeError(f"injected {name} failure") @@ -79,37 +246,54 @@ def record(name, token): _client_utils, "save_token", lambda host, port, token: record("persistence", token), ) stage = Usd.Stage.CreateNew(str(tmp_path / "scene.usda")) - client = kind( + observer = RecordingObserver(on_call=lambda name, token: record("observer", token)) + client = ManagedClient( stage, app_name="shared-credentials", token="configured", persist_token=True, - on_token_issued=lambda token: record("observer", token), + observer=observer, ) try: - callback = getattr(client, issuer)._on_token_issued - if failure is None: - callback("replacement") - else: - with pytest.raises(RuntimeError, match=f"injected {failure} failure"): + callback = client._sender._on_token_issued + if failure == "persistence": + with pytest.raises(RuntimeError, match="injected persistence failure"): callback("replacement") - - assert client._sender.token == client._receiver.token == "replacement" - expected = ["persistence"] if failure == "persistence" else ["persistence", "observer"] - assert calls == [(name, "replacement", "replacement", "replacement") for name in expected] + 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") finally: client.close() -def test_token_issued_while_loading_credentials_is_not_overwritten(monkeypatch): - sender = SimpleNamespace(token=None) - - def load_token(host, port): - sender.token = "issued-during-load" - return "old-stored-token" - - monkeypatch.setattr(_client_utils, "load_token", load_token) - _client_lifecycle.prepare_sender_token( - sender, None, host="test-host", port=7200, persist_token=True, +@pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient, UsdPublisher]) +def test_disconnected_update_does_no_token_io(kind, tmp_path, monkeypatch): + reads = [] + monkeypatch.setattr(_client_utils, "load_token", lambda host, port: reads.append(1)) + requests = [] + monkeypatch.setattr( + sender_module.EventSender, "request_connect", + lambda self, timeout=2.0: requests.append(True) or True, ) - assert sender.token == "issued-during-load" + stage = Usd.Stage.CreateNew(str(tmp_path / "scene.usda")) + client = kind(stage, app_name="token-io", port=1, persist_token=True) + try: + force_handshake(client) + reads.clear() + for _ in range(20): + client.update() + assert requests + assert reads == [] + finally: + client.close() @pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient]) @@ -134,11 +318,7 @@ def connect(*args, **kwargs): finished.set() monkeypatch.setattr(sender_module.socket, "create_connection", connect) - client._started = True - client._receiver.connected = True - client._receiver.layered_replay_active = True - if isinstance(client, SharedStageClient): - client._graph._ready = True + force_handshake(client) try: result = client.update() assert entered.wait(2) @@ -177,6 +357,8 @@ def flush(timeout=None): if isinstance(client, ManagedClient): client._receiver.connected = True client._receiver._synchronized_event.set() + else: + monkeypatch.setattr(client, "_is_synchronized", lambda: True) try: assert client.flush(timeout=1.0) assert calls == [("connect", 1.0), ("flush", 0.25)] @@ -184,29 +366,46 @@ def flush(timeout=None): client.close() -def test_publisher_accepts_host_owned_token_and_metadata_callbacks(): - tokens, metadata = [], [] - with UsdPublisher( - Usd.Stage.CreateInMemory(), - app_name="host-credentials", - client_id="host-id", - persist_token=False, - on_token_issued=tokens.append, - on_stage_metadata=metadata.append, - ) as client: - client.sender._on_token_issued("issued-token") - client.sender._on_stage_metadata({"metersPerUnit": 0.01}) - assert client.client_id == "host-id" - assert tokens == ["issued-token"] - assert metadata == [{"metersPerUnit": 0.01}] - - -def test_receiver_exposes_identity_and_current_delivery_sequence(): - client = UsdReceiver( - Usd.Stage.CreateInMemory(), app_name="viewer", client_id="viewer-id", persist_token=False +@pytest.mark.parametrize( + ("phase", "sender_connected", "edit_target_is_published", "expected"), + [ + ("ready", True, None, True), + ("ready", True, False, False), + ("ready", None, None, False), + ("replaying", True, True, False), + ], +) +def test_can_author_combines_readiness_role_and_edit_target( + phase, sender_connected, edit_target_is_published, expected, +): + status = client_types.ClientStatus( + phase=client_types.ClientPhase(phase), + connected=True, + synchronized=True, + sender_connected=sender_connected, + edit_target_is_published=edit_target_is_published, ) + assert status.can_author is expected + + +@pytest.mark.parametrize("kind", [ManagedClient, SharedStageClient, UsdReceiver, UsdPublisher]) +def test_close_delivers_notifications_queued_by_network_threads(kind, tmp_path): + 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 try: - assert client.client_id == "viewer-id" - assert client.applying_seq == client.dispatcher.applying_seq + 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 == [] 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()), + ] diff --git a/tests/unit/test_client_observer.py b/tests/unit/test_client_observer.py new file mode 100644 index 0000000..5ce05ca --- /dev/null +++ b/tests/unit/test_client_observer.py @@ -0,0 +1,132 @@ +"""ClientObserver wiring and typed notification payloads.""" + +from pxr import Usd + +from openusdconnect import ( + AppliedBatch, + ClientObserver, + ClientPhase, + ManagedClient, + PlaybackClaim, + 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 + + +def test_applied_batch_derives_sorted_unique_paths_once(): + events = [ + {"k": K_SET_VISIBILITY, "prim": "/World/B", "visible": True}, + {"k": K_SET_REFERENCE, "prim": "/World/A", "references": []}, + {"k": K_ENSURE_PRIM, "prim": "/World/B", "typeName": "Xform"}, + ] + batch = AppliedBatch(7, events) + assert batch.prim_paths == ("/World/A", "/World/B") + assert batch.prim_paths is batch.prim_paths + assert batch.imported_paths == ("/World/A",) + assert batch.seq == 7 + + +def test_only_overridden_methods_are_wired(): + 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 _event_frames(*paths): + return [ + encode_message({ + "type": "event", "seq": seq, + "event": {"k": K_ENSURE_PRIM, "prim": prim, "typeName": "Xform"}, + }) + for seq, prim in enumerate(paths, start=1) + ] + + +def test_receiver_update_delivers_applied_batches_with_the_drained_sequence(monkeypatch): + batches = [] + + class Record(ClientObserver): + def on_applied(self, batch): + batches.append((batch.seq, batch.prim_paths)) + + client = UsdReceiver( + Usd.Stage.CreateInMemory(), app_name="observer-batches", persist_token=False, + reconnect=False, observer=Record(), + ) + frames = _event_frames("/World", "/World/A") + monkeypatch.setattr(client.receiver, "drain_queue", lambda max_messages=None: frames) + client._started = True + try: + assert client.update().applied_events == 2 + assert batches == [(2, ("/World", "/World/A"))] + finally: + client.close() + + +def test_close_from_a_delivery_method_takes_effect_after_the_apply(monkeypatch): + class CloseOnApply(ClientObserver): + def on_applied(self, batch): + client.close() + + stage = Usd.Stage.CreateInMemory() + client = ManagedClient( + stage, app_name="close-on-apply", persist_token=False, reconnect=False, + observer=CloseOnApply(), + ) + monkeypatch.setattr( + client.receiver, "drain_queue", lambda max_messages=None: _event_frames("/World"), + ) + client._started = True + assert client.update().applied_events == 1 + assert client.status.phase is ClientPhase.CLOSED + assert stage.GetPrimAtPath("/World") + + +def test_stage_edits_made_in_on_resync_are_not_published(monkeypatch): + class Reset(ClientObserver): + def on_resync(self): + client.stage.DefinePrim("/FromResync", "Xform") + + client = ManagedClient( + Usd.Stage.CreateInMemory(), app_name="resync-edit", persist_token=False, + reconnect=False, observer=Reset(), + ) + monkeypatch.setattr( + client.receiver, "drain_queue", + lambda max_messages=None: [encode_message({"type": "resync"})], + ) + client._started = True + try: + client.update() + assert client.stage.GetPrimAtPath("/FromResync") + assert not client.status.has_unsent_changes + finally: + client.close() + + +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"), + PlaybackClaim(True, "a"), + PlaybackClaim(False, "b", "busy"), + ] diff --git a/tests/unit/test_dispatcher.py b/tests/unit/test_dispatcher.py index 965eee0..feabbaa 100644 --- a/tests/unit/test_dispatcher.py +++ b/tests/unit/test_dispatcher.py @@ -17,6 +17,7 @@ class _NullReceiver: layered_replay_active = False + sync_from = 1 origin = None def drain_queue(self): @@ -31,15 +32,16 @@ def request_replay_from(self, _seq_start): class _QueuedReceiver: layered_replay_active = False + sync_from = 1 origin = None def __init__(self, messages): self.messages = list(messages) self.replay_requests = [] - def drain_queue(self): - messages = self.messages - self.messages = [] + def drain_queue(self, max_messages=None): + count = len(self.messages) if max_messages is None else max_messages + messages, self.messages = self.messages[:count], self.messages[count:] return messages def request_replay_from(self, seq_start): @@ -111,6 +113,18 @@ def set_connectable_input(self, prim_path, info_id, inputs, input_types, time=No assert receiver.replay_requests == [] +def test_cursor_starts_at_the_receiver_continuation_point(): + receiver = _QueuedReceiver([_event(6, "/World/Continued")]) + receiver.sync_from = 6 + adapter = MockAdapter() + dispatcher = EventDispatcher(receiver=receiver, adapter=adapter) + + assert dispatcher.last_seq == 5 + assert dispatcher.drain_and_apply() == 1 + assert dispatcher.last_seq == 6 + assert receiver.replay_requests == [] + + def test_post_apply_callbacks_run_in_documented_order(): calls = [] event = { @@ -307,3 +321,14 @@ def test_sdf_spec_batches_use_full_layer_atomic_rollback(): dispatcher._apply([valid, invalid]) assert mirror.GetRootLayer().documentation == "original" + + +def test_budgeted_drain_reports_messages_taken(): + receiver = _QueuedReceiver([_event(seq, f"/World/P{seq}") for seq in (1, 2, 3)]) + dispatcher = EventDispatcher(receiver=receiver, adapter=MockAdapter()) + + assert dispatcher.drain_and_apply(max_messages=2) == 2 + assert dispatcher.drained_message_count == 2 + assert dispatcher.drain_and_apply(max_messages=2) == 1 + assert dispatcher.drained_message_count == 1 + assert dispatcher.last_seq == 3 diff --git a/tests/unit/test_emitter_notices.py b/tests/unit/test_emitter_notices.py index f85e971..cf19434 100644 --- a/tests/unit/test_emitter_notices.py +++ b/tests/unit/test_emitter_notices.py @@ -40,6 +40,21 @@ def _make_stage_and_emitter(): return stage, emitter +def test_removed_local_def_over_weaker_prim_leaves_no_pending_changes(): + stage = Usd.Stage.CreateInMemory() + stage.DefinePrim("/World/Child", "Xform") + stage.SetEditTarget(Usd.EditTarget(stage.GetSessionLayer())) + emitter = NoticeEmitter(stage) + stage.DefinePrim("/World", "Xform") + assert emitter.build_events_for_dirty() + + stage.RemovePrim("/World") + events = emitter.build_events_for_dirty() + + assert {"k": K_DELETE_PRIM, "prim": "/World"} in events + assert not emitter.has_local_changes + + class TestCreationDetection: """DefinePrim triggers ensure_prim + set_xform_trs events.""" diff --git a/tests/unit/test_mcp_session.py b/tests/unit/test_mcp_session.py index 15f8919..e162e35 100644 --- a/tests/unit/test_mcp_session.py +++ b/tests/unit/test_mcp_session.py @@ -9,6 +9,7 @@ from integrations.mcp.errors import ToolError from openusdconnect import usd_client from openusdconnect.checkpoints import MirrorCheckpoint +from openusdconnect.client_types import SyncUpdate class _FakeSender: @@ -30,17 +31,26 @@ def disconnect(self): self.is_connected = False +def _applied(count: int) -> SyncUpdate: + return SyncUpdate(applied_events=count, submitted_events=0) + + def _patch_net(monkeypatch, started, stopped): class _FakeReceiver: synchronized = True connected = True + auth_rejected = False + hello_rejected = False + rejection_reason = "" layered_replay_active = True server_instance = "test-server" replay_epoch = 0 + stopped = False def __init__(self, **kwargs): self.options = kwargs self.token = kwargs["token"] + self.sync_from = kwargs["sync_from"] self.joined = False def start(self): @@ -56,6 +66,12 @@ def join(self, timeout=None): assert self in stopped self.joined = True + def drain_queue(self, max_messages=None): + return [] + + def mark_replay_applied(self): + return False + monkeypatch.setattr(session_mod, "EventSender", _FakeSender) monkeypatch.setattr(usd_client, "ReceiverThread", _FakeReceiver) monkeypatch.setattr(session_mod.token_client, "load_token", lambda host, port: None) @@ -106,17 +122,14 @@ def test_playback_status_reflects_broadcast(monkeypatch): assert session.playback_status()["observed"] is False # nothing broadcast yet - session._on_playback_state( - {"playing": True, "time": 12.0, "rate": 2.0, "leader_client_id": "mcp-x"} - ) + notify = session.receiver.receiver.options["on_playback_state"] + notify({"playing": True, "time": 12.0, "rate": 2.0, "leader_client_id": "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 - session._on_playback_state( - {"playing": False, "time": 0.0, "rate": 1.0, "leader_client_id": "someone-else"} - ) + notify({"playing": False, "time": 0.0, "rate": 1.0, "leader_client_id": "someone-else"}) st2 = session.playback_status() assert st2["is_leader"] is False assert st2["leader_client_id"] == "someone-else" @@ -162,7 +175,7 @@ def test_concurrent_foreign_write_cannot_confirm_own_transaction(monkeypatch): def apply_foreign_write(): session.mirror_stage.DefinePrim("/Foreign", "Xform") session.receiver.dispatcher.last_seq += 1 - return 1 + return _applied(1) monkeypatch.setattr(session.receiver, "update", apply_foreign_write) try: @@ -196,7 +209,7 @@ def test_confirmation_requires_matching_applied_checkpoint( def apply(): session.receiver.dispatcher.last_seq = 1 session.receiver.receiver.synchronized = ready - return 0 + return _applied(0) monkeypatch.setattr(session.receiver, "update", apply) try: @@ -225,7 +238,7 @@ def flush(timeout): def apply(): updates.append(True) session.receiver.dispatcher.last_seq = 50 if len(updates) >= 2 else 1 - return 1 + return _applied(1) monkeypatch.setattr(session.sender, "flush", flush) monkeypatch.setattr(session.receiver, "update", apply) @@ -256,7 +269,7 @@ def apply(): acknowledged = True session.sender.acknowledged_checkpoint = MirrorCheckpoint("test-server", 0, 1) session.receiver.dispatcher.last_seq = 1 - return applied + return _applied(applied) monkeypatch.setattr(session.sender, "send_events", lambda events: True, raising=False) monkeypatch.setattr(session.sender, "flush", flush) @@ -283,7 +296,7 @@ def apply(): nonlocal elapsed elapsed += 0.001 session.receiver.dispatcher.last_seq += applied - return applied + return _applied(applied) def sleep(seconds): nonlocal elapsed @@ -313,7 +326,7 @@ def test_pending_ack_times_out_without_false_confirmation(monkeypatch): monkeypatch.setattr(session.sender, "send_events", lambda events: True, raising=False) monkeypatch.setattr(session.sender, "flush", lambda timeout: False) session.sender.acknowledged_checkpoint = MirrorCheckpoint("test-server", 0, 1) - monkeypatch.setattr(session.receiver, "update", lambda: 0) + monkeypatch.setattr(session.receiver, "update", lambda: _applied(0)) session.receiver.dispatcher.last_seq = 100 try: assert not session.send([{}])["mirror_synced"] @@ -335,7 +348,7 @@ def test_rejected_transaction_is_reported_as_tool_error(monkeypatch, after_apply def apply(): nonlocal applied applied = True - return 0 + return _applied(0) def reject(timeout): if after_apply and not applied: @@ -377,12 +390,14 @@ 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}) + options["on_playback_state"]( + {"playing": True, "time": 1.0, "rate": 1.0, "leader_client_id": ""} + ) assert session.playback_status()["playing"] is True dispatcher = session.receiver.dispatcher dispatcher.last_seq = 3 dispatcher._applying_seq = 7 - dispatcher.on_applied(["/World"]) + dispatcher.on_applied_events([{"k": "ensure_prim", "prim": "/World"}]) assert session._dirty == {"/World": 7} assert session.receiver.last_seq == 3 finally: diff --git a/tests/unit/test_sender.py b/tests/unit/test_sender.py index 631b18e..eda1473 100644 --- a/tests/unit/test_sender.py +++ b/tests/unit/test_sender.py @@ -46,9 +46,12 @@ def test_malformed_events_fail_before_encoding_or_outbox_ownership(self): assert sender.pending_transaction_count == 0 assert sender._next_txn_id == 1 - def test_handshake_and_send_events(self): + @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") + sender = EventSender( + "127.0.0.1", port, client_id="test-client", background_send=background_send + ) conn = None try: import threading @@ -288,10 +291,15 @@ def _fail(_endpoint, *, timeout): assert 0 < observed[0] <= 0.1 assert 0 < observed[1] <= 0.5 - def test_reconnect_replays_identical_bytes_until_duplicate_ack(self): + @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" + "127.0.0.1", + port, + client_id="test-client", + session_id="stable-session", + background_send=background_send, ) observed = [] first_closed = threading.Event() @@ -340,7 +348,8 @@ def _serve(): thread.join(timeout=2) srv.close() - def test_bounded_outbox_and_rejection_are_terminal(self): + @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", @@ -348,6 +357,7 @@ def test_bounded_outbox_and_rejection_are_terminal(self): client_id="test-client", session_id="bounded-session", max_pending_transactions=1, + background_send=background_send, ) def _serve(): @@ -387,13 +397,15 @@ def _serve(): thread.join(timeout=2) srv.close() - def test_rejection_closes_transport_and_quarantines_later_transactions(self): + @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, ) received = threading.Event() @@ -493,13 +505,15 @@ def _serve(): thread.join(timeout=2) srv.close() - def test_recoverable_rejection_reuses_boundary_before_later_transactions(self): + @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 = [] @@ -565,13 +579,15 @@ def test_nonrecoverable_rejection_cannot_be_retried(self): [{"k": "ensure_prim", "prim": "/World/X", "typeName": "Xform"}] ) - def test_abandon_rejected_session_never_replays_its_suffix(self): + @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 = [] diff --git a/tests/unit/test_sender_background.py b/tests/unit/test_sender_background.py new file mode 100644 index 0000000..ce8a0c0 --- /dev/null +++ b/tests/unit/test_sender_background.py @@ -0,0 +1,261 @@ +"""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 index aa7ff47..3882a4b 100644 --- a/tests/unit/test_sender_reconnect.py +++ b/tests/unit/test_sender_reconnect.py @@ -275,35 +275,51 @@ def test_background_handshake_honors_short_timeout(): sender.disconnect() -def test_cancel_after_hello_before_publication(monkeypatch): +@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") + sender = EventSender( + *server.getsockname(), client_id="late", background_send=background_send, + ) accept = sender._accept_handshake_response def blocked(*args): - result = accept(*args) entered.set() assert release.wait(3) - return result + 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: - conn.settimeout(2) - recv_framed(conn) - send_framed(conn, encode_message({"type": "hello_ok"})) + handshake(conn) assert entered.wait(2) sender.disconnect() release.set() _finish(sender) assert not sender.connected - assert not sender.recovery_required + + # 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() diff --git a/tests/unit/test_shared_stage_client.py b/tests/unit/test_shared_stage_client.py index 49b8693..dfcc7fd 100644 --- a/tests/unit/test_shared_stage_client.py +++ b/tests/unit/test_shared_stage_client.py @@ -6,7 +6,7 @@ from pxr import Ar, Sdf, Usd from openusdconnect import ClientPhase, RecoveryError -from openusdconnect.codec import ReceivedEvent, TransactionRejectionCode +from openusdconnect.codec import ReceivedEvent, TransactionRejectionCode, encode_message from openusdconnect.recovery import ( QuarantinedTransaction, RecoveryArtifact, @@ -15,6 +15,7 @@ ) from openusdconnect.sdf_spec_delta import serialize_spec_fields from openusdconnect.shared_stage_client import SharedStageClient +from tests.helpers import PeerTraffic def _create_root(path) -> Usd.Stage: @@ -36,6 +37,8 @@ def __init__(self, artifact: RecoveryArtifact): self.recovery_incident = make_recovery_incident(artifact) self.recovery_required = True self.pending_transaction_count = len(artifact.transactions) + self.pending_event_count = artifact.event_count + self.acknowledged_event_count = 0 self.abandoned_session_ids: list[str | None] = [] def abandon_rejected_session(self, *, session_id=None): @@ -184,6 +187,7 @@ class _StatusSender: acknowledged_event_count = 3 recovery_required = False recovery_incident = None + recovery_artifact = None sender = _StatusSender() client._sender = sender @@ -226,6 +230,34 @@ class _StatusSender: client.close() +def test_status_distinguishes_local_edit_targets_and_unsubmitted_changes(tmp_path): + stage = _create_root(tmp_path / "root.usda") + client = SharedStageClient(stage, app_name="authoring-scope", persist_token=False) + try: + assert client.status.edit_target_is_published + stage.SetEditTarget(stage.GetSessionLayer()) + stage.DefinePrim("/Local", "Xform") + assert not client.status.edit_target_is_published + assert not client.status.has_unsent_changes + + stage.SetEditTarget(stage.GetRootLayer()) + stage.DefinePrim("/Shared", "Xform") + assert client.status.edit_target_is_published + assert client.status.has_unsent_changes + assert client.status.prepared_events == 0 + assert client.status.pending_events == 0 + + client._started = True + result = client.update() + assert client.status.has_unsent_changes + assert result.submitted_events == 0 + assert client.status.prepared_events > 0 + client.close() + assert not client.status.has_unsent_changes + finally: + client.close() + + def test_unresolved_layer_events_apply_after_dependency_refresh(tmp_path): stage = _create_root(tmp_path / "root.usda") root = stage.GetRootLayer() @@ -278,7 +310,12 @@ def test_unresolved_layer_events_apply_after_dependency_refresh(tmp_path): "removed": False, } assert not client._apply_record(ReceivedEvent(seq=2, event=event, layer_key=child_key)) - assert client.deferred_event_count == 1 + assert client.status.deferred_events == 1 + client._receiver.connected = True + client._receiver._synchronized_event.set() + assert client.status.synchronized + assert client.status.deferred_events == 1 + assert client.status.deferred_layer_keys == (child_key,) late = Sdf.Layer.CreateNew(str(tmp_path / "late.usda")) Sdf.CreatePrimInLayer(late, "/Late") @@ -286,7 +323,8 @@ def test_unresolved_layer_events_apply_after_dependency_refresh(tmp_path): mapped = client.refresh_layer_graph() assert mapped == (child_key,) - assert client.deferred_event_count == 0 + assert client.status.deferred_events == 0 + assert client.status.deferred_layer_keys == () assert late.GetAttributeAtPath("/Late.value").default == 8 finally: client.close() @@ -317,7 +355,7 @@ def test_content_apply_failure_preserves_layer_and_tracker_until_retry( event={"k": "replace_sdf_layer_content", "fragment": source.ExportToString()}, )) if deferred: - client._pending_records.extend(records) + client._set_deferred(records) before = root.ExportToString() stage.SetEditTarget(stage.GetSessionLayer()) original_apply = client_module.apply_events @@ -347,13 +385,13 @@ def apply_records(): assert stage.GetEditTarget().GetLayer() == stage.GetSessionLayer() assert accepted == [] assert applied_batches == [[record.event for record in records]] - assert client.deferred_event_count == (len(records) if deferred else 0) + assert client.status.deferred_events == (len(records) if deferred else 0) monkeypatch.setattr(client_module, "apply_events", original_apply) assert apply_records() == len(records) assert accepted == [record.event for record in records] assert root.ExportToString() == source.ExportToString() - assert client.deferred_event_count == 0 + assert client.status.deferred_events == 0 finally: client.close() @@ -407,7 +445,7 @@ def test_update_restores_frozen_edits_when_replay_fails(tmp_path, monkeypatch): lambda: calls.append("restore"), ) - def _fail_replay(): + def _fail_replay(max_messages=None): calls.append("replay") raise RuntimeError("bad authoritative record") @@ -426,12 +464,12 @@ def test_repair_and_resume_targets_current_mapped_layer(tmp_path, monkeypatch): original_sender = client._sender repaired = [] - class _RepairSender: - recovery_artifact = _stale_artifact("layer:root") - transaction_failure = recovery_artifact.failure - + class _RepairSender(_RecoverySender): def repair_rejected_transaction(self, events, *, layer_key=""): + if not events: + raise ValueError("repair events must not be empty") repaired.append((events, layer_key)) + self.recovery_artifact = self.transaction_failure = None return 7 try: @@ -447,8 +485,9 @@ def repair_rejected_transaction(self, events, *, layer_key=""): ], } ) - client._sender = _RepairSender() + client._sender = _RepairSender(_stale_artifact("layer:root")) client._started = True + client._recovery_rebind_artifact = client._sender.recovery_artifact resumed = [] def _resume(): @@ -458,14 +497,18 @@ def _resume(): monkeypatch.setattr(client, "_connect_sender", _resume) events = [{"k": "replace_sdf_layer_content", "fragment": "#usda 1.0\n"}] - assert client.repair_and_resume(events, layer=stage.GetRootLayer()) == 7 - assert repaired == [(events, "layer:root")] - assert resumed == [True] - detached = Sdf.Layer.CreateAnonymous() with pytest.raises(RecoveryError, match="not mapped") as error: client.repair_and_resume(events, layer=detached) assert error.value.code == "invalid_repair_target" + with pytest.raises(ValueError, match="must not be empty"): + client.repair_and_resume([], layer=stage.GetRootLayer()) + assert client.status.recovery_stage_pending, "a failed repair keeps the incident" + + assert client.repair_and_resume(events, layer=stage.GetRootLayer()) == 7 + assert repaired == [(events, "layer:root")] + assert resumed == [True] + assert not client.status.recovery_stage_pending finally: client._sender = original_sender client.close() @@ -504,7 +547,7 @@ def _detach(_timeout): client._receiver.connected = True client._receiver._synchronized_event.set() - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", _detach) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", _detach) try: assessment = client.refresh_recovery_assessment() assert assessment.all_layers_detached @@ -538,7 +581,7 @@ def test_shared_use_server_refuses_a_quarantined_reachable_layer(tmp_path, monke sender = _RecoverySender(_stale_artifact("layer:child")) client._sender = sender client._started = True - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", lambda _timeout: None) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) try: assessment = client.refresh_recovery_assessment() assert assessment.recovery_artifact is sender.recovery_artifact @@ -574,7 +617,7 @@ def test_shared_assessment_reports_an_unavailable_source_layer(tmp_path, monkeyp sender = _RecoverySender(_stale_artifact("layer:missing")) client._sender = sender client._started = True - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", lambda _timeout: None) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) try: assessment = client.refresh_recovery_assessment() assert assessment.source_unavailable_layers == assessment.layers @@ -632,7 +675,7 @@ def test_shared_recovery_commands_distinguish_expected_policy_failures( ], } ) - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", lambda _timeout: None) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) assessment = client.refresh_recovery_assessment() with pytest.raises(RecoveryError) as not_synchronized: client.complete_recovery(assessment) @@ -660,7 +703,7 @@ def test_shared_use_server_keeps_incident_when_checkpoint_refresh_fails( def _timeout(_timeout): raise TimeoutError("injected checkpoint timeout") - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", _timeout) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", _timeout) try: with pytest.raises(TimeoutError, match="injected checkpoint timeout"): client.refresh_recovery_assessment() @@ -707,7 +750,7 @@ def _detach_child(_timeout): }, ) - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", _detach_child) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", _detach_child) try: assessment = client.refresh_recovery_assessment() assert [layer.rejected_layer_key for layer in assessment.detached_layers] == [ @@ -772,7 +815,7 @@ def _remap(_timeout): } ) - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", _remap) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", _remap) try: assessment = client.refresh_recovery_assessment() remapped = assessment.remapped_layers @@ -801,7 +844,7 @@ def test_shared_external_recovery_completes_a_structured_reachable_assessment( client._started = True client._receiver.connected = True client._receiver._synchronized_event.set() - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", lambda _timeout: None) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) try: assessment = client.refresh_recovery_assessment() assert not assessment.all_layers_detached @@ -837,7 +880,7 @@ def test_shared_external_recovery_rejects_an_assessment_from_another_incident( client._started = True client._receiver.connected = True client._receiver._synchronized_event.set() - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", lambda _timeout: None) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) try: assessment = client.refresh_recovery_assessment() sender.recovery_artifact = _stale_artifact("layer:child") @@ -868,7 +911,7 @@ def test_shared_external_recovery_rejects_a_stale_graph_assessment( client._started = True client._receiver.connected = True client._receiver._synchronized_event.set() - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", lambda _timeout: None) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) try: assessment = client.refresh_recovery_assessment() client._last_seq += 1 @@ -915,7 +958,7 @@ def _refresh(_timeout): client._receiver.connected = True client._receiver._synchronized_event.set() - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", _refresh) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", _refresh) try: with pytest.raises(RecoveryError, match="different clean stage") as error: client.recover_use_server(clean_stage=old_stage) @@ -947,6 +990,114 @@ def _refresh(_timeout): client.close() +@pytest.mark.parametrize("after_timeout", ["resume", "update", "local_edits", "different_incident"]) +def test_shared_rebind_recovery_resumes_after_replacement_replay_timeout( + tmp_path, + monkeypatch, + after_timeout, +): + old_stage = _create_root(tmp_path / "old-root.usda") + old_child = Sdf.Layer.CreateNew(str(tmp_path / "old-child.usda")) + Sdf.CreatePrimInLayer(old_child, "/Rejected") + old_child.Save() + old_stage.GetRootLayer().subLayerPaths.append("./old-child.usda") + client = SharedStageClient(old_stage, app_name="resume-recovery", persist_token=False) + with client._tracker.suppressed(): + _bind_child_graph(client) + original_sender = client._sender + sender = _RecoverySender(_stale_artifact("layer:child")) + client._sender = sender + client._started = True + + fresh_stage = _create_root(tmp_path / "fresh-root.usda") + fresh_child = Sdf.Layer.CreateNew(str(tmp_path / "fresh-child.usda")) + fresh_child.Save() + fresh_stage.GetRootLayer().subLayerPaths.append("./fresh-child.usda") + checkpoints = [] + + def refresh(_timeout): + checkpoints.append(client.stage) + client._receiver.connected = True + client._receiver._synchronized_event.set() + if len(checkpoints) == 2: + with client._tracker.suppressed(): + _bind_child_graph(client) + client.stage.DefinePrim("/Replayed", "Xform") + client._last_seq = 3 + raise TimeoutError("replacement replay timed out") + if len(checkpoints) == 3: + assert client.last_seq == (4 if after_timeout == "update" else 3) + assert client.stage.GetPrimAtPath("/Replayed") + client._last_seq += 1 + + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", refresh) + try: + with pytest.raises(TimeoutError, match="replacement replay timed out"): + client.recover_use_server(clean_stage=fresh_stage) + + assert client.stage is fresh_stage + assert client.status.recovery_stage_pending + assert client.status.phase is ClientPhase.RECOVERY_REQUIRED + assert sender.abandoned_session_ids == [] + + if after_timeout == "update": + source = Sdf.Layer.CreateAnonymous() + Sdf.CreatePrimInLayer(source, "/BetweenAttempts") + buffers = [encode_message({ + "type": "event", "seq": 4, "layer_key": "layer:child", + "event": { + "k": "replace_sdf_layer_content", "prim": "/", + "fragment": source.ExportToString(), + }, + })] + monkeypatch.setattr(client._receiver, "drain_queue", lambda max_messages=None: buffers) + monkeypatch.setattr(sender, "drain_acknowledged_event_count", lambda: 0, raising=False) + monkeypatch.setattr( + sender, "send_events", + lambda *_args, **_kwargs: pytest.fail("pending recovery must not publish"), + raising=False, + ) + sender.connected = True + assert client.update().applied_events == 1 + assert fresh_stage.GetPrimAtPath("/BetweenAttempts") + assert client.last_seq == 4 + assert client.status.recovery_stage_pending + assert client.status.phase is ClientPhase.RECOVERY_REQUIRED + + if after_timeout == "local_edits": + fresh_stage.DefinePrim("/Unsubmitted", "Xform") + with pytest.raises(RecoveryError) as error: + client.resume_recovery() + assert error.value.code == "local_changes_pending" + assert checkpoints == [old_stage, fresh_stage] + assert sender.recovery_required + elif after_timeout == "different_incident": + sender.recovery_artifact = _stale_artifact("layer:root") + assert not client.status.recovery_stage_pending + with pytest.raises(RecoveryError) as error: + client.resume_recovery() + assert error.value.code == "no_pending_recovery_stage" + with pytest.raises(RecoveryError) as error: + client.recover_use_server(clean_stage=fresh_stage) + assert error.value.code == "invalid_clean_stage" + assert checkpoints == [old_stage, fresh_stage] + else: + with pytest.raises(RecoveryError, match=r"resume_recovery\(\)") as error: + client.recover_use_server(clean_stage=fresh_stage) + assert error.value.code == "invalid_clean_stage" + result = client.resume_recovery() + assert checkpoints == [old_stage, fresh_stage, fresh_stage] + assert result.layers[0].source_layer is old_child + assert result.rejected_snapshots[0].GetPrimAtPath("/Rejected") + assert result.checkpoint_seq == (5 if after_timeout == "update" else 4) + assert client.stage is fresh_stage + assert not client.status.recovery_stage_pending + assert not sender.recovery_required + finally: + client._sender = original_sender + client.close() + + def test_shared_rebind_recovery_preflights_the_clean_stage(tmp_path, monkeypatch): old_stage = _create_root(tmp_path / "old-root.usda") client = SharedStageClient(old_stage, app_name="shared-invalid-clean", persist_token=False) @@ -968,7 +1119,7 @@ def test_shared_rebind_recovery_preflights_the_clean_stage(tmp_path, monkeypatch client._started = True client._receiver.connected = True client._receiver._synchronized_event.set() - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", lambda _timeout: None) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) clean_stage = _create_root(tmp_path / "clean-root.usda") clean_stage.SetEditTarget(Usd.EditTarget(clean_stage.GetSessionLayer())) @@ -1020,7 +1171,7 @@ def test_shared_rebind_recovery_rejects_a_detached_source_reused_by_clean_stage( client._started = True client._receiver.connected = True client._receiver._synchronized_event.set() - monkeypatch.setattr(client, "_refresh_recovery_checkpoint", lambda _timeout: None) + monkeypatch.setattr(client, "_replay_to_fresh_checkpoint", lambda _timeout: None) clean_stage = _create_root(tmp_path / "clean-root.usda") clean_stage.GetRootLayer().subLayerPaths.append("./child.usda") @@ -1037,3 +1188,37 @@ def test_shared_rebind_recovery_rejects_a_detached_source_reused_by_clean_stage( finally: client._sender = original_sender client.close() + + +def test_shared_budget_releases_local_edits_under_sustained_traffic(tmp_path, monkeypatch): + 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, + ) + 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() diff --git a/tests/unit/test_usd_client.py b/tests/unit/test_usd_client.py index d0583fb..6bbd114 100644 --- a/tests/unit/test_usd_client.py +++ b/tests/unit/test_usd_client.py @@ -22,17 +22,7 @@ make_recovery_incident, ) from openusdconnect.usd_client import UsdPublisher, UsdReceiver - - -def test_bidirectional_clients_share_one_update_result_contract(): - update = SyncUpdate( - applied_events=1, - submitted_events=2, - acknowledged_events_delta=3, - pending_events=4, - ) - assert update.acknowledged_events_delta == 3 - assert update.pending_events == 4 +from tests.helpers import PeerTraffic, RecordingObserver, force_handshake class _SenderStub: @@ -56,6 +46,7 @@ def __init__(self, results: list[bool]): self.connect_result = True self.repaired: list[tuple[list[dict], str]] = [] self.abandoned_session_ids: list[str | None] = [] + self.connect_requests = 0 def send_events(self, events: list[dict]) -> bool: self.batches.append(events) @@ -91,6 +82,10 @@ def connect(self, timeout=None) -> bool: self.connected = self.connect_result return self.connect_result + def request_connect(self, timeout=2.0) -> bool: + self.connect_requests += 1 + return True + def repair_rejected_transaction(self, events: list[dict], *, layer_key="") -> int: self.repaired.append((events, layer_key)) self.connected = False @@ -127,7 +122,7 @@ def test_receiver_parks_without_invalidating_dispatcher_adapter(): receiver._started = True assert receiver._dispatcher.adapter is adapter - assert receiver.update() == 0 + assert receiver.update().applied_events == 0 finally: receiver.close() @@ -148,7 +143,7 @@ def test_receiver_update_forwards_message_budget(monkeypatch): lambda *, max_messages=None: calls.append(max_messages) or 3, ) - assert receiver.update(max_messages=128) == 3 + assert receiver.update(max_messages=128).applied_events == 3 assert calls == [128] finally: receiver.close() @@ -231,14 +226,12 @@ def close(self): receiver._receiver.connected = True receiver._receiver._synchronized_event.set() try: - assert receiver.native_scene_rebuild_required assert receiver.status.phase is ClientPhase.RECOVERY_REQUIRED assert "must be rebuilt" in receiver.status.reason receiver.acknowledge_native_scene_rebuilt() assert state.acknowledgements == 1 - assert not receiver.native_scene_rebuild_required assert receiver.status.phase is ClientPhase.READY finally: receiver.close() @@ -266,7 +259,7 @@ def test_receiver_status_distinguishes_connecting_replay_and_ready(): assert receiver.status.phase is ClientPhase.CLOSED -def test_publisher_context_start_is_nonblocking_and_connect_is_explicit(): +def test_publisher_context_start_is_nonblocking_and_update_connects_in_background(): publisher = UsdPublisher( Usd.Stage.CreateInMemory(), app_name="lifecycle-publisher", @@ -275,11 +268,21 @@ def test_publisher_context_start_is_nonblocking_and_connect_is_explicit(): sender = _SenderStub([]) sender.connected = False publisher._sender = sender + assert publisher.status.phase is ClientPhase.OFFLINE with publisher as entered: assert entered is publisher assert sender.connect_timeouts == [] + assert publisher.status.phase is ClientPhase.CONNECTING + assert publisher.update().submitted_events == 0 + assert sender.connect_requests == 1 + assert sender.connect_timeouts == [] + + publisher.disconnect() assert publisher.status.phase is ClientPhase.OFFLINE + publisher.update() + assert sender.connect_requests == 1 + assert publisher.connect(timeout=0.25) assert sender.connect_timeouts == [0.25] assert publisher.status.phase is ClientPhase.READY @@ -306,10 +309,7 @@ def test_managed_status_exposes_partial_connection_and_event_counts(): sender.pending_event_count = 4 sender.acknowledged_event_count = 7 client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver.layered_replay_active = True - client._receiver._synchronized_event.set() + force_handshake(client, synchronized=True) try: status = client.status assert status.phase is ClientPhase.CONNECTING @@ -326,6 +326,149 @@ def test_managed_status_exposes_partial_connection_and_event_counts(): client.close() +@pytest.fixture +def ready_managed_client(monkeypatch): + client = ManagedClient( + Usd.Stage.CreateInMemory(), app_name="managed-api", persist_token=False, + ) + 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: + yield client, sender + finally: + client.close() + + +@pytest.mark.parametrize("options", [ + {"transform_coalesce_seconds": -1}, + {"extra_channels": [object()]}, +]) +def test_managed_invalid_configuration_preserves_stage(options): + stage = Usd.Stage.CreateInMemory() + stage.SetEditTarget(stage.GetSessionLayer()) + target = stage.GetEditTarget() + session = stage.GetSessionLayer().ExportToString() + with pytest.raises((TypeError, ValueError)): + ManagedClient(stage, app_name="invalid-options", persist_token=False, **options) + assert stage.GetEditTarget() == target + assert stage.GetSessionLayer().ExportToString() == session + + +def test_managed_metadata_only_changes_count_as_unsent_work(ready_managed_client): + client, _sender = ready_managed_client + with Usd.EditContext(client.stage, client.stage.GetRootLayer()): + client.stage.SetFramesPerSecond(48) + assert client.status.has_unsent_changes + # Evaluated even when the authoring layer does not own the opinion. + client.update() + assert not client.status.has_unsent_changes + + client.stage.GetRootLayer().framesPerSecond = 24 + assert client.status.has_unsent_changes + replacement = Usd.Stage.CreateInMemory() + replacement.SetFramesPerSecond(30) + client.rebind_stage(replacement, discard_unsent=True) + assert not client.status.has_unsent_changes + assert client.update().submitted_events == 0 + + +def test_managed_snapshot_waits_for_replay_without_losing_newer_edits( + ready_managed_client, monkeypatch, +): + client, sender = ready_managed_client + starts = [] + client._started = False + monkeypatch.setattr(client._receiver, "start", lambda: starts.append(True)) + client._receiver._synchronized_event.clear() + value = UsdGeom.Sphere.Define(client.stage, "/Local").GetRadiusAttr() + value.Set(1) + + assert client.publish_current_edit_target() == 0 + assert starts == [True] + assert sender.batches == [] + assert client.status.has_unsent_changes + with pytest.raises(RuntimeError, match="earlier publisher batch"): + client.publish_current_edit_target() + value.Set(2) + client._receiver._synchronized_event.set() + assert client.update().submitted_events > 0 + assert client.status.has_unsent_changes + assert client.update().submitted_events > 0 + assert not client.status.has_unsent_changes + values = [ + event["attrs"]["radius"] + for batch in sender.batches for event in batch + if event["k"] == "set_gprim_attrs" and "radius" in event["attrs"] + ] + assert values == [1, 2] + + +@pytest.mark.parametrize(("prepare", "park"), [(False, False), (True, True)]) +def test_managed_rebind_requires_explicit_unsent_discard( + ready_managed_client, prepare, park, +): + client, _sender = ready_managed_client + old_stage = client.stage + old_stage.DefinePrim("/Local", "Xform") + if prepare: + client.emitter.prepare_events_for_send() + replacement = None if park else Usd.Stage.CreateInMemory() + with pytest.raises(RuntimeError, match="unsent changes"): + client.rebind_stage(replacement) + assert client.stage is old_stage + assert client.status.has_unsent_changes + client.rebind_stage(replacement, discard_unsent=True) + assert client.stage is replacement + assert not client.status.has_unsent_changes + assert client.status.edit_target_is_published is not park + + +def test_managed_rebind_cannot_discard_submitted_work(ready_managed_client): + client, _sender = ready_managed_client + old_stage = client.stage + old_stage.DefinePrim("/Local", "Xform") + assert client.update().submitted_events > 0 + with pytest.raises(RuntimeError, match="submitted work pending"): + client.rebind_stage(None, discard_unsent=True) + assert client.stage is old_stage + + +def test_managed_parked_client_is_not_ready(ready_managed_client): + client, _sender = ready_managed_client + client.rebind_stage(None) + assert client.status.connected + assert not client.status.synchronized + assert client.status.phase is ClientPhase.PARKED + assert not client.status.edit_target_is_published + with pytest.raises(RuntimeError, match="no bound stage"): + client.wait_until_ready() + with pytest.raises(RuntimeError, match="no bound stage"): + client.submit_and_wait() + + +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": ""} + ) + 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): stage = Usd.Stage.CreateInMemory() @@ -363,11 +506,8 @@ def test_managed_use_server_preserves_and_clears_owned_authoring_layer(reconnect sender.recovery_incident = make_recovery_incident(artifact) sender.connect_result = reconnects client._sender = sender - client._started = True - client._refresh_recovery_checkpoint = lambda timeout: None - client._receiver.connected = True - client._receiver.layered_replay_active = True - client._receiver._synchronized_event.set() + client._replay_to_fresh_checkpoint = lambda timeout: None + force_handshake(client, synchronized=True) try: assert client.recovery_artifact is artifact result = client.recover_use_server(session_id="replacement-session") @@ -379,7 +519,7 @@ def test_managed_use_server_preserves_and_clears_owned_authoring_layer(reconnect assert sender.abandoned_session_ids == ["replacement-session"] assert sender.connected is reconnects assert sender.connect_timeouts - assert client.recovery_incident is None + assert client.status.recovery is None assert client.last_recovery_result is result client.dismiss_recovery_result() assert client.last_recovery_result is None @@ -397,10 +537,7 @@ def test_managed_client_rejects_an_edit_target_switch_before_publishing(): ) sender = _SenderStub([]) client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver.layered_replay_active = True - client._receiver._synchronized_event.set() + force_handshake(client, synchronized=True) try: stage.SetEditTarget(Usd.EditTarget(stage.GetRootLayer())) stage.DefinePrim("/World/WrongLayer", "Xform") @@ -429,11 +566,8 @@ def test_managed_use_server_refuses_an_edit_target_switch(): sender = _SenderStub([]) sender.recovery_required = True client._sender = sender - client._started = True - client._refresh_recovery_checkpoint = lambda timeout: None - client._receiver.connected = True - client._receiver.layered_replay_active = True - client._receiver._synchronized_event.set() + client._replay_to_fresh_checkpoint = lambda timeout: None + force_handshake(client, synchronized=True) try: with pytest.raises(RecoveryError, match="edit target changed") as error: client.recover_use_server() @@ -463,11 +597,8 @@ def _fail_abandonment(*, session_id=None): sender.abandon_rejected_session = _fail_abandonment client._sender = sender - client._started = True - client._refresh_recovery_checkpoint = lambda timeout: None - client._receiver.connected = True - client._receiver.layered_replay_active = True - client._receiver._synchronized_event.set() + client._replay_to_fresh_checkpoint = lambda timeout: None + force_handshake(client, synchronized=True) try: with pytest.raises(RuntimeError, match="injected abandonment failure"): client.recover_use_server() @@ -504,7 +635,7 @@ def test_managed_use_server_retains_result_when_emitter_reattach_fails(): sender.recovery_incident = make_recovery_incident(artifact) client._sender = sender client._started = True - client._refresh_recovery_checkpoint = lambda timeout: None + client._replay_to_fresh_checkpoint = lambda timeout: None def _fail_reattach(stage): raise RuntimeError("injected emitter reattach failure") @@ -540,23 +671,6 @@ def test_receiver_rejects_a_stage_that_already_contains_live_state( UsdReceiver(stage, app_name="test-receiver", persist_token=False) -def test_receiver_fails_closed_when_layered_replay_is_not_negotiated(): - receiver = UsdReceiver( - Usd.Stage.CreateInMemory(), - app_name="test-receiver", - persist_token=False, - reconnect=False, - ) - receiver._started = True - receiver._receiver.connected = True - receiver._receiver.layered_replay_active = False - - with pytest.raises(RuntimeError, match="required layered replay"): - receiver.update() - - assert not receiver.connected - - def test_publisher_retains_exact_batch_until_send_succeeds(): stage = Usd.Stage.CreateInMemory() publisher = UsdPublisher( @@ -566,6 +680,7 @@ def test_publisher_retains_exact_batch_until_send_succeeds(): ) sender = _SenderStub([False, True, True]) publisher._sender = sender + publisher.start() try: prim = stage.DefinePrim("/World/Thing", "Xform") value = prim.CreateAttribute( @@ -575,15 +690,15 @@ def test_publisher_retains_exact_batch_until_send_succeeds(): ) value.Set(1) - assert publisher.update() == 0 - assert publisher.prepared_event_count > 0 + assert publisher.update().submitted_events == 0 + assert publisher.status.prepared_events > 0 value.Set(2) - assert publisher.update() > 0 + assert publisher.update().submitted_events > 0 assert sender.batches[1] is sender.batches[0] - assert publisher.prepared_event_count == 0 + assert publisher.status.prepared_events == 0 - assert publisher.update() > 0 + assert publisher.update().submitted_events > 0 assert sender.batches[2] is not sender.batches[1] finally: publisher.close() @@ -606,6 +721,7 @@ def _publisher_with_transform(monkeypatch, results): ) sender = _SenderStub(results) publisher._sender = sender + publisher.start() return publisher, sender, translate, clock @@ -616,17 +732,18 @@ def test_publisher_coalesces_latest_default_time_transform_before_submission(mon ) try: translate.Set((1, 0, 0)) - assert publisher.update() == 3 # definition and op-order barriers publish immediately + # Definition and op-order barriers publish immediately. + assert publisher.update().submitted_events == 3 translate.Set((2, 0, 0)) - assert publisher.update() == 0 + assert publisher.update().submitted_events == 0 translate.Set((3, 0, 0)) - assert publisher.update() == 0 + assert publisher.update().submitted_events == 0 assert len(sender.batches) == 1 - assert publisher.prepared_event_count == 1 + assert publisher.status.prepared_events == 1 clock[0] = 0.11 - assert publisher.update() == 1 + assert publisher.update().submitted_events == 1 assert sender.batches[-1] == [ { "k": "set_xform_trs", @@ -646,19 +763,19 @@ def test_publisher_coalesced_batch_survives_failure_and_reconnect(monkeypatch): ) try: translate.Set((1, 0, 0)) - assert publisher.update() == 3 + assert publisher.update().submitted_events == 3 translate.Set((2, 0, 0)) - assert publisher.update() == 0 + assert publisher.update().submitted_events == 0 clock[0] = 0.11 - assert publisher.update() == 0 + assert publisher.update().submitted_events == 0 failed_batch = sender.batches[-1] sender.connected = False translate.Set((3, 0, 0)) - assert publisher.update() == 0 + assert publisher.update().submitted_events == 0 sender.connected = True - assert publisher.update() == 1 + assert publisher.update().submitted_events == 1 assert sender.batches[-1] is failed_batch assert sender.batches[-1][0]["t"] == [3.0, 0.0, 0.0] finally: @@ -672,9 +789,9 @@ def test_publisher_flush_forces_a_buffered_transform_to_sender(monkeypatch): ) try: translate.Set((1, 0, 0)) - assert publisher.update() == 3 + assert publisher.update().submitted_events == 3 translate.Set((2, 0, 0)) - assert publisher.update() == 0 + assert publisher.update().submitted_events == 0 assert publisher.flush(timeout=1.0) assert len(sender.batches) == 2 @@ -706,9 +823,7 @@ def test_managed_client_gates_new_edits_until_replay_is_applied_but_not_on_acks( ) sender = _SenderStub([True, True]) client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver.layered_replay_active = True + force_handshake(client) monkeypatch.setattr(client, "_connect_sender", lambda: None) try: prim = stage.DefinePrim("/World/Thing", "Xform") @@ -718,12 +833,12 @@ def test_managed_client_gates_new_edits_until_replay_is_applied_but_not_on_acks( replaying = client.update() assert replaying.submitted_events == 0 assert sender.batches == [] - assert not client.synchronized + assert not client.status.synchronized client._receiver._synchronized_event.set() first = client.update() assert first.submitted_events > 0 - assert client.synchronized + assert client.status.synchronized assert sender.pending_event_count == first.submitted_events value.Set(2) @@ -754,10 +869,7 @@ def test_managed_client_uses_the_same_pre_submission_transform_window(monkeypatc xformable.AddScaleOp(UsdGeom.XformOp.PrecisionDouble) sender = _SenderStub([True, True]) client._sender = sender - client._started = True - client._receiver.connected = True - client._receiver.layered_replay_active = True - client._receiver._synchronized_event.set() + force_handshake(client, synchronized=True) monkeypatch.setattr(client, "_connect_sender", lambda: None) try: translate.Set((1, 0, 0)) @@ -805,14 +917,17 @@ def test_publisher_does_not_consume_edits_while_disconnected(): sender = _SenderStub([True]) sender.connected = False publisher._sender = sender + publisher.start() try: stage.DefinePrim("/World/Thing", "Xform") - assert publisher.update() == 0 - assert publisher.prepared_event_count == 0 + assert publisher.update().submitted_events == 0 + assert publisher.status.prepared_events == 0 + assert publisher.status.has_unsent_changes + assert sender.connect_requests == 1 sender.connected = True - assert publisher.update() > 0 + assert publisher.update().submitted_events > 0 finally: publisher.close() @@ -873,10 +988,11 @@ def test_publisher_requires_a_prepared_batch_to_be_retried_before_snapshot(): persist_token=False, ) publisher._sender = _SenderStub([False]) + publisher.start() try: stage.DefinePrim("/World/Thing", "Xform") - assert publisher.update() == 0 - assert publisher.prepared_event_count > 0 + assert publisher.update().submitted_events == 0 + assert publisher.status.prepared_events > 0 with pytest.raises(RuntimeError, match=r"call update\(\)"): publisher.publish_current_edit_target() @@ -913,11 +1029,11 @@ def test_current_edit_target_publication_is_retained_for_update_retry(): publisher._sender = sender try: assert publisher.publish_current_edit_target() == 0 - assert publisher.prepared_event_count > 0 + assert publisher.status.prepared_events > 0 - assert publisher.update() > 0 + assert publisher.update().submitted_events > 0 assert sender.batches[1] is sender.batches[0] - assert publisher.prepared_event_count == 0 + assert publisher.status.prepared_events == 0 finally: publisher.close() @@ -963,3 +1079,25 @@ def test_app_name_is_required(): UsdPublisher(stage, app_name=" ", persist_token=False) with pytest.raises(ValueError, match="app_name"): 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() +