From 0385adfa5b6ec639e49a33f0d2a64f45ec26e7fb Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 4 Jun 2026 09:20:13 -0400 Subject: [PATCH 01/59] ql docs: add QLV2 overview and protocol reference --- QLV2_overview.md | 67 +++++++++ QL_V2.md | 376 +++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 443 insertions(+) create mode 100644 QLV2_overview.md create mode 100644 QL_V2.md diff --git a/QLV2_overview.md b/QLV2_overview.md new file mode 100644 index 00000000..cf267c23 --- /dev/null +++ b/QLV2_overview.md @@ -0,0 +1,67 @@ +# QuantumLink V2 + +QLV2 is designed around the shortcomings of QLV1. + +QLV1 worked, but it treated each message too much like a standalone encrypted blob. That made routing hard, pairing clunky, repeated messages expensive, reliability awkward, and application encoding too baked into the protocol. + +QLV2 moves QuantumLink toward authenticated encrypted sessions with reliable duplex byte streams. + +## Problems And Solutions + +### The protocol assumed specific product roles + +**Problem:** QLV1 was effectively shaped around the Prime/Envoy relationship. It was not a generic protocol for arbitrary peers; the wire model and pairing flow assumed who the participants were and what direction the relationship moved in. + +**Solution:** QLV2 models peers generically. A peer has a `PeerBundle` representing its public identity, including its QID, independent of whether it is Passport, Envoy, KeyOS, a mobile app, or something else. + +### Messages were not routable + +**Problem:** QLV1 messages were not cheaply introspectable. A peer or transport adapter could not look at a message and know where it was supposed to go. + +**Solution:** QLV2 records have public, but verified, routing headers. Known-peer records expose sender and recipient QIDs, and session records expose enough metadata to route them to the right session. + +This lets a single connection, like BLE, multiplex different senders and recipients on both ends. For example, iPhone/Android can have multiple apps using the same BLE connection, while KeyOS can still know which sender produced a record and which destination app should receive it. + +This also opens the door to QL-level packet forwarding, because peers can cheaply inspect a record and forward it to the right destination without needing to understand or decrypt the payload. + +### Encryption overhead was too high + +**Problem:** The minimum payload size for a QLV1 message is about 6.6KB. + +**Solution:** QLV2 amortizes the KEM and encryption setup cost into a session. Handshake messages are still large, but steady-state message overhead drops to roughly `35..42` bytes, depending on varint encoding size. + +### Key compromise could expose old messages + +**Problem:** QLV1 had no built-in key rotation model. If a long-term key was compromised, old messages were at risk. + +**Solution:** QLV2 uses Noise-style session handshakes so every session gets unique encryption keys. Compromising one session does not automatically compromise future sessions. + +### Reliability lived above the protocol + +**Problem:** QLV1 was fundamentally unreliable, so reliability had to be rebuilt by each higher-level API. Any message flow that needed dependable delivery had to invent its own retry/reliability behavior in userspace. + +**Solution:** QLV2 is built around reliable streams. One session can carry many duplex byte streams, and reliability is solved once at the QL layer instead of repeatedly in user/application space. + +### Encoding was baked into the protocol + +**Problem:** QLV1 was tied to a specific CBOR codec, even when the protocol did not really need that. This added serialization/deserialization cost and extra memory pressure. + +**Solution:** QLV2 uses binary framing internally and exposes byte streams to user space. Since QL is fundamentally moving bytes, the implementation can use zero-copy byte views where possible instead of copying payloads through a serialization layer. Applications can still layer whatever encoding they want on top: JSON, CBOR, XML, or something else. That choice belongs above QL. + +### RPC patterns had to be reinvented per API + +**Problem:** QLV1 mixed protocol transport with application workflow shape. If a feature needed request/response behavior, progress updates, downloads, uploads, or subscriptions, that behavior had to be manually modeled in its message types. The protocol did not provide reusable workflow primitives, so every feature had to encode its own control flow. + +**Solution:** QLV2 makes reliable byte streams the primitive. `ql-rpc` sits above QLV2 and gives those streams common RPC modalities: request/response, notification, upload, download, subscription, and duplex. QL stays focused on peer identity, session establishment, encryption, routing, and reliable byte transport. + +### Pairing lived in userspace + +**Problem:** QLV1 did not treat peer establishment as a protocol concern. Pairing had to be implemented as an application message flow, which made it directional, informal, and hard to generalize beyond the original Prime/Envoy convention. + +**Solution:** QLV2 lifts peer establishment into the protocol. Pairing becomes one protocol-supported way to establish a session, not a one-off userspace convention. + +QLV2 has different session establishment modalities for different starting states: + +- `XX`: first-time pairing using a small out-of-band `PairingToken` +- `IK`: the initiator already knows the responder +- `KK`: both peers already know each other diff --git a/QL_V2.md b/QL_V2.md new file mode 100644 index 00000000..0062c7e6 --- /dev/null +++ b/QL_V2.md @@ -0,0 +1,376 @@ +# QuantumLink V2 + +QuantumLink V2 is a peer-to-peer protocol for authenticated encrypted sessions carrying multiplexed duplex byte streams. + +It operates on whole QL records. Packetization, fragmentation, batching, and reassembly belong to the transport adapter, not to QLv2 itself. + +## Design goals +1. [Ephemeral peer sessions](#handshake): short-lived keys for encryption +2. [Forward secrecy](#security-properties): losing a long-term private key does not reveal old session data +3. [Minimal authenticated header](#record-and-frame-wire-format): keep routing visible, but authenticated +4. [QL-level reliability](#acknowledgment-and-retransmission): `ack` means received, decrypted, and accepted +5. [Duplex byte streams](#streams): avoid cross-stream head-of-line blocking and keep backpressure local +6. [Efficient wire format](#record-and-frame-wire-format): keep steady-state traffic compact +7. [Hardware-backed cryptography](#security-properties): allow platform-specific crypto implementations +8. Shared core state machine: keep implementation consistent across platforms + +## Non-goals + +QLv2 is not: + +- a packet framing format +- a generic reliability layer for arbitrary raw datagrams +- a globally ordered message bus + +## Core terms + +- `peer`: one QLv2 endpoint +- `QID`: a stable 16-byte peer identifier +- `peer bundle`: public peer information: `version`, `qid`, `capabilities`, and ML-KEM public key +- `pairing token`: an out-of-band secret that authorizes an `XX` pairing attempt +- `pairing_id`: the visible identifier derived from a pairing token and carried on `XX` records +- `session`: one live encrypted channel with directional keys and directional connection IDs +- `record`: one complete QLv2 wire unit +- `frame`: one logical item inside a session record +- `stream`: one duplex byte stream inside a session +- `route_id`: the application route carried once on the first initiator `StreamData` frame for a stream +- `stream origin`: the peer that opened the stream +- `origin lane`: bytes sent by the stream origin +- `return lane`: bytes sent back toward the stream origin + +## Record And Frame Wire Format + +QLv2 has two record types: + +- `handshake record`: used only during setup +- `session record`: used after the handshake completes + +Handshake records are large because they carry ML-KEM material. Session records are small and can carry multiple frames, including frames for different streams. + +All whole-record sizes below include the outer 2-byte record header: `version` plus `record type`. + +QLv2 uses QUIC-style variable-length integers for several steady-state fields. A varint is 1, 2, 4, or 8 bytes and can represent values in the range `0..2^62-1`. This keeps small values compact while allowing very large record and stream number spaces. + +Today, varints are used for: + +- session record `seq` +- `Ack.largest_acked` +- `Ack.block_count` +- `Ack.first_range_len` +- `Ack.gap` +- `Ack.range_len` +- `StreamData.stream_id` +- `StreamData.offset` +- `StreamData.route_id` when present +- `StreamData.bytes_len` +- `StreamWindow.stream_id` +- `StreamWindow.maximum_offset` +- `StreamClose.stream_id` + +### Handshake records + +QLv2 has two routed known-peer handshakes and one pairing handshake: + +- `IK` and `KK` carry a visible `sender` and `recipient` QID +- `XX` carries a visible `pairing_id` + +#### IK + +Used when the initiator already knows the responder bundle. + +| Record | Size | Purpose | +| --- | ---: | --- | +| `IK1` | 4785 bytes | start a handshake toward a known responder | +| `IK2` | 3195 bytes | complete `IK` and establish the session | + +#### KK + +Used when both peers already know each other. + +| Record | Size | Purpose | +| --- | ---: | --- | +| `KK1` | 3179 bytes | start a handshake between already-known peers | +| `KK2` | 3195 bytes | complete `KK` and establish the session | + +#### XX + +Used when the initiator has received an out of band pairing token, and neither peer knows each other. + +| Record | Size | Purpose | +| --- | ---: | --- | +| `XX1` | 1595 bytes | start pairing | +| `XX2` | 3201 bytes | send responder static identity and ciphertext | +| `XX3` | 3217 bytes | send initiator static identity and ciphertext | +| `XX4` | 1611 bytes | complete `XX` and establish the session | + +### Session records + +`session record size = 35..42 + sum(frame sizes)` + +There is no explicit AEAD nonce on the wire. The record `seq` is used to derive the nonce. + +| Fixed part | Size | Purpose | +| --- | ---: | --- | +| version | 1 byte | protocol version | +| record type | 1 byte | identifies a session record | +| `connection_id` | 16 bytes | route the record to the current session | +| `seq` | 1..8 bytes | varint record identity for ack and retransmit | +| AEAD auth tag | 16 bytes | authenticate the encrypted body | +| fixed overhead total | 35..42 bytes | overhead before any frames | + +The visible session header is authenticated as AEAD AAD but is not encrypted. + +### Session frames + +| Frame | Size | Purpose | +| --- | ---: | --- | +| `Ping` | 1 byte | keep the session alive when idle | +| `Unpair` | 1 byte | forget the currently bound peer and abort the session | +| `Ack` | `4+` bytes | acknowledge received session records with ACK ranges | +| `StreamWindow` | `3..17` bytes | extend per-stream send credit | +| `StreamClose` | `5..12` bytes | abort one stream lane or both lanes | +| `Close` | 3 bytes | close the whole session | +| `StreamData` | `5..34 + payload_len` bytes | carry stream bytes, optional opener route, and optional `fin` | + +`StreamData` is the main steady-state frame: + +`1 kind + varint(stream_id) + varint(offset) + 1 flags + optional varint(route_id) + varint(bytes_len) + payload_len` + +The flags byte carries: + +- `fin` +- `header present` + +Some useful minimum whole-record sizes for single-frame records: + +| Record | Size | Meaning | +| --- | ---: | --- | +| `Ping` only | 36 bytes | idle keepalive | +| `Unpair` only | 36 bytes | peer unpair | +| `Ack` only | 39 bytes | smallest selective ack | +| `Close` only | 38 bytes | session shutdown | +| empty `StreamData` without route header | 40 bytes | empty data or empty `fin` on an existing stream | +| empty opener `StreamData` with a 1-byte `route_id` | 41 bytes | open a new stream without payload bytes | + +## Handshake + +QLv2 currently supports three Noise-style handshake patterns: + +- `IK`: 2 messages, initiator already knows the responder bundle +- `KK`: 2 messages, both peers already know each other +- `XX`: 4 messages, peers authenticate through an out-of-band pairing token and exchange static identity during the handshake + +The handshake covers peer authentication and session establishment. + +Each successful handshake does five things: + +1. authenticate which peer we are talking to +2. derive a fresh transmit key and receive key +3. derive a directional transmit `connection_id` and receive `connection_id` +4. bind transport parameters into the transcript +5. produce a `handshake_hash` for the completed exchange + +Today the only transport parameter is: + +- initial per-stream receive window + +Future transport parameters could include session-wide byte credit or record-size limits. + +Each handshake attempt carries: + +- `handshake_id`: identifies one attempt and lets stale replies be ignored +- transport parameters + +`valid_until` is not currently part of the wire format. Handshake attempts instead expire by local timer. + +### Pattern summary + +- `IK` lets the responder learn the initiator during handshake completion. The initiator still needs the responder bundle before it can start. +- `KK` requires both peers to already know each other. +- `XX` requires the responder to be armed for pairing and to recognize the visible `pairing_id` derived from the expected pairing token. + +### Handshake rules + +- attempts are identified by `handshake_id` +- handshake messages are not retransmitted in place +- simultaneous starts must converge deterministically +- if `IK` and `KK` race, `IK` wins +- same-pattern races break ties by ordering the initial ephemeral public keys +- `XX` requires out-of-band authorization and uses visible `pairing_id` for lookup + +### Session establishment points + +- `IK` and `KK` complete after message 2 (1 RT) +- `XX` completes after 4 messages (2 RTT) + +## Session Model + +After the handshake, peers exchange encrypted session records. + +Each session record has: + +- one visible `connection_id` +- one visible `seq` +- one encrypted body containing one or more frames + +One session record may carry: + +- only control frames +- only stream data +- a mixture of frames for multiple streams + +This is the core steady-state model: records are the encrypted transport unit, frames are the logical items inside them. + +## Acknowledgment And Retransmission + +`Ack` is record-level, not stream-level. + +An `Ack` means the peer: + +- received that session record +- decrypted it with the current session key +- accepted its `seq` + +The ACK wire format is range-based, not bitmap-based. It carries: + +- `largest_acked` +- `block_count` +- `first_range_len` +- zero or more `(gap, range_len)` blocks + +Ranges are encoded from highest sequence numbers down to lowest sequence numbers. + +Receivers track a recent accepted record window so they can: + +- reject duplicates +- ignore records that are too old +- emit selective ACK ranges + +Pending ACK state is also range-based. If there are too many disjoint ranges, older low ranges may be dropped. An emitted ACK may also be truncated by the remaining record budget. + +Retransmission works at the frame level: + +- every emitted session record gets a fresh `seq` +- retransmit timers start only after the local transport confirms that it accepted the write +- if a record is considered lost, the FSM restores its frames +- those frames are packed into a new record with a new `seq` + +QLv2 does not resend the same logical record identity. + +There is no explicit `Nack` frame. Loss is inferred from timeout or from later ACK state that no longer includes a record. + +Pure ACK-only records are fire-and-forget: they are not themselves retransmitted. + +Example: + +`seq = 10` + +| Frame | Contents | +| --- | --- | +| `StreamData` | `stream_id=4 offset=0 bytes="hello"` | + +The sender receives more bytes for that stream before `seq = 10` is acked: + +| Pending new frame | Contents | +| --- | --- | +| `StreamData` | `stream_id=4 offset=5 bytes=" world"` | + +If `seq = 10` is considered lost, its frame is restored and packed again with a new record sequence: + +`seq = 11` + +| Frame | Contents | +| --- | --- | +| `StreamData` | `stream_id=4 offset=0 bytes="hello"` | +| `StreamData` | `stream_id=4 offset=5 bytes=" world"` | + +## Streams + +Streams are the application primitive. + +A stream has two independent lanes: + +- origin lane +- return lane + +Important properties: + +- either peer can open a stream +- stream IDs are split by parity derived from QID ordering, so both peers can open streams without collision +- stream IDs increase monotonically within each parity namespace and must not repeat within a session +- ordering is preserved within a stream lane +- different streams can make progress independently +- record loss on one stream does not block unrelated streams + +There is no separate open frame. + +Locally, opening a stream allocates: + +- a new `stream_id` +- an application `route_id` + +On the wire, the stream opener carries that `route_id` once, in the first initiator `StreamData` frame at `offset = 0`, using the optional `StreamHeader`. + +`StreamData` carries: + +- `stream_id` +- `offset` +- optional `StreamHeader { route_id }` +- `fin` +- bytes + +`StreamHeader` is only valid on the first initiator `StreamData` frame for a stream, at `offset = 0`. + +`fin` is graceful completion of one lane. It says "no more bytes on this lane" without aborting the other lane. + +## Flow Control + +Flow control is per stream. + +During the handshake, each peer advertises an initial per-stream receive window. That becomes the initial send credit the remote peer can use on each stream. + +`StreamWindow` extends that credit by advertising a larger absolute `maximum_offset`. + +In practice, a stream is writable only when both are true: + +- local send buffering has room +- peer-advertised stream credit allows more bytes + +Receive credit advances when the local application commits read bytes, not merely when bytes become readable. That is when the FSM emits a `StreamWindow` update. + +## Close And Liveness + +`StreamClose` aborts a stream early. Semantically it can target: + +- the origin lane +- the return lane +- both lanes + +`Close` aborts the whole session. + +`Unpair` is stronger than `Close`: + +- it forgets the currently bound peer locally +- it aborts the active session immediately +- it may emit one final outbound `Unpair` frame +- reconnect does not resume until a peer is paired again + +Idle sessions may send `Ping`. The peer does not answer with another ping; normal record acknowledgment is enough. + +Sessions also have local timers for: + +- handshake timeout +- delayed ack emission +- session record retransmit timeout +- keepalive ping interval +- peer silence timeout + +If peer silence exceeds the configured timeout, the session closes with timeout. + +## Security Properties + +The current handshake family is ML-KEM-based and post-quantum focused. + +Session payloads are encrypted and authenticated. The session header stays visible so the receiver can route the record, but it is still authenticated as AEAD AAD. + +QLv2 also provides forward secrecy in the following sense: even if an attacker later obtains a peer's long-term ML-KEM private key, they still cannot decrypt messages from earlier completed sessions. From 71702fbfe3489494f8ba2d44d566f0a0b5060f0e Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 4 Jun 2026 09:22:32 -0400 Subject: [PATCH 02/59] ql-wire: add wire-format definitions --- Cargo.lock | 197 ++++ Cargo.toml | 10 +- ql-wire/Cargo.toml | 27 + ql-wire/src/bytes.rs | 175 ++++ ql-wire/src/codec.rs | 245 +++++ ql-wire/src/crypto.rs | 47 + ql-wire/src/encrypted/ack.rs | 453 +++++++++ ql-wire/src/encrypted/builder.rs | 172 ++++ ql-wire/src/encrypted/close.rs | 55 ++ ql-wire/src/encrypted/mod.rs | 178 ++++ ql-wire/src/encrypted/route_id.rs | 55 ++ ql-wire/src/encrypted/stream_close.rs | 116 +++ ql-wire/src/encrypted/stream_data.rs | 135 +++ ql-wire/src/encrypted/stream_id.rs | 35 + ql-wire/src/encrypted/stream_window.rs | 29 + ql-wire/src/encrypted_message.rs | 97 ++ ql-wire/src/error.rs | 33 + ql-wire/src/handshake/ik.rs | 376 ++++++++ ql-wire/src/handshake/kk.rs | 352 +++++++ ql-wire/src/handshake/meta.rs | 48 + ql-wire/src/handshake/mod.rs | 590 ++++++++++++ ql-wire/src/handshake/pairing.rs | 83 ++ ql-wire/src/handshake/transport_params.rs | 38 + ql-wire/src/handshake/xx.rs | 612 +++++++++++++ ql-wire/src/header.rs | 121 +++ ql-wire/src/identity.rs | 192 ++++ ql-wire/src/lib.rs | 45 + ql-wire/src/nonce.rs | 13 + ql-wire/src/pq.rs | 159 ++++ ql-wire/src/qid.rs | 44 + ql-wire/src/record.rs | 254 +++++ ql-wire/src/testing.rs | 181 ++++ ql-wire/src/tests.rs | 1017 +++++++++++++++++++++ ql-wire/src/varint.rs | 181 ++++ 34 files changed, 6364 insertions(+), 1 deletion(-) create mode 100644 ql-wire/Cargo.toml create mode 100644 ql-wire/src/bytes.rs create mode 100644 ql-wire/src/codec.rs create mode 100644 ql-wire/src/crypto.rs create mode 100644 ql-wire/src/encrypted/ack.rs create mode 100644 ql-wire/src/encrypted/builder.rs create mode 100644 ql-wire/src/encrypted/close.rs create mode 100644 ql-wire/src/encrypted/mod.rs create mode 100644 ql-wire/src/encrypted/route_id.rs create mode 100644 ql-wire/src/encrypted/stream_close.rs create mode 100644 ql-wire/src/encrypted/stream_data.rs create mode 100644 ql-wire/src/encrypted/stream_id.rs create mode 100644 ql-wire/src/encrypted/stream_window.rs create mode 100644 ql-wire/src/encrypted_message.rs create mode 100644 ql-wire/src/error.rs create mode 100644 ql-wire/src/handshake/ik.rs create mode 100644 ql-wire/src/handshake/kk.rs create mode 100644 ql-wire/src/handshake/meta.rs create mode 100644 ql-wire/src/handshake/mod.rs create mode 100644 ql-wire/src/handshake/pairing.rs create mode 100644 ql-wire/src/handshake/transport_params.rs create mode 100644 ql-wire/src/handshake/xx.rs create mode 100644 ql-wire/src/header.rs create mode 100644 ql-wire/src/identity.rs create mode 100644 ql-wire/src/lib.rs create mode 100644 ql-wire/src/nonce.rs create mode 100644 ql-wire/src/pq.rs create mode 100644 ql-wire/src/qid.rs create mode 100644 ql-wire/src/record.rs create mode 100644 ql-wire/src/testing.rs create mode 100644 ql-wire/src/tests.rs create mode 100644 ql-wire/src/varint.rs diff --git a/Cargo.lock b/Cargo.lock index f144305f..c2c3b23c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -532,6 +532,17 @@ version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" +[[package]] +name = "core-models" +version = "0.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "657f625ff361906f779745d08375ae3cc9fef87a35fba5f22874cf773010daf4" +dependencies = [ + "hax-lib", + "pastey", + "rand 0.9.2", +] + [[package]] name = "cpufeatures" version = "0.2.17" @@ -1081,6 +1092,43 @@ version = "0.16.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" +[[package]] +name = "hax-lib" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "543f93241d32b3f00569201bfce9d7a93c92c6421b23c77864ac929dc947b9fc" +dependencies = [ + "hax-lib-macros", + "num-bigint", + "num-traits", +] + +[[package]] +name = "hax-lib-macros" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8755751e760b11021765bb04cb4a6c4e24742688d9f3aa14c2079638f537b0f" +dependencies = [ + "hax-lib-macros-types", + "proc-macro-error2", + "proc-macro2", + "quote", + "syn 2.0.106", +] + +[[package]] +name = "hax-lib-macros-types" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f177c9ae8ea456e2f71ff3c1ea47bf4464f772a05133fcbba56cd5ba169035a2" +dependencies = [ + "proc-macro2", + "quote", + "serde", + "serde_json", + "uuid", +] + [[package]] name = "hermit-abi" version = "0.5.2" @@ -1356,6 +1404,84 @@ version = "0.2.175" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6a82ae493e598baaea5209805c49bbf2ea7de956d50d7da0da1164f9c6d28543" +[[package]] +name = "libcrux-aesgcm" +version = "0.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "99f2a019dab4097585a7d4f5b9deebe46cd1e628b16a5bc4cb0ce35e1da334e6" +dependencies = [ + "libcrux-intrinsics", + "libcrux-platform", + "libcrux-secrets", + "libcrux-traits", +] + +[[package]] +name = "libcrux-intrinsics" +version = "0.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b1b5db005ff8001e026b73a6842ee81bbef8ec5ff0e1915a67ae65fd2a9fafa5" +dependencies = [ + "core-models", + "hax-lib", +] + +[[package]] +name = "libcrux-ml-kem" +version = "0.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "aca7de713c6dddcf7aaf76e8ef9dc0097c8d7ce23a8eadf04c8761734714e184" +dependencies = [ + "hax-lib", + "libcrux-intrinsics", + "libcrux-platform", + "libcrux-secrets", + "libcrux-sha3", + "libcrux-traits", + "rand 0.9.2", + "tls_codec", +] + +[[package]] +name = "libcrux-platform" +version = "0.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d9e21d7ed31a92ac539bd69a8c970b183ee883872d2d19ce27036e24cb8ecc4" +dependencies = [ + "libc", +] + +[[package]] +name = "libcrux-secrets" +version = "0.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ce650f3041b44ba40d4263852347d007cd2cd9d1cc856a6f6c8b2e10c3fd40b" +dependencies = [ + "hax-lib", +] + +[[package]] +name = "libcrux-sha3" +version = "0.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8c50f6e04a184511b782c5cc1eb6a227c6d36f2c935e93d698655a93a99696b5" +dependencies = [ + "hax-lib", + "libcrux-intrinsics", + "libcrux-platform", + "libcrux-traits", +] + +[[package]] +name = "libcrux-traits" +version = "0.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "812e4fa89f3f5e34b47f928b22b1b78395a0d4ec23b1f583db635f128159d65f" +dependencies = [ + "libcrux-secrets", + "rand 0.9.2", +] + [[package]] name = "libm" version = "0.2.15" @@ -1469,6 +1595,16 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "num-bigint" +version = "0.4.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a5e44f723f1133c9deac646763579fdb3ac745e418f2a7af9cd0c431da1f20b9" +dependencies = [ + "num-integer", + "num-traits", +] + [[package]] name = "num-bigint-dig" version = "0.8.6" @@ -1625,6 +1761,12 @@ version = "1.0.15" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" +[[package]] +name = "pastey" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4" + [[package]] name = "pbkdf2" version = "0.12.2" @@ -1810,6 +1952,28 @@ dependencies = [ "elliptic-curve", ] +[[package]] +name = "proc-macro-error-attr2" +version = "2.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96de42df36bb9bba5542fe9f1a054b8cc87e172759a1868aa05c1f3acc89dfc5" +dependencies = [ + "proc-macro2", + "quote", +] + +[[package]] +name = "proc-macro-error2" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11ec05c52be0a07b08061f7dd003e7d7092e0472bc731b4af7bb1ef876109802" +dependencies = [ + "proc-macro-error-attr2", + "proc-macro2", + "quote", + "syn 2.0.106", +] + [[package]] name = "proc-macro2" version = "1.0.101" @@ -1863,6 +2027,17 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "ql-wire" +version = "0.1.0" +dependencies = [ + "bytes", + "getrandom 0.2.16", + "libcrux-aesgcm", + "libcrux-ml-kem", + "sha2", +] + [[package]] name = "quantum-link-macros" version = "0.1.0" @@ -2436,6 +2611,27 @@ version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" +[[package]] +name = "tls_codec" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0de2e01245e2bb89d6f05801c564fa27624dbd7b1846859876c7dad82e90bf6b" +dependencies = [ + "tls_codec_derive", + "zeroize", +] + +[[package]] +name = "tls_codec_derive" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d2e76690929402faae40aebdda620a2c0e25dd6d3b9afe48867dfd95991f4bd" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.106", +] + [[package]] name = "tokio" version = "1.47.1" @@ -2517,6 +2713,7 @@ version = "1.18.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2f87b8aa10b915a06587d0dec516c282ff295b475d94abf425d62b57710070a2" dependencies = [ + "getrandom 0.3.3", "js-sys", "wasm-bindgen", ] diff --git a/Cargo.toml b/Cargo.toml index 0fd0e755..8aad9104 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,12 @@ [workspace] resolver = "2" -members = ["api", "backup-shard", "btp", "quantum-link-macros"] +members = [ + "api", + "backup-shard", + "btp", + "ql-wire", + "quantum-link-macros", +] [workspace.package] homepage = "https://github.com/Foundation-Devices/foundation-api" @@ -14,6 +20,7 @@ dcbor = { version = "0.23.3" } gstp = { version = "0.11.0" } chrono = "0.4" +bytes = "1" getrandom = { version = "0.2" } insta = { version = "1.43.2" } thiserror = { version = "2" } @@ -24,6 +31,7 @@ backup-shard = { path = "backup-shard" } btp = { path = "btp" } foundation-api = { path = "api" } quantum-link-macros = { path = "quantum-link-macros" } +ql-wire = { path = "ql-wire" } [patch.crates-io] pqcrypto-traits = { git = "https://github.com/Foundation-Devices/pqcrypto", rev = "ebadf71214f67cb970242fa1053b4acb65767737" } diff --git a/ql-wire/Cargo.toml b/ql-wire/Cargo.toml new file mode 100644 index 00000000..399846cc --- /dev/null +++ b/ql-wire/Cargo.toml @@ -0,0 +1,27 @@ +[package] +name = "ql-wire" +version = "0.1.0" +edition = "2021" +description = "QuantumLink protocol wire format" +license = "Proprietary" + +[features] +test-utils = [ + "dep:getrandom", + "dep:libcrux-aesgcm", + "dep:libcrux-ml-kem", + "dep:sha2", +] + +[dependencies] +bytes = { workspace = true } +getrandom = { workspace = true, optional = true } +libcrux-aesgcm = { version = "0.0.7", optional = true } +libcrux-ml-kem = { version = "0.0.7", optional = true } +sha2 = { version = "0.10", optional = true } + +[dev-dependencies] +getrandom = { workspace = true } +libcrux-aesgcm = "0.0.7" +libcrux-ml-kem = "0.0.7" +sha2 = "0.10" diff --git a/ql-wire/src/bytes.rs b/ql-wire/src/bytes.rs new file mode 100644 index 00000000..9fecf5ea --- /dev/null +++ b/ql-wire/src/bytes.rs @@ -0,0 +1,175 @@ +use core::ops::{Deref, DerefMut}; + +use bytes::{Buf, Bytes}; + +/// A mutable or immutable byte slice owner used by the wire parser. +pub trait ByteSlice: Deref + Sized { + /// Splits the current byte view at `mid`. + /// + /// Returns `Err(self)` when `mid` is out of bounds. + fn split_at(self, mid: usize) -> Result<(Self, Self), Self>; +} + +/// A mutable reference to bytes. +pub trait ByteSliceMut: ByteSlice + DerefMut {} + +impl ByteSliceMut for B where B: ByteSlice + DerefMut {} + +impl ByteSlice for &[u8] { + #[inline] + fn split_at(self, mid: usize) -> Result<(Self, Self), Self> { + if mid <= self.len() { + Ok(<[u8]>::split_at(self, mid)) + } else { + Err(self) + } + } +} + +impl ByteSlice for &mut [u8] { + #[inline] + fn split_at(self, mid: usize) -> Result<(Self, Self), Self> { + if mid <= self.len() { + Ok(<[u8]>::split_at_mut(self, mid)) + } else { + Err(self) + } + } +} + +impl ByteSlice for Bytes { + #[inline] + fn split_at(self, mid: usize) -> Result<(Self, Self), Self> { + if mid <= self.len() { + Ok((self.slice(..mid), self.slice(mid..))) + } else { + Err(self) + } + } +} + +/// A byte container that can expose a replayable [`Buf`] view for encoding. +pub trait BufView { + type Buf<'a>: Buf + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_>; + + fn is_empty(&self) -> bool { + self.buf().remaining() == 0 + } +} + +impl BufView for &T { + type Buf<'a> + = T::Buf<'a> + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + (*self).buf() + } +} + +impl BufView for &mut T { + type Buf<'a> + = T::Buf<'a> + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + (**self).buf() + } +} + +impl BufView for [u8] { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self + } +} + +impl BufView for [u8; N] { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self.as_slice() + } +} + +impl BufView for Vec { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self.as_slice() + } +} + +impl BufView for Bytes { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self.as_ref() + } +} + +#[cfg(test)] +mod tests { + use bytes::Buf; + + use super::{BufView, ByteSlice, ByteSliceMut}; + + #[test] + fn shared_slice_split_at() { + let bytes: &[u8] = b"abcdef"; + let (left, right) = ByteSlice::split_at(bytes, 2).unwrap(); + assert_eq!(left, b"ab"); + assert_eq!(right, b"cdef"); + } + + #[test] + fn mutable_slice_split_at() { + let mut bytes = *b"abcdef"; + let (left, right) = ByteSlice::split_at(&mut bytes[..], 2).unwrap(); + assert_eq!(left, b"ab"); + assert_eq!(right, b"cdef"); + } + + #[test] + fn mutable_split_trait_is_implemented() { + fn assert_split_mut(_value: T) {} + + let mut bytes = [0u8; 4]; + assert_split_mut(&mut bytes[..]); + } + + #[test] + fn split_at_rejects_out_of_bounds_index() { + let bytes: &[u8] = b"abcdef"; + assert!(ByteSlice::split_at(bytes, 7).is_err()); + } + + #[test] + fn slice_buf_view_is_contiguous() { + let bytes: &[u8] = b"abcdef"; + let mut buf = bytes.buf(); + assert_eq!(buf.remaining(), 6); + assert_eq!(buf.chunk(), b"abcdef"); + buf.advance(6); + assert!(!buf.has_remaining()); + } +} diff --git a/ql-wire/src/codec.rs b/ql-wire/src/codec.rs new file mode 100644 index 00000000..0245ef6d --- /dev/null +++ b/ql-wire/src/codec.rs @@ -0,0 +1,245 @@ +use bytes::BufMut; + +use crate::{ByteSlice, WireError}; + +pub trait WireEncode { + fn encoded_len(&self) -> usize; + + fn encode(&self, out: &mut W); + + fn encode_vec(&self) -> Vec { + let mut out = Vec::with_capacity(self.encoded_len()); + self.encode(&mut out); + debug_assert_eq!(out.len(), self.encoded_len()); + out + } +} + +pub trait WireDecode: Sized { + fn decode(reader: &mut Reader) -> Result; + + fn decode_bytes(bytes: B) -> Result { + let mut reader = Reader::new(bytes); + Self::decode(&mut reader) + } + + fn decode_exact(bytes: B) -> Result { + let mut reader = Reader::new(bytes); + let value = Self::decode(&mut reader)?; + if reader.is_empty() { + Ok(value) + } else { + Err(WireError::InvalidPayload) + } + } +} + +impl WireDecode for [u8; N] { + fn decode(reader: &mut Reader) -> Result { + let bytes = reader.take_bytes(N)?; + let mut out = [0u8; N]; + out.copy_from_slice(&bytes); + Ok(out) + } +} + +impl WireEncode for [u8; N] { + fn encoded_len(&self) -> usize { + N + } + + fn encode(&self, out: &mut W) { + out.put_slice(self); + } +} + +impl WireDecode for Box<[u8; N]> { + fn decode(reader: &mut Reader) -> Result { + let bytes = reader.take_bytes(N)?; + let mut out = Self::new_uninit(); + let src = bytes.as_ptr(); + let dst = out.as_mut_ptr().cast::(); + // SAFETY: `take_bytes(N)` guarantees the source has exactly `N` bytes. + unsafe { + std::ptr::copy_nonoverlapping(src, dst, N); + Ok(out.assume_init()) + } + } +} + +impl WireEncode for Box<[u8; N]> { + fn encoded_len(&self) -> usize { + N + } + + fn encode(&self, out: &mut W) { + out.put_slice(self.as_ref()); + } +} + +impl WireEncode for [u8] { + fn encoded_len(&self) -> usize { + self.len() + } + + fn encode(&self, out: &mut W) { + out.put_slice(self); + } +} + +impl WireDecode for u8 { + fn decode(reader: &mut Reader) -> Result { + Ok(reader.take_bytes(1)?[0]) + } +} + +impl WireEncode for u8 { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u8(*self); + } +} + +impl WireDecode for u16 { + fn decode(reader: &mut Reader) -> Result { + Ok(Self::from_be_bytes(reader.decode()?)) + } +} + +impl WireEncode for u16 { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u16(*self); + } +} + +impl WireDecode for u32 { + fn decode(reader: &mut Reader) -> Result { + Ok(Self::from_be_bytes(reader.decode()?)) + } +} + +impl WireEncode for u32 { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u32(*self); + } +} + +impl WireDecode for u64 { + fn decode(reader: &mut Reader) -> Result { + Ok(Self::from_be_bytes(reader.decode()?)) + } +} + +impl WireEncode for u64 { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u64(*self); + } +} + +impl WireDecode for bool { + fn decode(reader: &mut Reader) -> Result { + match reader.decode::()? { + 0 => Ok(false), + 1 => Ok(true), + _ => Err(WireError::InvalidPayload), + } + } +} + +impl WireEncode for bool { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u8(u8::from(*self)); + } +} + +impl WireEncode for Option { + fn encoded_len(&self) -> usize { + 1 + self.as_ref().map_or(0, WireEncode::encoded_len) + } + + fn encode(&self, out: &mut W) { + match self { + None => out.put_u8(0), + Some(inner) => { + out.put_u8(1); + inner.encode(out); + } + } + } +} + +impl> WireDecode for Option { + fn decode(reader: &mut Reader) -> Result { + match reader.decode::()? { + 0 => Ok(None), + 1 => Ok(Some(reader.decode::()?)), + _ => Err(WireError::InvalidPayload), + } + } +} + +#[derive(Clone)] +pub struct Reader { + remaining: Option, +} + +impl Reader { + pub fn new(bytes: B) -> Self { + Self { + remaining: Some(bytes), + } + } + + pub fn is_empty(&self) -> bool { + self.remaining.as_ref().unwrap().is_empty() + } + + pub fn remaining_len(&self) -> usize { + self.remaining.as_ref().unwrap().len() + } + + pub fn take_bytes(&mut self, len: usize) -> Result { + let remaining = self.remaining.take().unwrap(); + match remaining.split_at(len) { + Ok((head, tail)) => { + self.remaining = Some(tail); + Ok(head) + } + Err(remaining) => { + self.remaining = Some(remaining); + Err(WireError::InvalidPayload) + } + } + } + + pub fn take_rest(&mut self) -> B { + self.take_bytes(self.remaining_len()).unwrap() + } + + #[inline] + pub fn decode(&mut self) -> Result + where + T: WireDecode, + { + T::decode(self) + } +} diff --git a/ql-wire/src/crypto.rs b/ql-wire/src/crypto.rs new file mode 100644 index 00000000..96ace383 --- /dev/null +++ b/ql-wire/src/crypto.rs @@ -0,0 +1,47 @@ +use crate::{ + MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, Nonce, SessionKey, + ENCRYPTED_MESSAGE_AUTH_SIZE, +}; + +pub trait QlRandom { + fn fill_random_bytes(&self, out: &mut [u8]); +} + +pub trait QlHash { + fn sha256(&self, parts: &[&[u8]]) -> [u8; 32]; +} + +pub trait QlAead { + fn aes256_gcm_encrypt( + &self, + key: &SessionKey, + nonce: &Nonce, + aad: &[u8], + buffer: &mut [u8], + ) -> [u8; ENCRYPTED_MESSAGE_AUTH_SIZE]; + + fn aes256_gcm_decrypt( + &self, + key: &SessionKey, + nonce: &Nonce, + aad: &[u8], + buffer: &mut [u8], + auth_tag: &[u8; ENCRYPTED_MESSAGE_AUTH_SIZE], + ) -> bool; +} + +pub trait QlKem { + fn mlkem_generate_keypair(&self) -> MlKemKeyPair; + + fn mlkem_encapsulate(&self, public_key: &MlKemPublicKey) -> (MlKemCiphertext, SessionKey); + + fn mlkem_decapsulate( + &self, + private_key: &MlKemPrivateKey, + ciphertext: &MlKemCiphertext, + ) -> SessionKey; +} + +pub trait QlCrypto: QlRandom + QlHash + QlAead + QlKem {} + +impl QlCrypto for T where T: QlRandom + QlHash + QlAead + QlKem {} diff --git a/ql-wire/src/encrypted/ack.rs b/ql-wire/src/encrypted/ack.rs new file mode 100644 index 00000000..2eb34b33 --- /dev/null +++ b/ql-wire/src/encrypted/ack.rs @@ -0,0 +1,453 @@ +use std::{fmt, ops::RangeInclusive}; + +use crate::{codec, ByteSlice, RecordSeq, VarInt, WireEncode, WireError}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct RecordAck { + largest_acked: RecordSeq, + first_range_len: VarInt, + blocks: Box<[RecordAckBlock]>, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RecordAckBlock { + pub gap: VarInt, + pub range_len: VarInt, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum RecordAckRangeError { + Empty, + InvertedRange, + NotCanonical, +} + +impl RecordAck { + /// Build a record ACK from canonical ranges ordered from highest to lowest sequence number. + /// + /// Ranges must be: + /// - non-empty + /// - individually valid (`start <= end`) + /// - strictly descending + /// - separated by at least one missing sequence number + pub fn from_ranges(ranges: I) -> Result + where + I: IntoIterator>, + { + let mut builder = RecordAckBuilder::new(); + for range in ranges { + let pushed = builder.try_push_range(range, usize::MAX)?; + if !pushed { + unreachable!("record ack should fit inside usize::MAX"); + } + } + builder.build() + } + + pub fn ranges(&self) -> RecordAckRangeIter<'_> { + RecordAckRangeIter { + largest_acked: self.largest_acked.into_inner(), + first_range_len: Some(self.first_range_len), + previous_start: None, + blocks: self.blocks.iter(), + } + } + + pub fn contains(&self, seq: u64) -> bool { + let Ok(seq) = RecordSeq::from_u64(seq) else { + return false; + }; + self.ranges().any(|range| range.contains(&seq)) + } + + fn block_count_len(block_count: usize) -> usize { + VarInt::try_from(block_count).unwrap().encoded_len() + } +} + +impl RecordAckBlock { + fn encoded_len(&self) -> usize { + self.gap.encoded_len() + self.range_len.encoded_len() + } +} + +impl fmt::Display for RecordAckRangeError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Empty => f.write_str("record ack requires at least one acknowledged range"), + Self::InvertedRange => { + f.write_str("record ack range start must be less than or equal to end") + } + Self::NotCanonical => f.write_str( + "record ack ranges must be passed in descending, disjoint order with a gap between adjacent ranges", + ), + } + } +} + +impl std::error::Error for RecordAckRangeError {} + +pub struct RecordAckRangeIter<'a> { + largest_acked: u64, + first_range_len: Option, + previous_start: Option, + blocks: std::slice::Iter<'a, RecordAckBlock>, +} + +impl Iterator for RecordAckRangeIter<'_> { + type Item = RangeInclusive; + + fn next(&mut self) -> Option { + if let Some(first_range_len) = self.first_range_len.take() { + let end = self.largest_acked; + let start = end - first_range_len.into_inner(); + self.previous_start = Some(start); + return Some(RecordSeq::from_u64(start).unwrap()..=RecordSeq::from_u64(end).unwrap()); + } + + let block = self.blocks.next()?; + let previous_start = self + .previous_start + .expect("first ack range is always yielded"); + // gap is encoded as missing_count - 1, so decoding steps back by gap + 2. + let end = previous_start - block.gap.into_inner() - 2; + let start = end - block.range_len.into_inner(); + self.previous_start = Some(start); + Some(RecordSeq::from_u64(start).unwrap()..=RecordSeq::from_u64(end).unwrap()) + } +} + +impl WireEncode for RecordAck { + fn encoded_len(&self) -> usize { + self.largest_acked.encoded_len() + + Self::block_count_len(self.blocks.len()) + + self.first_range_len.encoded_len() + + self + .blocks + .iter() + .map(RecordAckBlock::encoded_len) + .sum::() + } + + fn encode(&self, out: &mut W) { + self.largest_acked.encode(out); + VarInt::try_from(self.blocks.len()).unwrap().encode(out); + self.first_range_len.encode(out); + for block in &self.blocks { + block.gap.encode(out); + block.range_len.encode(out); + } + } +} + +impl codec::WireDecode for RecordAck { + fn decode(reader: &mut codec::Reader) -> Result { + let largest_acked = reader.decode()?; + let block_count = usize::try_from(reader.decode::()?.into_inner()) + .map_err(|_| WireError::InvalidPayload)?; + let first_range_len = reader.decode::()?; + let mut blocks = Vec::with_capacity(block_count); + for _ in 0..block_count { + blocks.push(RecordAckBlock { + gap: reader.decode::()?, + range_len: reader.decode::()?, + }); + } + + let ack = Self { + largest_acked, + first_range_len, + blocks: blocks.into_boxed_slice(), + }; + + // validate + { + let mut previous_start = ack + .largest_acked + .into_inner() + .checked_sub(ack.first_range_len.into_inner()) + .ok_or(WireError::InvalidPayload)?; + + for block in &ack.blocks { + let end = previous_start + .checked_sub( + block + .gap + .into_inner() + .checked_add(2) + .ok_or(WireError::InvalidPayload)?, + ) + .ok_or(WireError::InvalidPayload)?; + previous_start = end + .checked_sub(block.range_len.into_inner()) + .ok_or(WireError::InvalidPayload)?; + } + } + Ok(ack) + } +} + +#[derive(Debug, Clone, Default, PartialEq, Eq)] +pub struct RecordAckBuilder { + largest_acked: Option, + first_range_len: Option, + blocks: Vec, + previous_start: Option, + wire_len: usize, +} + +impl RecordAckBuilder { + pub fn new() -> Self { + Self::default() + } + + pub fn try_push_range( + &mut self, + range: RangeInclusive, + max_wire_size: usize, + ) -> Result { + let start = range.start().into_inner(); + let end = range.end().into_inner(); + if start > end { + return Err(RecordAckRangeError::InvertedRange); + } + + let range_len = VarInt::from_u64(end - start).unwrap(); + if let Some(previous_start) = self.previous_start { + if end.saturating_add(1) >= previous_start { + return Err(RecordAckRangeError::NotCanonical); + } + + let gap = previous_start + .checked_sub(end) + .and_then(|delta| delta.checked_sub(2)) + .expect("canonical ack ranges stay separated by at least one sequence"); + let block = RecordAckBlock { + gap: VarInt::from_u64(gap).unwrap(), + range_len, + }; + let current_block_count_len = RecordAck::block_count_len(self.blocks.len()); + let next_block_count_len = RecordAck::block_count_len(self.blocks.len() + 1); + let next_wire_len = self.wire_len + + (next_block_count_len - current_block_count_len) + + block.encoded_len(); + if next_wire_len > max_wire_size { + return Ok(false); + } + + self.previous_start = Some(start); + self.wire_len = next_wire_len; + self.blocks.push(block); + return Ok(true); + } + + let largest_acked = RecordSeq::from_u64(end).unwrap(); + let wire_len = + largest_acked.encoded_len() + RecordAck::block_count_len(0) + range_len.encoded_len(); + if wire_len > max_wire_size { + return Ok(false); + } + + self.largest_acked = Some(largest_acked); + self.first_range_len = Some(range_len); + self.previous_start = Some(start); + self.wire_len = wire_len; + Ok(true) + } + + pub fn build(self) -> Result { + let Some(largest_acked) = self.largest_acked else { + return Err(RecordAckRangeError::Empty); + }; + + Ok(RecordAck { + largest_acked, + first_range_len: self.first_range_len.unwrap(), + blocks: self.blocks.into_boxed_slice(), + }) + } +} +#[cfg(test)] +mod tests { + use std::ops::RangeInclusive; + + use super::{RecordAck, RecordAckBlock, RecordAckBuilder, RecordAckRangeError}; + use crate::{RecordSeq, VarInt, WireDecode, WireEncode, WireError}; + + fn seq(value: u64) -> RecordSeq { + RecordSeq::from_u64(value).unwrap() + } + + fn ack_range(start: u64, end: u64) -> RangeInclusive { + seq(start)..=seq(end) + } + + fn varint(value: u64) -> VarInt { + VarInt::from_u64(value).unwrap() + } + + #[test] + fn encode_decode_round_trip() { + let ack = + RecordAck::from_ranges([ack_range(95, 100), ack_range(90, 92), ack_range(80, 80)]) + .unwrap(); + let encoded = ack.encode_vec(); + + assert_eq!(RecordAck::decode_exact(encoded.as_slice()).unwrap(), ack); + } + + #[test] + fn wire_fields_match_gap_encoding() { + let ack = + RecordAck::from_ranges([ack_range(95, 100), ack_range(90, 92), ack_range(80, 80)]) + .unwrap(); + + assert_eq!(ack.largest_acked, seq(100)); + assert_eq!(ack.first_range_len, varint(5)); + assert_eq!( + ack.blocks.as_ref(), + &[ + RecordAckBlock { + gap: varint(1), + range_len: varint(2), + }, + RecordAckBlock { + gap: varint(8), + range_len: varint(0), + } + ] + ); + } + + #[test] + fn builder_matches_from_ranges() { + let mut builder = RecordAckBuilder::new(); + assert!(builder + .try_push_range(ack_range(95, 100), usize::MAX) + .unwrap()); + assert!(builder + .try_push_range(ack_range(90, 92), usize::MAX) + .unwrap()); + assert!(builder + .try_push_range(ack_range(80, 80), usize::MAX) + .unwrap()); + + assert_eq!( + builder.build().unwrap(), + RecordAck::from_ranges([ack_range(95, 100), ack_range(90, 92), ack_range(80, 80)]) + .unwrap() + ); + } + + #[test] + fn builder_stops_when_budget_is_exhausted() { + let first_only = RecordAck::from_ranges([ack_range(95, 100)]).unwrap(); + let mut builder = RecordAckBuilder::new(); + + assert!(builder + .try_push_range(ack_range(95, 100), first_only.encoded_len()) + .unwrap()); + assert!(!builder + .try_push_range(ack_range(90, 92), first_only.encoded_len()) + .unwrap()); + assert_eq!(builder.build().unwrap(), first_only); + } + + #[test] + fn builder_rejects_non_canonical_ranges() { + let mut builder = RecordAckBuilder::new(); + assert!(builder + .try_push_range(ack_range(95, 100), usize::MAX) + .unwrap()); + assert_eq!( + builder.try_push_range(ack_range(90, 95), usize::MAX), + Err(RecordAckRangeError::NotCanonical) + ); + } + + #[test] + fn rejects_unsorted_ranges() { + assert_eq!( + RecordAck::from_ranges([ack_range(90, 92), ack_range(95, 100)]), + Err(RecordAckRangeError::NotCanonical) + ); + } + + #[test] + fn rejects_touching_ranges() { + assert_eq!( + RecordAck::from_ranges([ack_range(10, 12), ack_range(7, 9)]), + Err(RecordAckRangeError::NotCanonical) + ); + } + + #[test] + fn rejects_overlapping_ranges() { + assert_eq!( + RecordAck::from_ranges([ack_range(10, 12), ack_range(8, 11)]), + Err(RecordAckRangeError::NotCanonical) + ); + } + + #[test] + fn contains_matches_range_membership() { + let ack = RecordAck::from_ranges([ + ack_range(150, 163), + ack_range(105, 110), + ack_range(100, 100), + ]) + .unwrap(); + + assert!(ack.contains(100)); + assert!(ack.contains(107)); + assert!(ack.contains(163)); + assert!(!ack.contains(99)); + assert!(!ack.contains(104)); + assert!(!ack.contains(164)); + } + + #[test] + fn empty_ack_is_rejected() { + assert_eq!(RecordAck::from_ranges([]), Err(RecordAckRangeError::Empty)); + } + + #[test] + fn inverted_range_is_rejected() { + assert_eq!( + RecordAck::from_ranges([ack_range(5, 4)]), + Err(RecordAckRangeError::InvertedRange) + ); + } + + #[test] + fn decode_rejects_underflowing_ack_blocks() { + let encoded = vec![ + 42, // largest_acked + 1, // block_count + 0, // first_range_len + 41, // gap: implies a missing run larger than largest_acked + 0, // range_len + ]; + + assert_eq!( + RecordAck::decode_exact(encoded.as_slice()), + Err(WireError::InvalidPayload) + ); + } + + #[test] + fn decode_rejects_truncated_payload() { + assert_eq!( + RecordAck::decode_exact(&[][..]), + Err(WireError::InvalidPayload) + ); + + let encoded = RecordAck::from_ranges([ack_range(42, 42)]) + .unwrap() + .encode_vec(); + assert_eq!( + RecordAck::decode_exact(&encoded[..encoded.len() - 1]), + Err(WireError::InvalidPayload) + ); + } +} diff --git a/ql-wire/src/encrypted/builder.rs b/ql-wire/src/encrypted/builder.rs new file mode 100644 index 00000000..42933235 --- /dev/null +++ b/ql-wire/src/encrypted/builder.rs @@ -0,0 +1,172 @@ +use bytes::BufMut; + +use super::{RecordAck, SessionClose, SessionFrame, StreamClose, StreamData, StreamWindow}; +use crate::{ + BufView, ConnectionId, Nonce, QlCrypto, RecordSeq, RecordType, SessionHeader, SessionKey, + WireEncode, QL_WIRE_VERSION, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SessionRecordBuilder { + seq: RecordSeq, + prefix_len: usize, + max_capacity: usize, + bytes: Vec, +} + +impl SessionRecordBuilder { + pub const MIN_CAPACITY: usize = 1 + + 1 + + ConnectionId::SIZE + + RecordSeq::MAX_ENCODED_LEN + + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; + + pub fn new(seq: RecordSeq, max_capacity: usize) -> Self { + let prefix_len = + 1 + 1 + ConnectionId::SIZE + seq.encoded_len() + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; + assert!(max_capacity >= prefix_len); + Self { + seq, + prefix_len, + max_capacity, + bytes: Vec::new(), + } + } + + pub fn seq(&self) -> RecordSeq { + self.seq + } + + pub fn prefix_len(&self) -> usize { + self.prefix_len + } + + pub fn max_capacity(&self) -> usize { + self.max_capacity + } + + pub fn len(&self) -> usize { + self.bytes.len().saturating_sub(self.prefix_len) + } + + pub fn is_empty(&self) -> bool { + self.len() == 0 + } + + pub fn remaining_capacity(&self) -> usize { + self.max_capacity + .saturating_sub(self.bytes.len().max(self.prefix_len)) + } + + pub fn bytes(&self) -> &[u8] { + self.bytes.get(self.prefix_len..).unwrap_or_default() + } + + pub fn push_ping(&mut self) -> bool { + self.push_empty_frame(super::SessionFrameKind::Ping) + } + + pub fn push_unpair(&mut self) -> bool { + self.push_empty_frame(super::SessionFrameKind::Unpair) + } + + pub fn push_ack(&mut self, ack: &RecordAck) -> bool { + self.push_frame_payload(super::SessionFrameKind::Ack, ack) + } + + pub fn push_stream_data(&mut self, frame: &StreamData) -> bool { + self.push_frame_payload(super::SessionFrameKind::StreamData, frame) + } + + pub fn push_stream_window(&mut self, frame: &StreamWindow) -> bool { + self.push_frame_payload(super::SessionFrameKind::StreamWindow, frame) + } + + pub fn push_stream_close(&mut self, frame: &StreamClose) -> bool { + self.push_frame_payload(super::SessionFrameKind::StreamClose, frame) + } + + pub fn push_close(&mut self, close: &SessionClose) -> bool { + self.push_frame_payload(super::SessionFrameKind::Close, close) + } + + pub fn push_frame(&mut self, frame: &SessionFrame) -> bool { + match frame { + SessionFrame::Ping => self.push_ping(), + SessionFrame::Unpair => self.push_unpair(), + SessionFrame::Ack(frame) => self.push_ack(frame), + SessionFrame::StreamData(frame) => self.push_stream_data(frame), + SessionFrame::StreamWindow(frame) => self.push_stream_window(frame), + SessionFrame::StreamClose(frame) => self.push_stream_close(frame), + SessionFrame::Close(close) => self.push_close(close), + } + } + + pub fn encrypt( + mut self, + crypto: &impl QlCrypto, + connection_id: ConnectionId, + session_key: &SessionKey, + ) -> Vec { + self.ensure_prefix_capacity(0); + let header = SessionHeader { + connection_id, + seq: self.seq, + }; + let aad = header.aad(); + let nonce = Nonce::from_counter(self.seq.into_inner()); + let auth = crypto.aes256_gcm_encrypt( + session_key, + &nonce, + &aad, + &mut self.bytes[self.prefix_len..], + ); + + let mut prefix = &mut self.bytes[..self.prefix_len]; + prefix[0] = QL_WIRE_VERSION; + prefix[1] = RecordType::Session as u8; + prefix = &mut prefix[2..]; + header.encode(&mut prefix); + auth.encode(&mut prefix); + debug_assert!(prefix.is_empty()); + self.bytes + } + + fn push_wire_size(&mut self, wire_size: usize, encode: impl FnOnce(&mut Vec)) -> bool { + if !self.can_push_len(wire_size) { + return false; + } + self.ensure_prefix_capacity(wire_size); + let start = self.bytes.len(); + encode(&mut self.bytes); + debug_assert_eq!(self.bytes.len(), start + wire_size); + true + } + + fn push_empty_frame(&mut self, kind: super::SessionFrameKind) -> bool { + self.push_wire_size(1, |out| out.put_u8(kind as u8)) + } + + fn push_frame_payload( + &mut self, + kind: super::SessionFrameKind, + payload: &T, + ) -> bool { + let payload_wire_size = payload.encoded_len(); + self.push_wire_size(1 + payload_wire_size, |out| { + out.put_u8(kind as u8); + payload.encode(out); + }) + } + + fn can_push_len(&self, len: usize) -> bool { + len <= self.remaining_capacity() + } + + fn ensure_prefix_capacity(&mut self, additional_body_len: usize) { + if self.bytes.is_empty() { + self.bytes.reserve(self.prefix_len + additional_body_len); + self.bytes.resize(self.prefix_len, 0); + } + } +} diff --git a/ql-wire/src/encrypted/close.rs b/ql-wire/src/encrypted/close.rs new file mode 100644 index 00000000..e0860d7a --- /dev/null +++ b/ql-wire/src/encrypted/close.rs @@ -0,0 +1,55 @@ +use crate::{codec, codec::Reader, ByteSlice, WireEncode, WireError}; + +/// closes the whole session immediately with a close code. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SessionClose { + pub code: SessionCloseCode, +} + +impl SessionClose { + pub const WIRE_SIZE: usize = size_of::(); +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct SessionCloseCode(pub u16); + +impl SessionCloseCode { + pub const CANCELLED: Self = Self(0); + pub const PROTOCOL: Self = Self(1); + pub const TIMEOUT: Self = Self(2); +} + +impl WireEncode for SessionCloseCode { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl codec::WireDecode for SessionCloseCode { + fn decode(reader: &mut Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +impl codec::WireDecode for SessionClose { + fn decode(reader: &mut Reader) -> Result { + Ok(Self { + code: reader.decode()?, + }) + } +} + +impl WireEncode for SessionClose { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.code.encode(out); + } +} diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs new file mode 100644 index 00000000..563f9ded --- /dev/null +++ b/ql-wire/src/encrypted/mod.rs @@ -0,0 +1,178 @@ +use crate::{ + codec, encrypted_message::EncryptedMessage, BufView, ByteSlice, Nonce, QlCrypto, Reader, + SessionHeader, SessionKey, WireDecode, WireEncode, WireError, +}; + +mod ack; +mod builder; +mod close; +mod route_id; +mod stream_close; +mod stream_data; +mod stream_id; +mod stream_window; + +pub use ack::*; +pub use builder::*; +pub use close::*; +pub use route_id::*; +pub use stream_close::*; +pub use stream_data::*; +pub use stream_id::*; +pub use stream_window::*; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum SessionFrame { + // todo: do we need ping as explicit frame? + Ping, + Unpair, + Ack(RecordAck), + StreamData(StreamData), + StreamWindow(StreamWindow), + StreamClose(StreamClose), + Close(SessionClose), +} + +impl WireDecode for SessionFrame { + fn decode(reader: &mut Reader) -> Result { + let kind = reader.decode::()?; + let frame = match kind { + SessionFrameKind::Ping => Self::Ping, + SessionFrameKind::Unpair => Self::Unpair, + SessionFrameKind::Ack => Self::Ack(reader.decode::()?), + SessionFrameKind::StreamData => Self::StreamData(reader.decode::>()?), + SessionFrameKind::StreamWindow => Self::StreamWindow(reader.decode::()?), + SessionFrameKind::StreamClose => Self::StreamClose(reader.decode::()?), + SessionFrameKind::Close => Self::Close(reader.decode::()?), + }; + Ok(frame) + } +} + +impl SessionFrame { + fn kind(&self) -> SessionFrameKind { + match self { + Self::Ping => SessionFrameKind::Ping, + Self::Unpair => SessionFrameKind::Unpair, + Self::Ack(_) => SessionFrameKind::Ack, + Self::StreamData(_) => SessionFrameKind::StreamData, + Self::StreamWindow(_) => SessionFrameKind::StreamWindow, + Self::StreamClose(_) => SessionFrameKind::StreamClose, + Self::Close(_) => SessionFrameKind::Close, + } + } +} + +impl SessionFrame { + pub fn into_owned(self) -> SessionFrame> { + match self { + Self::Ping => SessionFrame::Ping, + Self::Unpair => SessionFrame::Unpair, + Self::Ack(frame) => SessionFrame::Ack(frame), + Self::StreamData(frame) => SessionFrame::StreamData(frame.into_owned()), + Self::StreamWindow(frame) => SessionFrame::StreamWindow(frame), + Self::StreamClose(frame) => SessionFrame::StreamClose(frame), + Self::Close(frame) => SessionFrame::Close(frame), + } + } +} + +impl WireEncode for SessionFrame { + fn encoded_len(&self) -> usize { + 1 + match self { + Self::Ping | Self::Unpair => 0, + Self::Ack(frame) => frame.encoded_len(), + Self::StreamData(frame) => frame.encoded_len(), + Self::StreamWindow(frame) => frame.encoded_len(), + Self::StreamClose(frame) => frame.encoded_len(), + Self::Close(frame) => frame.encoded_len(), + } + } + + fn encode(&self, out: &mut W) { + out.put_u8(self.kind() as u8); + match self { + Self::Ping | Self::Unpair => {} + Self::Ack(frame) => frame.encode(out), + Self::StreamData(frame) => frame.encode(out), + Self::StreamWindow(frame) => frame.encode(out), + Self::StreamClose(frame) => frame.encode(out), + Self::Close(frame) => frame.encode(out), + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +pub enum SessionFrameKind { + Ping = 1, + Ack = 2, + StreamData = 3, + StreamWindow = 4, + StreamClose = 5, + Close = 6, + Unpair = 7, +} + +impl TryFrom for SessionFrameKind { + type Error = WireError; + + fn try_from(value: u8) -> Result { + match value { + 1 => Ok(Self::Ping), + 2 => Ok(Self::Ack), + 3 => Ok(Self::StreamData), + 4 => Ok(Self::StreamWindow), + 5 => Ok(Self::StreamClose), + 6 => Ok(Self::Close), + 7 => Ok(Self::Unpair), + _ => Err(WireError::InvalidPayload), + } + } +} + +impl codec::WireDecode for SessionFrameKind { + fn decode(reader: &mut codec::Reader) -> Result { + reader.decode::()?.try_into() + } +} + +pub fn parse_session_frames(bytes: B) -> SessionFrameIter { + SessionFrameIter { + reader: Reader::new(bytes), + } +} + +pub fn decode_session_frames(bytes: &[u8]) -> Result>>, WireError> { + parse_session_frames(bytes) + .map(|frame| frame.map(SessionFrame::into_owned)) + .collect() +} + +#[derive(Clone)] +pub struct SessionFrameIter { + reader: Reader, +} + +impl Iterator for SessionFrameIter { + type Item = Result, WireError>; + + fn next(&mut self) -> Option { + if self.reader.is_empty() { + None + } else { + Some(self.reader.decode::>()) + } + } +} + +pub fn decrypt_record>( + crypto: &impl QlCrypto, + header: &SessionHeader, + encrypted: EncryptedMessage, + session_key: &SessionKey, +) -> Result { + let aad = header.aad(); + let nonce = Nonce::from_counter(header.seq.into_inner()); + encrypted.decrypt_in_place(crypto, session_key, &nonce, &aad) +} diff --git a/ql-wire/src/encrypted/route_id.rs b/ql-wire/src/encrypted/route_id.rs new file mode 100644 index 00000000..6b91a521 --- /dev/null +++ b/ql-wire/src/encrypted/route_id.rs @@ -0,0 +1,55 @@ +use crate::{ByteSlice, Reader, VarInt, VarIntBoundsExceeded, WireDecode, WireEncode, WireError}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct RouteId(pub VarInt); + +impl RouteId { + pub const MAX_ENCODED_LEN: usize = VarInt::MAX_SIZE; + + pub const fn from_u32(value: u32) -> Self { + Self(VarInt::from_u32(value)) + } + + pub fn from_u64(value: u64) -> Result { + Ok(Self(VarInt::from_u64(value)?)) + } + + pub const fn into_inner(self) -> u64 { + self.0.into_inner() + } +} + +impl WireEncode for RouteId { + fn encoded_len(&self) -> usize { + self.0.size() + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl WireDecode for RouteId { + fn decode(reader: &mut Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +impl From for RouteId { + fn from(value: VarInt) -> Self { + Self(value) + } +} + +impl From for RouteId { + fn from(value: u32) -> Self { + Self::from_u32(value) + } +} + +impl std::fmt::Display for RouteId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } +} diff --git a/ql-wire/src/encrypted/stream_close.rs b/ql-wire/src/encrypted/stream_close.rs new file mode 100644 index 00000000..20ddb879 --- /dev/null +++ b/ql-wire/src/encrypted/stream_close.rs @@ -0,0 +1,116 @@ +use super::StreamId; +use crate::{codec, ByteSlice, WireEncode, WireError}; + +/// aborts one or both lanes of a stream with a close code +/// +/// stream origin is the peer that opened the stream +/// origin lane carries bytes sent by the stream origin +/// return lane carries bytes sent back toward the stream origin +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct StreamClose { + pub stream_id: StreamId, + pub target: CloseTarget, + pub code: StreamCloseCode, +} + +impl StreamClose {} + +impl WireEncode for StreamClose { + fn encoded_len(&self) -> usize { + self.stream_id.encoded_len() + self.target.encoded_len() + self.code.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.stream_id.encode(out); + self.target.encode(out); + self.code.encode(out); + } +} + +impl codec::WireDecode for StreamClose { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + stream_id: reader.decode()?, + target: reader.decode()?, + code: reader.decode()?, + }) + } +} + +/// selects which stream lane a [`StreamClose`] applies to +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +pub enum CloseTarget { + /// close the lane sent by the stream origin + Origin = 1, + /// close the lane sent back toward the stream origin + Return = 2, + /// close both stream lanes + Both = 3, +} + +impl CloseTarget { + pub const fn to_wire(self) -> u8 { + self as u8 + } +} + +impl WireEncode for CloseTarget { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + self.to_wire().encode(out); + } +} + +impl TryFrom for CloseTarget { + type Error = WireError; + + fn try_from(value: u8) -> Result { + match value { + 1 => Ok(Self::Origin), + 2 => Ok(Self::Return), + 3 => Ok(Self::Both), + _ => Err(WireError::InvalidPayload), + } + } +} + +impl codec::WireDecode for CloseTarget { + fn decode(reader: &mut codec::Reader) -> Result { + reader.decode::()?.try_into() + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct StreamCloseCode(pub u16); + +impl StreamCloseCode { + /// the stream was aborted intentionally before graceful completion + pub const CANCELLED: Self = Self(0); +} + +impl codec::WireDecode for StreamCloseCode { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +impl WireEncode for StreamCloseCode { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl std::fmt::Display for StreamCloseCode { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } +} diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs new file mode 100644 index 00000000..9174fe5a --- /dev/null +++ b/ql-wire/src/encrypted/stream_data.rs @@ -0,0 +1,135 @@ +use bytes::Buf; + +use super::{RouteId, StreamId}; +use crate::{codec, BufView, ByteSlice, VarInt, WireDecode, WireEncode, WireError}; + +/// carries bytes for a stream and may finish that sending direction. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct StreamData { + pub stream_id: StreamId, + pub offset: VarInt, + pub header: Option, + pub fin: bool, + pub bytes: B, +} + +impl StreamData { + pub const MIN_WIRE_SIZE: usize = StreamId::MAX_ENCODED_LEN + + VarInt::MAX_SIZE + + size_of::() + + StreamHeader::MAX_WIRE_SIZE + + VarInt::MAX_SIZE; +} + +impl WireDecode for StreamData { + fn decode(reader: &mut codec::Reader) -> Result { + let stream_id = reader.decode()?; + let offset: VarInt = reader.decode()?; + let flags = reader.decode::()?; + let fin = (flags & flag::FIN) != 0; + let has_header = (flags & flag::HEADER) != 0; + let header = if has_header { + Some(reader.decode()?) + } else { + None + }; + let bytes_len = usize::try_from(reader.decode::()?.into_inner()) + .map_err(|_| WireError::InvalidPayload)?; + + Ok(Self { + stream_id, + offset, + header, + fin, + bytes: reader.take_bytes(bytes_len)?, + }) + } +} + +impl StreamData { + pub fn into_owned(self) -> StreamData> + where + B: ByteSlice, + { + StreamData { + stream_id: self.stream_id, + offset: self.offset, + header: self.header, + fin: self.fin, + bytes: self.bytes.to_vec(), + } + } +} + +impl WireEncode for StreamData { + fn encoded_len(&self) -> usize { + let bytes = self.bytes.buf(); + let bytes_len = bytes.remaining(); + self.stream_id.encoded_len() + + self.offset.encoded_len() + + size_of::() + + self.header.as_ref().map_or(0, WireEncode::encoded_len) + + VarInt::try_from(bytes_len).unwrap().encoded_len() + + bytes_len + } + + fn encode(&self, out: &mut W) { + debug_assert!( + self.offset.into_inner() == 0 || self.header.is_none(), + "stream header is only valid at offset 0" + ); + + self.stream_id.encode(out); + self.offset.encode(out); + let mut flags = 0; + if self.fin { + flags |= flag::FIN; + } + if self.header.is_some() { + flags |= flag::HEADER; + } + flags.encode(out); + if let Some(header) = &self.header { + header.encode(out); + } + let mut bytes = self.bytes.buf(); + VarInt::try_from(bytes.remaining()).unwrap().encode(out); + while bytes.has_remaining() { + let chunk = bytes.chunk(); + out.put_slice(chunk); + bytes.advance(chunk.len()); + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct StreamHeader { + pub route_id: RouteId, +} + +impl StreamHeader { + pub const MAX_WIRE_SIZE: usize = RouteId::MAX_ENCODED_LEN; +} + +impl WireDecode for StreamHeader { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + route_id: reader.decode()?, + }) + } +} + +impl WireEncode for StreamHeader { + fn encoded_len(&self) -> usize { + self.route_id.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.route_id.encode(out); + } +} + +mod flag { + pub const FIN: u8 = 0x01; + pub const HEADER: u8 = 0x02; +} diff --git a/ql-wire/src/encrypted/stream_id.rs b/ql-wire/src/encrypted/stream_id.rs new file mode 100644 index 00000000..07002259 --- /dev/null +++ b/ql-wire/src/encrypted/stream_id.rs @@ -0,0 +1,35 @@ +use crate::{ByteSlice, Reader, VarInt, WireDecode, WireEncode, WireError}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct StreamId(pub VarInt); + +impl StreamId { + pub const MAX_ENCODED_LEN: usize = VarInt::MAX_SIZE; + + pub const fn into_inner(self) -> u64 { + self.0.into_inner() + } +} + +impl WireEncode for StreamId { + fn encoded_len(&self) -> usize { + self.0.size() + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl WireDecode for StreamId { + fn decode(reader: &mut Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +impl std::fmt::Display for StreamId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } +} diff --git a/ql-wire/src/encrypted/stream_window.rs b/ql-wire/src/encrypted/stream_window.rs new file mode 100644 index 00000000..6a2274f9 --- /dev/null +++ b/ql-wire/src/encrypted/stream_window.rs @@ -0,0 +1,29 @@ +use super::StreamId; +use crate::{codec, ByteSlice, VarInt, WireEncode, WireError}; + +/// advertises the highest byte offset the peer may send on a stream. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct StreamWindow { + pub stream_id: StreamId, + pub maximum_offset: VarInt, +} + +impl WireEncode for StreamWindow { + fn encoded_len(&self) -> usize { + self.stream_id.encoded_len() + self.maximum_offset.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.stream_id.encode(out); + self.maximum_offset.encode(out); + } +} + +impl codec::WireDecode for StreamWindow { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + stream_id: reader.decode()?, + maximum_offset: reader.decode()?, + }) + } +} diff --git a/ql-wire/src/encrypted_message.rs b/ql-wire/src/encrypted_message.rs new file mode 100644 index 00000000..9e11d3d0 --- /dev/null +++ b/ql-wire/src/encrypted_message.rs @@ -0,0 +1,97 @@ +use crate::{ + codec, ByteSlice, Nonce, QlCrypto, SessionKey, WireDecode, WireEncode, WireError, + ENCRYPTED_MESSAGE_AUTH_SIZE, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct EncryptedMessage { + pub auth: [u8; ENCRYPTED_MESSAGE_AUTH_SIZE], + pub ciphertext: B, +} + +impl EncryptedMessage { + pub const AUTH_SIZE: usize = ENCRYPTED_MESSAGE_AUTH_SIZE; + pub const HEADER_LEN: usize = Self::AUTH_SIZE; + + pub fn into_owned(self) -> EncryptedMessage> + where + B: ByteSlice, + { + EncryptedMessage { + auth: self.auth, + ciphertext: self.ciphertext.to_vec(), + } + } +} + +impl WireDecode for EncryptedMessage { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + auth: reader.decode()?, + ciphertext: reader.take_rest(), + }) + } +} + +impl> EncryptedMessage { + pub fn decrypt( + &self, + crypto: &impl QlCrypto, + key: &SessionKey, + nonce: &Nonce, + aad: &[u8], + ) -> Result, WireError> { + let mut plaintext = self.ciphertext.as_ref().to_vec(); + if !crypto.aes256_gcm_decrypt(key, nonce, aad, &mut plaintext, &self.auth) { + return Err(WireError::DecryptFailed); + } + Ok(plaintext) + } +} + +impl> WireEncode for EncryptedMessage { + fn encoded_len(&self) -> usize { + Self::HEADER_LEN + self.ciphertext.as_ref().len() + } + + fn encode(&self, out: &mut W) { + self.auth.encode(out); + self.ciphertext.as_ref().encode(out); + } +} + +impl> EncryptedMessage { + pub fn decrypt_in_place( + mut self, + crypto: &impl QlCrypto, + key: &SessionKey, + nonce: &Nonce, + aad: &[u8], + ) -> Result { + let ciphertext = self.ciphertext.as_mut(); + if !crypto.aes256_gcm_decrypt(key, nonce, aad, ciphertext, &self.auth) { + return Err(WireError::DecryptFailed); + } + Ok(self.ciphertext) + } +} + +impl EncryptedMessage> { + pub fn encrypt( + crypto: &impl QlCrypto, + key: &SessionKey, + mut plaintext: Vec, + nonce: &Nonce, + aad: &[u8], + ) -> Self { + let auth = crypto.aes256_gcm_encrypt(key, nonce, aad, &mut plaintext); + Self { + auth, + ciphertext: plaintext, + } + } + + pub fn decode(bytes: &[u8]) -> Result { + Ok(EncryptedMessage::decode_exact(bytes)?.into_owned()) + } +} diff --git a/ql-wire/src/error.rs b/ql-wire/src/error.rs new file mode 100644 index 00000000..8da1eec0 --- /dev/null +++ b/ql-wire/src/error.rs @@ -0,0 +1,33 @@ +use core::fmt; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum WireError { + InvalidPayload, + InvalidHandshakeHeader, + InvalidHandshakeMeta, + InvalidPairingId, + InvalidRemoteBundle, + InvalidTransportParams, + Expired, + DecryptFailed, + InvalidState, +} + +impl fmt::Display for WireError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let message = match self { + Self::InvalidPayload => "invalid payload", + Self::InvalidHandshakeHeader => "invalid handshake header", + Self::InvalidHandshakeMeta => "invalid handshake meta", + Self::InvalidPairingId => "invalid pairing id", + Self::InvalidRemoteBundle => "invalid remote bundle", + Self::InvalidTransportParams => "invalid transport params", + Self::Expired => "expired", + Self::DecryptFailed => "decryption failed", + Self::InvalidState => "invalid state", + }; + f.write_str(message) + } +} + +impl std::error::Error for WireError {} diff --git a/ql-wire/src/handshake/ik.rs b/ql-wire/src/handshake/ik.rs new file mode 100644 index 00000000..628e30e7 --- /dev/null +++ b/ql-wire/src/handshake/ik.rs @@ -0,0 +1,376 @@ +use super::{ + decrypt_mlkem_ciphertext, decrypt_peer_bundle, encrypt_mlkem_ciphertext, encrypt_peer_bundle, + finalize_handshake, generate_ephemeral_keypair, init_ik_symmetric, initialize_handshake_meta, + mix_hash_ephemeral, mix_hash_routed_handshake, require_handshake_meta, + EncryptedMlKemCiphertext, EncryptedPeerBundle, EphemeralKeyPair, EphemeralPublicKey, + FinalizedHandshake, HandshakeHeader, Role, SymmetricState, TransportParams, +}; +use crate::{ + codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, + QlIdentity, WireEncode, WireError, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Ik1 { + pub header: HandshakeHeader, + pub meta: HandshakeMeta, + pub transport_params: TransportParams, + pub skem_ciphertext: MlKemCiphertext, + pub ephemeral: EphemeralPublicKey, + pub static_bundle: EncryptedPeerBundle, +} + +impl codec::WireDecode for Ik1 { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + meta: reader.decode()?, + transport_params: reader.decode()?, + skem_ciphertext: reader.decode()?, + ephemeral: reader.decode()?, + static_bundle: reader.decode()?, + }) + } +} + +impl WireEncode for Ik1 { + fn encoded_len(&self) -> usize { + HandshakeHeader::WIRE_SIZE + + HandshakeMeta::WIRE_SIZE + + TransportParams::WIRE_SIZE + + MlKemCiphertext::SIZE + + EphemeralPublicKey::WIRE_SIZE + + self.static_bundle.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.meta.encode(out); + self.transport_params.encode(out); + self.skem_ciphertext.encode(out); + self.ephemeral.encode(out); + self.static_bundle.encode(out); + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Ik2 { + pub header: HandshakeHeader, + pub meta: HandshakeMeta, + pub transport_params: TransportParams, + pub ekem_ciphertext: MlKemCiphertext, + pub skem_ciphertext: EncryptedMlKemCiphertext, +} + +impl Ik2 { + pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE + + HandshakeMeta::WIRE_SIZE + + TransportParams::WIRE_SIZE + + MlKemCiphertext::SIZE + + EncryptedMlKemCiphertext::WIRE_SIZE; +} + +impl codec::WireDecode for Ik2 { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + meta: reader.decode()?, + transport_params: reader.decode()?, + ekem_ciphertext: reader.decode()?, + skem_ciphertext: reader.decode()?, + }) + } +} + +impl WireEncode for Ik2 { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.meta.encode(out); + self.transport_params.encode(out); + self.ekem_ciphertext.encode(out); + self.skem_ciphertext.encode(out); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum IkStep { + Send1, + Recv1, + Send2, + Recv2, + Done, +} + +#[derive(Debug, Clone)] +pub struct IkHandshake { + role: Role, + step: IkStep, + symmetric: SymmetricState, + local: QlIdentity, + remote_bundle: Option, + local_ephemeral: Option, + remote_ephemeral: Option, + handshake_meta: Option, + local_transport_params: TransportParams, + remote_transport_params: Option, +} + +impl IkHandshake { + pub fn new_initiator( + crypto: &impl QlCrypto, + local: QlIdentity, + remote_bundle: PeerBundle, + local_transport_params: TransportParams, + ) -> Self { + let symmetric = init_ik_symmetric(crypto, &remote_bundle); + Self { + role: Role::Initiator, + step: IkStep::Send1, + symmetric, + local, + remote_bundle: Some(remote_bundle), + local_ephemeral: None, + remote_ephemeral: None, + handshake_meta: None, + local_transport_params, + remote_transport_params: None, + } + } + + pub fn new_responder( + crypto: &impl QlCrypto, + local: QlIdentity, + expected_remote: Option, + local_transport_params: TransportParams, + ) -> Self { + let symmetric = init_ik_symmetric(crypto, &local.bundle()); + Self { + role: Role::Responder, + step: IkStep::Recv1, + symmetric, + local, + remote_bundle: expected_remote, + local_ephemeral: None, + remote_ephemeral: None, + handshake_meta: None, + local_transport_params, + remote_transport_params: None, + } + } + + pub fn is_finished(&self) -> bool { + self.step == IkStep::Done + } + + fn outbound_header(&self) -> Result { + let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + Ok(HandshakeHeader { + sender: self.local.qid, + recipient: remote_bundle.qid, + }) + } + + fn ensure_inbound_recipient(&self, header: HandshakeHeader) -> Result<(), WireError> { + if header.recipient == self.local.qid { + Ok(()) + } else { + Err(WireError::InvalidPayload) + } + } + + fn ensure_known_remote_sender(&self, header: HandshakeHeader) -> Result<(), WireError> { + if let Some(remote_bundle) = self.remote_bundle.as_ref() { + if header.sender != remote_bundle.qid { + return Err(WireError::InvalidPayload); + } + } + Ok(()) + } + + pub fn write_1( + &mut self, + crypto: &impl QlCrypto, + meta: HandshakeMeta, + ) -> Result { + if self.step != IkStep::Send1 { + return Err(WireError::InvalidState); + } + initialize_handshake_meta(&mut self.handshake_meta, meta)?; + let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + let header = self.outbound_header()?; + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Ik1, + meta, + self.local_transport_params, + ); + let (skem_ciphertext, skem_secret) = + crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); + self.symmetric.mix_hash(crypto, skem_ciphertext.as_bytes()); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + let local_ephemeral = generate_ephemeral_keypair(crypto); + let public = local_ephemeral.public(); + mix_hash_ephemeral(&mut self.symmetric, crypto, &public); + + let static_bundle = encrypt_peer_bundle(crypto, &mut self.symmetric, &self.local.bundle())?; + + self.local_ephemeral = Some(local_ephemeral); + self.step = IkStep::Recv2; + Ok(Ik1 { + header, + meta, + transport_params: self.local_transport_params, + skem_ciphertext, + ephemeral: public, + static_bundle, + }) + } + + pub fn write_2( + &mut self, + crypto: &impl QlCrypto, + meta: HandshakeMeta, + ) -> Result { + if self.step != IkStep::Send2 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), meta)?; + let header = self.outbound_header()?; + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Ik2, + meta, + self.local_transport_params, + ); + let remote_ephemeral = self + .remote_ephemeral + .clone() + .ok_or(WireError::InvalidState)?; + let (ekem_ciphertext, ekem_secret) = + crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); + self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); + self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); + + let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + let (skem_ciphertext, skem_secret) = + crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); + let skem_ciphertext = + encrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &skem_ciphertext)?; + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + self.step = IkStep::Done; + Ok(Ik2 { + header, + meta, + transport_params: self.local_transport_params, + ekem_ciphertext, + skem_ciphertext, + }) + } + + pub fn read_1(&mut self, crypto: &impl QlCrypto, message: &Ik1) -> Result<(), WireError> { + if self.step != IkStep::Recv1 { + return Err(WireError::InvalidState); + } + initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; + self.ensure_inbound_recipient(message.header)?; + self.ensure_known_remote_sender(message.header)?; + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + message.header, + HandshakeKind::Ik1, + message.meta, + message.transport_params, + ); + self.symmetric + .mix_hash(crypto, message.skem_ciphertext.as_bytes()); + let skem_secret = + crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &message.skem_ciphertext); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + mix_hash_ephemeral(&mut self.symmetric, crypto, &message.ephemeral); + self.remote_ephemeral = Some(message.ephemeral.clone()); + + let remote_bundle = + decrypt_peer_bundle(crypto, &mut self.symmetric, &message.static_bundle)?; + if remote_bundle.qid != message.header.sender { + return Err(WireError::InvalidPayload); + } + match self.remote_bundle.as_ref() { + Some(expected) if expected != &remote_bundle => { + return Err(WireError::InvalidPayload); + } + Some(_) => {} + None => self.remote_bundle = Some(remote_bundle), + } + self.remote_transport_params = Some(message.transport_params); + self.step = IkStep::Send2; + Ok(()) + } + + pub fn read_2(&mut self, crypto: &impl QlCrypto, message: &Ik2) -> Result<(), WireError> { + if self.step != IkStep::Recv2 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; + self.ensure_inbound_recipient(message.header)?; + self.ensure_known_remote_sender(message.header)?; + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + message.header, + HandshakeKind::Ik2, + message.meta, + message.transport_params, + ); + let local_ephemeral = self + .local_ephemeral + .as_ref() + .ok_or(WireError::InvalidState)?; + self.symmetric + .mix_hash(crypto, message.ekem_ciphertext.as_bytes()); + let ekem_secret = + crypto.mlkem_decapsulate(&local_ephemeral.mlkem.private, &message.ekem_ciphertext); + self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); + + let skem_ciphertext = + decrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &message.skem_ciphertext)?; + let skem_secret = crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &skem_ciphertext); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + self.remote_transport_params = Some(message.transport_params); + self.step = IkStep::Done; + Ok(()) + } + + pub fn finalize(self, crypto: &impl QlCrypto) -> Result { + if !self.is_finished() { + return Err(WireError::InvalidState); + } + let remote_bundle = self.remote_bundle.ok_or(WireError::InvalidState)?; + let remote_transport_params = self + .remote_transport_params + .ok_or(WireError::InvalidState)?; + Ok(finalize_handshake( + crypto, + &self.symmetric, + self.role, + remote_bundle, + remote_transport_params, + )) + } +} diff --git a/ql-wire/src/handshake/kk.rs b/ql-wire/src/handshake/kk.rs new file mode 100644 index 00000000..2ad5ee2a --- /dev/null +++ b/ql-wire/src/handshake/kk.rs @@ -0,0 +1,352 @@ +use super::{ + decrypt_mlkem_ciphertext, encrypt_mlkem_ciphertext, finalize_handshake, + generate_ephemeral_keypair, init_kk_symmetric, initialize_handshake_meta, mix_hash_ephemeral, + mix_hash_routed_handshake, require_handshake_meta, EncryptedMlKemCiphertext, EphemeralKeyPair, + EphemeralPublicKey, FinalizedHandshake, HandshakeHeader, Role, SymmetricState, TransportParams, +}; +use crate::{ + codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, + QlIdentity, WireEncode, WireError, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Kk1 { + pub header: HandshakeHeader, + pub meta: HandshakeMeta, + pub transport_params: TransportParams, + pub skem_ciphertext: MlKemCiphertext, + pub ephemeral: EphemeralPublicKey, +} + +impl Kk1 { + pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE + + HandshakeMeta::WIRE_SIZE + + TransportParams::WIRE_SIZE + + MlKemCiphertext::SIZE + + EphemeralPublicKey::WIRE_SIZE; +} + +impl codec::WireDecode for Kk1 { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + meta: reader.decode()?, + transport_params: reader.decode()?, + skem_ciphertext: reader.decode()?, + ephemeral: reader.decode()?, + }) + } +} + +impl WireEncode for Kk1 { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.meta.encode(out); + self.transport_params.encode(out); + self.skem_ciphertext.encode(out); + self.ephemeral.encode(out); + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Kk2 { + pub header: HandshakeHeader, + pub meta: HandshakeMeta, + pub transport_params: TransportParams, + pub ekem_ciphertext: MlKemCiphertext, + pub skem_ciphertext: EncryptedMlKemCiphertext, +} + +impl Kk2 { + pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE + + HandshakeMeta::WIRE_SIZE + + TransportParams::WIRE_SIZE + + MlKemCiphertext::SIZE + + EncryptedMlKemCiphertext::WIRE_SIZE; +} + +impl codec::WireDecode for Kk2 { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + meta: reader.decode()?, + transport_params: reader.decode()?, + ekem_ciphertext: reader.decode()?, + skem_ciphertext: reader.decode()?, + }) + } +} + +impl WireEncode for Kk2 { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.meta.encode(out); + self.transport_params.encode(out); + self.ekem_ciphertext.encode(out); + self.skem_ciphertext.encode(out); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum KkStep { + Send1, + Recv1, + Send2, + Recv2, + Done, +} + +#[derive(Debug, Clone)] +pub struct KkHandshake { + role: Role, + step: KkStep, + symmetric: SymmetricState, + local: QlIdentity, + remote_bundle: PeerBundle, + local_ephemeral: Option, + remote_ephemeral: Option, + handshake_meta: Option, + local_transport_params: TransportParams, + remote_transport_params: Option, +} + +impl KkHandshake { + pub fn new_initiator( + crypto: &impl QlCrypto, + local: QlIdentity, + remote_bundle: PeerBundle, + local_transport_params: TransportParams, + ) -> Self { + let symmetric = init_kk_symmetric(crypto, &local.bundle(), &remote_bundle); + Self { + role: Role::Initiator, + step: KkStep::Send1, + symmetric, + local, + remote_bundle, + local_ephemeral: None, + remote_ephemeral: None, + handshake_meta: None, + local_transport_params, + remote_transport_params: None, + } + } + + pub fn new_responder( + crypto: &impl QlCrypto, + local: QlIdentity, + remote_bundle: PeerBundle, + local_transport_params: TransportParams, + ) -> Self { + let symmetric = init_kk_symmetric(crypto, &remote_bundle, &local.bundle()); + Self { + role: Role::Responder, + step: KkStep::Recv1, + symmetric, + local, + remote_bundle, + local_ephemeral: None, + remote_ephemeral: None, + handshake_meta: None, + local_transport_params, + remote_transport_params: None, + } + } + + pub fn is_finished(&self) -> bool { + self.step == KkStep::Done + } + + fn outbound_header(&self) -> HandshakeHeader { + HandshakeHeader { + sender: self.local.qid, + recipient: self.remote_bundle.qid, + } + } + + fn inbound_header(&self) -> HandshakeHeader { + HandshakeHeader { + sender: self.remote_bundle.qid, + recipient: self.local.qid, + } + } + + fn ensure_inbound_header(&self, header: HandshakeHeader) -> Result<(), WireError> { + if header == self.inbound_header() { + Ok(()) + } else { + Err(WireError::InvalidPayload) + } + } + + pub fn write_1( + &mut self, + crypto: &impl QlCrypto, + meta: HandshakeMeta, + ) -> Result { + if self.step != KkStep::Send1 { + return Err(WireError::InvalidState); + } + initialize_handshake_meta(&mut self.handshake_meta, meta)?; + let header = self.outbound_header(); + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Kk1, + meta, + self.local_transport_params, + ); + let (skem_ciphertext, skem_secret) = + crypto.mlkem_encapsulate(&self.remote_bundle.mlkem_public_key); + self.symmetric + .encrypt_and_hash(crypto, skem_ciphertext.as_bytes())?; + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + let local_ephemeral = generate_ephemeral_keypair(crypto); + let public = local_ephemeral.public(); + mix_hash_ephemeral(&mut self.symmetric, crypto, &public); + + self.local_ephemeral = Some(local_ephemeral); + self.step = KkStep::Recv2; + Ok(Kk1 { + header, + meta, + transport_params: self.local_transport_params, + skem_ciphertext, + ephemeral: public, + }) + } + + pub fn write_2( + &mut self, + crypto: &impl QlCrypto, + meta: HandshakeMeta, + ) -> Result { + if self.step != KkStep::Send2 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), meta)?; + let header = self.outbound_header(); + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Kk2, + meta, + self.local_transport_params, + ); + let remote_ephemeral = self + .remote_ephemeral + .clone() + .ok_or(WireError::InvalidState)?; + let (ekem_ciphertext, ekem_secret) = + crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); + self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); + self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); + + let (skem_ciphertext, skem_secret) = + crypto.mlkem_encapsulate(&self.remote_bundle.mlkem_public_key); + let skem_ciphertext = + encrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &skem_ciphertext)?; + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + self.step = KkStep::Done; + Ok(Kk2 { + header, + meta, + transport_params: self.local_transport_params, + ekem_ciphertext, + skem_ciphertext, + }) + } + + pub fn read_1(&mut self, crypto: &impl QlCrypto, message: &Kk1) -> Result<(), WireError> { + if self.step != KkStep::Recv1 { + return Err(WireError::InvalidState); + } + initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; + self.ensure_inbound_header(message.header)?; + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + message.header, + HandshakeKind::Kk1, + message.meta, + message.transport_params, + ); + self.symmetric + .decrypt_and_hash(crypto, message.skem_ciphertext.as_bytes())?; + let skem_secret = + crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &message.skem_ciphertext); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + mix_hash_ephemeral(&mut self.symmetric, crypto, &message.ephemeral); + self.remote_ephemeral = Some(message.ephemeral.clone()); + self.remote_transport_params = Some(message.transport_params); + self.step = KkStep::Send2; + Ok(()) + } + + pub fn read_2(&mut self, crypto: &impl QlCrypto, message: &Kk2) -> Result<(), WireError> { + if self.step != KkStep::Recv2 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; + self.ensure_inbound_header(message.header)?; + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + message.header, + HandshakeKind::Kk2, + message.meta, + message.transport_params, + ); + let local_ephemeral = self + .local_ephemeral + .as_ref() + .ok_or(WireError::InvalidState)?; + self.symmetric + .mix_hash(crypto, message.ekem_ciphertext.as_bytes()); + let ekem_secret = + crypto.mlkem_decapsulate(&local_ephemeral.mlkem.private, &message.ekem_ciphertext); + self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); + + let skem_ciphertext = + decrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &message.skem_ciphertext)?; + let skem_secret = crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &skem_ciphertext); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + self.remote_transport_params = Some(message.transport_params); + self.step = KkStep::Done; + Ok(()) + } + + pub fn finalize(self, crypto: &impl QlCrypto) -> Result { + if !self.is_finished() { + return Err(WireError::InvalidState); + } + let remote_transport_params = self + .remote_transport_params + .ok_or(WireError::InvalidState)?; + Ok(finalize_handshake( + crypto, + &self.symmetric, + self.role, + self.remote_bundle, + remote_transport_params, + )) + } +} diff --git a/ql-wire/src/handshake/meta.rs b/ql-wire/src/handshake/meta.rs new file mode 100644 index 00000000..8cb0cf97 --- /dev/null +++ b/ql-wire/src/handshake/meta.rs @@ -0,0 +1,48 @@ +use crate::{codec, ByteSlice, WireEncode, WireError}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct HandshakeId(pub u32); + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct HandshakeMeta { + pub handshake_id: HandshakeId, +} + +impl codec::WireDecode for HandshakeId { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +impl WireEncode for HandshakeId { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl HandshakeMeta { + pub const WIRE_SIZE: usize = size_of::(); +} + +impl WireEncode for HandshakeMeta { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.handshake_id.encode(out); + } +} + +impl codec::WireDecode for HandshakeMeta { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + handshake_id: reader.decode()?, + }) + } +} diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs new file mode 100644 index 00000000..a9b7cf87 --- /dev/null +++ b/ql-wire/src/handshake/mod.rs @@ -0,0 +1,590 @@ +use crate::{ + codec, ByteSlice, ConnectionId, HandshakeKind, MlKemCiphertext, MlKemKeyPair, MlKemPublicKey, + Nonce, PeerBundle, QlCrypto, SessionKey, WireDecode, WireEncode, WireError, + ENCRYPTED_MESSAGE_AUTH_SIZE, QID, +}; + +mod ik; +mod kk; +mod meta; +mod pairing; +mod transport_params; +mod xx; + +pub use ik::{Ik1, Ik2, IkHandshake}; +pub use kk::{Kk1, Kk2, KkHandshake}; +pub use meta::{HandshakeId, HandshakeMeta}; +pub use pairing::{PairingId, PairingToken}; +pub use transport_params::TransportParams; +pub use xx::{Xx1, Xx2, Xx3, Xx4, XxHandshake}; + +const SHA256_BLOCK_LEN: usize = 64; +const PROTOCOL_IK: &[u8] = b"ql-wire:pq-ik:v1"; +const PROTOCOL_KK: &[u8] = b"ql-wire:pq-kk:v1"; +const PROTOCOL_XX: &[u8] = b"ql-wire:pq-xx:v1"; +const CONNECTION_ID_DOMAIN: &[u8] = b"ql-wire:conn-id:v1"; +const HANDSHAKE_PREAMBLE_DOMAIN: &[u8] = b"ql-wire:handshake-preamble:v1"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct HandshakeHeader { + pub sender: QID, + pub recipient: QID, +} + +impl HandshakeHeader { + pub const WIRE_SIZE: usize = QID::SIZE * 2; +} + +impl WireEncode for HandshakeHeader { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.sender.encode(out); + self.recipient.encode(out); + } +} + +impl codec::WireDecode for HandshakeHeader { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + sender: reader.decode()?, + recipient: reader.decode()?, + }) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct EphemeralPublicKey { + pub mlkem_public_key: MlKemPublicKey, +} + +impl EphemeralPublicKey { + pub const WIRE_SIZE: usize = MlKemPublicKey::SIZE; +} + +impl WireEncode for EphemeralPublicKey { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.mlkem_public_key.encode(out); + } +} + +impl codec::WireDecode for EphemeralPublicKey { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + mlkem_public_key: reader.decode()?, + }) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct EncryptedMlKemCiphertext(pub Box<[u8; Self::WIRE_SIZE]>); + +impl EncryptedMlKemCiphertext { + pub const WIRE_SIZE: usize = MlKemCiphertext::SIZE + ENCRYPTED_MESSAGE_AUTH_SIZE; + + pub fn new(data: Box<[u8; Self::WIRE_SIZE]>) -> Self { + Self(data) + } + + pub fn as_bytes(&self) -> &[u8; Self::WIRE_SIZE] { + self.0.as_ref() + } +} + +impl WireEncode for EncryptedMlKemCiphertext { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.0.as_ref().encode(out); + } +} + +impl codec::WireDecode for EncryptedMlKemCiphertext { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self::new(reader.decode()?)) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct EncryptedPeerBundle(pub Box<[u8]>); + +impl EncryptedPeerBundle { + pub const MAX_WIRE_SIZE: usize = PeerBundle::MAX_WIRE_SIZE + ENCRYPTED_MESSAGE_AUTH_SIZE; + + pub fn as_bytes(&self) -> &[u8] { + self.0.as_ref() + } +} + +impl WireEncode for EncryptedPeerBundle { + fn encoded_len(&self) -> usize { + self.0.len() + } + + fn encode(&self, out: &mut W) { + self.as_bytes().encode(out); + } +} + +impl codec::WireDecode for EncryptedPeerBundle { + fn decode(reader: &mut codec::Reader) -> Result { + let data = reader.take_rest(); + if data.len() > Self::MAX_WIRE_SIZE { + return Err(WireError::InvalidPayload); + } + Ok(Self(data.to_vec().into_boxed_slice())) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct FinalizedHandshake { + pub tx_key: SessionKey, + pub rx_key: SessionKey, + pub tx_connection_id: ConnectionId, + pub rx_connection_id: ConnectionId, + pub handshake_hash: [u8; 32], + pub remote_bundle: PeerBundle, + /// Transport parameters advertised by the remote peer + pub remote_transport_params: TransportParams, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Role { + Initiator, + Responder, +} + +#[derive(Debug, Clone)] +struct EphemeralKeyPair { + mlkem: MlKemKeyPair, +} + +impl EphemeralKeyPair { + fn public(&self) -> EphemeralPublicKey { + EphemeralPublicKey { + mlkem_public_key: self.mlkem.public.clone(), + } + } +} + +#[derive(Debug, Clone)] +struct CipherState { + key: Option, + nonce: u64, +} + +impl CipherState { + fn new() -> Self { + Self { + key: None, + nonce: 0, + } + } + + fn initialize_key(&mut self, key: SessionKey) { + self.key = Some(key); + self.nonce = 0; + } + + fn has_key(&self) -> bool { + self.key.is_some() + } + + fn encrypt( + &mut self, + crypto: &impl QlCrypto, + aad: &[u8], + plaintext: &[u8], + ) -> Result, WireError> { + let key = self.key.as_ref().ok_or(WireError::InvalidState)?; + let nonce = Nonce::from_counter(self.nonce); + let mut ciphertext = Vec::with_capacity(plaintext.len() + ENCRYPTED_MESSAGE_AUTH_SIZE); + ciphertext.extend_from_slice(plaintext); + let auth = crypto.aes256_gcm_encrypt(key, &nonce, aad, &mut ciphertext); + self.nonce = self.nonce.wrapping_add(1); + ciphertext.extend_from_slice(&auth); + Ok(ciphertext) + } + + fn decrypt( + &mut self, + crypto: &impl QlCrypto, + aad: &[u8], + ciphertext: &[u8], + ) -> Result, WireError> { + if ciphertext.len() < ENCRYPTED_MESSAGE_AUTH_SIZE { + return Err(WireError::InvalidPayload); + } + let split = ciphertext.len() - ENCRYPTED_MESSAGE_AUTH_SIZE; + let (ciphertext, auth) = ciphertext.split_at(split); + let mut plaintext = ciphertext.to_vec(); + let key = self.key.as_ref().ok_or(WireError::InvalidState)?; + let nonce = Nonce::from_counter(self.nonce); + let mut auth_tag = [0u8; ENCRYPTED_MESSAGE_AUTH_SIZE]; + auth_tag.copy_from_slice(auth); + if !crypto.aes256_gcm_decrypt(key, &nonce, aad, &mut plaintext, &auth_tag) { + return Err(WireError::DecryptFailed); + } + self.nonce = self.nonce.wrapping_add(1); + Ok(plaintext) + } +} + +#[derive(Debug, Clone)] +struct SymmetricState { + chaining_key: [u8; 32], + handshake_hash: [u8; 32], + cipher: CipherState, +} + +impl SymmetricState { + fn new(crypto: &impl QlCrypto, protocol_name: &[u8]) -> Self { + let h = crypto.sha256(&[protocol_name]); + Self { + chaining_key: h, + handshake_hash: h, + cipher: CipherState::new(), + } + } + + fn mix_hash(&mut self, crypto: &impl QlCrypto, data: &[u8]) { + self.handshake_hash = crypto.sha256(&[&self.handshake_hash, data]); + } + + fn mix_key(&mut self, crypto: &impl QlCrypto, input_key_material: &[u8]) { + let (chaining_key, cipher_key) = hkdf2(crypto, &self.chaining_key, input_key_material); + self.chaining_key = chaining_key; + self.cipher.initialize_key(cipher_key); + } + + fn mix_key_and_hash(&mut self, crypto: &impl QlCrypto, input_key_material: &[u8]) { + let (chaining_key, hash_input, cipher_key) = + hkdf3(crypto, &self.chaining_key, input_key_material); + self.chaining_key = chaining_key; + self.mix_hash(crypto, &hash_input); + self.cipher.initialize_key(cipher_key); + } + + fn encrypt_and_hash( + &mut self, + crypto: &impl QlCrypto, + plaintext: &[u8], + ) -> Result, WireError> { + if self.cipher.has_key() { + let ciphertext = self + .cipher + .encrypt(crypto, &self.handshake_hash, plaintext)?; + self.mix_hash(crypto, &ciphertext); + Ok(ciphertext) + } else { + self.mix_hash(crypto, plaintext); + Ok(plaintext.to_vec()) + } + } + + fn decrypt_and_hash( + &mut self, + crypto: &impl QlCrypto, + ciphertext: &[u8], + ) -> Result, WireError> { + if self.cipher.has_key() { + let plaintext = self + .cipher + .decrypt(crypto, &self.handshake_hash, ciphertext)?; + self.mix_hash(crypto, ciphertext); + Ok(plaintext) + } else { + self.mix_hash(crypto, ciphertext); + Ok(ciphertext.to_vec()) + } + } + + fn split_for_role(&self, crypto: &impl QlCrypto, role: Role) -> (SessionKey, SessionKey) { + let temp_key = hmac_sha256(crypto, &self.chaining_key, &[&[]]); + let k1 = SessionKey::from_data(hmac_sha256(crypto, &temp_key, &[&[1]])); + let k2 = SessionKey::from_data(hmac_sha256(crypto, &temp_key, &[k1.as_bytes(), &[2]])); + match role { + Role::Initiator => (k1, k2), + Role::Responder => (k2, k1), + } + } +} + +fn init_kk_symmetric( + crypto: &impl QlCrypto, + initiator_bundle: &PeerBundle, + responder_bundle: &PeerBundle, +) -> SymmetricState { + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_KK); + symmetric.mix_hash(crypto, &initiator_bundle.encode_vec()); + symmetric.mix_hash(crypto, &responder_bundle.encode_vec()); + symmetric +} + +fn init_ik_symmetric(crypto: &impl QlCrypto, responder_bundle: &PeerBundle) -> SymmetricState { + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_IK); + symmetric.mix_hash(crypto, &responder_bundle.encode_vec()); + symmetric +} + +fn init_xx_symmetric(crypto: &impl QlCrypto) -> SymmetricState { + SymmetricState::new(crypto, PROTOCOL_XX) +} + +fn mix_psk_pairing_token( + symmetric: &mut SymmetricState, + crypto: &impl QlCrypto, + pairing_token: PairingToken, +) { + symmetric.mix_key_and_hash(crypto, &pairing_token.psk(crypto)); +} + +fn generate_ephemeral_keypair(crypto: &impl QlCrypto) -> EphemeralKeyPair { + EphemeralKeyPair { + mlkem: crypto.mlkem_generate_keypair(), + } +} + +fn mix_hash_ephemeral( + symmetric: &mut SymmetricState, + crypto: &impl QlCrypto, + public: &EphemeralPublicKey, +) { + symmetric.mix_hash(crypto, public.mlkem_public_key.as_bytes()); +} + +fn mix_hash_routed_handshake( + symmetric: &mut SymmetricState, + crypto: &impl QlCrypto, + header: HandshakeHeader, + kind: HandshakeKind, + meta: HandshakeMeta, + transport_params: TransportParams, +) { + mix_hash_handshake_preamble( + symmetric, + crypto, + &header.encode_vec(), + kind, + meta, + transport_params, + ); +} + +fn mix_hash_pairing_handshake( + symmetric: &mut SymmetricState, + crypto: &impl QlCrypto, + header: HandshakeHeader, + kind: HandshakeKind, + meta: HandshakeMeta, + pairing_id: PairingId, + transport_params: TransportParams, +) { + let mut preamble = header.encode_vec(); + pairing_id.encode(&mut preamble); + mix_hash_handshake_preamble(symmetric, crypto, &preamble, kind, meta, transport_params); +} + +fn mix_hash_handshake_preamble( + symmetric: &mut SymmetricState, + crypto: &impl QlCrypto, + header: &[u8], + kind: HandshakeKind, + meta: HandshakeMeta, + transport_params: TransportParams, +) { + symmetric.mix_hash(crypto, HANDSHAKE_PREAMBLE_DOMAIN); + symmetric.mix_hash(crypto, header); + symmetric.mix_hash(crypto, &[kind as u8]); + symmetric.mix_hash(crypto, &meta.encode_vec()); + symmetric.mix_hash(crypto, &transport_params.encode_vec()); +} + +fn initialize_handshake_meta( + expected: &mut Option, + meta: HandshakeMeta, +) -> Result<(), WireError> { + match expected { + Some(stored) if *stored != meta => Err(WireError::InvalidHandshakeMeta), + Some(_) => Ok(()), + None => { + *expected = Some(meta); + Ok(()) + } + } +} + +fn require_handshake_meta( + expected: Option<&HandshakeMeta>, + meta: HandshakeMeta, +) -> Result<(), WireError> { + match expected { + Some(stored) if *stored == meta => Ok(()), + _ => Err(WireError::InvalidHandshakeMeta), + } +} + +fn initialize_transport_params( + expected: &mut Option, + transport_params: TransportParams, +) -> Result<(), WireError> { + match expected { + Some(stored) if *stored != transport_params => Err(WireError::InvalidTransportParams), + Some(_) => Ok(()), + None => { + *expected = Some(transport_params); + Ok(()) + } + } +} + +fn require_transport_params( + expected: Option<&TransportParams>, + transport_params: TransportParams, +) -> Result<(), WireError> { + match expected { + Some(stored) if *stored == transport_params => Ok(()), + _ => Err(WireError::InvalidTransportParams), + } +} + +fn encrypt_peer_bundle( + crypto: &impl QlCrypto, + symmetric: &mut SymmetricState, + bundle: &PeerBundle, +) -> Result { + let ciphertext = symmetric.encrypt_and_hash(crypto, &bundle.encode_vec())?; + Ok(EncryptedPeerBundle(ciphertext.into_boxed_slice())) +} + +fn decrypt_peer_bundle( + crypto: &impl QlCrypto, + symmetric: &mut SymmetricState, + bundle: &EncryptedPeerBundle, +) -> Result { + let plaintext = symmetric.decrypt_and_hash(crypto, bundle.as_bytes())?; + let bundle = PeerBundle::decode_exact(plaintext.as_slice())?; + if !bundle.qid_matches_public_key(crypto) { + return Err(WireError::InvalidRemoteBundle); + } + Ok(bundle) +} + +fn encrypt_mlkem_ciphertext( + crypto: &impl QlCrypto, + symmetric: &mut SymmetricState, + ciphertext: &MlKemCiphertext, +) -> Result { + let encrypted = symmetric.encrypt_and_hash(crypto, ciphertext.as_bytes())?; + let out: Box<[u8; EncryptedMlKemCiphertext::WIRE_SIZE]> = + encrypted.try_into().map_err(|_| WireError::InvalidState)?; + Ok(EncryptedMlKemCiphertext::new(out)) +} + +fn decrypt_mlkem_ciphertext( + crypto: &impl QlCrypto, + symmetric: &mut SymmetricState, + ciphertext: &EncryptedMlKemCiphertext, +) -> Result { + let plaintext = symmetric.decrypt_and_hash(crypto, ciphertext.as_bytes())?; + let out: Box<[u8; MlKemCiphertext::SIZE]> = plaintext + .try_into() + .map_err(|_| WireError::InvalidPayload)?; + Ok(MlKemCiphertext::new(out)) +} + +fn finalize_handshake( + crypto: &impl QlCrypto, + symmetric: &SymmetricState, + role: Role, + remote_bundle: PeerBundle, + remote_transport_params: TransportParams, +) -> FinalizedHandshake { + let handshake_hash = symmetric.handshake_hash; + let (tx_key, rx_key) = symmetric.split_for_role(crypto, role); + let (initiator_rx, responder_rx) = derive_connection_ids(crypto, &handshake_hash); + let (tx_connection_id, rx_connection_id) = match role { + Role::Initiator => (responder_rx, initiator_rx), + Role::Responder => (initiator_rx, responder_rx), + }; + FinalizedHandshake { + tx_key, + rx_key, + tx_connection_id, + rx_connection_id, + handshake_hash, + remote_bundle, + remote_transport_params, + } +} + +fn derive_connection_ids( + crypto: &impl QlCrypto, + handshake_hash: &[u8; 32], +) -> (ConnectionId, ConnectionId) { + let initiator = crypto.sha256(&[CONNECTION_ID_DOMAIN, handshake_hash, b"initiator-rx"]); + let responder = crypto.sha256(&[CONNECTION_ID_DOMAIN, handshake_hash, b"responder-rx"]); + let mut initiator_rx = [0u8; ConnectionId::SIZE]; + let mut responder_rx = [0u8; ConnectionId::SIZE]; + initiator_rx.copy_from_slice(&initiator[..ConnectionId::SIZE]); + responder_rx.copy_from_slice(&responder[..ConnectionId::SIZE]); + ( + ConnectionId::from_data(initiator_rx), + ConnectionId::from_data(responder_rx), + ) +} + +fn hkdf2( + crypto: &impl QlCrypto, + chaining_key: &[u8; 32], + input_key_material: &[u8], +) -> ([u8; 32], SessionKey) { + let temp_key = hmac_sha256(crypto, chaining_key, &[input_key_material]); + let out1 = hmac_sha256(crypto, &temp_key, &[&[1]]); + let out2 = hmac_sha256(crypto, &temp_key, &[&out1, &[2]]); + (out1, SessionKey::from_data(out2)) +} + +fn hkdf3( + crypto: &impl QlCrypto, + chaining_key: &[u8; 32], + input_key_material: &[u8], +) -> ([u8; 32], [u8; 32], SessionKey) { + let temp_key = hmac_sha256(crypto, chaining_key, &[input_key_material]); + let out1 = hmac_sha256(crypto, &temp_key, &[&[1]]); + let out2 = hmac_sha256(crypto, &temp_key, &[&out1, &[2]]); + let out3 = hmac_sha256(crypto, &temp_key, &[&out2, &[3]]); + (out1, out2, SessionKey::from_data(out3)) +} + +fn hmac_sha256(crypto: &impl QlCrypto, key: &[u8], parts: &[&[u8]]) -> [u8; 32] { + let mut key_block = [0u8; SHA256_BLOCK_LEN]; + if key.len() > SHA256_BLOCK_LEN { + key_block[..32].copy_from_slice(&crypto.sha256(&[key])); + } else { + key_block[..key.len()].copy_from_slice(key); + } + + let mut ipad = [0x36u8; SHA256_BLOCK_LEN]; + let mut opad = [0x5cu8; SHA256_BLOCK_LEN]; + for (dst, src) in ipad.iter_mut().zip(key_block.iter()) { + *dst ^= *src; + } + for (dst, src) in opad.iter_mut().zip(key_block.iter()) { + *dst ^= *src; + } + + let mut inner_parts: Vec<&[u8]> = Vec::with_capacity(parts.len() + 1); + inner_parts.push(&ipad); + inner_parts.extend_from_slice(parts); + let inner = crypto.sha256(&inner_parts); + crypto.sha256(&[&opad, &inner]) +} diff --git a/ql-wire/src/handshake/pairing.rs b/ql-wire/src/handshake/pairing.rs new file mode 100644 index 00000000..237f066b --- /dev/null +++ b/ql-wire/src/handshake/pairing.rs @@ -0,0 +1,83 @@ +use std::fmt::{self, Display, Formatter}; + +use crate::{codec, ByteSlice, QlCrypto, WireEncode, WireError}; + +const PAIRING_ID_DOMAIN: &[u8] = b"ql-wire:pairing-id:v1"; +const PAIRING_PSK_DOMAIN: &[u8] = b"ql-wire:pairing-psk:v1"; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct PairingToken(pub [u8; Self::SIZE]); + +impl PairingToken { + pub const SIZE: usize = 16; + + pub fn id(&self, crypto: &impl QlCrypto) -> PairingId { + let hash = crypto.sha256(&[PAIRING_ID_DOMAIN, &self.0]); + let mut id = [0u8; PairingId::SIZE]; + id.copy_from_slice(&hash[..PairingId::SIZE]); + PairingId(id) + } + + pub(super) fn psk(&self, crypto: &impl QlCrypto) -> [u8; 32] { + crypto.sha256(&[PAIRING_PSK_DOMAIN, &self.0]) + } +} + +impl Display for PairingToken { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + for byte in self.0 { + write!(f, "{byte:02x}")?; + } + Ok(()) + } +} + +impl WireEncode for PairingToken { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl codec::WireDecode for PairingToken { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct PairingId(pub [u8; Self::SIZE]); + +impl PairingId { + pub const SIZE: usize = 16; +} + +impl Display for PairingId { + fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { + for byte in self.0 { + write!(f, "{byte:02x}")?; + } + Ok(()) + } +} + +impl WireEncode for PairingId { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl codec::WireDecode for PairingId { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} diff --git a/ql-wire/src/handshake/transport_params.rs b/ql-wire/src/handshake/transport_params.rs new file mode 100644 index 00000000..bfd0d427 --- /dev/null +++ b/ql-wire/src/handshake/transport_params.rs @@ -0,0 +1,38 @@ +use crate::{codec, ByteSlice, WireEncode, WireError}; + +/// Session parameters advertised in the handshake +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct TransportParams { + /// Initial per-stream receive credit granted to the remote peer + pub initial_stream_receive_window: u32, +} + +impl TransportParams { + pub const WIRE_SIZE: usize = size_of::(); +} + +impl WireEncode for TransportParams { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.initial_stream_receive_window.encode(out); + } +} + +impl Default for TransportParams { + fn default() -> Self { + Self { + initial_stream_receive_window: 16 * 1024, + } + } +} + +impl codec::WireDecode for TransportParams { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + initial_stream_receive_window: reader.decode()?, + }) + } +} diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs new file mode 100644 index 00000000..0b6452d4 --- /dev/null +++ b/ql-wire/src/handshake/xx.rs @@ -0,0 +1,612 @@ +use super::{ + decrypt_mlkem_ciphertext, decrypt_peer_bundle, encrypt_mlkem_ciphertext, encrypt_peer_bundle, + finalize_handshake, generate_ephemeral_keypair, init_xx_symmetric, initialize_handshake_meta, + initialize_transport_params, mix_hash_ephemeral, mix_hash_pairing_handshake, + mix_psk_pairing_token, require_handshake_meta, require_transport_params, + EncryptedMlKemCiphertext, EncryptedPeerBundle, EphemeralKeyPair, EphemeralPublicKey, + FinalizedHandshake, HandshakeHeader, Role, SymmetricState, TransportParams, +}; +use crate::{ + codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PairingId, PairingToken, + PeerBundle, QlCrypto, QlIdentity, WireEncode, WireError, QID, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Xx1 { + pub header: HandshakeHeader, + pub meta: HandshakeMeta, + pub pairing_id: PairingId, + pub transport_params: TransportParams, + pub ephemeral: EphemeralPublicKey, +} + +impl Xx1 { + pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE + + HandshakeMeta::WIRE_SIZE + + PairingId::SIZE + + TransportParams::WIRE_SIZE + + EphemeralPublicKey::WIRE_SIZE; +} + +impl codec::WireDecode for Xx1 { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + meta: reader.decode()?, + pairing_id: reader.decode()?, + transport_params: reader.decode()?, + ephemeral: reader.decode()?, + }) + } +} + +impl WireEncode for Xx1 { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.meta.encode(out); + self.pairing_id.encode(out); + self.transport_params.encode(out); + self.ephemeral.encode(out); + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Xx2 { + pub header: HandshakeHeader, + pub meta: HandshakeMeta, + pub pairing_id: PairingId, + pub transport_params: TransportParams, + pub ekem_ciphertext: MlKemCiphertext, + pub static_bundle: EncryptedPeerBundle, +} + +impl codec::WireDecode for Xx2 { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + meta: reader.decode()?, + pairing_id: reader.decode()?, + transport_params: reader.decode()?, + ekem_ciphertext: reader.decode()?, + static_bundle: reader.decode()?, + }) + } +} + +impl WireEncode for Xx2 { + fn encoded_len(&self) -> usize { + HandshakeHeader::WIRE_SIZE + + HandshakeMeta::WIRE_SIZE + + PairingId::SIZE + + TransportParams::WIRE_SIZE + + MlKemCiphertext::SIZE + + self.static_bundle.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.meta.encode(out); + self.pairing_id.encode(out); + self.transport_params.encode(out); + self.ekem_ciphertext.encode(out); + self.static_bundle.encode(out); + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Xx3 { + pub header: HandshakeHeader, + pub meta: HandshakeMeta, + pub pairing_id: PairingId, + pub transport_params: TransportParams, + pub skem_ciphertext: EncryptedMlKemCiphertext, + pub static_bundle: EncryptedPeerBundle, +} + +impl codec::WireDecode for Xx3 { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + meta: reader.decode()?, + pairing_id: reader.decode()?, + transport_params: reader.decode()?, + skem_ciphertext: reader.decode()?, + static_bundle: reader.decode()?, + }) + } +} + +impl WireEncode for Xx3 { + fn encoded_len(&self) -> usize { + HandshakeHeader::WIRE_SIZE + + HandshakeMeta::WIRE_SIZE + + PairingId::SIZE + + TransportParams::WIRE_SIZE + + EncryptedMlKemCiphertext::WIRE_SIZE + + self.static_bundle.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.meta.encode(out); + self.pairing_id.encode(out); + self.transport_params.encode(out); + self.skem_ciphertext.encode(out); + self.static_bundle.encode(out); + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Xx4 { + pub header: HandshakeHeader, + pub meta: HandshakeMeta, + pub pairing_id: PairingId, + pub transport_params: TransportParams, + pub skem_ciphertext: EncryptedMlKemCiphertext, +} + +impl Xx4 { + pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE + + HandshakeMeta::WIRE_SIZE + + PairingId::SIZE + + TransportParams::WIRE_SIZE + + EncryptedMlKemCiphertext::WIRE_SIZE; +} + +impl codec::WireDecode for Xx4 { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + meta: reader.decode()?, + pairing_id: reader.decode()?, + transport_params: reader.decode()?, + skem_ciphertext: reader.decode()?, + }) + } +} + +impl WireEncode for Xx4 { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.meta.encode(out); + self.pairing_id.encode(out); + self.transport_params.encode(out); + self.skem_ciphertext.encode(out); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum XxStep { + Send1, + Recv1, + Send2, + Recv2, + Send3, + Recv3, + Send4, + Recv4, + Done, +} + +#[derive(Debug, Clone)] +pub struct XxHandshake { + role: Role, + step: XxStep, + symmetric: SymmetricState, + local: QlIdentity, + remote_qid: QID, + pairing_token: PairingToken, + remote_bundle: Option, + local_ephemeral: Option, + remote_ephemeral: Option, + handshake_meta: Option, + local_transport_params: TransportParams, + remote_transport_params: Option, +} + +impl XxHandshake { + pub fn new_initiator( + crypto: &impl QlCrypto, + local: QlIdentity, + remote_qid: QID, + pairing_token: PairingToken, + local_transport_params: TransportParams, + ) -> Self { + Self { + role: Role::Initiator, + step: XxStep::Send1, + symmetric: init_xx_symmetric(crypto), + local, + remote_qid, + pairing_token, + remote_bundle: None, + local_ephemeral: None, + remote_ephemeral: None, + handshake_meta: None, + local_transport_params, + remote_transport_params: None, + } + } + + pub fn new_responder( + crypto: &impl QlCrypto, + local: QlIdentity, + remote_qid: QID, + pairing_token: PairingToken, + local_transport_params: TransportParams, + ) -> Self { + Self { + role: Role::Responder, + step: XxStep::Recv1, + symmetric: init_xx_symmetric(crypto), + local, + remote_qid, + pairing_token, + remote_bundle: None, + local_ephemeral: None, + remote_ephemeral: None, + handshake_meta: None, + local_transport_params, + remote_transport_params: None, + } + } + + pub fn is_finished(&self) -> bool { + self.step == XxStep::Done + } + + pub fn pairing_token(&self) -> PairingToken { + self.pairing_token + } + + pub fn pairing_id(&self, crypto: &impl QlCrypto) -> PairingId { + self.pairing_token.id(crypto) + } + + pub fn remote_qid(&self) -> QID { + self.remote_qid + } + + pub fn remote_bundle(&self) -> Option<&PeerBundle> { + self.remote_bundle.as_ref() + } + + fn header(&self) -> HandshakeHeader { + HandshakeHeader { + sender: self.local.qid, + recipient: self.remote_qid, + } + } + + fn ensure_inbound_header( + &self, + crypto: &impl QlCrypto, + header: HandshakeHeader, + pairing_id: PairingId, + ) -> Result<(), WireError> { + if header.sender != self.remote_qid || header.recipient != self.local.qid { + return Err(WireError::InvalidHandshakeHeader); + } + if pairing_id != self.pairing_token.id(crypto) { + return Err(WireError::InvalidPairingId); + } + Ok(()) + } + + fn ensure_remote_bundle(&self, bundle: &PeerBundle) -> Result<(), WireError> { + if bundle.qid == self.remote_qid { + Ok(()) + } else { + Err(WireError::InvalidRemoteBundle) + } + } + + pub fn write_1( + &mut self, + crypto: &impl QlCrypto, + meta: HandshakeMeta, + ) -> Result { + if self.step != XxStep::Send1 { + return Err(WireError::InvalidState); + } + initialize_handshake_meta(&mut self.handshake_meta, meta)?; + let header = self.header(); + let pairing_id = self.pairing_token.id(crypto); + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Xx1, + meta, + pairing_id, + self.local_transport_params, + ); + mix_psk_pairing_token(&mut self.symmetric, crypto, self.pairing_token); + + let local_ephemeral = generate_ephemeral_keypair(crypto); + let ephemeral = local_ephemeral.public(); + mix_hash_ephemeral(&mut self.symmetric, crypto, &ephemeral); + + self.local_ephemeral = Some(local_ephemeral); + self.step = XxStep::Recv2; + Ok(Xx1 { + header, + meta, + pairing_id, + transport_params: self.local_transport_params, + ephemeral, + }) + } + + pub fn read_1(&mut self, crypto: &impl QlCrypto, message: &Xx1) -> Result<(), WireError> { + if self.step != XxStep::Recv1 { + return Err(WireError::InvalidState); + } + initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; + self.ensure_inbound_header(crypto, message.header, message.pairing_id)?; + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + message.header, + HandshakeKind::Xx1, + message.meta, + message.pairing_id, + message.transport_params, + ); + mix_psk_pairing_token(&mut self.symmetric, crypto, self.pairing_token); + mix_hash_ephemeral(&mut self.symmetric, crypto, &message.ephemeral); + + self.remote_ephemeral = Some(message.ephemeral.clone()); + initialize_transport_params(&mut self.remote_transport_params, message.transport_params)?; + self.step = XxStep::Send2; + Ok(()) + } + + pub fn write_2( + &mut self, + crypto: &impl QlCrypto, + meta: HandshakeMeta, + ) -> Result { + if self.step != XxStep::Send2 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), meta)?; + let header = self.header(); + let pairing_id = self.pairing_token.id(crypto); + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Xx2, + meta, + pairing_id, + self.local_transport_params, + ); + + let remote_ephemeral = self + .remote_ephemeral + .as_ref() + .ok_or(WireError::InvalidState)?; + let (ekem_ciphertext, ekem_secret) = + crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); + self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); + self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); + + let static_bundle = encrypt_peer_bundle(crypto, &mut self.symmetric, &self.local.bundle())?; + + self.step = XxStep::Recv3; + Ok(Xx2 { + header, + meta, + pairing_id, + transport_params: self.local_transport_params, + ekem_ciphertext, + static_bundle, + }) + } + + pub fn read_2(&mut self, crypto: &impl QlCrypto, message: &Xx2) -> Result<(), WireError> { + if self.step != XxStep::Recv2 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; + self.ensure_inbound_header(crypto, message.header, message.pairing_id)?; + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + message.header, + HandshakeKind::Xx2, + message.meta, + message.pairing_id, + message.transport_params, + ); + + let local_ephemeral = self + .local_ephemeral + .as_ref() + .ok_or(WireError::InvalidState)?; + self.symmetric + .mix_hash(crypto, message.ekem_ciphertext.as_bytes()); + let ekem_secret = + crypto.mlkem_decapsulate(&local_ephemeral.mlkem.private, &message.ekem_ciphertext); + self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); + + let remote_bundle = + decrypt_peer_bundle(crypto, &mut self.symmetric, &message.static_bundle)?; + self.ensure_remote_bundle(&remote_bundle)?; + self.remote_bundle = Some(remote_bundle); + initialize_transport_params(&mut self.remote_transport_params, message.transport_params)?; + self.step = XxStep::Send3; + Ok(()) + } + + pub fn write_3( + &mut self, + crypto: &impl QlCrypto, + meta: HandshakeMeta, + ) -> Result { + if self.step != XxStep::Send3 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), meta)?; + let header = self.header(); + let pairing_id = self.pairing_token.id(crypto); + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Xx3, + meta, + pairing_id, + self.local_transport_params, + ); + + let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + let (skem_ciphertext, skem_secret) = + crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); + let skem_ciphertext = + encrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &skem_ciphertext)?; + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + let static_bundle = encrypt_peer_bundle(crypto, &mut self.symmetric, &self.local.bundle())?; + + self.step = XxStep::Recv4; + Ok(Xx3 { + header, + meta, + pairing_id, + transport_params: self.local_transport_params, + skem_ciphertext, + static_bundle, + }) + } + + pub fn read_3(&mut self, crypto: &impl QlCrypto, message: &Xx3) -> Result<(), WireError> { + if self.step != XxStep::Recv3 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; + self.ensure_inbound_header(crypto, message.header, message.pairing_id)?; + require_transport_params( + self.remote_transport_params.as_ref(), + message.transport_params, + )?; + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + message.header, + HandshakeKind::Xx3, + message.meta, + message.pairing_id, + message.transport_params, + ); + + let skem_ciphertext = + decrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &message.skem_ciphertext)?; + let skem_secret = crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &skem_ciphertext); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + let remote_bundle = + decrypt_peer_bundle(crypto, &mut self.symmetric, &message.static_bundle)?; + self.ensure_remote_bundle(&remote_bundle)?; + self.remote_bundle = Some(remote_bundle); + self.step = XxStep::Send4; + Ok(()) + } + + pub fn write_4( + &mut self, + crypto: &impl QlCrypto, + meta: HandshakeMeta, + ) -> Result { + if self.step != XxStep::Send4 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), meta)?; + let header = self.header(); + let pairing_id = self.pairing_token.id(crypto); + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Xx4, + meta, + pairing_id, + self.local_transport_params, + ); + + let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + let (skem_ciphertext, skem_secret) = + crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); + let skem_ciphertext = + encrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &skem_ciphertext)?; + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + self.step = XxStep::Done; + Ok(Xx4 { + header, + meta, + pairing_id, + transport_params: self.local_transport_params, + skem_ciphertext, + }) + } + + pub fn read_4(&mut self, crypto: &impl QlCrypto, message: &Xx4) -> Result<(), WireError> { + if self.step != XxStep::Recv4 { + return Err(WireError::InvalidState); + } + require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; + self.ensure_inbound_header(crypto, message.header, message.pairing_id)?; + require_transport_params( + self.remote_transport_params.as_ref(), + message.transport_params, + )?; + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + message.header, + HandshakeKind::Xx4, + message.meta, + message.pairing_id, + message.transport_params, + ); + + let skem_ciphertext = + decrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &message.skem_ciphertext)?; + let skem_secret = crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &skem_ciphertext); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + self.step = XxStep::Done; + Ok(()) + } + + pub fn finalize(self, crypto: &impl QlCrypto) -> Result { + if !self.is_finished() { + return Err(WireError::InvalidState); + } + let remote_bundle = self.remote_bundle.ok_or(WireError::InvalidState)?; + let remote_transport_params = self + .remote_transport_params + .ok_or(WireError::InvalidState)?; + Ok(finalize_handshake( + crypto, + &self.symmetric, + self.role, + remote_bundle, + remote_transport_params, + )) + } +} diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs new file mode 100644 index 00000000..88764c0a --- /dev/null +++ b/ql-wire/src/header.rs @@ -0,0 +1,121 @@ +use ::bytes::BufMut; + +use crate::{ + codec, ByteSlice, VarInt, VarIntBoundsExceeded, WireEncode, WireError, QL_WIRE_VERSION, +}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct SessionHeader { + pub connection_id: ConnectionId, + pub seq: RecordSeq, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct RecordSeq(pub VarInt); + +impl RecordSeq { + pub const MAX_ENCODED_LEN: usize = VarInt::MAX_SIZE; + + pub const fn from_u32(value: u32) -> Self { + Self(VarInt::from_u32(value)) + } + + pub fn from_u64(value: u64) -> Result { + Ok(Self(VarInt::from_u64(value)?)) + } + + pub const fn into_inner(self) -> u64 { + self.0.into_inner() + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct ConnectionId(pub [u8; Self::SIZE]); + +impl ConnectionId { + pub const SIZE: usize = 16; + + pub const fn from_data(data: [u8; Self::SIZE]) -> Self { + Self(data) + } + + pub const fn as_bytes(&self) -> &[u8; Self::SIZE] { + &self.0 + } +} + +impl codec::WireDecode for RecordSeq { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +impl WireEncode for RecordSeq { + fn encoded_len(&self) -> usize { + self.0.size() + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl codec::WireDecode for ConnectionId { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self::from_data(reader.decode()?)) + } +} + +impl WireEncode for ConnectionId { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl SessionHeader { + pub const MAX_ENCODED_LEN: usize = ConnectionId::SIZE + RecordSeq::MAX_ENCODED_LEN; + const AAD_DOMAIN: &[u8] = b"ql-wire:session-aad:v1"; + const AAD_RECORD_KIND_SESSION: u8 = 1; + + pub fn aad(&self) -> Vec { + let aad_len = Self::AAD_DOMAIN.len() + + size_of::() + + size_of::() + + ConnectionId::SIZE + + self.seq.encoded_len(); + let mut aad = Vec::with_capacity(aad_len); + aad.put_slice(Self::AAD_DOMAIN); + aad.put_u8(QL_WIRE_VERSION); + aad.put_u8(Self::AAD_RECORD_KIND_SESSION); + self.connection_id.encode(&mut aad); + self.seq.encode(&mut aad); + debug_assert_eq!(aad.len(), aad_len); + aad + } +} + +impl WireEncode for SessionHeader { + fn encoded_len(&self) -> usize { + ConnectionId::SIZE + self.seq.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.connection_id.encode(out); + self.seq.encode(out); + } +} + +impl codec::WireDecode for SessionHeader { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + connection_id: reader.decode()?, + seq: reader.decode()?, + }) + } +} diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs new file mode 100644 index 00000000..bdc54b2a --- /dev/null +++ b/ql-wire/src/identity.rs @@ -0,0 +1,192 @@ +use std::ops::Deref; + +use crate::{ + codec, ByteSlice, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash, VarInt, + WireEncode, WireError, QID, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PeerBundle { + pub version: u16, + pub qid: QID, + pub capabilities: u32, + pub mlkem_public_key: MlKemPublicKey, + pub name: QlName, +} + +impl PeerBundle { + pub const VERSION: u16 = 1; + pub const FIXED_WIRE_SIZE: usize = + size_of::() + QID::SIZE + size_of::() + MlKemPublicKey::SIZE; + pub const MAX_WIRE_SIZE: usize = Self::FIXED_WIRE_SIZE + VarInt::MAX_SIZE + QlName::MAX_LEN; + + pub fn qid_matches_public_key(&self, crypto: &impl QlHash) -> bool { + self.qid.matches_public_key(crypto, &self.mlkem_public_key) + } +} + +impl WireEncode for PeerBundle { + fn encoded_len(&self) -> usize { + Self::FIXED_WIRE_SIZE + self.name.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.version.encode(out); + self.qid.encode(out); + self.capabilities.encode(out); + self.mlkem_public_key.encode(out); + self.name.encode(out); + } +} + +impl codec::WireDecode for PeerBundle { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + version: reader.decode()?, + qid: reader.decode()?, + capabilities: reader.decode()?, + mlkem_public_key: reader.decode()?, + name: reader.decode()?, + }) + } +} + +#[derive(Debug, Clone)] +pub struct QlIdentity { + pub qid: QID, + pub mlkem_private_key: MlKemPrivateKey, + pub mlkem_public_key: MlKemPublicKey, + pub capabilities: u32, + pub name: QlName, +} + +impl QlIdentity { + pub const FIXED_WIRE_SIZE: usize = + QID::SIZE + MlKemPrivateKey::SIZE + MlKemPublicKey::SIZE + size_of::(); + pub const MAX_WIRE_SIZE: usize = Self::FIXED_WIRE_SIZE + VarInt::MAX_SIZE + QlName::MAX_LEN; + + pub fn new( + crypto: &impl QlHash, + mlkem_private_key: MlKemPrivateKey, + mlkem_public_key: MlKemPublicKey, + name: impl Into, + ) -> Result { + let name = QlName::new(name)?; + let qid = QID::derive(crypto, &mlkem_public_key); + Ok(Self { + qid, + mlkem_private_key, + mlkem_public_key, + capabilities: 0, + name, + }) + } + + #[must_use] + pub fn with_capabilities(mut self, capabilities: u32) -> Self { + self.capabilities = capabilities; + self + } + + pub fn with_name(mut self, name: impl Into) -> Result { + self.name = QlName::new(name)?; + Ok(self) + } + + pub fn bundle(&self) -> PeerBundle { + PeerBundle { + version: PeerBundle::VERSION, + qid: self.qid, + capabilities: self.capabilities, + mlkem_public_key: self.mlkem_public_key.clone(), + name: self.name.clone(), + } + } +} + +impl WireEncode for QlIdentity { + fn encoded_len(&self) -> usize { + Self::FIXED_WIRE_SIZE + self.name.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.qid.encode(out); + self.mlkem_private_key.as_bytes().encode(out); + self.mlkem_public_key.encode(out); + self.capabilities.encode(out); + self.name.encode(out); + } +} + +impl codec::WireDecode for QlIdentity { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + qid: reader.decode()?, + mlkem_private_key: MlKemPrivateKey::new(reader.decode()?), + mlkem_public_key: reader.decode()?, + capabilities: reader.decode()?, + name: reader.decode()?, + }) + } +} + +pub fn generate_identity( + crypto: &impl QlCrypto, + name: impl Into, +) -> Result { + let MlKemKeyPair { + private: mlkem_private_key, + public: mlkem_public_key, + } = crypto.mlkem_generate_keypair(); + QlIdentity::new(crypto, mlkem_private_key, mlkem_public_key, name) +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct QlName(String); + +impl QlName { + pub const MAX_LEN: usize = 256; + + pub fn new(name: impl Into) -> Result { + let name = name.into(); + if name.is_empty() || name.len() > Self::MAX_LEN { + return Err(WireError::InvalidPayload); + } + Ok(Self(name)) + } +} + +impl Deref for QlName { + type Target = str; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl WireEncode for QlName { + fn encoded_len(&self) -> usize { + let len = VarInt::try_from(self.0.len()).unwrap(); + len.encoded_len() + self.0.len() + } + + fn encode(&self, out: &mut W) { + VarInt::try_from(self.0.len()) + .expect("identity name length fits in varint") + .encode(out); + self.0.as_bytes().encode(out); + } +} + +impl codec::WireDecode for QlName { + fn decode(reader: &mut codec::Reader) -> Result { + let len = usize::try_from(reader.decode::()?.into_inner()) + .map_err(|_| WireError::InvalidPayload)?; + if len == 0 || len > Self::MAX_LEN { + return Err(WireError::InvalidPayload); + } + let bytes = reader.take_bytes(len)?; + let name = std::str::from_utf8(&bytes).map_err(|_| WireError::InvalidPayload)?; + QlName::new(name) + } +} diff --git a/ql-wire/src/lib.rs b/ql-wire/src/lib.rs new file mode 100644 index 00000000..1713745a --- /dev/null +++ b/ql-wire/src/lib.rs @@ -0,0 +1,45 @@ +//! +//! QuantumLink protocol wire format +//! + +#![allow(clippy::too_many_arguments)] + +mod bytes; +mod codec; +mod crypto; +mod encrypted; +mod encrypted_message; +mod error; +mod handshake; +mod header; +mod identity; +mod nonce; +mod pq; +mod qid; +mod record; +#[cfg(any(feature = "test-utils", test))] +mod testing; +mod varint; + +pub use bytes::*; +pub use codec::*; +pub use crypto::*; +pub use encrypted::*; +pub use encrypted_message::*; +pub use error::*; +pub use handshake::*; +pub use header::*; +pub use identity::*; +pub use nonce::*; +pub use pq::*; +pub use qid::*; +pub use record::*; +#[cfg(any(feature = "test-utils", test))] +pub use testing::*; +pub use varint::*; + +pub const QL_WIRE_VERSION: u8 = 1; +pub const ENCRYPTED_MESSAGE_AUTH_SIZE: usize = 16; + +#[cfg(test)] +mod tests; diff --git a/ql-wire/src/nonce.rs b/ql-wire/src/nonce.rs new file mode 100644 index 00000000..c7e6d793 --- /dev/null +++ b/ql-wire/src/nonce.rs @@ -0,0 +1,13 @@ +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct Nonce(pub [u8; Self::SIZE]); + +impl Nonce { + pub const SIZE: usize = 12; + + pub fn from_counter(counter: u64) -> Self { + let mut nonce = [0u8; Self::SIZE]; + nonce[4..].copy_from_slice(&counter.to_le_bytes()); + Self(nonce) + } +} diff --git a/ql-wire/src/pq.rs b/ql-wire/src/pq.rs new file mode 100644 index 00000000..327ef7c4 --- /dev/null +++ b/ql-wire/src/pq.rs @@ -0,0 +1,159 @@ +use crate::{codec, ByteSlice, WireEncode, WireError}; + +pub const ML_KEM_SUITE_TAG: &[u8] = b"ml-kem-1024"; + +// ql-wire fixes the protocol to ML-KEM-1024 on the wire, but the host +// platform is free to satisfy QlKem with any backend that produces the same +// serialized sizes. +const ML_KEM_1024_SHARED_SECRET_SIZE: usize = 32; +const ML_KEM_1024_PUBLIC_KEY_SIZE: usize = 1568; +const ML_KEM_1024_PRIVATE_KEY_SIZE: usize = 3168; +const ML_KEM_1024_CIPHERTEXT_SIZE: usize = 1568; + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct SessionKey([u8; Self::SIZE]); + +impl SessionKey { + pub const SIZE: usize = ML_KEM_1024_SHARED_SECRET_SIZE; + + pub const fn from_data(data: [u8; Self::SIZE]) -> Self { + Self(data) + } + + pub const fn data(&self) -> &[u8; Self::SIZE] { + &self.0 + } + + pub const fn as_bytes(&self) -> &[u8; Self::SIZE] { + &self.0 + } +} + +impl AsRef<[u8]> for SessionKey { + fn as_ref(&self) -> &[u8] { + &self.0 + } +} + +impl Drop for SessionKey { + fn drop(&mut self) { + self.0.fill(0); + } +} + +impl WireEncode for SessionKey { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl codec::WireDecode for SessionKey { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self::from_data(reader.decode()?)) + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct MlKemPublicKey(Box<[u8; MlKemPublicKey::SIZE]>); + +impl MlKemPublicKey { + pub const SIZE: usize = ML_KEM_1024_PUBLIC_KEY_SIZE; + + pub fn new(data: Box<[u8; Self::SIZE]>) -> Self { + Self(data) + } + + pub fn as_bytes(&self) -> &[u8; Self::SIZE] { + self.0.as_ref() + } +} + +impl Drop for MlKemPublicKey { + fn drop(&mut self) { + self.0.as_mut().fill(0); + } +} + +impl codec::WireDecode for MlKemPublicKey { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self::new(reader.decode()?)) + } +} + +impl WireEncode for MlKemPublicKey { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.as_ref().encode(out); + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct MlKemPrivateKey(Box<[u8; MlKemPrivateKey::SIZE]>); + +impl MlKemPrivateKey { + pub const SIZE: usize = ML_KEM_1024_PRIVATE_KEY_SIZE; + + pub fn new(data: Box<[u8; Self::SIZE]>) -> Self { + Self(data) + } + + pub fn as_bytes(&self) -> &[u8; Self::SIZE] { + self.0.as_ref() + } +} + +impl Drop for MlKemPrivateKey { + fn drop(&mut self) { + self.0.as_mut().fill(0); + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct MlKemCiphertext(Box<[u8; MlKemCiphertext::SIZE]>); + +impl MlKemCiphertext { + pub const SIZE: usize = ML_KEM_1024_CIPHERTEXT_SIZE; + + pub fn new(data: Box<[u8; Self::SIZE]>) -> Self { + Self(data) + } + + pub fn as_bytes(&self) -> &[u8; Self::SIZE] { + self.0.as_ref() + } +} + +impl Drop for MlKemCiphertext { + fn drop(&mut self) { + self.0.as_mut().fill(0); + } +} + +impl codec::WireDecode for MlKemCiphertext { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self::new(reader.decode()?)) + } +} + +impl WireEncode for MlKemCiphertext { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.as_ref().encode(out); + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct MlKemKeyPair { + pub private: MlKemPrivateKey, + pub public: MlKemPublicKey, +} diff --git a/ql-wire/src/qid.rs b/ql-wire/src/qid.rs new file mode 100644 index 00000000..55c6684f --- /dev/null +++ b/ql-wire/src/qid.rs @@ -0,0 +1,44 @@ +use crate::{codec, ByteSlice, MlKemPublicKey, QlHash, WireEncode, WireError, ML_KEM_SUITE_TAG}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct QID(pub [u8; Self::SIZE]); + +impl QID { + pub const SIZE: usize = 16; + + pub fn derive(crypto: &impl QlHash, mlkem_public_key: &MlKemPublicKey) -> Self { + let digest = crypto.sha256(&[ + b"quantum-link qid v1", + ML_KEM_SUITE_TAG, + mlkem_public_key.as_bytes(), + ]); + let mut qid = [0u8; Self::SIZE]; + qid.copy_from_slice(&digest[..Self::SIZE]); + Self(qid) + } + + pub fn matches_public_key( + &self, + crypto: &impl QlHash, + mlkem_public_key: &MlKemPublicKey, + ) -> bool { + *self == Self::derive(crypto, mlkem_public_key) + } +} + +impl WireEncode for QID { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl codec::WireDecode for QID { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} diff --git a/ql-wire/src/record.rs b/ql-wire/src/record.rs new file mode 100644 index 00000000..163a1bff --- /dev/null +++ b/ql-wire/src/record.rs @@ -0,0 +1,254 @@ +use crate::{ + codec, + encrypted_message::EncryptedMessage, + handshake::{Ik1, Ik2, Kk1, Kk2, Xx1, Xx2, Xx3, Xx4}, + ByteSlice, SessionHeader, WireDecode, WireEncode, WireError, QL_WIRE_VERSION, +}; + +pub fn encode_record(out: &mut W, record_type: RecordType, body: &T) +where + W: bytes::BufMut + ?Sized, + T: WireEncode + ?Sized, +{ + RecordHeader { + version: QL_WIRE_VERSION, + record_type, + } + .encode(out); + body.encode(out); +} + +pub fn encode_record_vec(record_type: RecordType, body: &T) -> Vec { + let mut out = Vec::with_capacity(RecordHeader::WIRE_SIZE + body.encoded_len()); + encode_record(&mut out, record_type, body); + out +} + +pub fn decode_record(bytes: B) -> Result<(RecordHeader, T), WireError> +where + T: WireDecode, + B: ByteSlice, +{ + let mut reader = codec::Reader::new(bytes); + Ok((reader.decode()?, reader.decode()?)) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RecordHeader { + pub version: u8, + pub record_type: RecordType, +} + +impl RecordHeader { + pub const WIRE_SIZE: usize = size_of::() + size_of::(); +} + +impl WireDecode for RecordHeader { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + version: reader.decode()?, + record_type: reader.decode()?, + }) + } +} + +impl WireEncode for RecordHeader { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + out.put_u8(self.version); + self.record_type.encode(out); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +pub enum RecordType { + Handshake = 1, + Session = 2, +} + +impl TryFrom for RecordType { + type Error = WireError; + + fn try_from(value: u8) -> Result { + match value { + 1 => Ok(Self::Handshake), + 2 => Ok(Self::Session), + _ => Err(WireError::InvalidPayload), + } + } +} + +impl WireDecode for RecordType { + fn decode(reader: &mut codec::Reader) -> Result { + reader.decode::()?.try_into() + } +} + +impl WireEncode for RecordType { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u8(*self as u8); + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum QlHandshakeRecord { + Ik1(Ik1), + Ik2(Ik2), + Kk1(Kk1), + Kk2(Kk2), + Xx1(Xx1), + Xx2(Xx2), + Xx3(Xx3), + Xx4(Xx4), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +pub enum HandshakeKind { + Ik1 = 1, + Ik2 = 2, + Kk1 = 3, + Kk2 = 4, + Xx1 = 5, + Xx2 = 6, + Xx3 = 7, + Xx4 = 8, +} + +impl TryFrom for HandshakeKind { + type Error = WireError; + + fn try_from(value: u8) -> Result { + match value { + 1 => Ok(Self::Ik1), + 2 => Ok(Self::Ik2), + 3 => Ok(Self::Kk1), + 4 => Ok(Self::Kk2), + 5 => Ok(Self::Xx1), + 6 => Ok(Self::Xx2), + 7 => Ok(Self::Xx3), + 8 => Ok(Self::Xx4), + _ => Err(WireError::InvalidPayload), + } + } +} + +impl WireDecode for HandshakeKind { + fn decode(reader: &mut codec::Reader) -> Result { + reader.decode::()?.try_into() + } +} + +impl WireEncode for HandshakeKind { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u8(*self as u8); + } +} + +impl QlHandshakeRecord { + pub fn kind(&self) -> HandshakeKind { + match self { + Self::Ik1(_) => HandshakeKind::Ik1, + Self::Ik2(_) => HandshakeKind::Ik2, + Self::Kk1(_) => HandshakeKind::Kk1, + Self::Kk2(_) => HandshakeKind::Kk2, + Self::Xx1(_) => HandshakeKind::Xx1, + Self::Xx2(_) => HandshakeKind::Xx2, + Self::Xx3(_) => HandshakeKind::Xx3, + Self::Xx4(_) => HandshakeKind::Xx4, + } + } +} + +impl WireEncode for QlHandshakeRecord { + fn encoded_len(&self) -> usize { + self.kind().encoded_len() + + match self { + Self::Ik1(message) => message.encoded_len(), + Self::Ik2(message) => message.encoded_len(), + Self::Kk1(message) => message.encoded_len(), + Self::Kk2(message) => message.encoded_len(), + Self::Xx1(message) => message.encoded_len(), + Self::Xx2(message) => message.encoded_len(), + Self::Xx3(message) => message.encoded_len(), + Self::Xx4(message) => message.encoded_len(), + } + } + + fn encode(&self, out: &mut W) { + self.kind().encode(out); + match self { + Self::Ik1(message) => message.encode(out), + Self::Ik2(message) => message.encode(out), + Self::Kk1(message) => message.encode(out), + Self::Kk2(message) => message.encode(out), + Self::Xx1(message) => message.encode(out), + Self::Xx2(message) => message.encode(out), + Self::Xx3(message) => message.encode(out), + Self::Xx4(message) => message.encode(out), + } + } +} + +impl WireDecode for QlHandshakeRecord { + fn decode(reader: &mut codec::Reader) -> Result { + let kind = reader.decode::()?; + match kind { + HandshakeKind::Ik1 => Ok(Self::Ik1(reader.decode()?)), + HandshakeKind::Ik2 => Ok(Self::Ik2(reader.decode()?)), + HandshakeKind::Kk1 => Ok(Self::Kk1(reader.decode()?)), + HandshakeKind::Kk2 => Ok(Self::Kk2(reader.decode()?)), + HandshakeKind::Xx1 => Ok(Self::Xx1(reader.decode()?)), + HandshakeKind::Xx2 => Ok(Self::Xx2(reader.decode()?)), + HandshakeKind::Xx3 => Ok(Self::Xx3(reader.decode()?)), + HandshakeKind::Xx4 => Ok(Self::Xx4(reader.decode()?)), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct QlSessionRecord { + pub header: SessionHeader, + pub payload: EncryptedMessage, +} + +impl> WireEncode for QlSessionRecord { + fn encoded_len(&self) -> usize { + self.header.encoded_len() + self.payload.encoded_len() + } + + fn encode(&self, out: &mut W) { + self.header.encode(out); + self.payload.encode(out); + } +} + +impl QlSessionRecord { + pub fn into_owned(self) -> QlSessionRecord> { + QlSessionRecord { + header: self.header, + payload: self.payload.into_owned(), + } + } +} + +impl WireDecode for QlSessionRecord { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + header: reader.decode()?, + payload: reader.decode()?, + }) + } +} diff --git a/ql-wire/src/testing.rs b/ql-wire/src/testing.rs new file mode 100644 index 00000000..a1223c12 --- /dev/null +++ b/ql-wire/src/testing.rs @@ -0,0 +1,181 @@ +use libcrux_aesgcm::AesGcm256Key; +use libcrux_ml_kem::mlkem1024; +use sha2::{Digest, Sha256}; + +use crate::{ + MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, Nonce, QlAead, QlCrypto, + QlHash, QlIdentity, QlKem, QlRandom, SessionKey, ENCRYPTED_MESSAGE_AUTH_SIZE, +}; + +#[derive(Debug, Default, Clone, Copy)] +pub struct SoftwareCrypto; + +#[derive(Debug, Default, Clone, Copy)] +pub struct NoopCrypto; + +pub fn test_identities(crypto: &impl QlCrypto) -> (QlIdentity, QlIdentity) { + ( + crate::generate_identity(crypto, "alice").unwrap(), + crate::generate_identity(crypto, "bob").unwrap(), + ) +} + +impl QlRandom for SoftwareCrypto { + fn fill_random_bytes(&self, out: &mut [u8]) { + getrandom::getrandom(out).unwrap(); + } +} + +impl QlHash for SoftwareCrypto { + fn sha256(&self, parts: &[&[u8]]) -> [u8; 32] { + let mut hasher = Sha256::new(); + for part in parts { + hasher.update(part); + } + hasher.finalize().into() + } +} + +impl QlAead for SoftwareCrypto { + fn aes256_gcm_encrypt( + &self, + key: &SessionKey, + nonce: &Nonce, + aad: &[u8], + buffer: &mut [u8], + ) -> [u8; ENCRYPTED_MESSAGE_AUTH_SIZE] { + let key: AesGcm256Key = (*key.data()).into(); + let plaintext = buffer.to_vec(); + let mut auth = [0u8; ENCRYPTED_MESSAGE_AUTH_SIZE]; + key.encrypt( + buffer, + (&mut auth).into(), + (&nonce.0).into(), + aad, + &plaintext, + ) + .unwrap(); + auth + } + + fn aes256_gcm_decrypt( + &self, + key: &SessionKey, + nonce: &Nonce, + aad: &[u8], + buffer: &mut [u8], + auth_tag: &[u8; ENCRYPTED_MESSAGE_AUTH_SIZE], + ) -> bool { + let key: AesGcm256Key = (*key.data()).into(); + let ciphertext = buffer.to_vec(); + key.decrypt(buffer, (&nonce.0).into(), aad, &ciphertext, auth_tag.into()) + .is_ok() + } +} + +impl QlKem for SoftwareCrypto { + fn mlkem_generate_keypair(&self) -> MlKemKeyPair { + let key_pair = mlkem1024::generate_key_pair(random_array(self)); + let mut public = [0u8; MlKemPublicKey::SIZE]; + public.copy_from_slice(key_pair.pk()); + let mut private = [0u8; MlKemPrivateKey::SIZE]; + private.copy_from_slice(key_pair.sk()); + + MlKemKeyPair { + private: MlKemPrivateKey::new(Box::new(private)), + public: MlKemPublicKey::new(Box::new(public)), + } + } + + fn mlkem_encapsulate(&self, public_key: &MlKemPublicKey) -> (MlKemCiphertext, SessionKey) { + let public_key = public_key.as_bytes().into(); + let (ciphertext_value, shared_value) = + mlkem1024::encapsulate(&public_key, random_array(self)); + let mut ciphertext = [0u8; MlKemCiphertext::SIZE]; + ciphertext.copy_from_slice(ciphertext_value.as_slice()); + let mut shared = [0u8; SessionKey::SIZE]; + shared.copy_from_slice(shared_value.as_slice()); + ( + MlKemCiphertext::new(Box::new(ciphertext)), + SessionKey::from_data(shared), + ) + } + + fn mlkem_decapsulate( + &self, + private_key: &MlKemPrivateKey, + ciphertext: &MlKemCiphertext, + ) -> SessionKey { + let private_key = private_key.as_bytes().into(); + let ciphertext = ciphertext.as_bytes().into(); + let shared = mlkem1024::decapsulate(&private_key, &ciphertext); + let mut out = [0u8; SessionKey::SIZE]; + out.copy_from_slice(shared.as_slice()); + SessionKey::from_data(out) + } +} + +impl QlRandom for NoopCrypto { + fn fill_random_bytes(&self, out: &mut [u8]) { + out.fill(0); + } +} + +impl QlHash for NoopCrypto { + fn sha256(&self, _parts: &[&[u8]]) -> [u8; 32] { + [0; 32] + } +} + +impl QlAead for NoopCrypto { + fn aes256_gcm_encrypt( + &self, + _key: &SessionKey, + _nonce: &Nonce, + _aad: &[u8], + _buffer: &mut [u8], + ) -> [u8; ENCRYPTED_MESSAGE_AUTH_SIZE] { + [0; ENCRYPTED_MESSAGE_AUTH_SIZE] + } + + fn aes256_gcm_decrypt( + &self, + _key: &SessionKey, + _nonce: &Nonce, + _aad: &[u8], + _buffer: &mut [u8], + _auth_tag: &[u8; ENCRYPTED_MESSAGE_AUTH_SIZE], + ) -> bool { + false + } +} + +impl QlKem for NoopCrypto { + fn mlkem_generate_keypair(&self) -> MlKemKeyPair { + MlKemKeyPair { + private: MlKemPrivateKey::new(Box::new([0; MlKemPrivateKey::SIZE])), + public: MlKemPublicKey::new(Box::new([0; MlKemPublicKey::SIZE])), + } + } + + fn mlkem_encapsulate(&self, _public_key: &MlKemPublicKey) -> (MlKemCiphertext, SessionKey) { + ( + MlKemCiphertext::new(Box::new([0; MlKemCiphertext::SIZE])), + SessionKey::from_data([0; SessionKey::SIZE]), + ) + } + + fn mlkem_decapsulate( + &self, + _private_key: &MlKemPrivateKey, + _ciphertext: &MlKemCiphertext, + ) -> SessionKey { + SessionKey::from_data([0; SessionKey::SIZE]) + } +} + +fn random_array(crypto: &impl QlRandom) -> [u8; L] { + let mut out = [0u8; L]; + crypto.fill_random_bytes(&mut out); + out +} diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs new file mode 100644 index 00000000..a09a36b8 --- /dev/null +++ b/ql-wire/src/tests.rs @@ -0,0 +1,1017 @@ +use std::ops::RangeInclusive; + +use super::*; + +fn decode_handshake_record(bytes: &[u8]) -> QlHandshakeRecord { + decode_record(bytes).unwrap().1 +} + +fn decode_session_record(bytes: &[u8]) -> QlSessionRecord> { + let (_, record) = decode_record::, _>(bytes).unwrap(); + record.into_owned() +} + +fn qid(byte: u8) -> QID { + QID([byte; QID::SIZE]) +} + +fn varint(value: u64) -> VarInt { + VarInt::from_u64(value).unwrap() +} + +fn record_seq(value: u64) -> RecordSeq { + RecordSeq(varint(value)) +} + +fn record_ack_range(start: u64, end: u64) -> RangeInclusive { + record_seq(start)..=record_seq(end) +} + +fn stream_id(value: u64) -> StreamId { + StreamId(varint(value)) +} + +fn handshake_meta(id: u32) -> HandshakeMeta { + HandshakeMeta { + handshake_id: HandshakeId(id), + } +} + +fn handshake_transport_params(window: u32) -> TransportParams { + TransportParams { + initial_stream_receive_window: window, + } +} + +fn handshake_header(sender: u8, recipient: u8) -> HandshakeHeader { + HandshakeHeader { + sender: qid(sender), + recipient: qid(recipient), + } +} + +fn pairing_token(byte: u8) -> PairingToken { + PairingToken([byte; PairingToken::SIZE]) +} + +fn pairing_id(byte: u8) -> PairingId { + PairingId([byte; PairingId::SIZE]) +} + +fn xx_header(sender: u8, recipient: u8) -> HandshakeHeader { + HandshakeHeader { + sender: qid(sender), + recipient: qid(recipient), + } +} + +fn encrypt_record( + crypto: &impl QlCrypto, + header: SessionHeader, + session_key: &SessionKey, + body: &[SessionFrame>], +) -> QlSessionRecord> { + let mut builder = SessionRecordBuilder::new(header.seq, usize::MAX); + for frame in body { + let pushed = builder.push_frame(frame); + debug_assert!(pushed); + } + decode_session_record( + builder + .encrypt(crypto, header.connection_id, session_key) + .as_slice(), + ) +} + +#[test] +fn peer_bundle_round_trip() { + let crypto = SoftwareCrypto; + let identity = generate_identity(&crypto, "alice") + .unwrap() + .with_capabilities(0x55aa_33cc); + let bundle = identity.bundle(); + + let encoded = bundle.encode_vec(); + let decoded = PeerBundle::decode_exact(encoded.as_slice()).unwrap(); + + assert_eq!(decoded, bundle); + assert_eq!(&*decoded.name, "alice"); +} + +#[test] +fn identity_name_validation() { + assert_eq!( + QlName::new("a".repeat(QlName::MAX_LEN)).unwrap().len(), + QlName::MAX_LEN + ); + assert!(matches!(QlName::new(""), Err(WireError::InvalidPayload))); + assert!(matches!( + QlName::new("a".repeat(QlName::MAX_LEN + 1)), + Err(WireError::InvalidPayload) + )); +} + +#[test] +fn qid_derives_from_mlkem_public_key() { + let crypto = SoftwareCrypto; + let public_key = MlKemPublicKey::new(Box::new([42; MlKemPublicKey::SIZE])); + let qid = QID::derive(&crypto, &public_key); + + let digest = crypto.sha256(&[ + b"quantum-link qid v1", + ML_KEM_SUITE_TAG, + public_key.as_bytes(), + ]); + let mut expected = [0u8; QID::SIZE]; + expected.copy_from_slice(&digest[..QID::SIZE]); + + assert_eq!(qid, QID(expected)); + assert!(qid.matches_public_key(&crypto, &public_key)); +} + +#[test] +fn qid_changes_when_mlkem_public_key_changes() { + let crypto = SoftwareCrypto; + let first = MlKemPublicKey::new(Box::new([1; MlKemPublicKey::SIZE])); + let second = MlKemPublicKey::new(Box::new([2; MlKemPublicKey::SIZE])); + + assert_ne!(QID::derive(&crypto, &first), QID::derive(&crypto, &second)); +} + +#[test] +fn peer_bundle_detects_tampered_qid() { + let crypto = SoftwareCrypto; + let identity = generate_identity(&crypto, "alice").unwrap(); + let mut bundle = identity.bundle(); + + bundle.qid = qid(9); + + assert!(!bundle.qid_matches_public_key(&crypto)); +} + +#[test] +fn handshake_record_round_trip_supports_ik_kk_and_xx() { + let ik = QlHandshakeRecord::Ik1(Ik1 { + header: handshake_header(1, 2), + meta: handshake_meta(1), + transport_params: handshake_transport_params(65_536), + skem_ciphertext: MlKemCiphertext::new(Box::new([7; MlKemCiphertext::SIZE])), + ephemeral: EphemeralPublicKey { + mlkem_public_key: MlKemPublicKey::new(Box::new([9; MlKemPublicKey::SIZE])), + }, + static_bundle: EncryptedPeerBundle(vec![13; 64].into_boxed_slice()), + }); + let ik_encoded = encode_record_vec(RecordType::Handshake, &ik); + assert_eq!( + RecordHeader::decode_bytes(ik_encoded.as_slice()).unwrap(), + RecordHeader { + version: QL_WIRE_VERSION, + record_type: RecordType::Handshake, + } + ); + assert_eq!(decode_handshake_record(ik_encoded.as_slice()), ik); + + let kk = QlHandshakeRecord::Kk1(Kk1 { + header: handshake_header(1, 2), + meta: handshake_meta(2), + transport_params: handshake_transport_params(131_072), + skem_ciphertext: MlKemCiphertext::new(Box::new([11; MlKemCiphertext::SIZE])), + ephemeral: EphemeralPublicKey { + mlkem_public_key: MlKemPublicKey::new(Box::new([15; MlKemPublicKey::SIZE])), + }, + }); + let kk_encoded = encode_record_vec(RecordType::Handshake, &kk); + assert_eq!( + RecordHeader::decode_bytes(kk_encoded.as_slice()).unwrap(), + RecordHeader { + version: QL_WIRE_VERSION, + record_type: RecordType::Handshake, + } + ); + assert_eq!(decode_handshake_record(kk_encoded.as_slice()), kk); + + let xx = QlHandshakeRecord::Xx1(Xx1 { + header: xx_header(1, 2), + meta: handshake_meta(3), + pairing_id: pairing_id(3), + transport_params: handshake_transport_params(196_608), + ephemeral: EphemeralPublicKey { + mlkem_public_key: MlKemPublicKey::new(Box::new([17; MlKemPublicKey::SIZE])), + }, + }); + let xx_encoded = encode_record_vec(RecordType::Handshake, &xx); + assert_eq!( + RecordHeader::decode_bytes(xx_encoded.as_slice()).unwrap(), + RecordHeader { + version: QL_WIRE_VERSION, + record_type: RecordType::Handshake, + } + ); + assert_eq!(decode_handshake_record(xx_encoded.as_slice()), xx); +} + +#[test] +fn ik_handshake_rejects_tampered_handshake_meta() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let mut initiator_state = IkHandshake::new_initiator( + &crypto, + initiator, + responder.bundle(), + TransportParams::default(), + ); + let mut responder_state = + IkHandshake::new_responder(&crypto, responder, None, TransportParams::default()); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(77)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let mut m2 = responder_state + .write_2(&crypto, handshake_meta(77)) + .unwrap(); + m2.meta.handshake_id = HandshakeId(78); + + assert_eq!( + initiator_state.read_2(&crypto, &m2), + Err(WireError::InvalidHandshakeMeta) + ); +} + +#[test] +fn kk_handshake_rejects_tampered_handshake_header() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let mut initiator_state = KkHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + TransportParams::default(), + ); + let mut responder_state = KkHandshake::new_responder( + &crypto, + responder, + initiator.bundle(), + TransportParams::default(), + ); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(88)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let mut m2 = responder_state + .write_2(&crypto, handshake_meta(88)) + .unwrap(); + m2.header = handshake_header(9, 1); + + assert_eq!( + initiator_state.read_2(&crypto, &m2), + Err(WireError::InvalidPayload) + ); +} + +#[test] +fn ik_handshake_rejects_tampered_transport_params() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let mut initiator_state = IkHandshake::new_initiator( + &crypto, + initiator, + responder.bundle(), + handshake_transport_params(4096), + ); + let mut responder_state = + IkHandshake::new_responder(&crypto, responder, None, handshake_transport_params(8192)); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(89)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let mut m2 = responder_state + .write_2(&crypto, handshake_meta(89)) + .unwrap(); + m2.transport_params.initial_stream_receive_window += 1; + + assert_eq!( + initiator_state.read_2(&crypto, &m2), + Err(WireError::DecryptFailed) + ); +} + +#[test] +fn ik_handshake_rejects_tampered_handshake_header() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let mut initiator_state = IkHandshake::new_initiator( + &crypto, + initiator, + responder.bundle(), + TransportParams::default(), + ); + let mut responder_state = + IkHandshake::new_responder(&crypto, responder, None, TransportParams::default()); + + let mut m1 = initiator_state + .write_1(&crypto, handshake_meta(90)) + .unwrap(); + m1.header.sender = qid(9); + + assert_eq!( + responder_state.read_1(&crypto, &m1), + Err(WireError::DecryptFailed) + ); +} + +#[test] +fn ik_handshake_rejects_bound_remote_bundle_mismatch() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + let bogus = generate_identity(&crypto, "bogus").unwrap(); + + let mut initiator_state = IkHandshake::new_initiator( + &crypto, + initiator, + responder.bundle(), + TransportParams::default(), + ); + let mut responder_state = IkHandshake::new_responder( + &crypto, + responder, + Some(bogus.bundle()), + TransportParams::default(), + ); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(91)) + .unwrap(); + + assert_eq!( + responder_state.read_1(&crypto, &m1), + Err(WireError::InvalidPayload) + ); +} + +#[test] +fn ik_handshake_round_trip_derives_matching_transport_and_learns_remote() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let initiator_params = handshake_transport_params(4096); + let responder_params = handshake_transport_params(8192); + let mut initiator_state = IkHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + initiator_params, + ); + let mut responder_state = + IkHandshake::new_responder(&crypto, responder.clone(), None, responder_params); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(11)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let m2 = responder_state + .write_2(&crypto, handshake_meta(11)) + .unwrap(); + initiator_state.read_2(&crypto, &m2).unwrap(); + + let initiator_final = initiator_state.finalize(&crypto).unwrap(); + let responder_final = responder_state.finalize(&crypto).unwrap(); + + assert_eq!( + initiator_final.handshake_hash, + responder_final.handshake_hash + ); + assert_eq!(initiator_final.tx_key, responder_final.rx_key); + assert_eq!(initiator_final.rx_key, responder_final.tx_key); + assert_eq!( + initiator_final.tx_connection_id, + responder_final.rx_connection_id + ); + assert_eq!( + initiator_final.rx_connection_id, + responder_final.tx_connection_id + ); + assert_eq!(initiator_final.remote_bundle, responder.bundle()); + assert_eq!(responder_final.remote_bundle, initiator.bundle()); + assert_eq!(initiator_final.remote_transport_params, responder_params); + assert_eq!(responder_final.remote_transport_params, initiator_params); +} + +#[test] +fn ik_handshake_round_trip_derives_matching_transport_with_bound_responder() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let initiator_params = handshake_transport_params(16_384); + let responder_params = handshake_transport_params(32_768); + let mut initiator_state = IkHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + initiator_params, + ); + let mut responder_state = IkHandshake::new_responder( + &crypto, + responder.clone(), + Some(initiator.bundle()), + responder_params, + ); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(12)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let m2 = responder_state + .write_2(&crypto, handshake_meta(12)) + .unwrap(); + initiator_state.read_2(&crypto, &m2).unwrap(); + + let initiator_final = initiator_state.finalize(&crypto).unwrap(); + let responder_final = responder_state.finalize(&crypto).unwrap(); + + assert_eq!( + initiator_final.handshake_hash, + responder_final.handshake_hash + ); + assert_eq!(initiator_final.tx_key, responder_final.rx_key); + assert_eq!(initiator_final.rx_key, responder_final.tx_key); + assert_eq!( + initiator_final.tx_connection_id, + responder_final.rx_connection_id + ); + assert_eq!( + initiator_final.rx_connection_id, + responder_final.tx_connection_id + ); + assert_eq!(initiator_final.remote_bundle, responder.bundle()); + assert_eq!(responder_final.remote_bundle, initiator.bundle()); + assert_eq!(initiator_final.remote_transport_params, responder_params); + assert_eq!(responder_final.remote_transport_params, initiator_params); +} + +#[test] +fn kk_handshake_round_trip_derives_matching_transport() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let initiator_params = handshake_transport_params(24_576); + let responder_params = handshake_transport_params(49_152); + let mut initiator_state = KkHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + initiator_params, + ); + let mut responder_state = KkHandshake::new_responder( + &crypto, + responder.clone(), + initiator.bundle(), + responder_params, + ); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(21)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let m2 = responder_state + .write_2(&crypto, handshake_meta(21)) + .unwrap(); + initiator_state.read_2(&crypto, &m2).unwrap(); + + let initiator_final = initiator_state.finalize(&crypto).unwrap(); + let responder_final = responder_state.finalize(&crypto).unwrap(); + + assert_eq!( + initiator_final.handshake_hash, + responder_final.handshake_hash + ); + assert_eq!(initiator_final.tx_key, responder_final.rx_key); + assert_eq!(initiator_final.rx_key, responder_final.tx_key); + assert_eq!( + initiator_final.tx_connection_id, + responder_final.rx_connection_id + ); + assert_eq!( + initiator_final.rx_connection_id, + responder_final.tx_connection_id + ); + assert_eq!(initiator_final.remote_bundle, responder.bundle()); + assert_eq!(responder_final.remote_bundle, initiator.bundle()); + assert_eq!(initiator_final.remote_transport_params, responder_params); + assert_eq!(responder_final.remote_transport_params, initiator_params); +} + +#[test] +fn kk_handshake_rejects_tampered_transport_params() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let mut initiator_state = KkHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + handshake_transport_params(12288), + ); + let mut responder_state = KkHandshake::new_responder( + &crypto, + responder, + initiator.bundle(), + handshake_transport_params(24576), + ); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(22)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let mut m2 = responder_state + .write_2(&crypto, handshake_meta(22)) + .unwrap(); + m2.transport_params.initial_stream_receive_window += 1; + + assert_eq!( + initiator_state.read_2(&crypto, &m2), + Err(WireError::DecryptFailed) + ); +} + +#[test] +fn xx_handshake_rejects_tampered_pairing_id() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + let token = pairing_token(7); + + let mut initiator_state = XxHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.qid, + token, + TransportParams::default(), + ); + let mut responder_state = XxHandshake::new_responder( + &crypto, + responder, + initiator.qid, + token, + TransportParams::default(), + ); + + let mut m1 = initiator_state + .write_1(&crypto, handshake_meta(31)) + .unwrap(); + m1.pairing_id = pairing_id(8); + + assert_eq!( + responder_state.read_1(&crypto, &m1), + Err(WireError::InvalidPairingId) + ); +} + +#[test] +fn xx_handshake_rejects_tampered_sender_or_recipient() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + let token = pairing_token(7); + + let mut initiator_state = XxHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.qid, + token, + TransportParams::default(), + ); + let mut responder_state = XxHandshake::new_responder( + &crypto, + responder.clone(), + initiator.qid, + token, + TransportParams::default(), + ); + + let mut m1 = initiator_state + .write_1(&crypto, handshake_meta(31)) + .unwrap(); + m1.header.sender = responder.qid; + + assert_eq!( + responder_state.read_1(&crypto, &m1), + Err(WireError::InvalidHandshakeHeader) + ); + + let mut initiator_state = XxHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.qid, + token, + TransportParams::default(), + ); + let mut responder_state = XxHandshake::new_responder( + &crypto, + responder.clone(), + initiator.qid, + token, + TransportParams::default(), + ); + + let mut m1 = initiator_state + .write_1(&crypto, handshake_meta(31)) + .unwrap(); + m1.header.recipient = initiator.qid; + + assert_eq!( + responder_state.read_1(&crypto, &m1), + Err(WireError::InvalidHandshakeHeader) + ); +} + +#[test] +fn xx_handshake_rejects_repeated_transport_param_change() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + let token = pairing_token(9); + + let mut initiator_state = XxHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.qid, + token, + handshake_transport_params(12_288), + ); + let mut responder_state = XxHandshake::new_responder( + &crypto, + responder, + initiator.qid, + token, + handshake_transport_params(24_576), + ); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(32)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let m2 = responder_state + .write_2(&crypto, handshake_meta(32)) + .unwrap(); + initiator_state.read_2(&crypto, &m2).unwrap(); + + let mut m3 = initiator_state + .write_3(&crypto, handshake_meta(32)) + .unwrap(); + m3.transport_params.initial_stream_receive_window += 1; + + assert_eq!( + responder_state.read_3(&crypto, &m3), + Err(WireError::InvalidTransportParams) + ); +} + +#[test] +fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + let token = pairing_token(10); + + let initiator_params = handshake_transport_params(28_672); + let responder_params = handshake_transport_params(57_344); + let mut initiator_state = XxHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.qid, + token, + initiator_params, + ); + let mut responder_state = XxHandshake::new_responder( + &crypto, + responder.clone(), + initiator.qid, + token, + responder_params, + ); + + assert_eq!(initiator_state.pairing_token(), token); + assert_eq!(responder_state.pairing_token(), token); + assert_eq!(initiator_state.pairing_id(&crypto), token.id(&crypto)); + assert_eq!(responder_state.pairing_id(&crypto), token.id(&crypto)); + assert!(initiator_state.remote_bundle().is_none()); + assert!(responder_state.remote_bundle().is_none()); + + let m1 = initiator_state + .write_1(&crypto, handshake_meta(33)) + .unwrap(); + responder_state.read_1(&crypto, &m1).unwrap(); + + let m2 = responder_state + .write_2(&crypto, handshake_meta(33)) + .unwrap(); + initiator_state.read_2(&crypto, &m2).unwrap(); + assert_eq!(initiator_state.remote_bundle(), Some(&responder.bundle())); + assert!(responder_state.remote_bundle().is_none()); + + let m3 = initiator_state + .write_3(&crypto, handshake_meta(33)) + .unwrap(); + responder_state.read_3(&crypto, &m3).unwrap(); + assert_eq!(responder_state.remote_bundle(), Some(&initiator.bundle())); + + let m4 = responder_state + .write_4(&crypto, handshake_meta(33)) + .unwrap(); + initiator_state.read_4(&crypto, &m4).unwrap(); + + let initiator_final = initiator_state.finalize(&crypto).unwrap(); + let responder_final = responder_state.finalize(&crypto).unwrap(); + + assert_eq!( + initiator_final.handshake_hash, + responder_final.handshake_hash + ); + assert_eq!(initiator_final.tx_key, responder_final.rx_key); + assert_eq!(initiator_final.rx_key, responder_final.tx_key); + assert_eq!( + initiator_final.tx_connection_id, + responder_final.rx_connection_id + ); + assert_eq!( + initiator_final.rx_connection_id, + responder_final.tx_connection_id + ); + assert_eq!(initiator_final.remote_bundle, responder.bundle()); + assert_eq!(responder_final.remote_bundle, initiator.bundle()); + assert_eq!(initiator_final.remote_transport_params, responder_params); + assert_eq!(responder_final.remote_transport_params, initiator_params); +} + +#[test] +fn encrypted_session_record_round_trip_uses_connection_id_header() { + let crypto = SoftwareCrypto; + let header = SessionHeader { + connection_id: ConnectionId::from_data([0x44; ConnectionId::SIZE]), + seq: record_seq(11), + }; + let body = vec![ + SessionFrame::Ping, + SessionFrame::Unpair, + SessionFrame::Ack( + RecordAck::from_ranges([record_ack_range(20, 23), record_ack_range(12, 13)]).unwrap(), + ), + SessionFrame::StreamWindow(StreamWindow { + stream_id: stream_id(9), + maximum_offset: varint(65_536), + }), + SessionFrame::StreamData(StreamData { + stream_id: stream_id(9), + offset: varint(1024), + header: None, + bytes: b"hello".to_vec(), + fin: true, + }), + SessionFrame::StreamClose(StreamClose { + stream_id: stream_id(9), + target: CloseTarget::Both, + code: StreamCloseCode::CANCELLED, + }), + SessionFrame::Close(SessionClose { + code: SessionCloseCode::TIMEOUT, + }), + ]; + let session_key = SessionKey::from_data([7; SessionKey::SIZE]); + let record = encrypt_record(&crypto, header, &session_key, &body); + + let bytes = encode_record_vec(RecordType::Session, &record); + assert_eq!( + RecordHeader::decode_bytes(bytes.as_slice()).unwrap(), + RecordHeader { + version: QL_WIRE_VERSION, + record_type: RecordType::Session, + } + ); + let decoded = decode_session_record(bytes.as_slice()); + assert_eq!(decoded.header, header); + let encrypted = decoded.payload; + + let decrypted = + encrypted::decrypt_record(&crypto, &header, encrypted.clone(), &session_key).unwrap(); + assert_eq!(decode_session_frames(&decrypted).unwrap(), body); + + let wrong_header = SessionHeader { + connection_id: ConnectionId::from_data([0x99; ConnectionId::SIZE]), + seq: header.seq, + }; + assert_eq!( + encrypted::decrypt_record(&crypto, &wrong_header, encrypted.clone(), &session_key), + Err(WireError::DecryptFailed) + ); + + let wrong_seq_header = SessionHeader { + connection_id: header.connection_id, + seq: record_seq(header.seq.into_inner() + 1), + }; + assert_eq!( + encrypted::decrypt_record(&crypto, &wrong_seq_header, encrypted, &session_key), + Err(WireError::DecryptFailed) + ); +} + +#[test] +fn session_varint_fields_expand_at_expected_boundaries() { + let short_header = SessionHeader { + connection_id: ConnectionId::from_data([0x11; ConnectionId::SIZE]), + seq: record_seq(63), + }; + let long_header = SessionHeader { + connection_id: ConnectionId::from_data([0x11; ConnectionId::SIZE]), + seq: record_seq(64), + }; + + assert_eq!(short_header.encode_vec().len(), ConnectionId::SIZE + 1); + assert_eq!(long_header.encode_vec().len(), ConnectionId::SIZE + 2); + + let frame = StreamData { + stream_id: stream_id(64), + offset: varint(16_384), + header: None, + fin: true, + bytes: b"abc".to_vec(), + }; + let encoded = frame.encode_vec(); + + assert_eq!( + StreamData::decode_exact(encoded.as_slice()) + .unwrap() + .into_owned(), + frame + ); +} + +#[test] +fn protocol_record_size_breakdown() { + fn print_size(label: &str, size: usize) { + println!("{label:<32}: {size} bytes"); + } + + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + + let mut ik_initiator = IkHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + TransportParams::default(), + ); + let mut ik_responder = + IkHandshake::new_responder(&crypto, responder.clone(), None, TransportParams::default()); + + let ik1 = ik_initiator.write_1(&crypto, handshake_meta(101)).unwrap(); + ik_responder.read_1(&crypto, &ik1).unwrap(); + + let ik2 = ik_responder.write_2(&crypto, handshake_meta(101)).unwrap(); + ik_initiator.read_2(&crypto, &ik2).unwrap(); + + let ik1 = QlHandshakeRecord::Ik1(ik1); + let ik2 = QlHandshakeRecord::Ik2(ik2); + + let mut kk_initiator = KkHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + TransportParams::default(), + ); + let mut kk_responder = KkHandshake::new_responder( + &crypto, + responder.clone(), + initiator.bundle(), + TransportParams::default(), + ); + + let kk1 = kk_initiator.write_1(&crypto, handshake_meta(201)).unwrap(); + kk_responder.read_1(&crypto, &kk1).unwrap(); + + let kk2 = kk_responder.write_2(&crypto, handshake_meta(201)).unwrap(); + kk_initiator.read_2(&crypto, &kk2).unwrap(); + + let kk1 = QlHandshakeRecord::Kk1(kk1); + let kk2 = QlHandshakeRecord::Kk2(kk2); + + let token = pairing_token(0x42); + let mut xx_initiator = XxHandshake::new_initiator( + &crypto, + initiator.clone(), + responder.qid, + token, + TransportParams::default(), + ); + let mut xx_responder = XxHandshake::new_responder( + &crypto, + responder.clone(), + initiator.qid, + token, + TransportParams::default(), + ); + + let xx1 = xx_initiator.write_1(&crypto, handshake_meta(301)).unwrap(); + xx_responder.read_1(&crypto, &xx1).unwrap(); + + let xx2 = xx_responder.write_2(&crypto, handshake_meta(301)).unwrap(); + xx_initiator.read_2(&crypto, &xx2).unwrap(); + + let xx3 = xx_initiator.write_3(&crypto, handshake_meta(301)).unwrap(); + xx_responder.read_3(&crypto, &xx3).unwrap(); + + let xx4 = xx_responder.write_4(&crypto, handshake_meta(301)).unwrap(); + xx_initiator.read_4(&crypto, &xx4).unwrap(); + + let xx1 = QlHandshakeRecord::Xx1(xx1); + let xx2 = QlHandshakeRecord::Xx2(xx2); + let xx3 = QlHandshakeRecord::Xx3(xx3); + let xx4 = QlHandshakeRecord::Xx4(xx4); + + let session = ik_initiator.finalize(&crypto).unwrap(); + let session_ping = encrypt_record( + &crypto, + SessionHeader { + connection_id: session.tx_connection_id, + seq: record_seq(1), + }, + &session.tx_key, + &[SessionFrame::Ping], + ); + let session_ack = encrypt_record( + &crypto, + SessionHeader { + connection_id: session.tx_connection_id, + seq: record_seq(2), + }, + &session.tx_key, + &[SessionFrame::Ack( + RecordAck::from_ranges([record_ack_range(6, 6), record_ack_range(1, 2)]).unwrap(), + )], + ); + let session_unpair = encrypt_record( + &crypto, + SessionHeader { + connection_id: session.tx_connection_id, + seq: record_seq(3), + }, + &session.tx_key, + &[SessionFrame::Unpair], + ); + let session_stream_empty = encrypt_record( + &crypto, + SessionHeader { + connection_id: session.tx_connection_id, + seq: record_seq(4), + }, + &session.tx_key, + &[SessionFrame::StreamData(StreamData { + stream_id: stream_id(1), + offset: varint(0), + header: None, + fin: false, + bytes: Vec::new(), + })], + ); + let session_close = encrypt_record( + &crypto, + SessionHeader { + connection_id: session.tx_connection_id, + seq: record_seq(5), + }, + &session.tx_key, + &[SessionFrame::Close(SessionClose { + code: SessionCloseCode::PROTOCOL, + })], + ); + + print_size("ql-wire peer bundle", initiator.bundle().encode_vec().len()); + print_size("ql-wire mlkem public key", MlKemPublicKey::SIZE); + print_size("ql-wire mlkem ciphertext", MlKemCiphertext::SIZE); + print_size("ql-wire pq ik1", ik1.encode_vec().len()); + print_size("ql-wire pq ik2", ik2.encode_vec().len()); + print_size("ql-wire pq kk1", kk1.encode_vec().len()); + print_size("ql-wire pq kk2", kk2.encode_vec().len()); + print_size("ql-wire pq xx1", xx1.encode_vec().len()); + print_size("ql-wire pq xx2", xx2.encode_vec().len()); + print_size("ql-wire pq xx3", xx3.encode_vec().len()); + print_size("ql-wire pq xx4", xx4.encode_vec().len()); + print_size("ql-wire session ping", session_ping.encode_vec().len()); + print_size("ql-wire session ack", session_ack.encode_vec().len()); + print_size("ql-wire session unpair", session_unpair.encode_vec().len()); + print_size( + "ql-wire session stream empty", + session_stream_empty.encode_vec().len(), + ); + print_size("ql-wire session close", session_close.encode_vec().len()); +} diff --git a/ql-wire/src/varint.rs b/ql-wire/src/varint.rs new file mode 100644 index 00000000..7a39bd16 --- /dev/null +++ b/ql-wire/src/varint.rs @@ -0,0 +1,181 @@ +use core::fmt; + +use bytes::BufMut; + +use crate::{ByteSlice, Reader, WireDecode, WireEncode, WireError}; + +/// An integer less than 2^62 encoded with QUIC variable-length integer rules. +#[derive(Default, Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)] +pub struct VarInt(pub(crate) u64); + +impl VarInt { + /// The largest representable value. + pub const MAX: Self = Self((1u64 << 62) - 1); + /// The largest encoded value length. + pub const MAX_SIZE: usize = 8; + pub const MIN_SIZE: usize = 1; + + /// Construct a `VarInt` infallibly from a `u32`. + pub const fn from_u32(x: u32) -> Self { + Self(x as u64) + } + + /// Construct a `VarInt` from a `u64`. + pub fn from_u64(x: u64) -> Result { + if x < (1u64 << 62) { + Ok(Self(x)) + } else { + Err(VarIntBoundsExceeded) + } + } + + /// Create a `VarInt` without checking the bounds. + /// + /// # Safety + /// + /// `x` must be less than 2^62. + pub const unsafe fn from_u64_unchecked(x: u64) -> Self { + Self(x) + } + + /// Extract the inner integer value. + pub const fn into_inner(self) -> u64 { + self.0 + } + + /// Return the number of bytes required to encode this value. + pub const fn size(self) -> usize { + let x = self.0; + if x < (1u64 << 6) { + 1 + } else if x < (1u64 << 14) { + 2 + } else if x < (1u64 << 30) { + 4 + } else { + 8 + } + } +} + +impl From for u64 { + fn from(value: VarInt) -> Self { + value.0 + } +} + +impl From for VarInt { + fn from(value: u8) -> Self { + Self(value.into()) + } +} + +impl From for VarInt { + fn from(value: u16) -> Self { + Self(value.into()) + } +} + +impl From for VarInt { + fn from(value: u32) -> Self { + Self(value.into()) + } +} + +impl TryFrom for VarInt { + type Error = VarIntBoundsExceeded; + + fn try_from(value: u64) -> Result { + Self::from_u64(value) + } +} + +impl TryFrom for VarInt { + type Error = VarIntBoundsExceeded; + + fn try_from(value: u128) -> Result { + Self::from_u64(value.try_into().map_err(|_| VarIntBoundsExceeded)?) + } +} + +impl TryFrom for VarInt { + type Error = VarIntBoundsExceeded; + + fn try_from(value: usize) -> Result { + Self::from_u64(value as u64) + } +} + +impl fmt::Debug for VarInt { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(f) + } +} + +impl fmt::Display for VarInt { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(f) + } +} + +#[derive(Debug, Copy, Clone, Eq, PartialEq)] +pub struct VarIntBoundsExceeded; + +impl fmt::Display for VarIntBoundsExceeded { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("value too large for varint encoding") + } +} + +impl std::error::Error for VarIntBoundsExceeded {} + +impl WireDecode for VarInt { + fn decode(reader: &mut Reader) -> Result { + let first = reader.decode::()?; + let tag = first >> 6; + let first = first & 0b0011_1111; + let value = match tag { + 0b00 => u64::from(first), + 0b01 => { + let mut buf = [0; 2]; + buf[0] = first; + buf[1] = reader.decode()?; + u64::from(u16::from_be_bytes(buf)) + } + 0b10 => { + let mut buf = [0; 4]; + buf[0] = first; + buf[1..].copy_from_slice(&reader.take_bytes(3)?); + u64::from(u32::from_be_bytes(buf)) + } + 0b11 => { + let mut buf = [0; 8]; + buf[0] = first; + buf[1..].copy_from_slice(&reader.take_bytes(7)?); + u64::from_be_bytes(buf) + } + _ => unreachable!(), + }; + + // SAFETY: the decoded value is guaranteed to fit in the 62-bit varint range. + Ok(unsafe { Self::from_u64_unchecked(value) }) + } +} + +impl WireEncode for VarInt { + fn encoded_len(&self) -> usize { + self.size() + } + + #[allow(clippy::cast_possible_truncation)] + fn encode(&self, out: &mut W) { + let x = self.into_inner(); + match self.size() { + 1 => out.put_u8(x as u8), + 2 => out.put_u16((0b01 << 14) | x as u16), + 4 => out.put_u32((0b10 << 30) | x as u32), + 8 => out.put_u64((0b11 << 62) | x), + _ => unreachable!("malformed varint"), + } + } +} From 1a09d0a5bd6ca907a18d0fa74a54aff13c9cf2ae Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 4 Jun 2026 09:24:20 -0400 Subject: [PATCH 03/59] ql-fsm: add synchronous protocol state machine --- Cargo.lock | 144 ++- Cargo.toml | 2 + ql-fsm/Cargo.toml | 15 + ql-fsm/src/error.rs | 124 +++ ql-fsm/src/fsm.rs | 267 +++++ ql-fsm/src/handshake/ik.rs | 119 +++ ql-fsm/src/handshake/kk.rs | 118 +++ ql-fsm/src/handshake/mod.rs | 140 +++ ql-fsm/src/handshake/xx.rs | 207 ++++ ql-fsm/src/lib.rs | 334 ++++++ ql-fsm/src/pairing.rs | 38 + ql-fsm/src/session/ack_tracker.rs | 266 +++++ ql-fsm/src/session/mod.rs | 1041 +++++++++++++++++++ ql-fsm/src/session/range_set.rs | 221 ++++ ql-fsm/src/session/remote_stream_history.rs | 60 ++ ql-fsm/src/session/state.rs | 140 +++ ql-fsm/src/session/stream_ops.rs | 147 +++ ql-fsm/src/session/stream_parity.rs | 44 + ql-fsm/src/session/stream_rx.rs | 428 ++++++++ ql-fsm/src/session/stream_tx.rs | 579 +++++++++++ ql-fsm/src/session/tests.rs | 869 ++++++++++++++++ ql-fsm/src/session/tracked.rs | 29 + ql-fsm/src/state.rs | 139 +++ ql-fsm/src/tests/handshake.rs | 388 +++++++ ql-fsm/src/tests/mod.rs | 351 +++++++ ql-fsm/src/tests/proptest.rs | 1001 ++++++++++++++++++ ql-fsm/src/tests/session.rs | 532 ++++++++++ 27 files changed, 7741 insertions(+), 2 deletions(-) create mode 100644 ql-fsm/Cargo.toml create mode 100644 ql-fsm/src/error.rs create mode 100644 ql-fsm/src/fsm.rs create mode 100644 ql-fsm/src/handshake/ik.rs create mode 100644 ql-fsm/src/handshake/kk.rs create mode 100644 ql-fsm/src/handshake/mod.rs create mode 100644 ql-fsm/src/handshake/xx.rs create mode 100644 ql-fsm/src/lib.rs create mode 100644 ql-fsm/src/pairing.rs create mode 100644 ql-fsm/src/session/ack_tracker.rs create mode 100644 ql-fsm/src/session/mod.rs create mode 100644 ql-fsm/src/session/range_set.rs create mode 100644 ql-fsm/src/session/remote_stream_history.rs create mode 100644 ql-fsm/src/session/state.rs create mode 100644 ql-fsm/src/session/stream_ops.rs create mode 100644 ql-fsm/src/session/stream_parity.rs create mode 100644 ql-fsm/src/session/stream_rx.rs create mode 100644 ql-fsm/src/session/stream_tx.rs create mode 100644 ql-fsm/src/session/tests.rs create mode 100644 ql-fsm/src/session/tracked.rs create mode 100644 ql-fsm/src/state.rs create mode 100644 ql-fsm/src/tests/handshake.rs create mode 100644 ql-fsm/src/tests/mod.rs create mode 100644 ql-fsm/src/tests/proptest.rs create mode 100644 ql-fsm/src/tests/session.rs diff --git a/Cargo.lock b/Cargo.lock index c2c3b23c..016c88c3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -293,6 +293,21 @@ dependencies = [ "thiserror", ] +[[package]] +name = "bit-set" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3" +dependencies = [ + "bit-vec", +] + +[[package]] +name = "bit-vec" +version = "0.8.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" + [[package]] name = "bitcoin-io" version = "0.1.3" @@ -326,9 +341,9 @@ dependencies = [ [[package]] name = "bitflags" -version = "2.9.3" +version = "2.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "34efbcccd345379ca2868b2b2c9d3782e9cc58ba87bc7d79d5b53d9c9ae6f25d" +checksum = "84d7ced0ae9557296835c32bf1b1e02b44c746701f898460fb000d7eaa84f00a" [[package]] name = "blake2" @@ -820,6 +835,22 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys", +] + +[[package]] +name = "fastrand" +version = "2.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f1f227452a390804cdb637b74a86990f2a7d7ba4b7d5693aac9b4dd6defd8d6" + [[package]] name = "ff" version = "0.13.1" @@ -878,6 +909,12 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "fnv" +version = "1.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" + [[package]] name = "form_urlencoded" version = "1.2.2" @@ -1488,6 +1525,12 @@ version = "0.2.15" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f9fbbcab51052fe104eb5e5d351cf728d30a5be1fe14d9be8a3b097481fb97de" +[[package]] +name = "linux-raw-sys" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df1d3c3b53da64cf5760482273a98e575c651a67eec7f77df96b5b642de8f039" + [[package]] name = "litemap" version = "0.8.0" @@ -1983,6 +2026,25 @@ dependencies = [ "unicode-ident", ] +[[package]] +name = "proptest" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4b45fcc2344c680f5025fe57779faef368840d0bd1f42f216291f0dc4ace4744" +dependencies = [ + "bit-set", + "bit-vec", + "bitflags", + "num-traits", + "rand 0.9.2", + "rand_chacha 0.9.0", + "rand_xorshift", + "regex-syntax", + "rusty-fork", + "tempfile", + "unarray", +] + [[package]] name = "provenance-mark" version = "0.16.0" @@ -2027,6 +2089,16 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "ql-fsm" +version = "0.1.0" +dependencies = [ + "bytes", + "indexmap", + "proptest", + "ql-wire", +] + [[package]] name = "ql-wire" version = "0.1.0" @@ -2047,6 +2119,12 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "quick-error" +version = "1.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a1d01941d82fa2ab50be1e79e6714289dd7cde78eba4c074bc5a4374f650dfe0" + [[package]] name = "quote" version = "1.0.40" @@ -2130,6 +2208,15 @@ dependencies = [ "getrandom 0.3.3", ] +[[package]] +name = "rand_xorshift" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "513962919efc330f829edb2535844d1b912b0fbe2ca165d613e4e8788bb05a5a" +dependencies = [ + "rand_core 0.9.3", +] + [[package]] name = "rand_xoshiro" version = "0.6.0" @@ -2262,12 +2349,37 @@ dependencies = [ "semver", ] +[[package]] +name = "rustix" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cd15f8a2c5551a84d56efdc1cd049089e409ac19a3072d5037a17fd70719ff3e" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys", +] + [[package]] name = "rustversion" version = "1.0.22" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" +[[package]] +name = "rusty-fork" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc6bf79ff24e648f6da1f8d1f011e9cac26491b619e6b9280f2b47f1774e6ee2" +dependencies = [ + "fnv", + "quick-error", + "tempfile", + "wait-timeout", +] + [[package]] name = "ryu" version = "1.0.20" @@ -2557,6 +2669,19 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "tempfile" +version = "3.23.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d31c77bdf42a745371d260a26ca7163f1e0924b64afa0b688e61b5a9fa02f16" +dependencies = [ + "fastrand", + "getrandom 0.3.3", + "once_cell", + "rustix", + "windows-sys", +] + [[package]] name = "thiserror" version = "2.0.17" @@ -2652,6 +2777,12 @@ version = "1.18.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1dccffe3ce07af9386bfd29e80c0ab1a8205a2fc34e4bcd40364df902cfa8f3f" +[[package]] +name = "unarray" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eaea85b334db583fe3274d12b4cd1880032beab409c0d774be044d4480ab9a94" + [[package]] name = "unicode-ident" version = "1.0.18" @@ -2724,6 +2855,15 @@ version = "0.9.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" +[[package]] +name = "wait-timeout" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ac3b126d3914f9849036f826e054cbabdc8519970b8998ddaf3b5bd3c65f11" +dependencies = [ + "libc", +] + [[package]] name = "wasi" version = "0.11.1+wasi-snapshot-preview1" diff --git a/Cargo.toml b/Cargo.toml index 8aad9104..84dde3e4 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -4,6 +4,7 @@ members = [ "api", "backup-shard", "btp", + "ql-fsm", "ql-wire", "quantum-link-macros", ] @@ -31,6 +32,7 @@ backup-shard = { path = "backup-shard" } btp = { path = "btp" } foundation-api = { path = "api" } quantum-link-macros = { path = "quantum-link-macros" } +ql-fsm = { path = "ql-fsm" } ql-wire = { path = "ql-wire" } [patch.crates-io] diff --git a/ql-fsm/Cargo.toml b/ql-fsm/Cargo.toml new file mode 100644 index 00000000..be937f01 --- /dev/null +++ b/ql-fsm/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "ql-fsm" +version = "0.1.0" +edition = "2021" +description = "QuantumLink Sans-IO protocol finite state machine" +license = "Proprietary" + +[dependencies] +bytes = { workspace = true } +indexmap = "2" +ql-wire = { workspace = true } + +[dev-dependencies] +proptest = "1.6" +ql-wire = { workspace = true, features = ["test-utils"] } diff --git a/ql-fsm/src/error.rs b/ql-fsm/src/error.rs new file mode 100644 index 00000000..9bf2a915 --- /dev/null +++ b/ql-fsm/src/error.rs @@ -0,0 +1,124 @@ +use std::{ + error::Error, + fmt::{Display, Formatter}, +}; + +use ql_wire::{PairingId, WireError}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ReceiveError { + InvalidRecordHeader(WireError), + InvalidRecordVersion, + InvalidHandshakeRecord(WireError), + InvalidSessionRecord(WireError), + InvalidSessionConnectionId, + InvalidSessionPayload(WireError), + InvalidIkHandshake(WireError), + InvalidKkHandshake(WireError), + InvalidXxHandshake(WireError), + InvalidRemoteBundle, + InvalidQid, + NoPeer, + NoSession, + NotPairingMode, + InvalidPairingId { + expected: PairingId, + actual: PairingId, + }, +} + +impl Display for ReceiveError { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + match self { + Self::InvalidRecordHeader(error) => write!(f, "invalid record header: {error}"), + Self::InvalidRecordVersion => f.write_str("invalid record version"), + Self::InvalidHandshakeRecord(error) => { + write!(f, "invalid handshake record: {error}") + } + Self::InvalidSessionRecord(error) => write!(f, "invalid session record: {error}"), + Self::InvalidSessionConnectionId => f.write_str("invalid session connection id"), + Self::InvalidSessionPayload(error) => write!(f, "invalid session payload: {error}"), + Self::InvalidIkHandshake(error) => write!(f, "invalid ik handshake: {error}"), + Self::InvalidKkHandshake(error) => write!(f, "invalid kk handshake: {error}"), + Self::InvalidXxHandshake(error) => write!(f, "invalid xx handshake: {error}"), + Self::InvalidRemoteBundle => f.write_str("invalid remote bundle"), + Self::InvalidQid => f.write_str("invalid qid"), + Self::NoPeer => f.write_str("no bound peer"), + Self::NoSession => f.write_str("no active session"), + Self::NotPairingMode => f.write_str("not in pairing mode"), + Self::InvalidPairingId { expected, actual } => { + write!( + f, + "invalid pairing id: expected {expected}, actual {actual}" + ) + } + } + } +} + +impl std::error::Error for ReceiveError {} + +impl From for ReceiveError { + fn from(_: NoSessionError) -> Self { + Self::NoSession + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct NoPeerError; + +impl Display for NoPeerError { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.write_str("no peer bound") + } +} + +impl Error for NoPeerError {} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct NoSessionError; + +impl Display for NoSessionError { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "no session") + } +} + +impl Error for NoSessionError {} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamError { + MissingStream, + NotWritable, + NoSession, +} + +impl Display for StreamError { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + let message = match self { + Self::MissingStream => "missing stream", + Self::NotWritable => "stream is not writable", + Self::NoSession => "no session", + }; + f.write_str(message) + } +} + +impl Error for StreamError {} + +impl From for StreamError { + fn from(_: NoSessionError) -> Self { + Self::NoSession + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct CommitReadError; + +impl Display for CommitReadError { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + write!(f, "invalid read commit") + } +} + +impl Error for CommitReadError {} diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs new file mode 100644 index 00000000..036a336e --- /dev/null +++ b/ql-fsm/src/fsm.rs @@ -0,0 +1,267 @@ +use std::{collections::VecDeque, time::Instant}; + +use bytes::Bytes; +use ql_wire::{self as wire, QlCrypto, RouteId, SessionCloseCode, StreamId, WireDecode}; + +use crate::{ + handshake, + session::{self, SessionEvent, TerminalFrame}, + state::LinkState, + Event, NoPeerError, NoSessionError, OutboundWrite, QlFsm, ReceiveError, StreamError, WriteId, +}; + +pub struct EventSink<'a> { + events: &'a mut VecDeque, + termination: Option, +} + +impl<'a> EventSink<'a> { + fn new(events: &'a mut VecDeque) -> Self { + Self { + events, + termination: None, + } + } +} + +impl session::EventSink for EventSink<'_> { + fn emit(&mut self, event: SessionEvent) { + match event { + SessionEvent::Unpaired => { + self.termination = Some(TerminalFrame::Unpair); + } + SessionEvent::Opened { + stream_id, + route_id, + } => { + self.events.push_back(Event::Opened { + stream_id, + route_id, + }); + } + SessionEvent::Readable(stream_id) => { + self.events.push_back(Event::Readable(stream_id)); + } + SessionEvent::Writable(stream_id) => { + self.events.push_back(Event::Writable(stream_id)); + } + SessionEvent::Finished(stream_id) => { + self.events.push_back(Event::Finished(stream_id)); + } + SessionEvent::OutboundFinished(stream_id) => { + self.events.push_back(Event::OutboundFinished(stream_id)); + } + SessionEvent::Closed(frame) => { + self.events.push_back(Event::Closed(frame)); + } + SessionEvent::WritableClosed(frame) => { + self.events.push_back(Event::WritableClosed(frame)); + } + SessionEvent::SessionClosed(close) => { + self.termination = Some(TerminalFrame::Close(close.clone())); + self.events.push_back(Event::SessionClosed(close)); + } + } + } +} + +pub fn handle_bind_peer(fsm: &mut QlFsm, peer: ql_wire::PeerBundle) { + fsm.state.handshake = None; + fsm.state.link = LinkState::Idle; + fsm.state.peer = Some(peer); +} + +pub fn unpair(fsm: &mut QlFsm) { + let had_peer = fsm.state.peer.is_some(); + fsm.state.handshake = None; + fsm.state.armed_pairing_token = None; + + if let Some(conn) = fsm.state.link.connected_mut() { + let mut emit = EventSink::new(&mut fsm.events); + conn.session.unpair(&mut emit); + } else { + fsm.state.link = LinkState::Idle; + } + + if had_peer { + emit_peer_status(fsm, crate::PeerStatus::Unpaired); + } + fsm.state.peer = None; +} + +pub fn handle_disarm_pairing(fsm: &mut QlFsm) { + fsm.state.armed_pairing_token = None; + handshake::handle_disarm_pairing(fsm); +} + +pub fn handle_connect_xx(fsm: &mut QlFsm, invite: crate::PairingInvite, crypto: &impl QlCrypto) { + handshake::handle_connect_xx(fsm, invite, crypto); +} + +pub fn handle_connect_ik(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Result<(), NoPeerError> { + handshake::handle_connect_ik(fsm, crypto) +} + +pub fn handle_connect_kk(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Result<(), NoPeerError> { + handshake::handle_connect_kk(fsm, crypto) +} + +pub fn receive( + fsm: &mut QlFsm, + mut bytes: Vec, + crypto: &impl QlCrypto, +) -> Result<(), ReceiveError> { + let mut reader = wire::Reader::new(bytes.as_mut_slice()); + let header = + wire::RecordHeader::decode(&mut reader).map_err(ReceiveError::InvalidRecordHeader)?; + + if header.version != wire::QL_WIRE_VERSION { + return Err(ReceiveError::InvalidRecordVersion); + } + + match header.record_type { + wire::RecordType::Handshake => { + let record = wire::QlHandshakeRecord::decode(&mut reader) + .map_err(ReceiveError::InvalidHandshakeRecord)?; + handshake::handle_handshake_record(fsm, crypto, &record) + } + wire::RecordType::Session => { + let termination = { + let QlFsm { state, events, .. } = fsm; + let conn = state.link.connected_mut_or_err()?; + let (decrypt_len, seq) = { + let record = wire::QlSessionRecord::decode(&mut reader) + .map_err(ReceiveError::InvalidSessionRecord)?; + if record.header.connection_id != conn.transport.rx_connection_id { + return Err(ReceiveError::InvalidSessionConnectionId); + } + let payload = wire::decrypt_record( + crypto, + &record.header, + record.payload, + &conn.transport.rx_key, + ) + .map_err(ReceiveError::InvalidSessionPayload)?; + (payload.len(), record.header.seq) + }; + + let len = bytes.len(); + let plaintext = Bytes::from(bytes).slice(len - decrypt_len..); + let frames = wire::parse_session_frames(plaintext); + + let mut emit = EventSink::new(events); + conn.session.receive(state.now, seq, frames, &mut emit); + emit.termination + }; + + if matches!(termination, Some(TerminalFrame::Unpair)) { + if fsm.state.peer.is_some() { + emit_peer_status(fsm, crate::PeerStatus::Unpaired); + } + fsm.state.handshake = None; + fsm.state.armed_pairing_token = None; + fsm.state.peer = None; + } + Ok(()) + } + } +} + +pub fn on_timer(fsm: &mut QlFsm) { + handshake::handle_timer(fsm); + + let QlFsm { state, events, .. } = fsm; + let Some(conn) = state.link.connected_mut() else { + return; + }; + + let mut emit = EventSink::new(events); + conn.session.on_timer(state.now, &mut emit); +} + +pub fn next_deadline(fsm: &QlFsm) -> Option { + [ + handshake::next_handshake_deadline(fsm), + fsm.state + .link + .connected() + .and_then(|state| state.session.next_deadline()), + ] + .into_iter() + .flatten() + .min() +} + +pub fn take_next_write(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Option { + if let Some(record) = fsm.state.handshake.take() { + let record = wire::encode_record_vec(ql_wire::RecordType::Handshake, &record); + return Some(OutboundWrite { + record, + write_id: None, + }); + } + + let QlFsm { state, .. } = fsm; + let conn = state.link.connected_mut()?; + + let (write_id, builder) = conn.session.take_next_write(state.now)?; + let record = builder.encrypt( + crypto, + conn.transport.tx_connection_id, + &conn.transport.tx_key, + ); + if conn.session.is_closed() && matches!(fsm.state.link, LinkState::Connected(_)) { + fsm.state.link = LinkState::Idle; + emit_peer_status(fsm, fsm.state.link.status()); + } + Some(OutboundWrite { + record, + write_id: write_id.map(WriteId), + }) +} + +pub fn complete_write(fsm: &mut QlFsm, write_id: WriteId, success: bool) { + let QlFsm { state, .. } = fsm; + if let Some(conn) = state.link.connected_mut() { + conn.session.complete_write(state.now, write_id.0, success); + } +} + +pub fn close_session(fsm: &mut QlFsm, code: SessionCloseCode) { + let QlFsm { state, events, .. } = fsm; + let Some(conn) = state.link.connected_mut() else { + return; + }; + let mut emit = EventSink::new(events); + conn.session.close(code, &mut emit); +} + +pub fn open_stream( + fsm: &mut QlFsm, + route_id: RouteId, +) -> Result, NoSessionError> { + let QlFsm { state, events, .. } = fsm; + let conn = state.link.connected_mut_or_err()?; + let inner = conn.session.open_stream(route_id, EventSink::new(events))?; + Ok(crate::StreamOps { inner }) +} + +pub fn stream(fsm: &mut QlFsm, stream_id: StreamId) -> Result, StreamError> { + let QlFsm { state, events, .. } = fsm; + let conn = state.link.connected_mut_or_err()?; + let inner = conn.session.stream(stream_id, EventSink::new(events))?; + Ok(crate::StreamOps { inner }) +} + +pub fn queue_ping(fsm: &mut QlFsm) -> Result<(), NoSessionError> { + let conn = fsm.state.link.connected_mut_or_err()?; + conn.session.queue_ping() +} + +pub fn poll_event(fsm: &mut QlFsm) -> Option { + fsm.events.pop_front() +} + +pub fn emit_peer_status(fsm: &mut QlFsm, status: crate::PeerStatus) { + fsm.events.push_back(Event::PeerStatusChanged(status)); +} diff --git a/ql-fsm/src/handshake/ik.rs b/ql-fsm/src/handshake/ik.rs new file mode 100644 index 00000000..7e6ebd1e --- /dev/null +++ b/ql-fsm/src/handshake/ik.rs @@ -0,0 +1,119 @@ +use ql_wire::{self as wire, Ik1, Ik2, PeerBundle, QlCrypto, QlHandshakeRecord}; + +use super::{ + emit_peer_status, enqueue_handshake, finish_handshake, reset_connected_session_if_needed, +}; +use crate::{ + state::{IkInitiatorState, LinkState, SessionTransport}, + QlFsm, ReceiveError, +}; + +pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle) { + let meta = super::next_handshake_meta(fsm); + let mut handshake = wire::IkHandshake::new_initiator( + crypto, + fsm.identity.clone(), + peer, + super::local_transport_params(fsm), + ); + let message = handshake.write_1(crypto, meta).unwrap(); + + fsm.state.link = LinkState::IkInitiator(IkInitiatorState { + handshake_id: meta.handshake_id, + initial_ephemeral: message.ephemeral.clone(), + handshake, + deadline: fsm.state.now + fsm.config.handshake_timeout, + }); + enqueue_handshake(fsm, QlHandshakeRecord::Ik1(message)); + emit_peer_status(fsm, fsm.state.link.status()); +} + +pub fn handle_ik1( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + message: &Ik1, +) -> Result<(), ReceiveError> { + if should_ignore_inbound(fsm, message) { + return Ok(()); + } + if message.header.recipient != fsm.identity.qid { + return Err(ReceiveError::InvalidQid); + } + if let Some(peer) = fsm.state.peer.as_ref() { + if message.header.sender != peer.qid { + return Err(ReceiveError::InvalidQid); + } + } + + reset_connected_session_if_needed(fsm); + + let mut handshake = wire::IkHandshake::new_responder( + crypto, + fsm.identity.clone(), + fsm.state.peer.clone(), + super::local_transport_params(fsm), + ); + handshake + .read_1(crypto, message) + .map_err(ReceiveError::InvalidIkHandshake)?; + let outbound = handshake + .write_2(crypto, message.meta) + .map_err(ReceiveError::InvalidIkHandshake)?; + let (transport, remote_bundle) = SessionTransport::from_finalized( + handshake + .finalize(crypto) + .map_err(ReceiveError::InvalidIkHandshake)?, + ); + finish_handshake(fsm, transport, remote_bundle)?; + fsm.state.handshake = None; + enqueue_handshake(fsm, QlHandshakeRecord::Ik2(outbound)); + Ok(()) +} + +pub fn handle_ik2( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + message: &Ik2, +) -> Result<(), ReceiveError> { + { + let LinkState::IkInitiator(state) = &mut fsm.state.link else { + return Ok(()); + }; + + if message.meta.handshake_id != state.handshake_id { + return Ok(()); + } + + state + .handshake + .read_2(crypto, message) + .map_err(ReceiveError::InvalidIkHandshake)?; + } + + let LinkState::IkInitiator(state) = fsm.state.link.take() else { + unreachable!("active IK initiator was checked above"); + }; + let (transport, remote_bundle) = SessionTransport::from_finalized( + state + .handshake + .finalize(crypto) + .map_err(ReceiveError::InvalidIkHandshake)?, + ); + finish_handshake(fsm, transport, remote_bundle) +} + +pub fn should_ignore_inbound(fsm: &QlFsm, message: &Ik1) -> bool { + match &fsm.state.link { + LinkState::Idle + | LinkState::Connected(_) + | LinkState::KkInitiator(_) + | LinkState::XxInitiator(_) + | LinkState::XxResponder(_) => false, + LinkState::IkInitiator(state) => { + if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(message.header.sender) { + return false; + } + super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) + } + } +} diff --git a/ql-fsm/src/handshake/kk.rs b/ql-fsm/src/handshake/kk.rs new file mode 100644 index 00000000..e78c8a6d --- /dev/null +++ b/ql-fsm/src/handshake/kk.rs @@ -0,0 +1,118 @@ +use ql_wire::{self as wire, Kk1, Kk2, PeerBundle, QlCrypto, QlHandshakeRecord}; + +use super::{ + emit_peer_status, enqueue_handshake, finish_handshake, reset_connected_session_if_needed, +}; +use crate::{ + state::{KkInitiatorState, LinkState, SessionTransport}, + QlFsm, ReceiveError, +}; + +pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle) { + let meta = super::next_handshake_meta(fsm); + let mut handshake = wire::KkHandshake::new_initiator( + crypto, + fsm.identity.clone(), + peer, + super::local_transport_params(fsm), + ); + let message = handshake.write_1(crypto, meta).unwrap(); + + fsm.state.link = LinkState::KkInitiator(KkInitiatorState { + handshake_id: meta.handshake_id, + initial_ephemeral: message.ephemeral.clone(), + handshake, + deadline: fsm.state.now + fsm.config.handshake_timeout, + }); + enqueue_handshake(fsm, QlHandshakeRecord::Kk1(message)); + emit_peer_status(fsm, fsm.state.link.status()); +} + +pub fn handle_kk1( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + message: &Kk1, +) -> Result<(), ReceiveError> { + if should_ignore_inbound(fsm, message) { + return Ok(()); + } + + let Some(peer) = fsm.state.peer.clone() else { + return Err(ReceiveError::NoPeer); + }; + if message.header.recipient != fsm.identity.qid || message.header.sender != peer.qid { + return Err(ReceiveError::InvalidQid); + } + + reset_connected_session_if_needed(fsm); + + let mut handshake = wire::KkHandshake::new_responder( + crypto, + fsm.identity.clone(), + peer, + super::local_transport_params(fsm), + ); + handshake + .read_1(crypto, message) + .map_err(ReceiveError::InvalidKkHandshake)?; + let outbound = handshake + .write_2(crypto, message.meta) + .map_err(ReceiveError::InvalidKkHandshake)?; + let (transport, remote_bundle) = SessionTransport::from_finalized( + handshake + .finalize(crypto) + .map_err(ReceiveError::InvalidKkHandshake)?, + ); + finish_handshake(fsm, transport, remote_bundle)?; + fsm.state.handshake = None; + enqueue_handshake(fsm, QlHandshakeRecord::Kk2(outbound)); + Ok(()) +} + +pub fn handle_kk2( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + message: &Kk2, +) -> Result<(), ReceiveError> { + { + let LinkState::KkInitiator(state) = &mut fsm.state.link else { + return Ok(()); + }; + + if message.meta.handshake_id != state.handshake_id { + return Ok(()); + } + + state + .handshake + .read_2(crypto, message) + .map_err(ReceiveError::InvalidKkHandshake)?; + } + + let LinkState::KkInitiator(state) = fsm.state.link.take() else { + unreachable!("active KK initiator was checked above"); + }; + let (transport, remote_bundle) = SessionTransport::from_finalized( + state + .handshake + .finalize(crypto) + .map_err(ReceiveError::InvalidKkHandshake)?, + ); + finish_handshake(fsm, transport, remote_bundle) +} + +pub fn should_ignore_inbound(fsm: &QlFsm, message: &Kk1) -> bool { + match &fsm.state.link { + LinkState::Idle + | LinkState::Connected(_) + | LinkState::XxInitiator(_) + | LinkState::XxResponder(_) => false, + LinkState::IkInitiator(_) => true, + LinkState::KkInitiator(state) => { + if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(message.header.sender) { + return false; + } + super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) + } + } +} diff --git a/ql-fsm/src/handshake/mod.rs b/ql-fsm/src/handshake/mod.rs new file mode 100644 index 00000000..1881f66e --- /dev/null +++ b/ql-fsm/src/handshake/mod.rs @@ -0,0 +1,140 @@ +mod ik; +mod kk; +mod xx; + +use ql_wire::{self as wire, EphemeralPublicKey, HandshakeMeta, QlCrypto, QlHandshakeRecord}; + +use crate::{ + fsm::emit_peer_status, + session::{SessionConfig, SessionFsm, StreamParity}, + state::{ConnectedState, LinkState, SessionTransport}, + Event, NoPeerError, QlFsm, ReceiveError, +}; + +pub fn handle_connect_ik(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Result<(), NoPeerError> { + let peer = fsm.state.peer.clone().ok_or(NoPeerError)?; + prepare_for_outbound_connect(fsm); + ik::start_initiator(fsm, crypto, peer); + Ok(()) +} + +pub fn handle_connect_kk(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Result<(), NoPeerError> { + let peer = fsm.state.peer.clone().ok_or(NoPeerError)?; + prepare_for_outbound_connect(fsm); + kk::start_initiator(fsm, crypto, peer); + Ok(()) +} + +pub fn handle_connect_xx(fsm: &mut QlFsm, invite: crate::PairingInvite, crypto: &impl QlCrypto) { + prepare_for_outbound_connect(fsm); + xx::start_initiator(fsm, crypto, invite.token, invite.qid); +} + +pub fn next_handshake_meta(fsm: &mut QlFsm) -> HandshakeMeta { + let handshake_id = wire::HandshakeId(fsm.state.next_control_id); + fsm.state.next_control_id = fsm.state.next_control_id.wrapping_add(1); + HandshakeMeta { handshake_id } +} + +pub fn enqueue_handshake(fsm: &mut QlFsm, record: QlHandshakeRecord) { + debug_assert!(fsm.state.handshake.is_none()); + fsm.state.handshake = Some(record); +} + +pub fn handle_disarm_pairing(fsm: &mut QlFsm) { + xx::disarm_pairing(fsm); +} + +fn local_transport_params(fsm: &QlFsm) -> wire::TransportParams { + wire::TransportParams { + initial_stream_receive_window: fsm.config.session_stream_receive_buffer_size, + } +} + +pub fn prepare_for_outbound_connect(fsm: &mut QlFsm) { + fsm.state.handshake = None; + reset_connected_session_if_needed(fsm); +} + +pub fn handle_handshake_record( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + record: &QlHandshakeRecord, +) -> Result<(), ReceiveError> { + match record { + QlHandshakeRecord::Ik1(message) => ik::handle_ik1(fsm, crypto, message), + QlHandshakeRecord::Ik2(message) => ik::handle_ik2(fsm, crypto, message), + QlHandshakeRecord::Kk1(message) => kk::handle_kk1(fsm, crypto, message), + QlHandshakeRecord::Kk2(message) => kk::handle_kk2(fsm, crypto, message), + QlHandshakeRecord::Xx1(message) => xx::handle_xx1(fsm, crypto, message), + QlHandshakeRecord::Xx2(message) => xx::handle_xx2(fsm, crypto, message), + QlHandshakeRecord::Xx3(message) => xx::handle_xx3(fsm, crypto, message), + QlHandshakeRecord::Xx4(message) => xx::handle_xx4(fsm, crypto, message), + } +} + +pub fn handle_timer(fsm: &mut QlFsm) { + let Some(deadline) = fsm.state.link.handshake_deadline() else { + return; + }; + if deadline > fsm.state.now { + return; + } + + fsm.state.link = LinkState::Idle; + fsm.state.handshake = None; + emit_peer_status(fsm, fsm.state.link.status()); +} + +pub fn next_handshake_deadline(fsm: &QlFsm) -> Option { + fsm.state.link.handshake_deadline() +} + +pub fn finish_handshake( + fsm: &mut QlFsm, + transport: SessionTransport, + remote_bundle: wire::PeerBundle, +) -> Result<(), ReceiveError> { + let qid = remote_bundle.qid; + if let Some(peer) = fsm.state.peer.as_ref() { + if peer != &remote_bundle { + return Err(ReceiveError::InvalidRemoteBundle); + } + } else { + fsm.state.peer = Some(remote_bundle); + fsm.events.push_back(Event::NewPeer); + } + + let config = &fsm.config; + let session = SessionFsm::new( + SessionConfig { + local_parity: StreamParity::for_local(fsm.identity.qid, qid), + record_max_size: config.session_record_max_size, + ack_delay: config.session_record_ack_delay, + retransmit_timeout: config.session_record_retransmit_timeout, + keepalive_interval: config.session_keepalive_interval, + peer_timeout: config.session_peer_timeout, + stream_send_buffer_size: config.session_stream_send_buffer_size, + stream_receive_buffer_size: config.session_stream_receive_buffer_size, + accepted_record_window: config.session_accepted_record_window, + pending_ack_range_limit: config.session_pending_ack_range_limit, + initial_peer_stream_receive_window: transport + .remote_transport_params + .initial_stream_receive_window, + }, + fsm.state.now, + ); + fsm.state.link = LinkState::Connected(ConnectedState { transport, session }); + emit_peer_status(fsm, fsm.state.link.status()); + Ok(()) +} + +pub fn reset_connected_session_if_needed(fsm: &mut QlFsm) { + if matches!(fsm.state.link, LinkState::Connected(_)) { + fsm.state.link = LinkState::Idle; + } +} + +fn local_start_wins(local: &EphemeralPublicKey, inbound: &EphemeralPublicKey) -> bool { + local.mlkem_public_key.as_bytes() <= inbound.mlkem_public_key.as_bytes() +} diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs new file mode 100644 index 00000000..c9a289e0 --- /dev/null +++ b/ql-fsm/src/handshake/xx.rs @@ -0,0 +1,207 @@ +use ql_wire::{self as wire, PairingToken, QlCrypto, QlHandshakeRecord, Xx1, Xx2, Xx3, Xx4, QID}; + +use super::{ + emit_peer_status, enqueue_handshake, finish_handshake, reset_connected_session_if_needed, +}; +use crate::{ + state::{LinkState, SessionTransport, XxInitiatorState, XxResponderState}, + QlFsm, ReceiveError, +}; + +pub fn start_initiator( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + token: PairingToken, + remote_qid: QID, +) { + let meta = super::next_handshake_meta(fsm); + let mut handshake = wire::XxHandshake::new_initiator( + crypto, + fsm.identity.clone(), + remote_qid, + token, + super::local_transport_params(fsm), + ); + let message = handshake.write_1(crypto, meta).unwrap(); + + fsm.state.link = LinkState::XxInitiator(XxInitiatorState { + handshake_id: meta.handshake_id, + initial_ephemeral: message.ephemeral.clone(), + handshake, + deadline: fsm.state.now + fsm.config.handshake_timeout, + }); + enqueue_handshake(fsm, QlHandshakeRecord::Xx1(message)); + emit_peer_status(fsm, fsm.state.link.status()); +} + +pub fn handle_xx1( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + message: &Xx1, +) -> Result<(), ReceiveError> { + if should_ignore_inbound(fsm, crypto, message) { + return Ok(()); + } + match fsm.state.armed_pairing_token { + Some(expected) if expected.id(crypto) != message.pairing_id => { + Err(ReceiveError::InvalidPairingId { + expected: expected.id(crypto), + actual: message.pairing_id, + }) + } + Some(_) + if message.header.recipient != fsm.identity.qid + || message.header.sender == fsm.identity.qid => + { + Err(ReceiveError::InvalidQid) + } + Some(token) => { + reset_connected_session_if_needed(fsm); + + let mut handshake = wire::XxHandshake::new_responder( + crypto, + fsm.identity.clone(), + message.header.sender, + token, + super::local_transport_params(fsm), + ); + handshake + .read_1(crypto, message) + .map_err(ReceiveError::InvalidXxHandshake)?; + let outbound = handshake + .write_2(crypto, message.meta) + .map_err(ReceiveError::InvalidXxHandshake)?; + fsm.state.link = LinkState::XxResponder(XxResponderState { + handshake, + handshake_meta: message.meta, + deadline: fsm.state.now + fsm.config.handshake_timeout, + }); + fsm.state.handshake = None; + enqueue_handshake(fsm, QlHandshakeRecord::Xx2(outbound)); + Ok(()) + } + None => Err(ReceiveError::NotPairingMode), + } +} + +pub fn handle_xx2( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + message: &Xx2, +) -> Result<(), ReceiveError> { + { + let LinkState::XxInitiator(state) = &mut fsm.state.link else { + return Ok(()); + }; + + if message.meta.handshake_id != state.handshake_id { + return Ok(()); + } + + state + .handshake + .read_2(crypto, message) + .map_err(ReceiveError::InvalidXxHandshake)?; + let outbound = state + .handshake + .write_3(crypto, message.meta) + .map_err(ReceiveError::InvalidXxHandshake)?; + fsm.state.handshake = None; + enqueue_handshake(fsm, QlHandshakeRecord::Xx3(outbound)); + } + + Ok(()) +} + +pub fn handle_xx3( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + message: &Xx3, +) -> Result<(), ReceiveError> { + let LinkState::XxResponder(state) = &mut fsm.state.link else { + return Ok(()); + }; + + if message.meta.handshake_id != state.handshake_meta.handshake_id { + return Ok(()); + } + + state + .handshake + .read_3(crypto, message) + .map_err(ReceiveError::InvalidXxHandshake)?; + let handshake_meta = state.handshake_meta; + let LinkState::XxResponder(mut state) = fsm.state.link.take() else { + unreachable!("active XX responder was checked above"); + }; + let outbound = state + .handshake + .write_4(crypto, handshake_meta) + .map_err(ReceiveError::InvalidXxHandshake)?; + fsm.state.handshake = None; + enqueue_handshake(fsm, QlHandshakeRecord::Xx4(outbound)); + let (transport, remote_bundle) = SessionTransport::from_finalized( + state + .handshake + .finalize(crypto) + .map_err(ReceiveError::InvalidXxHandshake)?, + ); + finish_handshake(fsm, transport, remote_bundle) +} + +pub fn handle_xx4( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + message: &Xx4, +) -> Result<(), ReceiveError> { + { + let LinkState::XxInitiator(state) = &mut fsm.state.link else { + return Ok(()); + }; + + if message.meta.handshake_id != state.handshake_id { + return Ok(()); + } + + state + .handshake + .read_4(crypto, message) + .map_err(ReceiveError::InvalidXxHandshake)?; + } + + let LinkState::XxInitiator(state) = fsm.state.link.take() else { + unreachable!("active XX initiator was checked above"); + }; + let (transport, remote_bundle) = SessionTransport::from_finalized( + state + .handshake + .finalize(crypto) + .map_err(ReceiveError::InvalidXxHandshake)?, + ); + finish_handshake(fsm, transport, remote_bundle) +} + +pub fn disarm_pairing(fsm: &mut QlFsm) { + if matches!(fsm.state.link, LinkState::XxResponder(_)) { + fsm.state.link = LinkState::Idle; + fsm.state.handshake = None; + } +} + +pub fn should_ignore_inbound(fsm: &QlFsm, crypto: &impl QlCrypto, message: &Xx1) -> bool { + match &fsm.state.link { + LinkState::Idle | LinkState::Connected(_) => false, + LinkState::IkInitiator(_) | LinkState::KkInitiator(_) | LinkState::XxResponder(_) => true, + LinkState::XxInitiator(state) => { + if state.handshake.pairing_id(crypto) != message.pairing_id { + return false; + } + if message.header.recipient != fsm.identity.qid + || message.header.sender != state.handshake.remote_qid() + { + return false; + } + super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) + } + } +} diff --git a/ql-fsm/src/lib.rs b/ql-fsm/src/lib.rs new file mode 100644 index 00000000..3067efdb --- /dev/null +++ b/ql-fsm/src/lib.rs @@ -0,0 +1,334 @@ +//! sync finite state machine for QuantumLink protocol +//! +//! a caller drives `QlFsm` inside its own event loop +//! +//! inputs to that loop usually include +//! - app actions like `bind_peer`, `connect_ik`, `connect_kk`, `connect_xx`, `open_stream`, or +//! `stream` +//! - inbound transport bytes passed to `receive` +//! - a deadline expiring, handled by calling `on_timer` +//! - transport write results passed to `complete_write` +//! +//! outputs from `QlFsm` are +//! - outbound session and handshake records from `take_next_write` +//! - queued `QlFsmEvent`s returned by `poll_event` after `connect_ik`, `connect_kk`, +//! `connect_xx`, `receive`, and `on_timer` +//! +//! call `next_deadline` after handling current inputs and any queued outputs +//! use it to decide how long the outer loop can wait before `on_timer` must run +//! another input may arrive before that deadline, which is fine + +mod error; +mod fsm; +mod handshake; +mod pairing; +mod session; +pub(crate) mod state; +#[cfg(test)] +mod tests; + +use std::{ + collections::VecDeque, + time::{Duration, Instant}, +}; + +pub use bytes::Bytes; +pub use error::*; +pub use pairing::PairingInvite; +use ql_wire::{ + PairingToken, PeerBundle, QlCrypto, QlIdentity, RouteId, SessionClose, SessionCloseCode, + StreamClose, StreamId, +}; +pub use session::{SessionEvent, StreamReadIter, StreamWriter}; + +use crate::state::{LinkState, QlFsmState}; + +/// connection state for the bound peer +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum PeerStatus { + /// no active encrypted session + Disconnected, + /// we are driving the handshake + Initiator, + /// the encrypted session is up + Connected, + /// the bound peer was forgotten immediately + /// + /// unpair is abortive and best-effort. the binding is removed immediately + /// and one final write may remain: a record containing only `SessionFrame::Unpair` + Unpaired, +} + +/// events emitted by `QlFsm` +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum Event { + /// a peer was learned during handshake completion + NewPeer, + /// the peer changed lifecycle state + PeerStatusChanged(PeerStatus), + /// a stream was opened + Opened { + stream_id: StreamId, + route_id: RouteId, + }, + /// a stream has bytes ready to read + Readable(StreamId), + /// a stream has room for more local writes + Writable(StreamId), + /// the peer finished writing this stream and no more bytes remain to read + Finished(StreamId), + /// our local FIN was acknowledged by the peer at the session layer + OutboundFinished(StreamId), + /// a stream was closed + Closed(StreamClose), + /// local writes on this stream are closed + WritableClosed(StreamClose), + /// the encrypted session was closed + /// + /// session close is abortive and best-effort. the session ends immediately + /// one final write remains: a record containing only `SessionFrame::Close` + /// the FSM does not wait for an ack for that record + SessionClosed(SessionClose), +} + +/// handle for a session write returned by `QlFsm::take_next_write` +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct WriteId(pub(crate) u64); + +/// outbound record produced by `QlFsm` +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct OutboundWrite { + /// wire bytes to hand to the transport + pub record: Vec, + /// write handle that must be completed exactly once + pub write_id: Option, +} + +pub struct StreamOps<'a> { + inner: session::StreamOps<'a, fsm::EventSink<'a>>, +} + +impl StreamOps<'_> { + /// returns this stream's identifier + pub fn stream_id(&self) -> StreamId { + self.inner.stream_id() + } + + /// returns the readable stream bytes as owned `Bytes` views without consuming them + pub fn read(&self) -> StreamReadIter<'_> { + self.inner.read() + } + + /// returns how many bytes can be read from the stream + pub fn readable_bytes(&self) -> usize { + self.inner.readable_bytes() + } + + /// marks previously read bytes as consumed + pub fn commit_read(&mut self, len: usize) -> Result<(), CommitReadError> { + self.inner.commit_read(len) + } + + /// returns a writer if the local write side is still open + pub fn writer(&mut self) -> Option> { + self.inner.writer() + } + + /// closes the origin lane, return lane, or both lanes of the stream + pub fn close(&mut self, target: ql_wire::CloseTarget, code: ql_wire::StreamCloseCode) { + self.inner.close(target, code); + } +} + +/// timing and buffering knobs for `QlFsm` +#[derive(Debug, Clone, Copy)] +pub struct QlFsmConfig { + /// overall time limit for one handshake attempt + pub handshake_timeout: Duration, + /// delay before sending a pure record ack + pub session_record_ack_delay: Duration, + /// how long to wait before resending unacked session records + pub session_record_retransmit_timeout: Duration, + /// idle delay before sending a keepalive ping + pub session_keepalive_interval: Duration, + /// how long to wait before declaring the peer dead + pub session_peer_timeout: Duration, + /// maximum total wire size for one session record, including header and auth tag + pub session_record_max_size: usize, + /// maximum bytes buffered locally for one stream send side + pub session_stream_send_buffer_size: usize, + /// maximum bytes buffered locally for one stream receive side + pub session_stream_receive_buffer_size: u32, + /// how many accepted record sequence numbers to retain for duplicate detection + pub session_accepted_record_window: u64, + /// maximum disjoint pending ACK ranges to retain before dropping the oldest low ranges + pub session_pending_ack_range_limit: usize, +} + +impl Default for QlFsmConfig { + fn default() -> Self { + let s = session::SessionConfig::default(); + Self { + handshake_timeout: Duration::from_secs(5), + session_record_ack_delay: s.ack_delay, + session_record_retransmit_timeout: s.retransmit_timeout, + session_keepalive_interval: s.keepalive_interval, + session_peer_timeout: s.peer_timeout, + session_record_max_size: s.record_max_size, + session_stream_send_buffer_size: s.stream_send_buffer_size, + session_stream_receive_buffer_size: s.stream_receive_buffer_size, + session_accepted_record_window: s.accepted_record_window, + session_pending_ack_range_limit: s.pending_ack_range_limit, + } + } +} + +/// synchronous driver for peer binding, handshake, and encrypted streams +pub struct QlFsm { + config: QlFsmConfig, + identity: QlIdentity, + state: QlFsmState, + events: VecDeque, +} + +impl QlFsm { + /// creates a new `QlFsm` + pub fn new(config: QlFsmConfig, identity: QlIdentity, now: Instant) -> Self { + Self { + config, + identity, + state: QlFsmState { + next_control_id: 1, + peer: None, + armed_pairing_token: None, + handshake: None, + link: LinkState::Idle, + now, + }, + events: VecDeque::new(), + } + } + + /// binds the remote peer + pub fn bind_peer(&mut self, peer: PeerBundle) { + fsm::handle_bind_peer(self, peer); + } + + /// returns the currently bound peer, if any + pub fn peer(&self) -> Option<&PeerBundle> { + self.state.peer.as_ref() + } + + /// arms acceptance of inbound xx pairings for a single token + pub fn arm_pairing(&mut self, token: PairingToken) { + self.state.armed_pairing_token = Some(token); + } + + pub fn pairing_token(&self) -> Option<&PairingToken> { + self.state.armed_pairing_token.as_ref() + } + + /// disarms inbound xx pairing and rejects any in-flight inbound xx responder state + pub fn disarm_pairing(&mut self) { + fsm::handle_disarm_pairing(self); + } + + /// starts an outbound xx handshake using a pairing invite + pub fn connect_xx(&mut self, now: Instant, invite: PairingInvite, crypto: &impl QlCrypto) { + self.state.now = now; + fsm::handle_connect_xx(self, invite, crypto); + } + + /// starts an IK handshake with the currently bound peer + pub fn connect_ik(&mut self, now: Instant, crypto: &impl QlCrypto) -> Result<(), NoPeerError> { + self.state.now = now; + fsm::handle_connect_ik(self, crypto) + } + + /// starts a KK handshake with the currently bound peer + pub fn connect_kk(&mut self, now: Instant, crypto: &impl QlCrypto) -> Result<(), NoPeerError> { + self.state.now = now; + fsm::handle_connect_kk(self, crypto) + } + + /// handles one inbound wire message + pub fn receive( + &mut self, + now: Instant, + bytes: Vec, + crypto: &impl QlCrypto, + ) -> Result<(), ReceiveError> { + self.state.now = now; + fsm::receive(self, bytes, crypto) + } + + /// returns the next queued event, if any + pub fn poll_event(&mut self) -> Option { + fsm::poll_event(self) + } + + /// advances time-based state + pub fn on_timer(&mut self, now: Instant) { + self.state.now = now; + fsm::on_timer(self); + } + + /// returns the next timer deadline, if any + pub fn next_deadline(&self) -> Option { + fsm::next_deadline(self) + } + + pub fn has_shutdown_work(&self) -> bool { + self.state + .link + .connected() + .is_some_and(|state| state.session.has_shutdown_work()) + } + + /// returns the next outbound record + /// + /// if `write_id` is `Some`, call `complete_write` exactly once + /// + /// if it is `None`, the record is fire-and-forget + pub fn take_next_write( + &mut self, + now: Instant, + crypto: &impl QlCrypto, + ) -> Option { + self.state.now = now; + fsm::take_next_write(self, crypto) + } + + /// completes a `SessionWriteId` from `take_next_write` with the transport outcome + /// + /// call this at most once for each returned `SessionWriteId` + pub fn complete_write(&mut self, now: Instant, write_id: WriteId, success: bool) { + self.state.now = now; + fsm::complete_write(self, write_id, success); + } + + /// closes the current encrypted session locally + pub fn close_session(&mut self, code: SessionCloseCode) { + fsm::close_session(self, code); + } + + /// forgets the bound peer locally and may emit one final outbound `SessionFrame::Unpair` + pub fn unpair(&mut self) { + fsm::unpair(self); + } + + /// opens a new outgoing stream + pub fn open_stream(&mut self, route_id: RouteId) -> Result, NoSessionError> { + fsm::open_stream(self, route_id) + } + + /// returns a facade for an open stream + pub fn stream(&mut self, stream_id: StreamId) -> Result, StreamError> { + fsm::stream(self, stream_id) + } + + /// queues a ping on the active session + pub fn queue_ping(&mut self) -> Result<(), NoSessionError> { + fsm::queue_ping(self) + } +} diff --git a/ql-fsm/src/pairing.rs b/ql-fsm/src/pairing.rs new file mode 100644 index 00000000..4b8361b8 --- /dev/null +++ b/ql-fsm/src/pairing.rs @@ -0,0 +1,38 @@ +use ql_wire::{ByteSlice, PairingToken, Reader, WireDecode, WireEncode, WireError, QID}; + +/// Out-of-band invite consumed by the initiator of an XX pairing +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct PairingInvite { + pub qid: QID, + pub token: PairingToken, +} + +impl PairingInvite { + pub const VERSION: u8 = 1; + pub const WIRE_SIZE: usize = size_of::() + QID::SIZE + PairingToken::SIZE; +} + +impl WireEncode for PairingInvite { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + Self::VERSION.encode(out); + self.qid.encode(out); + self.token.encode(out); + } +} + +impl WireDecode for PairingInvite { + fn decode(reader: &mut Reader) -> Result { + if reader.decode::()? != Self::VERSION { + return Err(WireError::InvalidPayload); + } + + Ok(Self { + qid: reader.decode()?, + token: reader.decode()?, + }) + } +} diff --git a/ql-fsm/src/session/ack_tracker.rs b/ql-fsm/src/session/ack_tracker.rs new file mode 100644 index 00000000..a75b5c63 --- /dev/null +++ b/ql-fsm/src/session/ack_tracker.rs @@ -0,0 +1,266 @@ +use std::{ops::RangeInclusive, time::Instant}; + +use ql_wire::{RecordAck, RecordAckBuilder, RecordSeq}; + +use super::range_set::RangeSet; + +#[derive(Debug, Clone)] +pub struct AckTracker { + accepted_records: RangeSet, + pending_ack: RangeSet, + ack_state: AckState, + accepted_record_window: u64, + pending_ack_range_limit: usize, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PendingAck { + pub ack: RecordAck, + pub due_at: Instant, + pub includes_all_pending: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ReceiveOutcome { + New, + Duplicate, + TooOld, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum AckState { + Idle, + Dirty { due_at: Instant }, +} + +impl AckTracker { + pub fn new(accepted_record_window: u64, pending_ack_range_limit: usize) -> Self { + Self { + accepted_records: RangeSet::new(), + pending_ack: RangeSet::new(), + ack_state: AckState::Idle, + accepted_record_window: accepted_record_window.max(1), + pending_ack_range_limit: pending_ack_range_limit.max(1), + } + } + + pub fn insert(&mut self, seq: RecordSeq) -> ReceiveOutcome { + let seq = seq.into_inner(); + let largest_accepted = self.accepted_records.max(); + if largest_accepted.is_some_and(|largest| seq < self.accepted_cutoff(largest)) { + return ReceiveOutcome::TooOld; + } + if self.accepted_records.contains(seq) { + self.pending_ack.insert(single_range(seq)); + self.trim_pending_ack_ranges(); + return ReceiveOutcome::Duplicate; + } + + self.accepted_records.insert(single_range(seq)); + self.trim_accepted_records(); + + self.pending_ack.insert(single_range(seq)); + self.trim_pending_ack_ranges(); + + ReceiveOutcome::New + } + + pub fn ack_deadline(&self) -> Option { + match self.ack_state { + AckState::Idle => None, + AckState::Dirty { due_at } => Some(due_at), + } + } + + pub fn schedule_ack(&mut self, due_at: Instant) { + self.ack_state = match self.ack_state { + AckState::Dirty { due_at: old } => AckState::Dirty { + due_at: due_at.min(old), + }, + AckState::Idle => AckState::Dirty { due_at }, + }; + } + + pub fn pending_ack(&self, max_wire_size: usize) -> Option { + let due_at = self.ack_deadline()?; + if max_wire_size == 0 || self.pending_ack.range_count() == 0 { + return None; + } + + let total_range_count = self.pending_ack.range_count(); + let mut ack = RecordAckBuilder::new(); + let mut selected_range_count = 0usize; + + for range in self.pending_ack.iter_rev() { + let pushed = ack + .try_push_range(to_ack_range(range), max_wire_size) + .unwrap(); + if !pushed { + break; + } + selected_range_count += 1; + } + + (selected_range_count != 0).then(|| PendingAck { + ack: ack.build().unwrap(), + due_at, + includes_all_pending: total_range_count == selected_range_count, + }) + } + + pub fn on_ack_emitted(&mut self, pending_ack: &PendingAck) { + self.retire_acked_ranges(&pending_ack.ack); + if pending_ack.includes_all_pending || self.pending_ack.range_count() == 0 { + self.ack_state = AckState::Idle; + } + } + + pub fn retire_acked_ranges(&mut self, ack: &RecordAck) { + for range in ack.ranges() { + self.pending_ack.remove(from_ack_range(range)); + } + if self.pending_ack.range_count() == 0 { + self.ack_state = AckState::Idle; + } + } + + pub fn clear_ack_state(&mut self) { + self.ack_state = AckState::Idle; + } + + pub fn restore_acked_ranges(&mut self, ack: &RecordAck, due_at: Instant) { + for range in ack.ranges() { + self.pending_ack.insert(from_ack_range(range)); + } + self.trim_pending_ack_ranges(); + self.schedule_ack(due_at); + } + + fn accepted_cutoff(&self, largest_accepted: u64) -> u64 { + largest_accepted + .saturating_add(1) + .saturating_sub(self.accepted_record_window) + } + + fn trim_accepted_records(&mut self) { + let Some(largest_accepted) = self.accepted_records.max() else { + return; + }; + let cutoff = self.accepted_cutoff(largest_accepted); + self.accepted_records.remove(0..cutoff); + } + + fn trim_pending_ack_ranges(&mut self) { + while self.pending_ack.range_count() > self.pending_ack_range_limit { + self.pending_ack.pop_min(); + } + } +} + +fn single_range(seq: u64) -> std::ops::Range { + seq..seq.checked_add(1).unwrap() +} + +fn to_ack_range(range: std::ops::Range) -> RangeInclusive { + let end = range.end.checked_sub(1).unwrap(); + RecordSeq::from_u64(range.start).unwrap()..=RecordSeq::from_u64(end).unwrap() +} + +fn from_ack_range(range: RangeInclusive) -> std::ops::Range { + let start = range.start().into_inner(); + let end = range.end().into_inner().checked_add(1).unwrap(); + start..end +} + +#[cfg(test)] +mod tests { + use std::time::{Duration, Instant}; + + use ql_wire::RecordSeq; + + use super::{AckTracker, PendingAck, ReceiveOutcome}; + + fn seq(value: u64) -> RecordSeq { + RecordSeq::from_u64(value).unwrap() + } + + fn ack_ranges(pending_ack: &PendingAck) -> Vec<(u64, u64)> { + pending_ack + .ack + .ranges() + .map(|range| (range.start().into_inner(), range.end().into_inner())) + .collect() + } + + #[test] + fn contiguous_records_emit_one_ack_range() { + let now = Instant::now(); + let mut ack_tracker = AckTracker::new(128, 8); + + assert_eq!(ack_tracker.insert(seq(10)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(11)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(12)), ReceiveOutcome::New); + + ack_tracker.schedule_ack(now); + let pending_ack = ack_tracker.pending_ack(usize::MAX).unwrap(); + assert_eq!(ack_ranges(&pending_ack), vec![(10, 12)]); + } + + #[test] + fn sparse_records_emit_descending_ack_ranges() { + let now = Instant::now(); + let mut ack_tracker = AckTracker::new(128, 8); + + assert_eq!(ack_tracker.insert(seq(10)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(15)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(16)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(12)), ReceiveOutcome::New); + + ack_tracker.schedule_ack(now + Duration::from_millis(5)); + let pending_ack = ack_tracker.pending_ack(usize::MAX).unwrap(); + assert_eq!(ack_ranges(&pending_ack), vec![(15, 16), (12, 12), (10, 10)]); + } + + #[test] + fn accepted_record_window_evicts_old_sequences() { + let mut ack_tracker = AckTracker::new(4, 8); + + assert_eq!(ack_tracker.insert(seq(10)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(15)), ReceiveOutcome::New); + + assert_eq!(ack_tracker.insert(seq(10)), ReceiveOutcome::TooOld); + } + + #[test] + fn pending_ack_range_limit_drops_oldest_low_ranges() { + let now = Instant::now(); + let mut ack_tracker = AckTracker::new(128, 2); + + assert_eq!(ack_tracker.insert(seq(1)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(3)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(5)), ReceiveOutcome::New); + + ack_tracker.schedule_ack(now); + let pending_ack = ack_tracker.pending_ack(usize::MAX).unwrap(); + assert_eq!(ack_ranges(&pending_ack), vec![(5, 5), (3, 3)]); + } + + #[test] + fn retire_acked_ranges_removes_only_exact_snapshot() { + let now = Instant::now(); + let mut ack_tracker = AckTracker::new(128, 8); + + assert_eq!(ack_tracker.insert(seq(1)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(3)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(seq(5)), ReceiveOutcome::New); + ack_tracker.schedule_ack(now); + + let first_ack = ack_tracker.pending_ack(4).unwrap(); + assert_eq!(ack_ranges(&first_ack), vec![(5, 5)]); + ack_tracker.on_ack_emitted(&first_ack); + ack_tracker.retire_acked_ranges(&first_ack.ack); + + let second_ack = ack_tracker.pending_ack(usize::MAX).unwrap(); + assert_eq!(ack_ranges(&second_ack), vec![(3, 3), (1, 1)]); + } +} diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs new file mode 100644 index 00000000..55187757 --- /dev/null +++ b/ql-fsm/src/session/mod.rs @@ -0,0 +1,1041 @@ +pub use self::{state::TerminalFrame, stream_ops::*, stream_parity::*, stream_rx::*}; + +mod ack_tracker; +mod range_set; +mod remote_stream_history; +mod state; +mod stream_ops; +mod stream_parity; +mod stream_rx; +mod stream_tx; +mod tracked; + +#[cfg(test)] +mod tests; + +use std::time::{Duration, Instant}; + +use bytes::Bytes; +use indexmap::IndexMap; +use ql_wire::{ + CloseTarget, RecordAck, RecordSeq, RouteId, SessionClose, SessionCloseCode, SessionFrame, + SessionRecordBuilder, StreamClose, StreamData, StreamHeader, StreamId, StreamWindow, VarInt, + WireError, +}; + +use self::{ + ack_tracker::{AckTracker, PendingAck, ReceiveOutcome}, + remote_stream_history::RemoteStreamHistory, + state::{InboundState, OutboundState, SessionPhase, SessionState, StreamRole, StreamState}, + stream_tx::StreamTxRange, + tracked::{TrackedFrame, TrackedRecord, TrackedStreamData}, +}; +use crate::{NoSessionError, StreamError}; + +#[derive(Debug, Clone, Copy)] +pub struct SessionConfig { + pub local_parity: StreamParity, + pub record_max_size: usize, + pub ack_delay: Duration, + pub retransmit_timeout: Duration, + pub keepalive_interval: Duration, + pub peer_timeout: Duration, + pub stream_send_buffer_size: usize, + pub stream_receive_buffer_size: u32, + pub initial_peer_stream_receive_window: u32, + pub accepted_record_window: u64, + pub pending_ack_range_limit: usize, +} + +impl Default for SessionConfig { + fn default() -> Self { + Self { + local_parity: StreamParity::Even, + record_max_size: 8 * 1024, + ack_delay: Duration::from_millis(5), + retransmit_timeout: Duration::from_millis(150), + keepalive_interval: Duration::from_secs(10), + peer_timeout: Duration::from_secs(30), + stream_send_buffer_size: 16 * 1024, + stream_receive_buffer_size: 16 * 1024, + initial_peer_stream_receive_window: 16 * 1024, + accepted_record_window: 4096, + pending_ack_range_limit: 64, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum SessionEvent { + Opened { + stream_id: StreamId, + route_id: RouteId, + }, + Readable(StreamId), + Writable(StreamId), + Finished(StreamId), + OutboundFinished(StreamId), + Closed(StreamClose), + WritableClosed(StreamClose), + SessionClosed(SessionClose), + Unpaired, +} + +pub trait EventSink { + fn emit(&mut self, event: SessionEvent); +} + +impl EventSink for F +where + F: FnMut(SessionEvent), +{ + fn emit(&mut self, event: SessionEvent) { + self(event); + } +} + +pub struct SessionFsm { + config: SessionConfig, + state: SessionState, +} + +impl SessionFsm { + pub fn new(mut config: SessionConfig, now: Instant) -> Self { + config.record_max_size = config + .record_max_size + .max(SessionRecordBuilder::MIN_CAPACITY); + config.stream_send_buffer_size = config.stream_send_buffer_size.max(1); + config.stream_receive_buffer_size = config.stream_receive_buffer_size.max(1); + config.accepted_record_window = config.accepted_record_window.max(1); + config.pending_ack_range_limit = config.pending_ack_range_limit.max(1); + Self { + config, + state: SessionState { + last_activity_at: now, + last_inbound_at: now, + phase: SessionPhase::Open, + next_stream_ordinal: 0, + next_record_seq: RecordSeq::from_u32(0), + next_write_id: 0, + tracked_records: IndexMap::default(), + ack_tracker: AckTracker::new( + config.accepted_record_window, + config.pending_ack_range_limit, + ), + pending_ping: false, + streams: IndexMap::default(), + next_stream_index: 0, + remote_stream_history: RemoteStreamHistory::new(config.local_parity.remote()), + }, + } + } + + pub fn open_stream( + &mut self, + route_id: RouteId, + sink: E, + ) -> Result, NoSessionError> + where + E: EventSink, + { + self.ensure_session_open()?; + let stream_id = self + .config + .local_parity + .make_stream_id(self.state.next_stream_ordinal); + self.state.next_stream_ordinal = self.state.next_stream_ordinal.saturating_add(1); + self.state.streams.insert( + stream_id, + StreamState::new( + StreamRole::Initiator, + Some(route_id), + self.config.stream_receive_buffer_size, + self.config.initial_peer_stream_receive_window, + ), + ); + let stream_index = self.state.streams.len() - 1; + Ok(StreamOps::new(self, stream_id, stream_index, sink)) + } + + pub fn stream( + &mut self, + stream_id: StreamId, + sink: E, + ) -> Result, StreamError> + where + E: EventSink, + { + self.ensure_session_open()?; + let Some(stream_index) = self.state.streams.get_index_of(&stream_id) else { + return Err(StreamError::MissingStream); + }; + + Ok(StreamOps::new(self, stream_id, stream_index, sink)) + } + + pub fn queue_ping(&mut self) -> Result<(), NoSessionError> { + self.ensure_session_open()?; + self.state.pending_ping = true; + Ok(()) + } + + pub fn close(&mut self, code: SessionCloseCode, sink: &mut impl EventSink) { + if self.state.phase != SessionPhase::Open { + return; + } + + self.begin_termination(TerminalFrame::Close(SessionClose { code }), sink); + } + + pub fn unpair(&mut self, sink: &mut impl EventSink) { + if self.state.phase != SessionPhase::Open { + return; + } + + self.begin_termination(TerminalFrame::Unpair, sink); + } + + pub fn is_closed(&self) -> bool { + self.state.phase == SessionPhase::Closed + } + + pub fn receive(&mut self, now: Instant, seq: RecordSeq, frames: I, sink: &mut impl EventSink) + where + I: IntoIterator, WireError>>, + { + if self.state.phase != SessionPhase::Open { + return; + } + + self.state.last_activity_at = now; + self.state.last_inbound_at = now; + self.collect_timeouts(now); + + match self.state.ack_tracker.insert(seq) { + ReceiveOutcome::TooOld => return, + ReceiveOutcome::Duplicate => { + self.schedule_ack(now, true); + return; + } + ReceiveOutcome::New => {} + } + + let mut ack_eliciting = false; + + for frame in frames { + let Ok(frame) = frame else { + self.close(SessionCloseCode::PROTOCOL, sink); + return; + }; + ack_eliciting |= !matches!(frame, SessionFrame::Ack(_)); + match frame { + SessionFrame::Ping => {} + SessionFrame::Unpair => { + self.unpair(sink); + return; + } + SessionFrame::Ack(ack) => self.process_record_ack(&ack, sink), + SessionFrame::StreamData(frame) => { + if self.handle_stream_data(frame, sink).is_err() { + self.close(SessionCloseCode::PROTOCOL, sink); + return; + } + } + SessionFrame::StreamWindow(frame) => self.handle_stream_window(&frame, sink), + SessionFrame::StreamClose(frame) => { + if self.handle_stream_close(&frame, sink).is_err() { + self.close(SessionCloseCode::PROTOCOL, sink); + return; + } + } + SessionFrame::Close(close) => { + self.close(close.code, sink); + return; + } + } + } + + if ack_eliciting { + self.schedule_ack(now, false); + } + } + + pub fn complete_write(&mut self, now: Instant, write_id: u64, success: bool) { + if !self.state.phase.is_open() { + return; + } + if success { + let Some(record) = self.state.tracked_records.get_mut(&write_id) else { + return; + }; + if record.sent_at.is_some() { + return; + } + self.state.last_activity_at = now; + record.sent_at = Some(now); + } else { + if self + .state + .tracked_records + .get(&write_id) + .is_some_and(|record| record.sent_at.is_some()) + { + return; + } + let Some(record) = self.state.tracked_records.shift_remove(&write_id) else { + return; + }; + restore_tracked_record( + now, + &mut self.state.ack_tracker, + &mut self.state.pending_ping, + &mut self.state.streams, + record, + ); + } + } + + pub fn on_timer(&mut self, now: Instant, sink: &mut impl EventSink) { + if !self.state.phase.is_open() { + return; + } + self.collect_timeouts(now); + if !self.config.peer_timeout.is_zero() + && self.state.last_inbound_at + self.config.peer_timeout <= now + { + self.close(SessionCloseCode::TIMEOUT, sink); + return; + } + if self.state.phase == SessionPhase::Open + && !self.config.keepalive_interval.is_zero() + && self.state.last_activity_at + self.config.keepalive_interval <= now + { + self.state.pending_ping = true; + } + } + + pub fn next_deadline(&self) -> Option { + if !self.state.phase.is_open() { + return None; + } + let ack_deadline = self.state.ack_tracker.ack_deadline(); + let retransmit_deadline = self + .state + .tracked_records + .values() + .filter_map(|record| { + record + .sent_at + .map(|sent_at| sent_at + self.config.retransmit_timeout) + }) + .min(); + let is_open = self.state.phase.is_open(); + let keepalive_deadline = + (is_open && !self.config.keepalive_interval.is_zero() && !self.state.pending_ping) + .then_some(self.state.last_activity_at + self.config.keepalive_interval); + let peer_timeout_deadline = (is_open && !self.config.peer_timeout.is_zero()) + .then_some(self.state.last_inbound_at + self.config.peer_timeout); + [ + ack_deadline, + retransmit_deadline, + keepalive_deadline, + peer_timeout_deadline, + ] + .into_iter() + .flatten() + .min() + } + + pub fn has_shutdown_work(&self) -> bool { + matches!(self.state.phase, SessionPhase::Terminating(_)) + || self.state.ack_tracker.ack_deadline().is_some() + || !self.state.tracked_records.is_empty() + } + + pub fn take_next_write(&mut self, now: Instant) -> Option<(Option, SessionRecordBuilder)> { + match &self.state.phase { + SessionPhase::Terminating(frame) => { + let seq = self.state.next_record_seq; + next_seq(&mut self.state.next_record_seq); + let mut builder = SessionRecordBuilder::new(seq, self.config.record_max_size); + match frame { + TerminalFrame::Close(close) => { + assert!(builder.push_close(close), "builder has capacity"); + } + TerminalFrame::Unpair => { + assert!(builder.push_unpair(), "builder has capacity"); + } + } + self.state.phase = SessionPhase::Closed; + return Some((None, builder)); + } + SessionPhase::Closed => { + return None; + } + SessionPhase::Open => {} + } + self.collect_timeouts(now); + + let (builder, outbound) = self.build_next_record(now)?; + + let should_track = outbound.ping_included + || !outbound.window_updates.is_empty() + || !outbound.frames.is_empty(); + let write_id = should_track.then(|| { + let write_id = self.state.next_write_id; + self.state.next_write_id = self.state.next_write_id.wrapping_add(1); + self.state.tracked_records.insert(write_id, outbound); + write_id + }); + + Some((write_id, builder)) + } + + fn build_next_record(&mut self, now: Instant) -> Option<(SessionRecordBuilder, TrackedRecord)> { + let seq = self.state.next_record_seq; + let mut builder = SessionRecordBuilder::new(seq, self.config.record_max_size); + let mut outbound = TrackedRecord { + seq, + frames: Vec::new(), + ack: None, + ping_included: false, + window_updates: Vec::new(), + sent_at: None, + }; + + self.push_next_pending_stream_close(&mut builder, &mut outbound); + + if self.state.pending_ping && builder.push_ping() { + self.state.pending_ping = false; + outbound.ping_included = true; + } + + self.push_next_pending_stream_window(&mut builder, &mut outbound); + + self.push_next_stream_data(&mut builder, &mut outbound); + + if let Some(pending_ack) = self.pending_ack(builder.remaining_capacity()) { + if (!builder.is_empty() || pending_ack.due_at <= now) + && builder.push_ack(&pending_ack.ack) + { + self.state.ack_tracker.on_ack_emitted(&pending_ack); + outbound.ack = Some(pending_ack.ack); + } + } + + if builder.is_empty() { + return None; + } + + next_seq(&mut self.state.next_record_seq); + Some((builder, outbound)) + } + + fn begin_termination(&mut self, frame: TerminalFrame, sink: &mut impl EventSink) { + match &frame { + TerminalFrame::Close(close) => sink.emit(SessionEvent::SessionClosed(close.clone())), + TerminalFrame::Unpair => sink.emit(SessionEvent::Unpaired), + } + + self.state.phase = SessionPhase::Terminating(frame); + self.state.tracked_records.clear(); + self.state.ack_tracker.clear_ack_state(); + self.clear_streams(); + } + + fn push_next_pending_stream_close( + &mut self, + builder: &mut SessionRecordBuilder, + outbound: &mut TrackedRecord, + ) { + let len = self.state.streams.len(); + if len == 0 { + return; + } + + let start = self.state.next_stream_index % len; + for offset in 0..len { + let index = (start + offset) % len; + let stream = self.state.streams.get_index_mut(index).unwrap().1; + let Some(close) = stream.pending_close.as_ref() else { + continue; + }; + if !builder.push_stream_close(close) { + break; + } + + outbound.frames.push(TrackedFrame::StreamClose( + stream.pending_close.take().unwrap(), + )); + } + } + + fn push_next_pending_stream_window( + &mut self, + builder: &mut SessionRecordBuilder, + outbound: &mut TrackedRecord, + ) { + let len = self.state.streams.len(); + if len == 0 { + return; + } + + let start = self.state.next_stream_index % len; + for offset in 0..len { + let index = (start + offset) % len; + let (&stream_id, stream) = self.state.streams.get_index_mut(index).unwrap(); + if !stream.pending_window { + continue; + } + let frame = StreamWindow { + stream_id, + maximum_offset: VarInt::from_u64(stream.recv_limit()).unwrap(), + }; + if !builder.push_stream_window(&frame) { + break; + } + + stream.pending_window = false; + stream.advertised_max_offset = frame.maximum_offset.into_inner(); + outbound + .window_updates + .push((stream_id, frame.maximum_offset.into_inner())); + } + } + + fn push_next_stream_data( + &mut self, + builder: &mut SessionRecordBuilder, + outbound: &mut TrackedRecord, + ) { + const OVERHEAD: usize = 1 + StreamData::>::MIN_WIRE_SIZE; + + let len = self.state.streams.len(); + if len == 0 { + return; + } + + let start = self.state.next_stream_index % len; + let mut next_index = start; + + for offset in 0..len { + let Some(max_payload) = builder.remaining_capacity().checked_sub(OVERHEAD) else { + break; + }; + + let index = (start + offset) % len; + let (&stream_id, stream) = self.state.streams.get_index_mut(index).unwrap(); + if matches!(stream.outbound_state, OutboundState::Closed) { + continue; + } + let Some(candidate) = stream.tx.poll_transmit(max_payload, stream.peer_max_offset) + else { + continue; + }; + let offset = + VarInt::from_u64(candidate.offset).expect("stream offsets must fit ql-wire varint"); + let frame = StreamData { + stream_id, + offset, + header: if matches!(stream.role, StreamRole::Initiator) && candidate.offset == 0 { + stream.route_id.map(|route_id| StreamHeader { route_id }) + } else { + None + }, + fin: candidate.fin, + bytes: stream.tx.ranged_bytes(candidate), + }; + let res = builder.push_stream_data(&frame); + assert!(res, "builder has capacity"); + + if candidate.fin { + stream.outbound_state = OutboundState::Finished; + } + outbound + .frames + .push(TrackedFrame::StreamData(TrackedStreamData { + stream_id, + offset: candidate.offset, + len: candidate.len, + fin: candidate.fin, + })); + next_index = (index + 1) % len; + } + + self.state.next_stream_index = next_index; + } + + fn ensure_session_open(&self) -> Result<(), NoSessionError> { + if self.state.phase == SessionPhase::Open { + Ok(()) + } else { + Err(NoSessionError) + } + } + + fn process_record_ack(&mut self, ack: &RecordAck, sink: &mut impl EventSink) { + let stream_send_buffer_size = self.config.stream_send_buffer_size; + let acked_records = self + .state + .tracked_records + .extract_if(.., |_, record| { + record.sent_at.is_some() && ack.contains(record.seq.into_inner()) + }) + .map(|(_, record)| record) + .collect::>(); + + for record in acked_records { + for frame in &record.frames { + acknowledge_tracked_frame( + &mut self.state.streams, + stream_send_buffer_size, + frame, + sink, + ); + } + } + self.reap_reapable_streams(); + } + + fn schedule_ack(&mut self, now: Instant, immediate: bool) { + self.state.ack_tracker.schedule_ack(if immediate { + now + } else { + now + self.config.ack_delay + }); + } + + fn pending_ack(&self, remaining_capacity: usize) -> Option { + let max_ack_wire_size = remaining_capacity.checked_sub(1)?; + self.state.ack_tracker.pending_ack(max_ack_wire_size) + } + + fn collect_timeouts(&mut self, now: Instant) { + let retransmit_timeout = self.config.retransmit_timeout; + for (_, record) in self.state.tracked_records.extract_if(.., |_, record| { + record + .sent_at + .is_some_and(|sent_at| sent_at + retransmit_timeout <= now) + }) { + restore_tracked_record( + now, + &mut self.state.ack_tracker, + &mut self.state.pending_ping, + &mut self.state.streams, + record, + ); + } + } + + fn handle_stream_data( + &mut self, + frame: StreamData, + sink: &mut impl EventSink, + ) -> Result<(), ()> { + let StreamData { + stream_id, + offset, + header, + fin, + bytes, + } = frame; + let stream = match self.state.streams.get_mut(&stream_id) { + Some(stream) => stream, + None => match self.create_remote_stream(stream_id)? { + Some(stream) => stream, + None => return Ok(()), + }, + }; + + let frame_offset = offset.into_inner(); + let Some(frame_end) = frame_offset.checked_add(bytes.len() as u64) else { + return Err(()); + }; + let readable_before = stream.readable_bytes(); + let was_finished = matches!(stream.inbound_state, InboundState::Finished); + + let opened_route = match (stream.role, stream.route_id, header, frame_offset) { + (StreamRole::Responder, None, Some(header), 0) => { + stream.route_id = Some(header.route_id); + Some(header.route_id) + } + (StreamRole::Initiator, _, Some(_), _) + | (StreamRole::Responder, None, Some(_), _) + | (StreamRole::Responder, None, None, 0) => return Err(()), + _ => None, + }; + + match stream.inbound_state { + InboundState::Open => {} + InboundState::Discarding | InboundState::Closed(_) => return Ok(()), + InboundState::Finished => { + // finished stream should always have a final offset + let Some(final_offset) = stream.rx.final_offset() else { + debug_assert!(false, "finished stream must retain final offset"); + return Ok(()); + }; + + // retransmitted data for an already-finished stream is fine as long as it stays + // within the finalized byte range and any repeated FIN lands on that same offset. + if (!frame.fin || frame_end == final_offset) && frame_end <= final_offset { + if let Some(route_id) = opened_route { + sink.emit(SessionEvent::Opened { + stream_id, + route_id, + }); + if readable_before > 0 { + sink.emit(SessionEvent::Readable(stream_id)); + } else { + sink.emit(SessionEvent::Finished(stream_id)); + } + } + return Ok(()); + } + + return Err(()); + } + } + + let outcome = stream.rx.insert(frame_offset, fin, bytes).map_err(|_| ())?; + + if outcome.became_complete { + stream.inbound_state = InboundState::Finished; + } + + if let Some(route_id) = opened_route { + sink.emit(SessionEvent::Opened { + stream_id, + route_id, + }); + } + + if stream.route_id.is_some() && readable_before == 0 && stream.readable_bytes() > 0 { + sink.emit(SessionEvent::Readable(stream_id)); + } + + if stream.route_id.is_some() + && !was_finished + && matches!(stream.inbound_state, InboundState::Finished) + && stream.readable_bytes() == 0 + { + sink.emit(SessionEvent::Finished(stream_id)); + } + + self.try_reap_stream(stream_id); + Ok(()) + } + + fn handle_stream_window(&mut self, frame: &StreamWindow, sink: &mut impl EventSink) { + let Some(stream) = self.state.streams.get_mut(&frame.stream_id) else { + return; + }; + + let was_full = stream.send_capacity(self.config.stream_send_buffer_size) == 0; + let maximum_offset = frame.maximum_offset.into_inner(); + if maximum_offset > stream.peer_max_offset { + stream.peer_max_offset = maximum_offset; + } + if was_full && stream.send_capacity(self.config.stream_send_buffer_size) > 0 { + sink.emit(SessionEvent::Writable(frame.stream_id)); + } + } + + fn handle_stream_close( + &mut self, + frame: &StreamClose, + sink: &mut impl EventSink, + ) -> Result<(), ()> { + let stream_id = frame.stream_id; + let stream = match self.state.streams.get_mut(&stream_id) { + Some(stream) => stream, + None => match self.create_remote_stream(stream_id)? { + Some(stream) => stream, + None => return Ok(()), + }, + }; + + if Self::target_affects_inbound(stream.role, frame.target) + && !matches!( + stream.inbound_state, + InboundState::Closed(_) | InboundState::Discarding + ) + { + stream.inbound_state = InboundState::Closed(frame.clone()); + stream.reset_recv(); + sink.emit(SessionEvent::Closed(frame.clone())); + } + if Self::target_affects_outbound(stream.role, frame.target) + && !matches!(stream.outbound_state, OutboundState::Closed) + { + stream.outbound_state = OutboundState::Closed; + stream.tx.clear(); + stream.pending_close = None; + sink.emit(SessionEvent::WritableClosed(frame.clone())); + } + self.try_reap_stream(frame.stream_id); + Ok(()) + } + + fn apply_local_close_to_stream(stream: &mut StreamState, target: CloseTarget) { + if Self::target_affects_inbound(stream.role, target) { + stream.inbound_state = InboundState::Discarding; + stream.reset_recv(); + } + if Self::target_affects_outbound(stream.role, target) { + stream.outbound_state = OutboundState::Closed; + stream.tx.clear(); + } + } + + fn target_affects_inbound(role: StreamRole, target: CloseTarget) -> bool { + matches!(target, CloseTarget::Both) || role.inbound_target() == target + } + + fn target_affects_outbound(role: StreamRole, target: CloseTarget) -> bool { + matches!(target, CloseTarget::Both) || role.outbound_target() == target + } + + fn stream_is_reapable(&self, stream_id: StreamId, stream: &StreamState) -> bool { + let tracked_refs_stream = self.state.tracked_records.values().any(|record| { + record.window_updates.iter().any(|(id, _)| *id == stream_id) + || record.frames.iter().any(|frame| match frame { + TrackedFrame::StreamData(frame) => frame.stream_id == stream_id, + TrackedFrame::StreamClose(frame) => frame.stream_id == stream_id, + }) + }); + if tracked_refs_stream { + return false; + } + + if !stream.tx.is_empty() + || stream.pending_close.is_some() + || stream.pending_window + || stream.readable_bytes() > 0 + || stream.rx.buffered_end_offset() > stream.rx.start_offset() + { + return false; + } + + matches!( + stream.inbound_state, + InboundState::Finished | InboundState::Closed(_) | InboundState::Discarding + ) && matches!( + stream.outbound_state, + OutboundState::Finished | OutboundState::Closed + ) + } + + fn reap_reapable_streams(&mut self) { + let mut index = 0usize; + while index < self.state.streams.len() { + let stream_id = *self.state.streams.get_index(index).unwrap().0; + let len_before = self.state.streams.len(); + self.try_reap_stream(stream_id); + if self.state.streams.len() == len_before { + index += 1; + } + } + } + + fn try_reap_stream(&mut self, stream_id: StreamId) { + let Some(index) = self.state.streams.get_index_of(&stream_id) else { + return; + }; + self.try_reap_stream_at(stream_id, index); + } + + fn try_reap_stream_at(&mut self, stream_id: StreamId, index: usize) { + let Some((indexed_stream_id, stream)) = self.state.streams.get_index(index) else { + return; + }; + debug_assert_eq!(*indexed_stream_id, stream_id); + if !self.stream_is_reapable(stream_id, stream) { + return; + } + self.reap_stream_at(index); + } + + fn reap_stream_at(&mut self, index: usize) { + self.state.streams.shift_remove_index(index); + + if self.state.streams.is_empty() { + self.state.next_stream_index = 0; + return; + } + if index < self.state.next_stream_index { + self.state.next_stream_index -= 1; + } + if self.state.next_stream_index >= self.state.streams.len() { + self.state.next_stream_index %= self.state.streams.len(); + } + } + + fn clear_streams(&mut self) { + self.state.next_stream_index = 0; + self.state.streams.clear(); + } + + fn create_remote_stream( + &mut self, + stream_id: StreamId, + ) -> Result, ()> { + match classify_missing_stream( + self.config.local_parity, + self.state.next_stream_ordinal, + stream_id, + &mut self.state.remote_stream_history, + ) { + MissingStreamAction::Create => {} + MissingStreamAction::Ignore => return Ok(None), + MissingStreamAction::FailProtocol => { + return Err(()); + } + } + + let stream = self + .state + .streams + .entry(stream_id) + .insert_entry(StreamState::new( + StreamRole::Responder, + None, + self.config.stream_receive_buffer_size, + self.config.initial_peer_stream_receive_window, + )); + + Ok(Some(stream.into_mut())) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum MissingStreamAction { + Create, + Ignore, + FailProtocol, +} + +fn classify_missing_stream( + local_parity: StreamParity, + next_stream_ordinal: u32, + stream_id: StreamId, + remote_stream_history: &mut RemoteStreamHistory, +) -> MissingStreamAction { + if !local_parity.remote().matches(stream_id) { + return if local_stream_was_opened(local_parity, next_stream_ordinal, stream_id) { + MissingStreamAction::Ignore + } else { + MissingStreamAction::FailProtocol + }; + } + + if remote_stream_history.observe(stream_id) { + MissingStreamAction::Ignore + } else { + MissingStreamAction::Create + } +} + +fn local_stream_was_opened( + local_parity: StreamParity, + next_stream_ordinal: u32, + stream_id: StreamId, +) -> bool { + local_parity.matches(stream_id) + && stream_id.into_inner() + < local_parity + .make_stream_id(next_stream_ordinal) + .into_inner() +} + +fn restore_tracked_record( + now: Instant, + ack_tracker: &mut AckTracker, + pending_ping: &mut bool, + streams: &mut IndexMap, + record: TrackedRecord, +) { + if let Some(ack) = &record.ack { + ack_tracker.restore_acked_ranges(ack, now); + } + if record.ping_included { + *pending_ping = true; + } + for (stream_id, maximum_offset) in record.window_updates { + if let Some(stream) = streams.get_mut(&stream_id) { + if stream.recv_limit() >= maximum_offset { + stream.pending_window = true; + } + } + } + for frame in record.frames { + requeue_tracked_frame(streams, frame); + } +} + +fn requeue_tracked_frame(streams: &mut IndexMap, frame: TrackedFrame) { + match frame { + TrackedFrame::StreamClose(close) => restore_stream_close(streams, close), + TrackedFrame::StreamData(frame) => restore_stream_data(streams, frame), + } +} + +fn restore_stream_close(streams: &mut IndexMap, close: StreamClose) { + if let Some(stream) = streams.get_mut(&close.stream_id) { + stream.pending_close = Some(close); + } +} + +fn restore_stream_data(streams: &mut IndexMap, frame: TrackedStreamData) { + if let Some(stream) = streams.get_mut(&frame.stream_id) { + if matches!(stream.outbound_state, OutboundState::Closed) { + return; + } + stream.tx.retransmit(stream_tx::StreamTxRange { + offset: frame.offset, + len: frame.len, + fin: frame.fin, + }); + if frame.fin && matches!(stream.outbound_state, OutboundState::Finished) { + stream.outbound_state = OutboundState::FinQueued; + } + } +} + +fn acknowledge_tracked_frame( + streams: &mut IndexMap, + stream_send_buffer_size: usize, + frame: &TrackedFrame, + sink: &mut impl EventSink, +) { + match frame { + TrackedFrame::StreamClose(_) => {} + TrackedFrame::StreamData(frame) => { + let stream_id = frame.stream_id; + if let Some(stream) = streams.get_mut(&stream_id) { + let was_full = stream.send_capacity(stream_send_buffer_size) == 0; + let had_unacked_fin = frame.fin && stream.tx.has_unacked_fin(); + stream.tx.ack(StreamTxRange { + offset: frame.offset, + len: frame.len, + fin: frame.fin, + }); + if was_full && stream.send_capacity(stream_send_buffer_size) > 0 { + sink.emit(SessionEvent::Writable(stream_id)); + } + if had_unacked_fin && !stream.tx.has_unacked_fin() { + sink.emit(SessionEvent::OutboundFinished(stream_id)); + } + } + } + } +} + +#[inline] +#[track_caller] +fn next_seq(seq: &mut RecordSeq) { + *seq = seq + .into_inner() + .checked_add(1) + .and_then(|next| RecordSeq::from_u64(next).ok()) + .expect("record sequence overflow"); +} diff --git a/ql-fsm/src/session/range_set.rs b/ql-fsm/src/session/range_set.rs new file mode 100644 index 00000000..53d66269 --- /dev/null +++ b/ql-fsm/src/session/range_set.rs @@ -0,0 +1,221 @@ +use std::{ + cmp, + collections::BTreeMap, + ops::{ + Bound::{Excluded, Included}, + Range, + }, +}; + +/// A set of `u64` values optimized for long runs and random insert/delete. +#[derive(Debug, Default, Clone, PartialEq, Eq)] +pub struct RangeSet(BTreeMap); + +impl RangeSet { + pub fn new() -> Self { + Self::default() + } + + pub fn insert(&mut self, mut x: Range) -> bool { + if x.is_empty() { + return false; + } + + if let Some((start, end)) = self.before(x.start) { + if end >= x.end { + return false; + } else if end >= x.start { + self.0.remove(&start); + x.start = start; + } + } + + while let Some((next_start, next_end)) = self.after(x.start) { + if next_start > x.end { + break; + } + self.0.remove(&next_start); + x.end = cmp::max(next_end, x.end); + } + + self.0.insert(x.start, x.end); + true + } + + pub fn remove(&mut self, x: Range) -> bool { + if x.is_empty() { + return false; + } + + let before = match self.before(x.start) { + Some((start, end)) if end > x.start => { + self.0.remove(&start); + if start < x.start { + self.0.insert(start, x.start); + } + if end > x.end { + self.0.insert(x.end, end); + } + if end >= x.end { + return true; + } + true + } + Some(_) | None => false, + }; + + let mut after = false; + while let Some((start, end)) = self.after(x.start) { + if start >= x.end { + break; + } + after = true; + self.0.remove(&start); + if end > x.end { + self.0.insert(x.end, end); + break; + } + } + + before || after + } + + pub fn min(&self) -> Option { + self.0.first_key_value().map(|(&start, _)| start) + } + + pub fn max(&self) -> Option { + self.0 + .last_key_value() + .map(|(_, &end)| end.checked_sub(1).unwrap()) + } + + pub fn contains(&self, x: u64) -> bool { + self.before(x).is_some_and(|(_, end)| end > x) + } + + pub fn range_count(&self) -> usize { + self.0.len() + } + + pub fn iter(&self) -> Iter<'_> { + Iter(self.0.iter()) + } + + pub fn iter_rev(&self) -> RevIter<'_> { + RevIter(self.0.iter().rev()) + } + + pub fn peek_min(&self) -> Option> { + let (&start, &end) = self.0.iter().next()?; + Some(start..end) + } + + pub fn pop_min(&mut self) -> Option> { + let result = self.peek_min()?; + self.0.remove(&result.start); + Some(result) + } + + #[cfg(test)] + pub fn peek_max(&self) -> Option> { + let (&start, &end) = self.0.iter().next_back()?; + Some(start..end) + } + + #[cfg(test)] + pub fn pop_max(&mut self) -> Option> { + let result = self.peek_max()?; + self.0.remove(&result.start); + Some(result) + } + + /// find closest range to `x` that begins at or before it + fn before(&self, x: u64) -> Option<(u64, u64)> { + self.0 + .range((Included(0), Included(x))) + .next_back() + .map(|(&start, &end)| (start, end)) + } + + /// find the closest range to `x` that begins after it + fn after(&self, x: u64) -> Option<(u64, u64)> { + self.0 + .range((Excluded(x), Included(u64::MAX))) + .next() + .map(|(&start, &end)| (start, end)) + } +} + +pub struct Iter<'a>(std::collections::btree_map::Iter<'a, u64, u64>); + +impl Iterator for Iter<'_> { + type Item = Range; + + fn next(&mut self) -> Option { + self.0.next().map(|(&start, &end)| start..end) + } +} + +pub struct RevIter<'a>(std::iter::Rev>); + +impl Iterator for RevIter<'_> { + type Item = Range; + + fn next(&mut self) -> Option { + self.0.next().map(|(&start, &end)| start..end) + } +} + +#[cfg(test)] +mod tests { + use super::RangeSet; + + #[test] + fn insert_merges_overlaps() { + let mut set = RangeSet::new(); + assert!(set.insert(10..20)); + assert!(set.insert(30..40)); + assert!(set.insert(15..35)); + assert_eq!(set.iter().collect::>(), vec![10..40]); + } + + #[test] + fn remove_splits_ranges() { + let mut set = RangeSet::new(); + set.insert(10..40); + assert!(set.remove(20..30)); + assert_eq!(set.iter().collect::>(), vec![10..20, 30..40]); + } + + #[test] + fn reverse_iteration_visits_highest_range_first() { + let mut set = RangeSet::new(); + set.insert(10..20); + set.insert(30..40); + set.insert(50..60); + + assert_eq!( + set.iter_rev().collect::>(), + vec![50..60, 30..40, 10..20] + ); + assert_eq!(set.peek_max(), Some(50..60)); + assert_eq!(set.pop_max(), Some(50..60)); + assert_eq!(set.iter().collect::>(), vec![10..20, 30..40]); + } + + #[test] + fn contains_and_max_reflect_current_membership() { + let mut set = RangeSet::new(); + set.insert(10..20); + set.insert(30..31); + + assert!(!set.contains(9)); + assert!(set.contains(10)); + assert!(set.contains(19)); + assert!(!set.contains(20)); + assert_eq!(set.min(), Some(10)); + assert_eq!(set.max(), Some(30)); + assert_eq!(set.range_count(), 2); + } +} diff --git a/ql-fsm/src/session/remote_stream_history.rs b/ql-fsm/src/session/remote_stream_history.rs new file mode 100644 index 00000000..76c1e8bb --- /dev/null +++ b/ql-fsm/src/session/remote_stream_history.rs @@ -0,0 +1,60 @@ +use ql_wire::StreamId; + +use super::{range_set::RangeSet, stream_parity::StreamParity}; + +#[derive(Debug)] +pub struct RemoteStreamHistory { + parity: StreamParity, + seen: RangeSet, +} + +impl RemoteStreamHistory { + pub fn new(parity: StreamParity) -> Self { + Self { + parity, + seen: RangeSet::new(), + } + } + + /// returns true when this remote stream id was already observed before + /// panics if `stream_id` is wrong stream parity + #[allow(clippy::range_plus_one)] + pub fn observe(&mut self, stream_id: StreamId) -> bool { + let ordinal = self + .stream_ordinal(stream_id) + .expect("remote stream history used with wrong stream parity"); + !self.seen.insert(ordinal..ordinal + 1) + } + + fn stream_ordinal(&self, stream_id: StreamId) -> Option { + let delta = stream_id + .into_inner() + .checked_sub(u64::from(self.parity.first_stream_id()))?; + if delta % 2 != 0 { + return None; + } + Some(delta / 2) + } +} + +#[cfg(test)] +mod tests { + use super::RemoteStreamHistory; + use crate::session::stream_parity::StreamParity; + + #[test] + fn observe() { + let parity = StreamParity::Even; + let mut history = RemoteStreamHistory::new(parity); + + assert!(!history.observe(parity.make_stream_id(2))); + assert!(!history.observe(parity.make_stream_id(5))); + assert!(!history.observe(parity.make_stream_id(0))); + assert!(!history.observe(parity.make_stream_id(4))); + assert!(history.observe(parity.make_stream_id(2))); + assert!(!history.observe(parity.make_stream_id(1))); + assert!(history.observe(parity.make_stream_id(5))); + assert!(!history.observe(parity.make_stream_id(3))); + assert!(history.observe(parity.make_stream_id(0))); + } +} diff --git a/ql-fsm/src/session/state.rs b/ql-fsm/src/session/state.rs new file mode 100644 index 00000000..b63140a1 --- /dev/null +++ b/ql-fsm/src/session/state.rs @@ -0,0 +1,140 @@ +use std::time::Instant; + +use indexmap::IndexMap; +use ql_wire::{CloseTarget, RecordSeq, RouteId, SessionClose, StreamClose, StreamId}; + +use super::{ + ack_tracker::AckTracker, remote_stream_history::RemoteStreamHistory, stream_rx::StreamRx, + stream_tx::StreamTx, tracked::TrackedRecord, +}; + +pub struct SessionState { + pub last_activity_at: Instant, + pub last_inbound_at: Instant, + pub phase: SessionPhase, + pub next_stream_ordinal: u32, + pub next_record_seq: RecordSeq, + pub next_write_id: u64, + pub tracked_records: IndexMap, + pub ack_tracker: AckTracker, + pub pending_ping: bool, + pub streams: IndexMap, + pub next_stream_index: usize, + pub remote_stream_history: RemoteStreamHistory, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum SessionPhase { + Open, + Terminating(TerminalFrame), + Closed, +} + +impl SessionPhase { + pub fn is_open(&self) -> bool { + self == &Self::Open + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum TerminalFrame { + Close(SessionClose), + Unpair, +} + +#[derive(Debug)] +pub struct StreamState { + pub role: StreamRole, + pub route_id: Option, + pub rx: StreamRx, + pub tx: StreamTx, + pub pending_close: Option, + pub peer_max_offset: u64, + pub outbound_state: OutboundState, + pub inbound_state: InboundState, + pub advertised_max_offset: u64, + pub pending_window: bool, +} + +impl StreamState { + pub fn new( + role: StreamRole, + route_id: Option, + receive_buffer_size: u32, + initial_peer_stream_receive_window: u32, + ) -> Self { + let receive_buffer_size = receive_buffer_size as usize; + Self { + role, + route_id, + tx: StreamTx::new(), + pending_close: None, + peer_max_offset: u64::from(initial_peer_stream_receive_window), + outbound_state: OutboundState::Open, + inbound_state: InboundState::Open, + rx: StreamRx::new(receive_buffer_size), + advertised_max_offset: receive_buffer_size as u64, + pending_window: false, + } + } + + pub fn is_writable(&self) -> bool { + matches!(self.outbound_state, OutboundState::Open) + } + + pub fn send_capacity(&self, send_buffer_size: usize) -> usize { + send_buffer_size.saturating_sub(self.tx.buffered_len()) + } + + pub fn readable_bytes(&self) -> usize { + self.rx.readable_len() + } + + pub fn recv_limit(&self) -> u64 { + self.rx + .start_offset() + .saturating_add(self.rx.max_buffered() as u64) + } + + pub fn reset_recv(&mut self) { + self.rx = StreamRx::with_start_offset(self.rx.start_offset(), self.rx.max_buffered()); + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamRole { + Initiator, + Responder, +} + +impl StreamRole { + pub fn outbound_target(self) -> CloseTarget { + match self { + Self::Initiator => CloseTarget::Origin, + Self::Responder => CloseTarget::Return, + } + } + + pub fn inbound_target(self) -> CloseTarget { + match self { + Self::Initiator => CloseTarget::Return, + Self::Responder => CloseTarget::Origin, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum OutboundState { + Open, + FinQueued, + Finished, + Closed, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum InboundState { + Open, + Finished, + Closed(StreamClose), + Discarding, +} diff --git a/ql-fsm/src/session/stream_ops.rs b/ql-fsm/src/session/stream_ops.rs new file mode 100644 index 00000000..548189b7 --- /dev/null +++ b/ql-fsm/src/session/stream_ops.rs @@ -0,0 +1,147 @@ +use ql_wire::{CloseTarget, StreamClose, StreamCloseCode, StreamId}; + +use super::{ + state::{InboundState, StreamState}, + stream_rx::StreamReadIter, + EventSink, SessionEvent, SessionFsm, +}; +use crate::CommitReadError; + +pub struct StreamOps<'a, E> { + session: &'a mut SessionFsm, + emit: E, + stream_id: StreamId, + stream_index: usize, + reap_on_drop: bool, +} + +impl<'a, E: EventSink> StreamOps<'a, E> { + pub(super) fn new( + session: &'a mut SessionFsm, + stream_id: StreamId, + stream_index: usize, + emit: E, + ) -> Self { + Self { + session, + emit, + stream_id, + stream_index, + reap_on_drop: false, + } + } + + /// returns this stream's identifier + pub fn stream_id(&self) -> StreamId { + self.stream_id + } + + /// returns the readable stream bytes as owned `Bytes` views without consuming them + pub fn read(&self) -> StreamReadIter<'_> { + self.stream().rx.bytes() + } + + /// returns how many bytes can be read from the stream + pub fn readable_bytes(&self) -> usize { + self.stream().readable_bytes() + } + + /// marks previously read bytes as consumed + pub fn commit_read(&mut self, len: usize) -> Result<(), CommitReadError> { + let stream_id = self.stream_id; + let emit_finished = { + let stream = self.stream_mut(); + if len > stream.readable_bytes() { + return Err(CommitReadError); + } + stream.rx.consume(len); + if stream.recv_limit() > stream.advertised_max_offset { + stream.pending_window = true; + } + stream.route_id.is_some() + && matches!(stream.inbound_state, InboundState::Finished) + && stream.readable_bytes() == 0 + }; + if emit_finished { + self.emit.emit(SessionEvent::Finished(stream_id)); + } + self.reap_on_drop = true; + Ok(()) + } + + /// returns a writer if the local write side is still open + pub fn writer(&mut self) -> Option> { + let send_buffer_size = self.session.config.stream_send_buffer_size; + let stream = self.stream_mut(); + if !stream.is_writable() { + return None; + } + Some(StreamWriter::new(stream, send_buffer_size)) + } + + /// closes the origin lane, return lane, or both lanes of the stream + pub fn close(&mut self, target: CloseTarget, code: StreamCloseCode) { + let stream_id = self.stream_id; + let stream = self.stream_mut(); + SessionFsm::apply_local_close_to_stream(stream, target); + stream.pending_close = Some(StreamClose { + stream_id, + target, + code, + }); + self.reap_on_drop = true; + } + + fn stream(&self) -> &StreamState { + &self.session.state.streams[self.stream_index] + } + + fn stream_mut(&mut self) -> &mut StreamState { + &mut self.session.state.streams[self.stream_index] + } +} + +impl Drop for StreamOps<'_, E> { + fn drop(&mut self) { + if !self.reap_on_drop { + return; + } + + self.session + .try_reap_stream_at(self.stream_id, self.stream_index); + } +} + +pub struct StreamWriter<'a> { + stream: &'a mut StreamState, + send_buffer_size: usize, +} + +impl<'a> StreamWriter<'a> { + pub(super) fn new(stream: &'a mut StreamState, send_buffer_size: usize) -> Self { + Self { + stream, + send_buffer_size, + } + } + + /// returns how many bytes can still be buffered for local writes + pub fn capacity(&self) -> usize { + self.stream.send_capacity(self.send_buffer_size) + } + + /// appends as many bytes as possible and returns the accepted count + pub fn write(&mut self, bytes: &mut bytes::Bytes) -> usize { + let accepted = bytes.len().min(self.capacity()); + if accepted > 0 { + self.stream.tx.append(bytes.split_to(accepted)); + } + accepted + } + + /// marks the local write side as finished + pub fn finish(self) { + self.stream.tx.queue_fin(); + self.stream.outbound_state = super::state::OutboundState::FinQueued; + } +} diff --git a/ql-fsm/src/session/stream_parity.rs b/ql-fsm/src/session/stream_parity.rs new file mode 100644 index 00000000..70f60776 --- /dev/null +++ b/ql-fsm/src/session/stream_parity.rs @@ -0,0 +1,44 @@ +use ql_wire::{StreamId, QID}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamParity { + Even, + Odd, +} + +impl StreamParity { + pub fn for_local(local: QID, peer: QID) -> Self { + match local.0.cmp(&peer.0) { + std::cmp::Ordering::Less | std::cmp::Ordering::Equal => Self::Even, + std::cmp::Ordering::Greater => Self::Odd, + } + } + + pub const fn first_stream_id(self) -> u32 { + match self { + Self::Even => 0, + Self::Odd => 1, + } + } + + pub const fn matches(self, stream_id: StreamId) -> bool { + match self { + Self::Even => stream_id.into_inner() % 2 == 0, + Self::Odd => stream_id.into_inner() % 2 == 1, + } + } + + pub const fn remote(self) -> Self { + match self { + Self::Even => Self::Odd, + Self::Odd => Self::Even, + } + } + + pub fn make_stream_id(self, ordinal: u32) -> StreamId { + StreamId(ql_wire::VarInt::from_u32( + self.first_stream_id() + .saturating_add(ordinal.saturating_mul(2)), + )) + } +} diff --git a/ql-fsm/src/session/stream_rx.rs b/ql-fsm/src/session/stream_rx.rs new file mode 100644 index 00000000..0f5a8eab --- /dev/null +++ b/ql-fsm/src/session/stream_rx.rs @@ -0,0 +1,428 @@ +use std::collections::{btree_map, BTreeMap}; + +use bytes::{Buf, Bytes}; + +/// reassembles one stream direction from out-of-order byte ranges. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct StreamRx { + start_offset: u64, + chunks: BTreeMap, + final_offset: Option, + max_buffered: usize, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct InsertOutcome { + pub newly_readable_bytes: usize, + pub became_complete: bool, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamRxError { + OffsetOverflow, + OutOfWindow, + InconsistentFinalOffset, + FinalOffsetBeforeBufferedData, + BeyondFinalOffset, +} + +impl StreamRx { + pub fn new(max_buffered: usize) -> Self { + Self::with_start_offset(0, max_buffered) + } + + pub fn with_start_offset(start_offset: u64, max_buffered: usize) -> Self { + Self { + start_offset, + chunks: BTreeMap::new(), + final_offset: None, + max_buffered, + } + } + + pub fn start_offset(&self) -> u64 { + self.start_offset + } + + pub fn buffered_end_offset(&self) -> u64 { + self.chunks + .last_key_value() + .map_or(self.start_offset, |(&offset, bytes)| { + offset + bytes.len() as u64 + }) + } + + pub fn final_offset(&self) -> Option { + self.final_offset + } + + pub fn max_buffered(&self) -> usize { + self.max_buffered + } + + pub fn readable_len(&self) -> usize { + let mut cursor = self.start_offset; + for (&offset, bytes) in self.chunks.range(self.start_offset..) { + if offset > cursor { + break; + } + + let end = offset + bytes.len() as u64; + if end > cursor { + cursor = end; + } + } + + usize::try_from(cursor - self.start_offset).expect("readable prefix exceeds usize") + } + + pub fn bytes(&self) -> StreamReadIter<'_> { + StreamReadIter { + inner: self.chunks.range(self.start_offset..), + cursor: self.start_offset, + remaining: self.readable_len(), + } + } + + pub fn is_complete(&self) -> bool { + matches!(self.final_offset, Some(final_offset) + if final_offset == self.buffered_end_offset() + && final_offset == self.start_offset + self.readable_len() as u64) + } + + pub fn insert( + &mut self, + offset: u64, + fin: bool, + mut bytes: Bytes, + ) -> Result { + let end = offset + .checked_add(bytes.len() as u64) + .ok_or(StreamRxError::OffsetOverflow)?; + + let was_complete = self.is_complete(); + let old_readable = self.readable_len(); + + if fin { + self.set_or_validate_final_offset(end)?; + } + if let Some(final_offset) = self.final_offset { + if end > final_offset { + return Err(StreamRxError::BeyondFinalOffset); + } + } + + if bytes.is_empty() || end <= self.start_offset { + return Ok(self.insert_outcome(was_complete, old_readable)); + } + + let effective_offset = offset.max(self.start_offset); + let trim_front = + usize::try_from(effective_offset - offset).expect("front trim exceeds usize"); + bytes.advance(trim_front); + if bytes.is_empty() { + return Ok(self.insert_outcome(was_complete, old_readable)); + } + + let effective_end = effective_offset + bytes.len() as u64; + self.ensure_within_window(effective_end)?; + self.insert_chunk(effective_offset, bytes); + + Ok(self.insert_outcome(was_complete, old_readable)) + } + + pub fn consume(&mut self, len: usize) { + let readable = self.readable_len(); + debug_assert!(len <= readable, "consume beyond readable bytes"); + if len > readable { + return; + } + + let new_start = self.start_offset.saturating_add(len as u64); + while let Some((&offset, bytes)) = self.chunks.first_key_value() { + let end = offset + bytes.len() as u64; + if end <= new_start { + self.chunks.pop_first(); + continue; + } + if offset < new_start { + let (offset, mut bytes) = self.chunks.pop_first().unwrap(); + bytes.advance(usize::try_from(new_start - offset).expect("trim exceeds usize")); + self.chunks.insert(new_start, bytes); + } + break; + } + + self.start_offset = new_start; + } + + fn insert_outcome(&self, was_complete: bool, old_readable: usize) -> InsertOutcome { + InsertOutcome { + newly_readable_bytes: self.readable_len().saturating_sub(old_readable), + became_complete: !was_complete && self.is_complete(), + } + } + + fn set_or_validate_final_offset(&mut self, final_offset: u64) -> Result<(), StreamRxError> { + if let Some(existing) = self.final_offset { + return if existing == final_offset { + Ok(()) + } else { + Err(StreamRxError::InconsistentFinalOffset) + }; + } + + let buffered_end = self.buffered_end_offset(); + if final_offset < buffered_end { + return Err(StreamRxError::FinalOffsetBeforeBufferedData); + } + + self.final_offset = Some(final_offset); + Ok(()) + } + + fn ensure_within_window(&self, end: u64) -> Result<(), StreamRxError> { + let attempted = end.saturating_sub(self.start_offset); + if attempted > self.max_buffered as u64 { + return Err(StreamRxError::OutOfWindow); + } + Ok(()) + } + + fn insert_chunk(&mut self, mut offset: u64, mut bytes: Bytes) { + if bytes.is_empty() { + return; + } + + if let Some((&existing_offset, existing)) = self.chunks.range(..offset).next_back() { + let existing_end = existing_offset + existing.len() as u64; + if existing_end > offset { + let overlap = + usize::try_from((existing_end - offset).min(bytes.len() as u64)).unwrap(); + bytes.advance(overlap); + offset += overlap as u64; + } + } + + if bytes.is_empty() { + return; + } + + let end = offset + bytes.len() as u64; + let overlapping = self + .chunks + .range(offset..end) + .map(|(&chunk_offset, _)| chunk_offset) + .collect::>(); + + for chunk_offset in overlapping { + let chunk_end = chunk_offset + self.chunks[&chunk_offset].len() as u64; + + if chunk_offset > offset { + let len = usize::try_from(chunk_offset - offset).expect("gap exceeds usize"); + self.chunks.insert(offset, bytes.slice(..len)); + bytes.advance(len); + offset = chunk_offset; + } + + let overlap = usize::try_from((chunk_end - offset).min(bytes.len() as u64)).unwrap(); + bytes.advance(overlap); + offset += overlap as u64; + + if bytes.is_empty() { + return; + } + } + + self.chunks.insert(offset, bytes); + } +} + +#[derive(Debug, Clone)] +pub struct StreamReadIter<'a> { + inner: btree_map::Range<'a, u64, Bytes>, + cursor: u64, + remaining: usize, +} + +impl Iterator for StreamReadIter<'_> { + type Item = Bytes; + + fn next(&mut self) -> Option { + while self.remaining > 0 { + let (&offset, bytes) = self.inner.next()?; + if offset > self.cursor { + self.remaining = 0; + return None; + } + + let skip = usize::try_from(self.cursor.saturating_sub(offset)) + .expect("read cursor exceeds usize"); + if skip >= bytes.len() { + continue; + } + + let len = (bytes.len() - skip).min(self.remaining); + self.remaining -= len; + self.cursor += len as u64; + return Some(bytes.slice(skip..skip + len)); + } + + None + } +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use super::{InsertOutcome, StreamRx, StreamRxError}; + + pub fn copy_readable(rx: &StreamRx) -> Vec { + let readable = rx.readable_len(); + let mut out = Vec::with_capacity(readable); + for chunk in rx.bytes() { + out.extend_from_slice(&chunk); + } + out + } + + fn bytes(bytes: &'static [u8]) -> Bytes { + Bytes::from_static(bytes) + } + + #[test] + fn contiguous_insert_becomes_readable_and_complete() { + let mut rx = StreamRx::new(64); + + let outcome = rx.insert(0, true, bytes(b"hello")).unwrap(); + + assert_eq!( + outcome, + InsertOutcome { + newly_readable_bytes: 5, + became_complete: true, + } + ); + assert_eq!(rx.readable_len(), 5); + assert_eq!(copy_readable(&rx), b"hello"); + assert_eq!(rx.final_offset, Some(5)); + assert!(rx.is_complete()); + } + + #[test] + fn out_of_order_insert_tracks_gap_until_prefix_is_filled() { + let mut rx = StreamRx::new(64); + + let first = rx.insert(5, true, bytes(b" world")).unwrap(); + assert_eq!( + first, + InsertOutcome { + newly_readable_bytes: 0, + became_complete: false, + } + ); + assert_eq!(rx.readable_len(), 0); + + let second = rx.insert(0, false, bytes(b"hello")).unwrap(); + assert_eq!( + second, + InsertOutcome { + newly_readable_bytes: 11, + became_complete: true, + } + ); + assert_eq!(copy_readable(&rx), b"hello world"); + assert!(rx.is_complete()); + } + + #[test] + fn duplicate_insert_is_ignored_if_bytes_match() { + let mut rx = StreamRx::new(64); + + rx.insert(0, false, bytes(b"hello")).unwrap(); + let duplicate = rx.insert(0, false, bytes(b"hello")).unwrap(); + + assert_eq!( + duplicate, + InsertOutcome { + newly_readable_bytes: 0, + became_complete: false, + } + ); + assert_eq!(copy_readable(&rx), b"hello"); + } + + #[test] + fn consume_advances_start_offset_and_trims_old_prefix() { + let mut rx = StreamRx::new(64); + + rx.insert(0, false, bytes(b"abcd")).unwrap(); + rx.consume(2); + assert_eq!(rx.start_offset(), 2); + assert_eq!(copy_readable(&rx), b"cd"); + + let outcome = rx.insert(1, true, bytes(b"bcde")).unwrap(); + assert_eq!( + outcome, + InsertOutcome { + newly_readable_bytes: 1, + became_complete: true, + } + ); + assert_eq!(copy_readable(&rx), b"cde"); + assert_eq!(rx.final_offset, Some(5)); + assert!(rx.is_complete()); + } + + #[test] + fn insert_can_fill_multiple_gaps_without_rebuilding_state() { + let mut rx = StreamRx::new(64); + + rx.insert(0, false, bytes(b"ab")).unwrap(); + rx.insert(4, false, bytes(b"ef")).unwrap(); + rx.insert(8, true, bytes(b"ij")).unwrap(); + + let outcome = rx.insert(2, false, bytes(b"cdefgh")).unwrap(); + + assert_eq!( + outcome, + InsertOutcome { + newly_readable_bytes: 8, + became_complete: true, + } + ); + + assert_eq!(copy_readable(&rx), b"abcdefghij"); + assert!(rx.is_complete()); + } + + #[test] + fn heavily_fragmented_inserts_stay_valid() { + let mut rx = StreamRx::new(64); + + rx.insert(1, false, bytes(b"b")).unwrap(); + rx.insert(3, false, bytes(b"d")).unwrap(); + rx.insert(5, false, bytes(b"f")).unwrap(); + rx.insert(7, false, bytes(b"h")).unwrap(); + rx.insert(9, true, bytes(b"j")).unwrap(); + + let outcome = rx.insert(0, false, bytes(b"abcdefghi")).unwrap(); + assert_eq!( + outcome, + InsertOutcome { + newly_readable_bytes: 10, + became_complete: true, + } + ); + assert_eq!(copy_readable(&rx), b"abcdefghij"); + assert!(rx.is_complete()); + } + + #[test] + fn out_of_window_insert_is_rejected() { + let mut rx = StreamRx::new(4); + let error = rx.insert(5, false, bytes(b"a")).unwrap_err(); + assert_eq!(error, StreamRxError::OutOfWindow); + } +} diff --git a/ql-fsm/src/session/stream_tx.rs b/ql-fsm/src/session/stream_tx.rs new file mode 100644 index 00000000..15533922 --- /dev/null +++ b/ql-fsm/src/session/stream_tx.rs @@ -0,0 +1,579 @@ +use std::{collections::VecDeque, ops::Range}; + +use bytes::{Buf, Bytes}; +use ql_wire::BufView; + +use super::range_set::RangeSet; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct StreamTx { + chunks: VecDeque, + buffered_len: usize, + base_offset: u64, + unsent: u64, + acked: RangeSet, + retransmits: RangeSet, + final_offset: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct TrackedFinalOffset { + offset: u64, + state: SendState, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum SendState { + Unsent, + Sent, + Lost, + Acked, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct StreamTxRange { + pub offset: u64, + pub len: usize, + pub fin: bool, +} + +#[derive(Debug, Clone, Copy)] +pub struct StreamTxBytes<'a> { + inner: &'a VecDeque, + offset: usize, + len: usize, +} + +pub struct StreamTxBuf<'a> { + inner: std::collections::vec_deque::Iter<'a, Bytes>, + skip: usize, + remaining: usize, + current: &'a [u8], +} + +impl BufView for StreamTxBytes<'_> { + type Buf<'a> + = StreamTxBuf<'a> + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + let mut buf = StreamTxBuf { + inner: self.inner.iter(), + skip: self.offset, + remaining: self.len, + current: &[], + }; + buf.refill(); + buf + } +} + +impl StreamTxBuf<'_> { + fn refill(&mut self) { + if self.remaining == 0 { + self.current = &[]; + return; + } + + for chunk in self.inner.by_ref() { + if self.skip >= chunk.len() { + self.skip -= chunk.len(); + continue; + } + + let chunk = &chunk[self.skip..]; + self.skip = 0; + if chunk.is_empty() { + continue; + } + + let len = chunk.len().min(self.remaining); + self.current = &chunk[..len]; + return; + } + + self.current = &[]; + } +} + +impl Buf for StreamTxBuf<'_> { + fn remaining(&self) -> usize { + self.remaining + } + + fn chunk(&self) -> &[u8] { + self.current + } + + fn advance(&mut self, cnt: usize) { + let remaining = self.remaining; + assert!( + cnt <= remaining, + "cannot advance past remaining bytes: {cnt} > {remaining}", + ); + + self.remaining -= cnt; + let mut cnt = cnt; + while cnt > 0 { + if cnt < self.current.len() { + self.current = &self.current[cnt..]; + return; + } + + cnt -= self.current.len(); + self.refill(); + } + + if self.remaining == 0 { + self.current = &[]; + } + } +} + +impl StreamTx { + pub fn new() -> Self { + Self { + chunks: VecDeque::new(), + buffered_len: 0, + base_offset: 0, + unsent: 0, + acked: RangeSet::new(), + retransmits: RangeSet::new(), + final_offset: None, + } + } + + pub fn buffered_len(&self) -> usize { + self.buffered_len + } + + pub fn end_offset(&self) -> u64 { + self.base_offset + self.buffered_len as u64 + } + + pub fn is_empty(&self) -> bool { + self.buffered_len == 0 && self.final_offset.is_none() + } + + pub fn append(&mut self, bytes: Bytes) { + if bytes.is_empty() { + return; + } + + self.buffered_len += bytes.len(); + self.chunks.push_back(bytes); + } + + pub fn queue_fin(&mut self) { + self.final_offset = Some(TrackedFinalOffset { + offset: self.end_offset(), + state: SendState::Unsent, + }); + } + + pub fn has_unacked_fin(&self) -> bool { + self.final_offset + .is_some_and(|final_offset| final_offset.state != SendState::Acked) + } + + pub fn poll_transmit( + &mut self, + max_payload: usize, + peer_max_offset: u64, + ) -> Option { + let budget_end = |start: u64| { + start + .saturating_add(max_payload as u64) + .min(peer_max_offset) + }; + + // prefer the lowest lost bytes before sending new bytes + if let Some(range) = self.retransmits.peek_min() { + let mut end = range.end.min(budget_end(range.start)); + + // extend only when lost bytes end where unsent bytes begin + if end == range.end && range.end == self.unsent { + end = self.end_offset().min(budget_end(range.start)); + } + + if end > range.start { + let range = self.retransmits.pop_min().unwrap(); + if end < range.end { + self.retransmits.insert(end..range.end); + } + + // mark any new bytes in this frame as sent + self.unsent = self.unsent.max(end); + return Some(StreamTxRange { + offset: range.start, + len: usize::try_from(end - range.start).unwrap(), + fin: self.poll_fin(end), + }); + } + } + + // send bytes that have not been sent yet + if self.unsent < self.end_offset() { + let end = self.end_offset().min(budget_end(self.unsent)); + if end > self.unsent { + let start = self.unsent; + self.unsent = end; + return Some(StreamTxRange { + offset: start, + len: usize::try_from(end - start).unwrap(), + fin: self.poll_fin(end), + }); + } + } + + // send a fin after all data has been sent + let final_offset = + self.final_offset + .as_mut() + .filter(|TrackedFinalOffset { offset, state }| { + (*state == SendState::Lost || *state == SendState::Unsent) + && *offset <= peer_max_offset + })?; + final_offset.state = SendState::Sent; + Some(StreamTxRange { + offset: final_offset.offset, + len: 0, + fin: true, + }) + } + + pub fn ranged_bytes(&self, range: StreamTxRange) -> StreamTxBytes<'_> { + let offset = usize::try_from(range.offset - self.base_offset).unwrap(); + let len = range.len.min(self.buffered_len.saturating_sub(offset)); + StreamTxBytes { + inner: &self.chunks, + offset, + len, + } + } + + pub fn retransmit(&mut self, range: StreamTxRange) { + if let Some(range) = self.clamp_sent_range(range.offset, range.len) { + Self::insert_not_acked(&self.acked, &mut self.retransmits, range); + } + if range.fin { + self.mark_fin_lost(); + } + } + + pub fn ack(&mut self, range: StreamTxRange) { + if let Some(range) = self.clamp_buffered_range(range.offset, range.len) { + self.acked.insert(range.clone()); + self.retransmits.remove(range); + self.trim_acked_prefix(); + } + if range.fin { + if let Some(final_offset) = self.final_offset.as_mut() { + final_offset.state = SendState::Acked; + } + } + self.trim_acked_fin(); + } + + pub fn clear(&mut self) { + self.chunks.clear(); + self.buffered_len = 0; + self.unsent = self.base_offset; + self.acked = RangeSet::new(); + self.retransmits = RangeSet::new(); + self.final_offset = None; + } + + fn clamp_buffered_range(&self, offset: u64, len: usize) -> Option> { + if len == 0 { + return None; + } + let start = offset.max(self.base_offset); + let end = offset.saturating_add(len as u64).min(self.end_offset()); + (start < end).then_some(start..end) + } + + fn clamp_sent_range(&self, offset: u64, len: usize) -> Option> { + if len == 0 { + return None; + } + let start = offset.max(self.base_offset); + let end = offset.saturating_add(len as u64).min(self.unsent); + (start < end).then_some(start..end) + } + + fn insert_not_acked(acked_set: &RangeSet, target: &mut RangeSet, range: Range) { + let mut cursor = range.start; + for acked in acked_set.iter() { + if acked.end <= cursor { + continue; + } + if acked.start >= range.end { + break; + } + if cursor < acked.start { + target.insert(cursor..acked.start.min(range.end)); + } + cursor = cursor.max(acked.end); + if cursor >= range.end { + break; + } + } + if cursor < range.end { + target.insert(cursor..range.end); + } + } + + fn poll_fin(&mut self, offset: u64) -> bool { + let Some(final_offset) = self.final_offset.as_mut() else { + return false; + }; + if matches!(final_offset.state, SendState::Lost | SendState::Unsent) + && final_offset.offset == offset + { + final_offset.state = SendState::Sent; + true + } else { + false + } + } + + fn mark_fin_lost(&mut self) { + if let Some(final_offset) = self.final_offset.as_mut() { + if final_offset.state != SendState::Acked { + final_offset.state = SendState::Lost; + } + } + } + + fn trim_acked_prefix(&mut self) { + while self.acked.min() == Some(self.base_offset) { + let prefix = self.acked.pop_min().unwrap(); + let mut to_advance = usize::try_from(prefix.end - prefix.start).unwrap(); + self.buffered_len -= to_advance; + while to_advance > 0 { + let front = self + .chunks + .front_mut() + .expect("expected buffered chunks for acked prefix"); + if front.len() <= to_advance { + to_advance -= front.len(); + self.chunks.pop_front(); + } else { + front.advance(to_advance); + to_advance = 0; + } + } + self.base_offset = prefix.end; + } + } + + fn trim_acked_fin(&mut self) { + if self.final_offset.is_some_and(|final_offset| { + final_offset.state == SendState::Acked + && final_offset.offset == self.base_offset + && self.buffered_len == 0 + }) { + self.final_offset = None; + } + } +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use super::{StreamTx, StreamTxRange}; + + #[test] + fn append_tracks_unsent_bytes() { + let mut tx = StreamTx::new(); + tx.append(Bytes::from_static(b"abc")); + tx.append(Bytes::from_static(b"de")); + + assert_eq!( + tx.poll_transmit(8, u64::MAX), + Some(StreamTxRange { + offset: 0, + len: 5, + fin: false, + }) + ); + } + + #[test] + fn lost_range_is_selected_before_unsent_bytes() { + let mut tx = StreamTx::new(); + tx.append(Bytes::from_static(b"abcdef")); + + let first = tx.poll_transmit(3, u64::MAX).unwrap(); + tx.retransmit(first); + + assert_eq!( + tx.poll_transmit(3, u64::MAX), + Some(StreamTxRange { + offset: 0, + len: 3, + fin: false, + }) + ); + } + + #[test] + fn lost_range_coalesces_contiguous_unsent_bytes() { + let mut tx = StreamTx::new(); + tx.append(Bytes::from_static(b"abc")); + + let first = tx.poll_transmit(3, u64::MAX).unwrap(); + tx.retransmit(first); + tx.append(Bytes::from_static(b"def")); + + assert_eq!( + tx.poll_transmit(6, u64::MAX), + Some(StreamTxRange { + offset: 0, + len: 6, + fin: false, + }) + ); + assert_eq!(tx.poll_transmit(6, u64::MAX), None); + } + + #[test] + fn lost_range_coalesces_only_new_bytes_that_fit() { + let mut tx = StreamTx::new(); + tx.append(Bytes::from_static(b"abc")); + + let first = tx.poll_transmit(3, u64::MAX).unwrap(); + tx.retransmit(first); + tx.append(Bytes::from_static(b"def")); + + assert_eq!( + tx.poll_transmit(5, u64::MAX), + Some(StreamTxRange { + offset: 0, + len: 5, + fin: false, + }) + ); + assert_eq!( + tx.poll_transmit(6, u64::MAX), + Some(StreamTxRange { + offset: 5, + len: 1, + fin: false, + }) + ); + } + + #[test] + fn non_contiguous_lost_range_does_not_coalesce_unsent_bytes() { + let mut tx = StreamTx::new(); + tx.append(Bytes::from_static(b"abcdef")); + + let first = tx.poll_transmit(3, u64::MAX).unwrap(); + let _second = tx.poll_transmit(3, u64::MAX).unwrap(); + tx.retransmit(first); + tx.append(Bytes::from_static(b"ghi")); + + assert_eq!( + tx.poll_transmit(6, u64::MAX), + Some(StreamTxRange { + offset: 0, + len: 3, + fin: false, + }) + ); + assert_eq!( + tx.poll_transmit(6, u64::MAX), + Some(StreamTxRange { + offset: 6, + len: 3, + fin: false, + }) + ); + } + + #[test] + fn acked_prefix_is_trimmed() { + let mut tx = StreamTx::new(); + tx.append(Bytes::from_static(b"abcdef")); + + let first = tx.poll_transmit(3, u64::MAX).unwrap(); + tx.ack(first); + + assert_eq!( + tx.poll_transmit(3, u64::MAX), + Some(StreamTxRange { + offset: 3, + len: 3, + fin: false, + }) + ); + } + + #[test] + fn empty_fin_is_tracked_separately() { + let mut tx = StreamTx::new(); + tx.queue_fin(); + + let range = tx.poll_transmit(16, u64::MAX).unwrap(); + assert_eq!( + range, + StreamTxRange { + offset: 0, + len: 0, + fin: true, + } + ); + + tx.ack(range); + assert!(tx.is_empty()); + } + + #[test] + fn subrange_updates_split_merged_in_flight_segments() { + let mut tx = StreamTx::new(); + tx.append(Bytes::from_static(b"abcdefghijkl")); + + let _first = tx.poll_transmit(4, u64::MAX).unwrap(); + let second = tx.poll_transmit(4, u64::MAX).unwrap(); + let _third = tx.poll_transmit(4, u64::MAX).unwrap(); + + tx.retransmit(second); + + assert_eq!( + tx.poll_transmit(4, u64::MAX), + Some(StreamTxRange { + offset: 4, + len: 4, + fin: false, + }) + ); + } + + #[test] + fn acked_subrange_is_not_reopened_by_stale_timeout() { + let mut tx = StreamTx::new(); + tx.append(Bytes::from_static(b"abcdefghijklmnop")); + + let _first = tx.poll_transmit(4, u64::MAX).unwrap(); + let second = tx.poll_transmit(4, u64::MAX).unwrap(); + let third = tx.poll_transmit(4, u64::MAX).unwrap(); + let _fourth = tx.poll_transmit(4, u64::MAX).unwrap(); + + tx.ack(second); + tx.retransmit(second); + tx.retransmit(third); + + assert_eq!( + tx.poll_transmit(4, u64::MAX), + Some(StreamTxRange { + offset: 8, + len: 4, + fin: false, + }) + ); + } +} diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs new file mode 100644 index 00000000..f1f29879 --- /dev/null +++ b/ql-fsm/src/session/tests.rs @@ -0,0 +1,869 @@ +use std::time::{Duration, Instant}; + +use bytes::Bytes; +use ql_wire::{ + decode_session_frames, parse_session_frames, CloseTarget, RecordAck, RecordSeq, RouteId, + SessionFrame, SessionRecordBuilder, StreamClose, StreamCloseCode, StreamData, StreamHeader, + StreamId, VarInt, QID, +}; + +use super::{SessionConfig, SessionEvent, SessionFsm}; +use crate::session::stream_parity::StreamParity; + +fn seq(value: u64) -> RecordSeq { + RecordSeq::from_u64(value).unwrap() +} + +fn stream_id(value: u64) -> StreamId { + StreamId(VarInt::from_u64(value).unwrap()) +} + +fn offset(value: u64) -> VarInt { + VarInt::from_u64(value).unwrap() +} + +fn route_id(value: u64) -> RouteId { + RouteId::from_u64(value).unwrap() +} + +fn record_ack(seq: RecordSeq) -> RecordAck { + RecordAck::from_ranges([seq..=seq]).unwrap() +} + +const REFUSED: StreamCloseCode = StreamCloseCode(1); +const TIMEOUT: StreamCloseCode = StreamCloseCode(2); + +fn header(value: u64) -> StreamHeader { + StreamHeader { + route_id: route_id(value), + } +} + +fn opened(stream_id: StreamId) -> SessionEvent { + SessionEvent::Opened { + stream_id, + route_id: route_id(1), + } +} + +fn open_stream_id(fsm: &mut SessionFsm) -> StreamId { + fsm.open_stream(route_id(1), |_| {}).unwrap().stream_id() +} + +fn write_stream_bytes(fsm: &mut SessionFsm, stream_id: StreamId, bytes: &[u8]) -> usize { + let mut bytes = Bytes::copy_from_slice(bytes); + let mut stream = fsm.stream(stream_id, |_| {}).unwrap(); + let mut writer = stream.writer().unwrap(); + writer.write(&mut bytes) +} + +fn read_stream_all(fsm: &mut SessionFsm, stream_id: StreamId) -> Vec { + let mut stream = fsm.stream(stream_id, |_| {}).unwrap(); + let out = stream.read().flatten().collect::>(); + stream.commit_read(out.len()).unwrap(); + out +} + +fn read_stream_all_with_events( + fsm: &mut SessionFsm, + stream_id: StreamId, + events: &mut Vec, +) -> Vec { + let mut stream = fsm.stream(stream_id, |event| events.push(event)).unwrap(); + let out = stream.read().flatten().collect::>(); + stream.commit_read(out.len()).unwrap(); + out +} + +fn next_outbound( + fsm: &mut SessionFsm, + now: Instant, +) -> Option<(RecordSeq, Vec>>)> { + let (write_id, builder) = fsm.take_next_write(now)?; + if let Some(write_id) = write_id { + fsm.complete_write(now, write_id, true); + } + Some(( + builder.seq(), + decode_session_frames(builder.bytes()).unwrap(), + )) +} + +fn drain_outbound( + fsm: &mut SessionFsm, + now: Instant, + limit: usize, +) -> Vec<(RecordSeq, Vec>>)> { + let mut records = Vec::new(); + for _ in 0..limit { + let Some(record) = next_outbound(fsm, now) else { + return records; + }; + records.push(record); + } + + panic!("session did not quiesce within outbound limit"); +} + +fn receive_events( + fsm: &mut SessionFsm, + now: Instant, + seq: RecordSeq, + record: &[SessionFrame>], +) -> Vec { + let mut builder = SessionRecordBuilder::new(seq, usize::MAX); + for frame in record { + assert!(builder.push_frame(frame)); + } + let bytes = Bytes::from(builder.bytes().to_vec()); + let frames = parse_session_frames(bytes); + let mut events = Vec::new(); + let mut emit = |event| events.push(event); + fsm.receive(now, seq, frames, &mut emit); + events +} + +#[test] +fn outbound_record_seq_increments_monotonically() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = open_stream_id(&mut fsm); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"one"), 3); + let (first_seq, _) = next_outbound(&mut fsm, now).unwrap(); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"two"), 3); + let (second_seq, _) = next_outbound(&mut fsm, now + Duration::from_millis(1)).unwrap(); + + assert_eq!(first_seq, seq(0)); + assert_eq!(second_seq, seq(1)); +} + +#[test] +fn retransmit_uses_new_record_seq() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = open_stream_id(&mut fsm); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"retry"), 5); + let (first_seq, first) = next_outbound(&mut fsm, now).unwrap(); + + let mut emit = |_| {}; + fsm.on_timer(now + Duration::from_millis(200), &mut emit); + let (retried_seq, retried) = next_outbound(&mut fsm, now + Duration::from_millis(200)).unwrap(); + + assert_ne!(first_seq, retried_seq); + assert_eq!(first, retried); +} + +#[test] +fn lost_record_on_one_stream_does_not_block_another_stream() { + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + record_max_size: 80 + SessionRecordBuilder::MIN_CAPACITY, + ..SessionConfig::default() + }, + now, + ); + let stream_id_a = open_stream_id(&mut fsm); + let stream_id_b = open_stream_id(&mut fsm); + let payload_a = vec![b'a'; 40]; + let payload_b = vec![b'b'; 40]; + + assert_eq!(write_stream_bytes(&mut fsm, stream_id_a, &payload_a), 40); + assert_eq!(write_stream_bytes(&mut fsm, stream_id_b, &payload_b), 40); + + let (first_seq, first) = next_outbound(&mut fsm, now).unwrap(); + let (second_seq, _second) = next_outbound(&mut fsm, now + Duration::from_millis(1)).unwrap(); + assert_ne!(first_seq, second_seq); + assert!(first.iter().any( + |frame| matches!(frame, SessionFrame::StreamData(frame) if frame.stream_id == stream_id_a) + )); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id_b, b"b-2"), 3); + let (_third_seq, third) = next_outbound(&mut fsm, now + Duration::from_millis(2)).unwrap(); + + let stream_ids: Vec<_> = third + .iter() + .filter_map(|frame| match frame { + SessionFrame::StreamData(frame) => Some(frame.stream_id), + _ => None, + }) + .collect(); + assert_eq!(stream_ids, vec![stream_id_b]); +} + +#[test] +fn ack_reopens_write_capacity() { + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + stream_send_buffer_size: 4, + ..SessionConfig::default() + }, + now, + ); + let stream_id = open_stream_id(&mut fsm); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"abcd"), 4); + let (record_seq, _record) = next_outbound(&mut fsm, now).unwrap(); + + let mut events = Vec::new(); + let mut emit = |event| events.push(event); + fsm.receive( + now + Duration::from_millis(1), + seq(9), + std::iter::once(Ok(SessionFrame::Ack(record_ack(record_seq)))), + &mut emit, + ); + + assert!(events.contains(&SessionEvent::Writable(stream_id))); + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"z"), 1); +} + +#[test] +fn ack_of_fin_emits_outbound_finished_once() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = open_stream_id(&mut fsm); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"done"), 4); + fsm.stream(stream_id, |_| {}) + .unwrap() + .writer() + .unwrap() + .finish(); + + let (record_seq, record) = next_outbound(&mut fsm, now).unwrap(); + assert!(matches!( + record.as_slice(), + [SessionFrame::StreamData(StreamData { + stream_id: id, + fin: true, + .. + })] if *id == stream_id + )); + + let mut events = Vec::new(); + { + let mut emit = |event| events.push(event); + fsm.receive( + now + Duration::from_millis(1), + seq(9), + std::iter::once(Ok(SessionFrame::Ack(record_ack(record_seq)))), + &mut emit, + ); + } + assert_eq!(events, vec![SessionEvent::OutboundFinished(stream_id)]); + + { + let mut emit = |event| events.push(event); + fsm.receive( + now + Duration::from_millis(2), + seq(10), + std::iter::once(Ok(SessionFrame::Ack(record_ack(record_seq)))), + &mut emit, + ); + } + assert_eq!(events, vec![SessionEvent::OutboundFinished(stream_id)]); +} + +#[test] +fn commit_stream_read_is_what_advances_stream_window() { + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + local_parity: StreamParity::Even, + ack_delay: Duration::ZERO, + ..SessionConfig::default() + }, + now, + ); + let stream_id = stream_id(1); + let data = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: offset(0), + header: Some(header(1)), + fin: false, + bytes: b"hi".to_vec(), + })]; + let events = receive_events(&mut fsm, now, seq(7), &data); + assert_eq!( + events, + vec![opened(stream_id), SessionEvent::Readable(stream_id)] + ); + + let (write_id, builder) = fsm.take_next_write(now + Duration::from_millis(1)).unwrap(); + let first = decode_session_frames(builder.bytes()).unwrap(); + assert!(write_id.is_none()); + assert!(matches!(first.as_slice(), [SessionFrame::Ack(_)])); + + let read = fsm + .stream(stream_id, |_| {}) + .unwrap() + .read() + .map(|chunk| chunk.len()) + .sum::(); + assert_eq!(read, 2); + + assert!(next_outbound(&mut fsm, now + Duration::from_millis(2)).is_none()); + + fsm.stream(stream_id, |_| {}) + .unwrap() + .commit_read(2) + .unwrap(); + let (_second_seq, second) = next_outbound(&mut fsm, now + Duration::from_millis(3)).unwrap(); + assert!(matches!( + second.as_slice(), + [SessionFrame::StreamWindow(window)] if window.stream_id == stream_id + )); +} + +#[test] +fn pure_ack_only_records_are_fire_and_forget() { + let now = Instant::now(); + let config = SessionConfig { + ack_delay: Duration::ZERO, + ..SessionConfig::default() + }; + let retransmit_timeout = config.retransmit_timeout; + let mut fsm = SessionFsm::new(config, now); + let stream_id = stream_id(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: offset(0), + header: Some(header(1)), + fin: false, + bytes: b"hi".to_vec(), + })]; + + let _ = receive_events(&mut fsm, now, seq(7), &record); + + let (write_id, builder) = fsm.take_next_write(now + Duration::from_millis(1)).unwrap(); + let ack = decode_session_frames(builder.bytes()).unwrap(); + assert!(write_id.is_none()); + assert!(matches!(ack.as_slice(), [SessionFrame::Ack(_)])); + + let mut emit = |_| {}; + fsm.on_timer( + now + retransmit_timeout + Duration::from_millis(1), + &mut emit, + ); + assert!(fsm + .take_next_write(now + retransmit_timeout + Duration::from_millis(1)) + .is_none()); +} + +#[test] +fn inbound_stream_data_emits_opened_and_readable() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = stream_id(1); + let record = vec![SessionFrame::StreamData(ql_wire::StreamData { + stream_id, + offset: offset(0), + header: Some(header(1)), + fin: true, + bytes: b"hello".to_vec(), + })]; + + let events = receive_events(&mut fsm, now, seq(0), &record); + assert_eq!( + events, + vec![opened(stream_id), SessionEvent::Readable(stream_id)] + ); + let mut events = Vec::new(); + assert_eq!( + read_stream_all_with_events(&mut fsm, stream_id, &mut events), + b"hello".to_vec() + ); + assert_eq!(events, vec![SessionEvent::Finished(stream_id)]); +} + +#[test] +fn inbound_empty_fin_emits_finished_immediately() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = stream_id(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: offset(0), + header: Some(header(1)), + fin: true, + bytes: Vec::new(), + })]; + + let events = receive_events(&mut fsm, now, seq(0), &record); + assert_eq!( + events, + vec![opened(stream_id), SessionEvent::Finished(stream_id)] + ); +} + +#[test] +fn remote_stream_close_is_reliable_and_retried() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = open_stream_id(&mut fsm); + + fsm.stream(stream_id, |_| {}) + .unwrap() + .close(CloseTarget::Both, StreamCloseCode::CANCELLED); + + let (write_id, builder) = fsm.take_next_write(now).unwrap(); + fsm.complete_write(now, write_id.expect("stream close should be tracked"), true); + let first = decode_session_frames(builder.bytes()).unwrap(); + assert!(matches!( + first.as_slice(), + [SessionFrame::StreamClose(StreamClose { stream_id: id, .. })] if *id == stream_id + )); + + let mut emit = |_| {}; + fsm.on_timer(now + Duration::from_millis(200), &mut emit); + let (_retried_seq, retried) = + next_outbound(&mut fsm, now + Duration::from_millis(200)).unwrap(); + assert_eq!(first, retried); +} + +#[test] +fn stream_ids_follow_even_odd_xid_ordering() { + let now = Instant::now(); + let even = StreamParity::for_local(QID([1; QID::SIZE]), QID([2; QID::SIZE])); + let odd = StreamParity::for_local(QID([2; QID::SIZE]), QID([1; QID::SIZE])); + + let even_id = SessionFsm::new( + SessionConfig { + local_parity: even, + ..SessionConfig::default() + }, + now, + ) + .open_stream(route_id(1), |_| {}) + .unwrap() + .stream_id(); + let odd_id = SessionFsm::new( + SessionConfig { + local_parity: odd, + ..SessionConfig::default() + }, + now, + ) + .open_stream(route_id(1), |_| {}) + .unwrap() + .stream_id(); + + assert_eq!(even_id.into_inner() % 2, 0); + assert_eq!(odd_id.into_inner() % 2, 1); +} + +#[test] +fn duplicate_stream_data_is_not_redelivered() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = stream_id(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: offset(0), + header: Some(header(1)), + fin: false, + bytes: b"hi".to_vec(), + })]; + let _ = receive_events(&mut fsm, now, seq(1), &record); + let _ = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); + + assert_eq!(read_stream_all(&mut fsm, stream_id), b"hi".to_vec()); +} + +#[test] +fn duplicate_remote_close_after_reap_is_ignored() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let close = StreamClose { + stream_id: stream_id(1), + target: CloseTarget::Both, + code: StreamCloseCode(9), + }; + let record = vec![SessionFrame::StreamClose(close.clone())]; + + let first = receive_events(&mut fsm, now, seq(1), &record); + assert_eq!( + first, + vec![ + SessionEvent::Closed(close.clone()), + SessionEvent::WritableClosed(close), + ] + ); + + let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); + assert!(second.is_empty()); +} + +#[test] +fn late_remote_stream_data_after_close_is_ignored() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = stream_id(1); + let close = vec![SessionFrame::StreamClose(StreamClose { + stream_id, + target: CloseTarget::Both, + code: StreamCloseCode(9), + })]; + let data = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: offset(0), + header: Some(header(1)), + fin: false, + bytes: b"hello".to_vec(), + })]; + + let first = receive_events(&mut fsm, now, seq(1), &close); + assert_eq!( + first, + vec![ + SessionEvent::Closed(StreamClose { + stream_id, + target: CloseTarget::Both, + code: StreamCloseCode(9), + }), + SessionEvent::WritableClosed(StreamClose { + stream_id, + target: CloseTarget::Both, + code: StreamCloseCode(9), + }), + ] + ); + + let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &data); + assert!(second.is_empty()); +} + +#[test] +fn duplicate_finished_remote_data_after_reap_is_ignored() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = stream_id(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: offset(0), + header: Some(header(1)), + fin: true, + bytes: b"hello".to_vec(), + })]; + + let first = receive_events(&mut fsm, now, seq(1), &record); + assert_eq!( + first, + vec![opened(stream_id), SessionEvent::Readable(stream_id)] + ); + let mut events = Vec::new(); + assert_eq!( + read_stream_all_with_events(&mut fsm, stream_id, &mut events), + b"hello".to_vec() + ); + assert_eq!(events, vec![SessionEvent::Finished(stream_id)]); + + let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); + assert!(second.is_empty()); +} + +#[test] +fn duplicate_finished_remote_data_before_read_is_ignored() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = stream_id(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: offset(0), + header: Some(header(1)), + fin: true, + bytes: b"hello".to_vec(), + })]; + + let first = receive_events(&mut fsm, now, seq(1), &record); + assert_eq!( + first, + vec![opened(stream_id), SessionEvent::Readable(stream_id)] + ); + + let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); + assert!(second.is_empty()); + let mut events = Vec::new(); + assert_eq!( + read_stream_all_with_events(&mut fsm, stream_id, &mut events), + b"hello".to_vec() + ); + assert_eq!(events, vec![SessionEvent::Finished(stream_id)]); +} + +#[test] +fn out_of_order_remote_stream_first_observations_still_open_once_each() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let close3 = vec![SessionFrame::StreamClose(StreamClose { + stream_id: stream_id(3), + target: CloseTarget::Both, + code: REFUSED, + })]; + let close1 = vec![SessionFrame::StreamClose(StreamClose { + stream_id: stream_id(1), + target: CloseTarget::Both, + code: TIMEOUT, + })]; + + let first = receive_events(&mut fsm, now, seq(1), &close3); + assert_eq!( + first, + vec![ + SessionEvent::Closed(StreamClose { + stream_id: stream_id(3), + target: CloseTarget::Both, + code: REFUSED, + }), + SessionEvent::WritableClosed(StreamClose { + stream_id: stream_id(3), + target: CloseTarget::Both, + code: REFUSED, + }), + ] + ); + + let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &close1); + assert_eq!( + second, + vec![ + SessionEvent::Closed(StreamClose { + stream_id: stream_id(1), + target: CloseTarget::Both, + code: TIMEOUT, + }), + SessionEvent::WritableClosed(StreamClose { + stream_id: stream_id(1), + target: CloseTarget::Both, + code: TIMEOUT, + }), + ] + ); + + let third = receive_events(&mut fsm, now + Duration::from_millis(2), seq(3), &close3); + assert!(third.is_empty()); +} + +#[test] +fn invalid_remote_stream_close_closes_session() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + + let invalid = vec![SessionFrame::StreamClose(StreamClose { + stream_id: stream_id(0), + target: CloseTarget::Both, + code: StreamCloseCode(9), + })]; + let events = receive_events(&mut fsm, now, seq(1), &invalid); + + assert_eq!( + events, + vec![SessionEvent::SessionClosed(ql_wire::SessionClose { + code: ql_wire::SessionCloseCode::PROTOCOL, + })] + ); +} + +#[test] +fn close_does_not_ack_rejected_record_seq() { + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + ack_delay: Duration::ZERO, + ..SessionConfig::default() + }, + now, + ); + + let invalid = vec![SessionFrame::StreamData(StreamData { + stream_id: stream_id(0), + offset: offset(0), + header: Some(header(1)), + fin: false, + bytes: b"bad".to_vec(), + })]; + let events = receive_events(&mut fsm, now, seq(7), &invalid); + assert_eq!( + events, + vec![SessionEvent::SessionClosed(ql_wire::SessionClose { + code: ql_wire::SessionCloseCode::PROTOCOL, + })] + ); + + let valid_after_close = vec![SessionFrame::Ping]; + let events = receive_events( + &mut fsm, + now + Duration::from_millis(1), + seq(8), + &valid_after_close, + ); + assert!(events.is_empty()); + + let (_seq, outbound) = next_outbound(&mut fsm, now + Duration::from_millis(2)).unwrap(); + assert!(matches!(outbound.as_slice(), [SessionFrame::Close(_)])); +} + +#[test] +fn inbound_unpair_emits_final_unpair_frame() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + + let events = receive_events(&mut fsm, now, seq(1), &[SessionFrame::Unpair]); + assert_eq!(events, vec![SessionEvent::Unpaired]); + assert!(!fsm.is_closed()); + + let (_seq, outbound) = next_outbound(&mut fsm, now + Duration::from_millis(1)).unwrap(); + assert!(matches!(outbound.as_slice(), [SessionFrame::Unpair])); + assert!(fsm.is_closed()); +} + +#[test] +fn terminating_session_ignores_inbound_frames() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + + let mut events = Vec::new(); + fsm.unpair(&mut |event| events.push(event)); + assert_eq!(events, vec![SessionEvent::Unpaired]); + + let ignored = receive_events( + &mut fsm, + now + Duration::from_millis(1), + seq(1), + &[SessionFrame::Ping], + ); + assert!(ignored.is_empty()); + + let (_seq, outbound) = next_outbound(&mut fsm, now + Duration::from_millis(2)).unwrap(); + assert!(matches!(outbound.as_slice(), [SessionFrame::Unpair])); + assert!(fsm.is_closed()); +} + +#[test] +fn initial_peer_stream_receive_window_limits_first_send() { + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + initial_peer_stream_receive_window: 3, + ..SessionConfig::default() + }, + now, + ); + let stream_id = open_stream_id(&mut fsm); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"hello"), 5); + let (_first_seq, first) = next_outbound(&mut fsm, now).unwrap(); + assert!(matches!( + first.as_slice(), + [SessionFrame::StreamData(frame)] if frame.stream_id == stream_id && frame.bytes.as_slice() == b"hel" + )); + + let events = receive_events( + &mut fsm, + now + Duration::from_millis(1), + seq(9), + &[SessionFrame::StreamWindow(ql_wire::StreamWindow { + stream_id, + maximum_offset: offset(5), + })], + ); + assert!(events.is_empty()); + + let (_second_seq, second) = next_outbound(&mut fsm, now + Duration::from_millis(2)).unwrap(); + assert!(second.iter().any(|frame| { + matches!( + frame, + SessionFrame::StreamData(frame) + if frame.stream_id == stream_id + && frame.offset == offset(3) + && frame.bytes.as_slice() == b"lo" + ) + })); +} + +#[test] +fn sparse_out_of_order_ack_ranges_page_and_quiesce() { + let now = Instant::now(); + let sender_config = SessionConfig { + local_parity: StreamParity::Even, + record_max_size: SessionRecordBuilder::MIN_CAPACITY + 40, + ack_delay: Duration::from_millis(5), + retransmit_timeout: Duration::from_millis(25), + stream_send_buffer_size: 8 * 1024, + initial_peer_stream_receive_window: 8 * 1024, + ..SessionConfig::default() + }; + let receiver_config = SessionConfig { + local_parity: StreamParity::Odd, + record_max_size: SessionRecordBuilder::MIN_CAPACITY + 10, + ack_delay: Duration::from_millis(1), + retransmit_timeout: Duration::from_millis(25), + pending_ack_range_limit: 512, + initial_peer_stream_receive_window: 8 * 1024, + ..SessionConfig::default() + }; + let mut sender = SessionFsm::new(sender_config, now); + let mut receiver = SessionFsm::new(receiver_config, now); + + let stream_id = open_stream_id(&mut sender); + let payload = vec![b'x'; 2048]; + assert_eq!( + write_stream_bytes(&mut sender, stream_id, &payload), + payload.len() + ); + + let originals = drain_outbound(&mut sender, now, 4096); + assert!(originals.len() >= 64); + + for (seq, record) in originals + .iter() + .filter(|(seq, _)| seq.into_inner() % 2 == 1) + { + let _ = receive_events(&mut receiver, now, *seq, record); + } + + let first_ack_time = now + receiver_config.ack_delay; + let first_acks = drain_outbound(&mut receiver, first_ack_time, originals.len()); + assert!(first_acks.len() > 1); + assert!(first_acks + .iter() + .all(|(_, frames)| matches!(frames.as_slice(), [SessionFrame::Ack(_)]))); + + for (seq, record) in &first_acks { + let _ = receive_events(&mut sender, first_ack_time, *seq, record); + } + + let retransmit_time = now + sender_config.retransmit_timeout + Duration::from_millis(1); + let mut emit = |_| {}; + sender.on_timer(retransmit_time, &mut emit); + let retransmits = drain_outbound(&mut sender, retransmit_time, originals.len()); + assert!(!retransmits.is_empty()); + + for (seq, record) in &retransmits { + let _ = receive_events(&mut receiver, retransmit_time, *seq, record); + } + + let second_ack_time = retransmit_time + receiver_config.ack_delay; + let second_acks = drain_outbound(&mut receiver, second_ack_time, retransmits.len() + 16); + assert!(!second_acks.is_empty()); + assert!(second_acks + .iter() + .all(|(_, frames)| matches!(frames.as_slice(), [SessionFrame::Ack(_)]))); + + for (seq, record) in &second_acks { + let _ = receive_events(&mut sender, second_ack_time, *seq, record); + } + + let final_now = second_ack_time + sender_config.retransmit_timeout + Duration::from_millis(1); + let mut sender_emit = |_| {}; + sender.on_timer(final_now, &mut sender_emit); + let mut receiver_emit = |_| {}; + receiver.on_timer(final_now, &mut receiver_emit); + assert!(next_outbound(&mut sender, final_now).is_none()); + assert!(next_outbound(&mut receiver, final_now).is_none()); +} diff --git a/ql-fsm/src/session/tracked.rs b/ql-fsm/src/session/tracked.rs new file mode 100644 index 00000000..84317951 --- /dev/null +++ b/ql-fsm/src/session/tracked.rs @@ -0,0 +1,29 @@ +//! outbound record tracking state for ack and retransmit handling + +use std::time::Instant; + +use ql_wire::{RecordAck, RecordSeq, StreamClose, StreamId}; + +#[derive(Debug, Clone)] +pub struct TrackedRecord { + pub seq: RecordSeq, + pub frames: Vec, + pub ack: Option, + pub ping_included: bool, + pub window_updates: Vec<(StreamId, u64)>, + pub sent_at: Option, +} + +#[derive(Debug, Clone)] +pub enum TrackedFrame { + StreamData(TrackedStreamData), + StreamClose(StreamClose), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct TrackedStreamData { + pub stream_id: StreamId, + pub offset: u64, + pub len: usize, + pub fin: bool, +} diff --git a/ql-fsm/src/state.rs b/ql-fsm/src/state.rs new file mode 100644 index 00000000..8268bc16 --- /dev/null +++ b/ql-fsm/src/state.rs @@ -0,0 +1,139 @@ +use std::time::Instant; + +use ql_wire::{ + ConnectionId, EphemeralPublicKey, HandshakeId, HandshakeMeta, IkHandshake, KkHandshake, + PairingToken, PeerBundle, QlHandshakeRecord, SessionKey, TransportParams, XxHandshake, +}; + +use crate::{session::SessionFsm, NoSessionError, PeerStatus}; + +pub struct QlFsmState { + pub next_control_id: u32, + pub peer: Option, + pub armed_pairing_token: Option, + pub handshake: Option, + pub link: LinkState, + pub now: Instant, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct SessionTransport { + pub tx_key: SessionKey, + pub rx_key: SessionKey, + pub tx_connection_id: ConnectionId, + pub rx_connection_id: ConnectionId, + pub remote_transport_params: TransportParams, +} + +impl SessionTransport { + pub fn from_finalized(finalized: ql_wire::FinalizedHandshake) -> (Self, PeerBundle) { + ( + Self { + tx_key: finalized.tx_key, + rx_key: finalized.rx_key, + tx_connection_id: finalized.tx_connection_id, + rx_connection_id: finalized.rx_connection_id, + remote_transport_params: finalized.remote_transport_params, + }, + finalized.remote_bundle, + ) + } +} + +#[allow(clippy::large_enum_variant)] +pub enum LinkState { + Idle, + IkInitiator(IkInitiatorState), + KkInitiator(KkInitiatorState), + XxInitiator(XxInitiatorState), + XxResponder(XxResponderState), + Connected(ConnectedState), +} + +pub struct ConnectedState { + pub transport: SessionTransport, + pub session: SessionFsm, +} + +#[derive(Debug, Clone)] +pub struct IkInitiatorState { + pub handshake: IkHandshake, + pub handshake_id: HandshakeId, + pub deadline: Instant, + pub initial_ephemeral: EphemeralPublicKey, +} + +#[derive(Debug, Clone)] +pub struct KkInitiatorState { + pub handshake: KkHandshake, + pub handshake_id: HandshakeId, + pub deadline: Instant, + pub initial_ephemeral: EphemeralPublicKey, +} + +#[derive(Debug, Clone)] +pub struct XxInitiatorState { + pub handshake: XxHandshake, + pub handshake_id: HandshakeId, + pub deadline: Instant, + pub initial_ephemeral: EphemeralPublicKey, +} + +#[derive(Debug, Clone)] +pub struct XxResponderState { + pub handshake: XxHandshake, + pub handshake_meta: HandshakeMeta, + pub deadline: Instant, +} + +impl LinkState { + pub fn take(&mut self) -> Self { + std::mem::replace(self, Self::Idle) + } + + pub fn status(&self) -> PeerStatus { + match self { + Self::Idle | Self::XxResponder(_) => PeerStatus::Disconnected, + Self::IkInitiator(_) | Self::KkInitiator(_) | Self::XxInitiator(_) => { + PeerStatus::Initiator + } + Self::Connected(_) => PeerStatus::Connected, + } + } + + #[inline] + pub fn connected(&self) -> Option<&ConnectedState> { + match self { + Self::Connected(state) => Some(state), + _ => None, + } + } + + #[inline] + pub fn connected_mut(&mut self) -> Option<&mut ConnectedState> { + match self { + Self::Connected(state) => Some(state), + _ => None, + } + } + + #[inline] + pub fn connected_mut_or_err(&mut self) -> Result<&mut ConnectedState, NoSessionError> { + self.connected_mut().ok_or(NoSessionError) + } + + pub fn handshake_deadline(&self) -> Option { + match self { + Self::Idle | Self::Connected(_) => None, + Self::IkInitiator(state) => Some(state.deadline), + Self::KkInitiator(state) => Some(state.deadline), + Self::XxInitiator(state) => Some(state.deadline), + Self::XxResponder(state) => Some(state.deadline), + } + } + + #[cfg(test)] + pub fn transport(&self) -> Option<&SessionTransport> { + self.connected().map(|state| &state.transport) + } +} diff --git a/ql-fsm/src/tests/handshake.rs b/ql-fsm/src/tests/handshake.rs new file mode 100644 index 00000000..4de4f06e --- /dev/null +++ b/ql-fsm/src/tests/handshake.rs @@ -0,0 +1,388 @@ +use std::time::Duration; + +use ql_wire::QlHandshakeRecord; + +use super::*; +use crate::{state::LinkState, Event, NoPeerError, PeerStatus, ReceiveError}; + +#[test] +fn ik_connect_round_trip_establishes_transport() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.connect_ik(Side::A).unwrap(); + harness.pump(); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +#[test] +fn kk_connect_round_trip_establishes_transport() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.connect_kk(Side::A).unwrap(); + harness.pump(); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +#[test] +fn xx_connect_round_trip_establishes_transport_when_armed() { + let mut harness = Harness::paired(QlFsmConfig::default(), false, false); + let token = pairing_token(1); + + harness.b.fsm.arm_pairing(token); + harness.connect_xx(Side::A, token); + + let xx1 = harness.next_outbound(Side::A).unwrap(); + harness.deliver(Side::B, xx1); + let xx2 = harness.next_outbound(Side::B).unwrap(); + harness.deliver(Side::A, xx2); + let xx3 = harness.next_outbound(Side::A).unwrap(); + harness.deliver(Side::B, xx3); + + let xx4 = harness.next_outbound(Side::B).unwrap(); + harness.deliver(Side::A, xx4); + + assert_eq!(harness.a.fsm.peer(), Some(&harness.b.fsm.identity.bundle())); + assert_eq!(harness.b.fsm.peer(), Some(&harness.a.fsm.identity.bundle())); + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +#[test] +fn ik_connect_learns_remote_initial_stream_receive_window() { + let mut harness = Harness::paired_known_with_configs( + QlFsmConfig { + session_stream_receive_buffer_size: 9, + ..QlFsmConfig::default() + }, + QlFsmConfig { + session_stream_receive_buffer_size: 3, + ..QlFsmConfig::default() + }, + ); + + harness.connect_ik(Side::A).unwrap(); + harness.pump(); + + assert_eq!( + harness + .a + .fsm + .state + .link + .transport() + .unwrap() + .remote_transport_params + .initial_stream_receive_window, + 3 + ); + assert_eq!( + harness + .b + .fsm + .state + .link + .transport() + .unwrap() + .remote_transport_params + .initial_stream_receive_window, + 9 + ); +} + +#[test] +fn connect_methods_require_bound_peer() { + let time = Harness::paired_known(QlFsmConfig::default()).time(); + let identity = generate_identity(&SoftwareCrypto, "identity").unwrap(); + let mut fsm = QlFsm::new(QlFsmConfig::default(), identity, time); + let crypto = SoftwareCrypto; + + assert_eq!(fsm.connect_ik(time, &crypto), Err(NoPeerError)); + assert_eq!(fsm.connect_kk(time, &crypto), Err(NoPeerError)); + + fsm.connect_xx( + time, + PairingInvite { + qid: ql_wire::QID([2; ql_wire::QID::SIZE]), + token: pairing_token(2), + }, + &crypto, + ); +} + +#[test] +fn connect_ik_emits_initiator_status() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.connect_ik(Side::A).unwrap(); + + assert_eq!( + harness.drain_events(Side::A), + vec![Event::PeerStatusChanged(PeerStatus::Initiator)] + ); +} + +#[test] +fn inbound_xx1_rejects_when_not_in_pairing_mode() { + let mut harness = Harness::paired(QlFsmConfig::default(), false, false); + let token = pairing_token(3); + + harness.connect_xx(Side::A, token); + let xx1 = harness.next_outbound(Side::A).unwrap(); + let time = harness.time(); + let Node { fsm, crypto } = &mut harness.b; + let err = fsm.receive(time, xx1, crypto); + + assert_eq!(err, Err(ReceiveError::NotPairingMode)); + assert!(matches!(harness.b.fsm.state.link, LinkState::Idle)); + assert!(harness.drain_events(Side::B).is_empty()); + assert!(harness.next_outbound(Side::B).is_none()); +} + +#[test] +fn inbound_xx1_rejects_mismatched_pairing_id_with_expected_and_actual() { + let mut harness = Harness::paired(QlFsmConfig::default(), false, false); + let expected = pairing_token(4); + let actual = pairing_token(7); + + harness.b.fsm.arm_pairing(expected); + harness.connect_xx(Side::A, actual); + let xx1 = harness.next_outbound(Side::A).unwrap(); + + let time = harness.time(); + let Node { fsm, crypto } = &mut harness.b; + let err = fsm.receive(time, xx1, crypto); + + assert_eq!( + err, + Err(ReceiveError::InvalidPairingId { + expected: expected.id(&SoftwareCrypto), + actual: actual.id(&SoftwareCrypto), + }) + ); +} + +#[test] +fn disarm_pairing_rejects_inflight_inbound_xx_responder() { + let mut harness = Harness::paired(QlFsmConfig::default(), false, false); + let token = pairing_token(5); + + harness.b.fsm.arm_pairing(token); + harness.connect_xx(Side::A, token); + let xx1 = harness.next_outbound(Side::A).unwrap(); + harness.deliver(Side::B, xx1); + let xx2 = harness.next_outbound(Side::B).unwrap(); + harness.deliver(Side::A, xx2); + let xx3 = harness.next_outbound(Side::A).unwrap(); + harness.b.fsm.disarm_pairing(); + harness.deliver(Side::B, xx3); + + assert!(matches!(harness.b.fsm.state.link, LinkState::Idle)); + assert!(harness.next_outbound(Side::B).is_none()); +} + +#[test] +fn simultaneous_xx_connect_converges() { + let mut harness = Harness::paired(QlFsmConfig::default(), false, false); + let token = pairing_token(6); + + harness.a.fsm.arm_pairing(token); + harness.b.fsm.arm_pairing(token); + harness.connect_xx(Side::A, token); + harness.connect_xx(Side::B, token); + + for _ in 0..2 { + if let Some(record) = harness.next_outbound(Side::A) { + harness.deliver(Side::B, record); + } + if let Some(record) = harness.next_outbound(Side::B) { + harness.deliver(Side::A, record); + } + } + harness.pump(); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +#[test] +fn connect_ik_replaces_in_flight_attempt_and_ignores_stale_reply() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.connect_ik(Side::A).unwrap(); + harness.drain_events(Side::A); + let first = harness.next_outbound(Side::A).unwrap(); + let first_id = handshake_id(&first); + + harness.connect_ik(Side::A).unwrap(); + let second = harness.next_outbound(Side::A).unwrap(); + let second_id = handshake_id(&second); + + assert_ne!(first_id, second_id); + + harness.deliver(Side::B, first); + let stale_reply = harness.next_outbound(Side::B).unwrap(); + assert_eq!(handshake_id(&stale_reply), first_id); + + harness.deliver(Side::A, stale_reply); + assert!(matches!( + harness.a.fsm.state.link, + LinkState::IkInitiator(_) + )); + + harness.deliver(Side::B, second); + harness.pump(); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +#[test] +fn connect_kk_replaces_in_flight_attempt_and_ignores_stale_reply() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.connect_kk(Side::A).unwrap(); + let first = harness.next_outbound(Side::A).unwrap(); + let first_id = handshake_id(&first); + + harness.connect_kk(Side::A).unwrap(); + let second = harness.next_outbound(Side::A).unwrap(); + let second_id = handshake_id(&second); + + assert_ne!(first_id, second_id); + + harness.deliver(Side::B, first); + let stale_reply = harness.next_outbound(Side::B).unwrap(); + assert_eq!(handshake_id(&stale_reply), first_id); + + harness.deliver(Side::A, stale_reply); + assert!(matches!( + harness.a.fsm.state.link, + LinkState::KkInitiator(_) + )); + + harness.deliver(Side::B, second); + harness.pump(); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +#[test] +fn inbound_ik1_auto_binds_unbound_responder() { + let mut harness = Harness::paired(QlFsmConfig::default(), true, false); + + harness.connect_ik(Side::A).unwrap(); + harness.pump(); + + let expected_peer = harness.a.fsm.identity.bundle(); + assert_eq!(harness.b.fsm.peer(), Some(&expected_peer)); + assert_eq!( + harness.drain_events(Side::B), + vec![ + Event::NewPeer, + Event::PeerStatusChanged(PeerStatus::Connected), + ] + ); + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +#[test] +fn handshake_timeout_drops_single_ik_attempt_without_resend() { + let config = QlFsmConfig { + handshake_timeout: Duration::from_millis(60), + ..QlFsmConfig::default() + }; + let mut harness = Harness::paired_known(config); + + harness.connect_ik(Side::A).unwrap(); + harness.drain_events(Side::A); + let first = harness.next_outbound(Side::A).unwrap(); + let (_, first) = ql_wire::decode_record::(first.as_slice()).unwrap(); + assert!(matches!(first, ql_wire::QlHandshakeRecord::Ik1(_))); + assert!(harness.next_outbound(Side::A).is_none()); + + harness.advance(config.handshake_timeout); + harness.on_timer(Side::A); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Idle)); + assert_eq!( + harness.take_event(Side::A), + Some(Event::PeerStatusChanged(PeerStatus::Disconnected)) + ); + assert!(harness.next_outbound(Side::A).is_none()); +} + +#[test] +fn handshake_timeout_clears_queued_kk_output() { + let config = QlFsmConfig { + handshake_timeout: Duration::from_millis(60), + ..QlFsmConfig::default() + }; + let mut harness = Harness::paired_known(config); + + harness.connect_kk(Side::A).unwrap(); + + harness.advance(config.handshake_timeout); + harness.on_timer(Side::A); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Idle)); + assert!(harness.next_outbound(Side::A).is_none()); +} + +#[test] +fn bind_peer_clears_queued_handshake_output() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.connect_ik(Side::A).unwrap(); + harness.drain_events(Side::A); + harness + .a + .fsm + .bind_peer(generate_identity(&SoftwareCrypto, "peer").unwrap().bundle()); + + assert!(harness.drain_events(Side::A).is_empty()); + assert!(harness.next_outbound(Side::A).is_none()); +} + +#[test] +fn simultaneous_ik_connect_converges() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.connect_ik(Side::A).unwrap(); + harness.connect_ik(Side::B).unwrap(); + harness.pump(); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +#[test] +fn simultaneous_ik_and_kk_connect_prefers_ik() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.connect_ik(Side::A).unwrap(); + harness.connect_kk(Side::B).unwrap(); + harness.pump(); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!(harness.b.fsm.state.link, LinkState::Connected(_))); +} + +fn handshake_id(record: &[u8]) -> ql_wire::HandshakeId { + let (_, record) = ql_wire::decode_record(record).unwrap(); + match record { + ql_wire::QlHandshakeRecord::Ik1(message) => message.meta.handshake_id, + ql_wire::QlHandshakeRecord::Ik2(message) => message.meta.handshake_id, + ql_wire::QlHandshakeRecord::Kk1(message) => message.meta.handshake_id, + ql_wire::QlHandshakeRecord::Kk2(message) => message.meta.handshake_id, + ql_wire::QlHandshakeRecord::Xx1(message) => message.meta.handshake_id, + ql_wire::QlHandshakeRecord::Xx2(message) => message.meta.handshake_id, + ql_wire::QlHandshakeRecord::Xx3(message) => message.meta.handshake_id, + ql_wire::QlHandshakeRecord::Xx4(message) => message.meta.handshake_id, + } +} diff --git a/ql-fsm/src/tests/mod.rs b/ql-fsm/src/tests/mod.rs new file mode 100644 index 00000000..c4d005e4 --- /dev/null +++ b/ql-fsm/src/tests/mod.rs @@ -0,0 +1,351 @@ +mod handshake; +mod proptest; +mod session; + +use std::time::{Duration, Instant}; + +use ql_wire::{ + self, generate_identity, test_identities, ConnectionId, PairingToken, QlCrypto, SessionKey, + SoftwareCrypto, TransportParams, QID, +}; + +use crate::{ + session::{SessionConfig, SessionFsm, StreamParity}, + state::{ConnectedState, LinkState, SessionTransport}, + Event, NoPeerError, OutboundWrite, PairingInvite, QlFsm, QlFsmConfig, WriteId, +}; + +type TestCrypto = SoftwareCrypto; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum Side { + A, + B, +} + +impl Side { + fn idx(self) -> usize { + match self { + Self::A => 0, + Self::B => 1, + } + } +} + +struct Node { + fsm: QlFsm, + crypto: TestCrypto, +} + +struct Harness { + now: Instant, + a: Node, + b: Node, +} + +struct DecodedSessionWrite { + record: Vec, + write_id: Option, + header: ql_wire::SessionHeader, + frames: Vec>>, +} + +impl Harness { + fn paired_known(config: QlFsmConfig) -> Self { + Self::paired_with_configs(config, config, true, true) + } + + fn paired(config: QlFsmConfig, know_a: bool, know_b: bool) -> Self { + Self::paired_with_configs(config, config, know_a, know_b) + } + + fn paired_known_with_configs(config_a: QlFsmConfig, config_b: QlFsmConfig) -> Self { + Self::paired_with_configs(config_a, config_b, true, true) + } + + fn paired_with_configs( + config_a: QlFsmConfig, + config_b: QlFsmConfig, + know_a: bool, + know_b: bool, + ) -> Self { + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + let now = Instant::now(); + + let mut harness = Self { + now, + a: Node { + fsm: QlFsm::new(config_a, identity_a.clone(), now), + crypto: SoftwareCrypto, + }, + b: Node { + fsm: QlFsm::new(config_b, identity_b.clone(), now), + crypto: SoftwareCrypto, + }, + }; + + if know_a { + harness.a.fsm.bind_peer(identity_b.bundle()); + } + if know_b { + harness.b.fsm.bind_peer(identity_a.bundle()); + } + + harness + } + + fn connected(config: QlFsmConfig) -> Self { + let mut harness = Self::paired_known(config); + let a_to_b_key = SessionKey::from_data([7; SessionKey::SIZE]); + let b_to_a_key = SessionKey::from_data([9; SessionKey::SIZE]); + let a_to_b_conn = ConnectionId::from_data([0xA1; ConnectionId::SIZE]); + let b_to_a_conn = ConnectionId::from_data([0xB2; ConnectionId::SIZE]); + + harness.a.fsm.state.link = LinkState::Connected(ConnectedState { + transport: SessionTransport { + tx_key: a_to_b_key.clone(), + rx_key: b_to_a_key.clone(), + tx_connection_id: a_to_b_conn, + rx_connection_id: b_to_a_conn, + remote_transport_params: TransportParams { + initial_stream_receive_window: harness + .b + .fsm + .config + .session_stream_receive_buffer_size, + }, + }, + session: SessionFsm::new(session_config(&harness, true), harness.now), + }); + harness.b.fsm.state.link = LinkState::Connected(ConnectedState { + transport: SessionTransport { + tx_key: b_to_a_key, + rx_key: a_to_b_key, + tx_connection_id: b_to_a_conn, + rx_connection_id: a_to_b_conn, + remote_transport_params: TransportParams { + initial_stream_receive_window: harness + .a + .fsm + .config + .session_stream_receive_buffer_size, + }, + }, + session: SessionFsm::new(session_config(&harness, false), harness.now), + }); + harness + } + + fn time(&self) -> Instant { + self.now + } + + fn advance(&mut self, duration: Duration) { + self.now += duration; + } + + fn node(&self, side: Side) -> &Node { + match side { + Side::A => &self.a, + Side::B => &self.b, + } + } + + fn node_mut(&mut self, side: Side) -> &mut Node { + match side { + Side::A => &mut self.a, + Side::B => &mut self.b, + } + } + + fn next_outbound(&mut self, side: Side) -> Option> { + let write = self.next_write(side)?; + if let Some(id) = write.write_id { + self.confirm_write(side, id); + } + Some(write.record) + } + + fn next_write(&mut self, side: Side) -> Option { + let time = self.time(); + let Node { fsm, crypto } = self.node_mut(side); + fsm.take_next_write(time, crypto) + } + + fn next_decoded_outbound(&mut self, side: Side) -> Option { + let write = self.next_write(side)?; + if let Some(id) = write.write_id { + self.confirm_write(side, id); + } + Some(self.decode_session_write(write, side)) + } + + fn next_decoded_write(&mut self, side: Side) -> Option { + let write = self.next_write(side)?; + Some(self.decode_session_write(write, side)) + } + + fn connect_ik(&mut self, side: Side) -> Result<(), NoPeerError> { + let time = self.time(); + let Node { fsm, crypto } = self.node_mut(side); + fsm.connect_ik(time, crypto) + } + + fn connect_kk(&mut self, side: Side) -> Result<(), NoPeerError> { + let time = self.time(); + let Node { fsm, crypto } = self.node_mut(side); + fsm.connect_kk(time, crypto) + } + + fn connect_xx(&mut self, side: Side, token: PairingToken) { + let time = self.time(); + let remote_qid = self.remote_qid(side); + let Node { fsm, crypto } = self.node_mut(side); + fsm.connect_xx( + time, + PairingInvite { + qid: remote_qid, + token, + }, + crypto, + ); + } + + fn remote_qid(&self, side: Side) -> QID { + match side { + Side::A => self.b.fsm.identity.qid, + Side::B => self.a.fsm.identity.qid, + } + } + + fn deliver(&mut self, side: Side, record: Vec) { + let time = self.time(); + let Node { fsm, crypto } = self.node_mut(side); + fsm.receive(time, record, crypto).unwrap(); + } + + fn confirm_write(&mut self, side: Side, write_id: WriteId) { + let time = self.time(); + self.node_mut(side).fsm.complete_write(time, write_id, true); + } + + fn reject_write(&mut self, side: Side, write_id: WriteId) { + let time = self.time(); + self.node_mut(side) + .fsm + .complete_write(time, write_id, false); + } + + fn decode_session_write(&self, write: OutboundWrite, side: Side) -> DecodedSessionWrite { + let peer = self.node(match side { + Side::A => Side::B, + Side::B => Side::A, + }); + let crypto = &peer.crypto; + let session_key = &peer.fsm.state.link.transport().unwrap().rx_key; + let (header, frames) = decrypt_record(crypto, &write.record, session_key); + DecodedSessionWrite { + record: write.record, + write_id: write.write_id, + header, + frames, + } + } + + fn on_timer(&mut self, side: Side) { + let time = self.time(); + self.node_mut(side).fsm.on_timer(time); + } + + fn take_event(&mut self, side: Side) -> Option { + self.node_mut(side).fsm.poll_event() + } + + fn drain_events(&mut self, side: Side) -> Vec { + let mut events = Vec::new(); + while let Some(event) = self.take_event(side) { + events.push(event); + } + events + } + + fn pump(&mut self) { + for _ in 0..128 { + let mut progressed = false; + + while let Some(record) = self.next_outbound(Side::A) { + progressed = true; + self.deliver(Side::B, record); + } + + while let Some(record) = self.next_outbound(Side::B) { + progressed = true; + self.deliver(Side::A, record); + } + + if !progressed { + return; + } + } + + panic!("pump did not quiesce"); + } +} + +fn pairing_token(byte: u8) -> PairingToken { + PairingToken([byte; PairingToken::SIZE]) +} + +fn session_config(harness: &Harness, a: bool) -> SessionConfig { + let (local, peer, config) = if a { + ( + harness.a.fsm.identity.qid, + harness.a.fsm.state.peer.as_ref().unwrap().qid, + harness.a.fsm.config, + ) + } else { + ( + harness.b.fsm.identity.qid, + harness.b.fsm.state.peer.as_ref().unwrap().qid, + harness.b.fsm.config, + ) + }; + + SessionConfig { + local_parity: StreamParity::for_local(local, peer), + record_max_size: config.session_record_max_size, + ack_delay: config.session_record_ack_delay, + retransmit_timeout: config.session_record_retransmit_timeout, + keepalive_interval: config.session_keepalive_interval, + peer_timeout: config.session_peer_timeout, + stream_send_buffer_size: config.session_stream_send_buffer_size, + stream_receive_buffer_size: config.session_stream_receive_buffer_size, + accepted_record_window: config.session_accepted_record_window, + pending_ack_range_limit: config.session_pending_ack_range_limit, + initial_peer_stream_receive_window: if a { + harness.b.fsm.config.session_stream_receive_buffer_size + } else { + harness.a.fsm.config.session_stream_receive_buffer_size + }, + } +} + +fn decrypt_record( + crypto: &impl QlCrypto, + record: &[u8], + session_key: &SessionKey, +) -> (ql_wire::SessionHeader, Vec>>) { + let (_header, record) = + ql_wire::decode_record::, _>(record).unwrap(); + let plaintext = ql_wire::decrypt_record( + crypto, + &record.header, + record.payload.into_owned(), + session_key, + ) + .unwrap(); + ( + record.header, + ql_wire::decode_session_frames(&plaintext).unwrap(), + ) +} diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs new file mode 100644 index 00000000..bc97ca77 --- /dev/null +++ b/ql-fsm/src/tests/proptest.rs @@ -0,0 +1,1001 @@ +use std::{ + collections::{BTreeMap, BTreeSet}, + time::Duration, +}; + +extern crate proptest as proptest_crate; + +use bytes::Bytes; +use proptest_crate::{collection::vec, prelude::*, test_runner::TestCaseResult}; +use ql_wire::{CloseTarget, StreamCloseCode, StreamId, WireError}; + +use super::*; + +fn test_route_id() -> ql_wire::RouteId { + ql_wire::RouteId::from_u32(1) +} +use crate::{state::LinkState, Event, PeerStatus, ReceiveError, WriteId}; + +const SLOT_COUNT: usize = 4; + +#[derive(Clone, Debug)] +enum Action { + ConnectIk(Side), + ConnectKk(Side), + AdvanceMs(u8), + OnTimer(Side), + OnTimerBoth, + Pump, + TakeNext(Side), + ConfirmTaken { + side: Side, + index: usize, + }, + RejectTaken { + side: Side, + index: usize, + }, + CaptureNext(Side), + DeliverNext(Side), + DropNext(Side), + DeliverQueued { + side: Side, + index: usize, + }, + DuplicateQueued { + side: Side, + index: usize, + }, + DropQueued { + side: Side, + index: usize, + }, + OpenStream { + side: Side, + slot: usize, + }, + Write { + side: Side, + slot: usize, + bytes: Vec, + }, + Finish { + side: Side, + slot: usize, + }, + Close { + side: Side, + slot: usize, + }, +} + +impl Action { + fn confirm_taken(side: Side, index: usize) -> Self { + Self::ConfirmTaken { side, index } + } + + fn reject_taken(side: Side, index: usize) -> Self { + Self::RejectTaken { side, index } + } + + fn deliver_queued(side: Side, index: usize) -> Self { + Self::DeliverQueued { side, index } + } + + fn duplicate_queued(side: Side, index: usize) -> Self { + Self::DuplicateQueued { side, index } + } + + fn drop_queued(side: Side, index: usize) -> Self { + Self::DropQueued { side, index } + } + + fn open_stream(side: Side, slot: usize) -> Self { + Self::OpenStream { side, slot } + } + + fn write(side: Side, slot: usize, bytes: Vec) -> Self { + Self::Write { side, slot, bytes } + } + + fn finish(side: Side, slot: usize) -> Self { + Self::Finish { side, slot } + } + + fn close(side: Side, slot: usize) -> Self { + Self::Close { side, slot } + } +} + +#[derive(Clone, Debug)] +struct TakenWrite { + record: Vec, + write_id: Option, +} + +#[derive(Default)] +struct SideEventState { + opened: BTreeSet, + finished: BTreeSet, + outbound_finished: BTreeSet, + writable_closed: BTreeSet, + closed: BTreeSet, + peer_statuses: Vec, + last_peer_status: Option, + session_epoch: usize, + session_closed_epoch: Option, +} + +impl SideEventState { + fn note_peer_status(&mut self, status: PeerStatus) { + if status == PeerStatus::Connected && self.last_peer_status != Some(PeerStatus::Connected) { + self.session_epoch = self.session_epoch.saturating_add(1); + } + self.peer_statuses.push(status); + self.last_peer_status = Some(status); + } +} + +struct Runner { + harness: Harness, + slots: [[Option; SLOT_COUNT]; 2], + taken: [Vec; 2], + pending: [Vec>; 2], + receive_errors: Vec<(Side, ReceiveError)>, + events: [SideEventState; 2], + known_streams: BTreeSet, + expected: [BTreeMap>; 2], + received: [BTreeMap>; 2], + finished_by: [BTreeSet; 2], + closed_by: [BTreeSet; 2], +} + +impl Runner { + fn handshake() -> Self { + let config = QlFsmConfig { + handshake_timeout: Duration::from_millis(60), + session_record_ack_delay: Duration::from_millis(5), + session_record_retransmit_timeout: Duration::from_millis(15), + session_peer_timeout: Duration::from_millis(80), + ..QlFsmConfig::default() + }; + + Self { + harness: Harness::paired_known(config), + slots: [[None; SLOT_COUNT]; 2], + taken: [Vec::new(), Vec::new()], + pending: [Vec::new(), Vec::new()], + receive_errors: Vec::new(), + events: [SideEventState::default(), SideEventState::default()], + known_streams: BTreeSet::new(), + expected: [BTreeMap::new(), BTreeMap::new()], + received: [BTreeMap::new(), BTreeMap::new()], + finished_by: [BTreeSet::new(), BTreeSet::new()], + closed_by: [BTreeSet::new(), BTreeSet::new()], + } + } + + fn connected() -> Self { + let config = QlFsmConfig { + session_record_ack_delay: Duration::from_millis(5), + session_record_retransmit_timeout: Duration::from_millis(15), + session_peer_timeout: Duration::from_secs(5), + ..QlFsmConfig::default() + }; + Self::connected_with_config(config) + } + + fn connected_with_config(config: QlFsmConfig) -> Self { + let connected_events = || SideEventState { + last_peer_status: Some(PeerStatus::Connected), + session_epoch: 1, + ..SideEventState::default() + }; + + Self { + harness: Harness::connected(config), + slots: [[None; SLOT_COUNT]; 2], + taken: [Vec::new(), Vec::new()], + pending: [Vec::new(), Vec::new()], + receive_errors: Vec::new(), + events: [connected_events(), connected_events()], + known_streams: BTreeSet::new(), + expected: [BTreeMap::new(), BTreeMap::new()], + received: [BTreeMap::new(), BTreeMap::new()], + finished_by: [BTreeSet::new(), BTreeSet::new()], + closed_by: [BTreeSet::new(), BTreeSet::new()], + } + } + + fn run(&mut self, actions: &[Action]) -> TestCaseResult { + for action in actions { + self.apply(action); + self.observe_and_assert()?; + } + + self.cleanup()?; + self.observe_and_assert()?; + self.assert_terminal_semantics()?; + self.assert_quiesced() + } + + #[allow(clippy::cognitive_complexity, clippy::too_many_lines)] + fn apply(&mut self, action: &Action) { + match action { + Action::ConnectIk(side) => { + let _ = self.harness.connect_ik(*side); + } + Action::ConnectKk(side) => { + let _ = self.harness.connect_kk(*side); + } + Action::AdvanceMs(ms) => { + self.harness + .advance(Duration::from_millis(u64::from(*ms) + 1)); + } + Action::OnTimer(side) => self.harness.on_timer(*side), + Action::OnTimerBoth => { + self.harness.on_timer(Side::A); + self.harness.on_timer(Side::B); + } + Action::Pump => self.capture_all_outbound(), + Action::TakeNext(side) => { + if let Some(write) = take_unconfirmed_outbound(&mut self.harness, *side) { + self.taken[side.idx()].push(write); + } + } + Action::ConfirmTaken { side, index } => { + if let Some(write) = take_taken(&mut self.taken[side.idx()], *index) { + confirm_taken(&mut self.harness, *side, &write); + self.pending[side.idx()].push(write.record); + } + } + Action::RejectTaken { side, index } => { + if let Some(write) = take_taken(&mut self.taken[side.idx()], *index) { + reject_taken(&mut self.harness, *side, &write); + } + } + Action::CaptureNext(side) => { + if let Some(record) = take_confirmed_outbound(&mut self.harness, *side) { + self.pending[side.idx()].push(record); + } + } + Action::DeliverNext(side) => { + if let Some(record) = take_confirmed_outbound(&mut self.harness, *side) { + self.deliver_to(opposite(*side), record); + } + } + Action::DropNext(side) => { + let _ = take_confirmed_outbound(&mut self.harness, *side); + } + Action::DeliverQueued { side, index } => { + if let Some(record) = take_pending(&mut self.pending[side.idx()], *index) { + self.deliver_to(opposite(*side), record); + } + } + Action::DuplicateQueued { side, index } => { + if let Some(record) = peek_pending(&self.pending[side.idx()], *index) { + self.deliver_to(opposite(*side), record); + } + } + Action::DropQueued { side, index } => { + let _ = take_pending(&mut self.pending[side.idx()], *index); + } + Action::OpenStream { side, slot } => { + let stream_id = self + .harness + .node_mut(*side) + .fsm + .open_stream(test_route_id()) + .ok() + .map(|stream| stream.stream_id()); + if let Some(stream_id) = stream_id { + self.slots[side.idx()][*slot] = Some(stream_id); + self.known_streams.insert(stream_id); + } + } + Action::Write { side, slot, bytes } => { + if let Some(stream_id) = self.slots[side.idx()][*slot] { + let mut chunk = Bytes::copy_from_slice(bytes); + let accepted = self.harness.node_mut(*side).fsm.stream(stream_id).map_or( + 0, + |mut stream| { + stream + .writer() + .map_or(0, |mut writer| writer.write(&mut chunk)) + }, + ); + if accepted != 0 { + self.expected[opposite(*side).idx()] + .entry(stream_id) + .or_default() + .extend_from_slice(&bytes[..accepted]); + } + } + } + Action::Finish { side, slot } => { + if let Some(stream_id) = self.slots[side.idx()][*slot] { + let finished = self + .harness + .node_mut(*side) + .fsm + .stream(stream_id) + .is_ok_and(|mut stream| { + stream.writer().is_some_and(|writer| { + writer.finish(); + true + }) + }); + if finished { + self.finished_by[side.idx()].insert(stream_id); + } + } + } + Action::Close { side, slot } => { + if let Some(stream_id) = self.slots[side.idx()][*slot] { + let closed = self + .harness + .node_mut(*side) + .fsm + .stream(stream_id) + .is_ok_and(|mut stream| { + stream.close(CloseTarget::Both, StreamCloseCode::CANCELLED); + true + }); + if closed { + self.closed_by[side.idx()].insert(stream_id); + self.slots[side.idx()][*slot] = None; + } + } + } + } + } + + fn observe_and_assert(&mut self) -> TestCaseResult { + self.drain_reads(Side::A); + self.drain_reads(Side::B); + let events_a = self.harness.drain_events(Side::A); + let events_b = self.harness.drain_events(Side::B); + self.process_events(Side::A, events_a)?; + self.process_events(Side::B, events_b)?; + self.assert_prefix_invariants()?; + self.assert_legal_link_state()?; + self.assert_receive_errors() + } + + fn cleanup(&mut self) -> TestCaseResult { + let tick = self + .harness + .a + .fsm + .config + .session_record_retransmit_timeout + .max(self.harness.a.fsm.config.session_record_ack_delay) + + Duration::from_millis(1); + + self.reject_all_taken(); + + for _ in 0..12 { + self.capture_all_outbound(); + self.flush_pending_in_order(); + self.capture_all_outbound(); + self.flush_pending_in_order(); + self.observe_and_assert()?; + self.harness.advance(tick); + self.harness.on_timer(Side::A); + self.harness.on_timer(Side::B); + self.capture_all_outbound(); + self.flush_pending_in_order(); + self.observe_and_assert()?; + self.reject_all_taken(); + } + + Ok(()) + } + + fn drain_reads(&mut self, side: Side) { + for stream_id in self.known_streams.clone() { + let appended = drain_stream(&mut self.harness.node_mut(side).fsm, stream_id); + if !appended.is_empty() { + self.received[side.idx()] + .entry(stream_id) + .or_default() + .extend_from_slice(&appended); + } + } + } + + fn process_events(&mut self, side: Side, events: Vec) -> TestCaseResult { + for event in events { + match event { + Event::NewPeer => {} + Event::PeerStatusChanged(status) => { + if status == PeerStatus::Unpaired { + let state = &mut self.events[side.idx()]; + prop_assert!( + state.session_epoch > 0, + "side {side:?} emitted Unpaired without a connected session" + ); + prop_assert!( + state.session_closed_epoch != Some(state.session_epoch), + "side {side:?} emitted duplicate terminal event in session epoch {}", + state.session_epoch + ); + state.session_closed_epoch = Some(state.session_epoch); + } + self.events[side.idx()].note_peer_status(status); + } + Event::Opened { stream_id, .. } => { + prop_assert!( + self.known_streams.contains(&stream_id), + "side {side:?} emitted Opened for unknown stream {stream_id:?}" + ); + prop_assert!( + self.events[side.idx()].opened.insert(stream_id), + "side {side:?} emitted duplicate Opened for {stream_id:?}" + ); + } + Event::Readable(stream_id) | Event::Writable(stream_id) => { + prop_assert!( + self.known_streams.contains(&stream_id), + "side {side:?} emitted readiness for unknown stream {stream_id:?}" + ); + } + Event::Finished(stream_id) => { + prop_assert!( + self.known_streams.contains(&stream_id), + "side {side:?} emitted Finished for unknown stream {stream_id:?}" + ); + prop_assert!( + self.events[side.idx()].finished.insert(stream_id), + "side {side:?} emitted duplicate Finished for {stream_id:?}" + ); + prop_assert!( + !self.events[side.idx()].closed.contains(&stream_id), + "side {side:?} emitted Finished after Closed for {stream_id:?}" + ); + } + Event::OutboundFinished(stream_id) => { + prop_assert!( + self.known_streams.contains(&stream_id), + "side {side:?} emitted OutboundFinished for unknown stream {stream_id:?}" + ); + prop_assert!( + self.events[side.idx()].outbound_finished.insert(stream_id), + "side {side:?} emitted duplicate OutboundFinished for {stream_id:?}" + ); + } + Event::Closed(frame) => { + prop_assert!( + self.known_streams.contains(&frame.stream_id), + "side {side:?} emitted Closed for unknown stream {:?}", + frame.stream_id + ); + prop_assert!( + self.events[side.idx()].closed.insert(frame.stream_id), + "side {side:?} emitted duplicate Closed for {:?}", + frame.stream_id + ); + } + Event::WritableClosed(frame) => { + let stream_id = frame.stream_id; + prop_assert!( + self.known_streams.contains(&stream_id), + "side {side:?} emitted WritableClosed for unknown stream {stream_id:?}" + ); + prop_assert!( + self.events[side.idx()].writable_closed.insert(stream_id), + "side {side:?} emitted duplicate WritableClosed for {stream_id:?}" + ); + } + Event::SessionClosed(_) => { + let state = &mut self.events[side.idx()]; + prop_assert!( + state.session_epoch > 0, + "side {side:?} emitted SessionClosed without a connected session" + ); + prop_assert!( + state.session_closed_epoch != Some(state.session_epoch), + "side {side:?} emitted duplicate SessionClosed in session epoch {}", + state.session_epoch + ); + state.session_closed_epoch = Some(state.session_epoch); + } + } + } + + Ok(()) + } + + fn assert_prefix_invariants(&self) -> TestCaseResult { + for side in [Side::A, Side::B] { + for (stream_id, received) in &self.received[side.idx()] { + let expected = self.expected[side.idx()] + .get(stream_id) + .map_or(&[][..], Vec::as_slice); + prop_assert!( + expected.starts_with(received), + "side {side:?} observed non-prefix bytes on {stream_id:?}: received={received:?} expected={expected:?}" + ); + } + } + + Ok(()) + } + + fn assert_legal_link_state(&self) -> TestCaseResult { + let a_connected = matches!(self.harness.a.fsm.state.link, LinkState::Connected(_)); + let b_connected = matches!(self.harness.b.fsm.state.link, LinkState::Connected(_)); + + prop_assert!( + !a_connected || self.harness.a.fsm.peer().is_some(), + "side A reached Connected without a bound peer" + ); + prop_assert!( + !b_connected || self.harness.b.fsm.peer().is_some(), + "side B reached Connected without a bound peer" + ); + + Ok(()) + } + + fn assert_receive_errors(&self) -> TestCaseResult { + for (side, error) in &self.receive_errors { + prop_assert!( + matches!( + error, + ReceiveError::NoSession + | ReceiveError::NoPeer + | ReceiveError::InvalidRemoteBundle + | ReceiveError::InvalidSessionPayload(WireError::InvalidPayload) + | ReceiveError::InvalidSessionPayload(WireError::DecryptFailed) + | ReceiveError::InvalidIkHandshake(WireError::InvalidPayload) + | ReceiveError::InvalidIkHandshake(WireError::InvalidState) + | ReceiveError::InvalidKkHandshake(WireError::InvalidPayload) + | ReceiveError::InvalidKkHandshake(WireError::InvalidState) + | ReceiveError::InvalidXxHandshake(WireError::InvalidPayload) + | ReceiveError::InvalidXxHandshake(WireError::InvalidState) + | ReceiveError::InvalidXxHandshake(WireError::DecryptFailed) + ), + "unexpected receive error on side {side:?}: {error:?}" + ); + } + + Ok(()) + } + + fn assert_terminal_semantics(&self) -> TestCaseResult { + let a_connected = matches!(self.harness.a.fsm.state.link, LinkState::Connected(_)); + let b_connected = matches!(self.harness.b.fsm.state.link, LinkState::Connected(_)); + let connected = [a_connected, b_connected]; + + for side in [Side::A, Side::B] { + for stream_id in &self.events[side.idx()].finished { + if self.inbound_aborted(side, *stream_id) { + continue; + } + let expected = self.expected[side.idx()] + .get(stream_id) + .map_or(&[][..], Vec::as_slice); + let received = self.received[side.idx()] + .get(stream_id) + .map_or(&[][..], Vec::as_slice); + prop_assert_eq!( + received, + expected, + "side {:?} finished {:?} without receiving all expected bytes", + side, + stream_id + ); + } + + for stream_id in &self.finished_by[side.idx()] { + prop_assert!( + self.events[opposite(side).idx()].finished.contains(stream_id) + || self.events[opposite(side).idx()].closed.contains(stream_id) + || !connected[opposite(side).idx()], + "side {side:?} finished {stream_id:?} but side {:?} saw neither Finished nor Closed", + opposite(side) + ); + } + + for stream_id in &self.closed_by[side.idx()] { + prop_assert!( + self.events[opposite(side).idx()].closed.contains(stream_id) + || !connected[opposite(side).idx()], + "side {side:?} closed {stream_id:?} but side {:?} saw no Closed event", + opposite(side) + ); + } + } + + Ok(()) + } + + fn assert_expected_delivered(&self, side: Side) -> TestCaseResult { + for (stream_id, expected) in &self.expected[side.idx()] { + let received = self.received[side.idx()] + .get(stream_id) + .map_or(&[][..], Vec::as_slice); + prop_assert_eq!( + received, + expected, + "side {:?} did not receive full payload for {:?}", + side, + stream_id + ); + } + + Ok(()) + } + + fn assert_no_stream_events(&self) -> TestCaseResult { + prop_assert!( + self.known_streams.is_empty() + && self.events.iter().all(|events| { + events.opened.is_empty() + && events.finished.is_empty() + && events.outbound_finished.is_empty() + && events.closed.is_empty() + && events.writable_closed.is_empty() + }), + "handshake-only property observed stream activity" + ); + Ok(()) + } + + fn assert_no_taken_writes(&self) -> TestCaseResult { + prop_assert!( + self.taken.iter().all(Vec::is_empty), + "cleanup left taken writes queued" + ); + Ok(()) + } + + fn assert_quiesced(&mut self) -> TestCaseResult { + self.reject_all_taken(); + + for _ in 0..8 { + self.capture_all_outbound(); + if self.pending.iter().all(Vec::is_empty) { + break; + } + self.flush_pending_in_order(); + self.observe_and_assert()?; + } + + self.capture_all_outbound(); + prop_assert!( + self.pending.iter().all(Vec::is_empty) && self.taken.iter().all(Vec::is_empty), + "cleanup did not quiesce: taken_a={} taken_b={} pending_a={} pending_b={}", + self.taken[Side::A.idx()].len(), + self.taken[Side::B.idx()].len(), + self.pending[Side::A.idx()].len(), + self.pending[Side::B.idx()].len() + ); + + Ok(()) + } + + fn capture_all_outbound(&mut self) { + for side in [Side::A, Side::B] { + while let Some(record) = take_confirmed_outbound(&mut self.harness, side) { + self.pending[side.idx()].push(record); + } + } + } + + fn flush_pending_in_order(&mut self) { + for side in [Side::A, Side::B] { + while let Some(record) = pop_front_pending(&mut self.pending[side.idx()]) { + self.deliver_to(opposite(side), record); + } + } + } + + fn reject_all_taken(&mut self) { + for side in [Side::A, Side::B] { + while let Some(write) = self.taken[side.idx()].pop() { + reject_taken(&mut self.harness, side, &write); + } + } + } + + fn deliver_to(&mut self, side: Side, record: Vec) { + if let Err(error) = deliver_to(&mut self.harness, side, record) { + self.receive_errors.push((side, error)); + } + } + + fn inbound_aborted(&self, side: Side, stream_id: StreamId) -> bool { + self.events[side.idx()].closed.contains(&stream_id) + || self.closed_by[side.idx()].contains(&stream_id) + } +} + +fn take_unconfirmed_outbound(harness: &mut Harness, side: Side) -> Option { + let write = harness.next_write(side)?; + Some(TakenWrite { + record: write.record, + write_id: write.write_id, + }) +} + +fn take_confirmed_outbound(harness: &mut Harness, side: Side) -> Option> { + let write = take_unconfirmed_outbound(harness, side)?; + confirm_taken(harness, side, &write); + Some(write.record) +} + +fn confirm_taken(harness: &mut Harness, side: Side, write: &TakenWrite) { + if let Some(write_id) = write.write_id { + harness.confirm_write(side, write_id); + } +} + +fn reject_taken(harness: &mut Harness, side: Side, write: &TakenWrite) { + if let Some(write_id) = write.write_id { + harness.reject_write(side, write_id); + } +} + +fn deliver_to(harness: &mut Harness, side: Side, record: Vec) -> Result<(), ReceiveError> { + let time = harness.time(); + let Node { fsm, crypto } = harness.node_mut(side); + fsm.receive(time, record, crypto) +} + +fn take_pending(pending: &mut Vec>, index: usize) -> Option> { + if pending.is_empty() { + return None; + } + + Some(pending.remove(index % pending.len())) +} + +fn peek_pending(pending: &[Vec], index: usize) -> Option> { + if pending.is_empty() { + return None; + } + + Some(pending[index % pending.len()].clone()) +} + +fn pop_front_pending(pending: &mut Vec>) -> Option> { + if pending.is_empty() { + None + } else { + Some(pending.remove(0)) + } +} + +fn take_taken(taken: &mut Vec, index: usize) -> Option { + if taken.is_empty() { + return None; + } + + Some(taken.remove(index % taken.len())) +} + +fn drain_stream(fsm: &mut QlFsm, stream_id: StreamId) -> Vec { + let mut out = Vec::new(); + let Ok(mut stream) = fsm.stream(stream_id) else { + return out; + }; + + loop { + let mut read = 0usize; + for chunk in stream.read() { + out.extend_from_slice(&chunk); + read += chunk.len(); + } + + if read == 0 { + break; + } + + stream.commit_read(read).unwrap(); + } + + out +} + +fn opposite(side: Side) -> Side { + match side { + Side::A => Side::B, + Side::B => Side::A, + } +} + +fn side_strategy() -> impl Strategy { + prop_oneof![Just(Side::A), Just(Side::B)] +} + +fn side_action(f: fn(Side) -> Action) -> impl Strategy { + side_strategy().prop_map(f) +} + +fn side_usize_action( + values: impl Strategy, + f: fn(Side, usize) -> Action, +) -> impl Strategy { + (side_strategy(), values).prop_map(move |(side, value)| f(side, value)) +} + +fn side_usize_vec_action( + values: impl Strategy, + bytes: impl Strategy>, + f: fn(Side, usize, Vec) -> Action, +) -> impl Strategy { + (side_strategy(), values, bytes).prop_map(move |(side, value, bytes)| f(side, value, bytes)) +} + +fn handshake_action_strategy() -> impl Strategy { + let queue_index = 0usize..6; + prop_oneof![ + side_action(Action::ConnectIk), + side_action(Action::ConnectKk), + (0u8..40).prop_map(Action::AdvanceMs), + side_action(Action::OnTimer), + Just(Action::OnTimerBoth), + Just(Action::Pump), + side_action(Action::TakeNext), + side_usize_action(queue_index.clone(), Action::confirm_taken), + side_usize_action(queue_index.clone(), Action::reject_taken), + side_action(Action::CaptureNext), + side_action(Action::DeliverNext), + side_action(Action::DropNext), + side_usize_action(queue_index.clone(), Action::deliver_queued), + side_usize_action(queue_index.clone(), Action::duplicate_queued), + side_usize_action(queue_index, Action::drop_queued), + ] +} + +fn connected_action_strategy() -> impl Strategy { + let bytes = vec(any::(), 0..24); + let slot = 0usize..SLOT_COUNT; + let queue_index = 0usize..6; + prop_oneof![ + (0u8..30).prop_map(Action::AdvanceMs), + side_action(Action::OnTimer), + Just(Action::OnTimerBoth), + Just(Action::Pump), + side_action(Action::TakeNext), + side_usize_action(queue_index.clone(), Action::confirm_taken), + side_usize_action(queue_index.clone(), Action::reject_taken), + side_action(Action::CaptureNext), + side_action(Action::DeliverNext), + side_action(Action::DropNext), + side_usize_action(queue_index.clone(), Action::deliver_queued), + side_usize_action(queue_index.clone(), Action::duplicate_queued), + side_usize_action(queue_index, Action::drop_queued), + side_usize_action(slot.clone(), Action::open_stream), + side_usize_vec_action(slot.clone(), bytes, Action::write), + side_usize_action(slot.clone(), Action::finish), + side_usize_action(slot, Action::close), + ] +} + +fn write_tracking_action_strategy() -> impl Strategy { + let bytes = vec(any::(), 0..16); + let slot = 0usize..SLOT_COUNT; + let queue_index = 0usize..6; + prop_oneof![ + side_usize_action(slot.clone(), Action::open_stream), + side_usize_vec_action(slot, bytes, Action::write), + side_action(Action::TakeNext), + side_usize_action(queue_index.clone(), Action::confirm_taken), + side_usize_action(queue_index.clone(), Action::reject_taken), + side_usize_action(queue_index.clone(), Action::deliver_queued), + side_usize_action(queue_index.clone(), Action::duplicate_queued), + side_usize_action(queue_index, Action::drop_queued), + Just(Action::Pump), + side_action(Action::OnTimer), + Just(Action::OnTimerBoth), + (0u8..20).prop_map(Action::AdvanceMs), + ] +} + +fn packet_loss_recovery_action_strategy() -> impl Strategy { + let queue_index = 0usize..16; + prop_oneof![ + (0u8..20).prop_map(Action::AdvanceMs), + side_action(Action::OnTimer), + Just(Action::OnTimerBoth), + Just(Action::Pump), + side_usize_action(queue_index.clone(), Action::deliver_queued), + side_usize_action(queue_index.clone(), Action::duplicate_queued), + side_usize_action(queue_index, Action::drop_queued), + ] +} + +fn terminal_action_strategy() -> impl Strategy { + let bytes = vec(any::(), 0..16); + let slot = 0usize..SLOT_COUNT; + let queue_index = 0usize..6; + prop_oneof![ + side_usize_action(slot.clone(), Action::open_stream), + side_usize_vec_action(slot.clone(), bytes, Action::write), + side_usize_action(slot.clone(), Action::finish), + side_usize_action(slot, Action::close), + side_action(Action::TakeNext), + side_usize_action(queue_index.clone(), Action::confirm_taken), + side_usize_action(queue_index.clone(), Action::reject_taken), + side_usize_action(queue_index.clone(), Action::deliver_queued), + side_usize_action(queue_index.clone(), Action::duplicate_queued), + side_usize_action(queue_index, Action::drop_queued), + Just(Action::Pump), + side_action(Action::OnTimer), + Just(Action::OnTimerBoth), + (0u8..20).prop_map(Action::AdvanceMs), + ] +} + +proptest_crate::proptest! { + #![proptest_config(ProptestConfig { + cases: 24, + max_shrink_iters: 10_000, + .. ProptestConfig::default() + })] + + #[test] + fn randomized_handshake_actions_quiesce(actions in vec(handshake_action_strategy(), 1..64)) { + let mut runner = Runner::handshake(); + runner.run(&actions)?; + runner.assert_no_stream_events()?; + } + + #[test] + fn randomized_stream_actions_preserve_integrity(actions in vec(connected_action_strategy(), 1..80)) { + let mut runner = Runner::connected(); + runner.run(&actions)?; + } + + #[test] + fn randomized_write_tracking_actions_quiesce(actions in vec(write_tracking_action_strategy(), 1..80)) { + let mut runner = Runner::connected(); + runner.run(&actions)?; + runner.assert_no_taken_writes()?; + } + + #[test] + fn randomized_session_packet_loss_recovers( + payload in vec(any::(), 512..2048), + actions in vec(packet_loss_recovery_action_strategy(), 1..96), + ) { + let config = QlFsmConfig { + session_record_ack_delay: Duration::from_millis(1), + session_record_retransmit_timeout: Duration::from_millis(10), + session_record_max_size: 96, + session_pending_ack_range_limit: 512, + ..QlFsmConfig::default() + }; + let mut runner = Runner::connected_with_config(config); + + runner.apply(&Action::open_stream(Side::A, 0)); + runner.observe_and_assert()?; + + runner.apply(&Action::write(Side::A, 0, payload)); + runner.observe_and_assert()?; + + runner.apply(&Action::finish(Side::A, 0)); + runner.observe_and_assert()?; + + for action in &actions { + runner.apply(action); + runner.observe_and_assert()?; + } + + runner.cleanup()?; + runner.observe_and_assert()?; + runner.assert_expected_delivered(Side::B)?; + runner.assert_terminal_semantics()?; + runner.assert_quiesced()?; + } + + #[test] + fn randomized_terminal_actions_preserve_terminal_semantics(actions in vec(terminal_action_strategy(), 1..80)) { + let mut runner = Runner::connected(); + runner.run(&actions)?; + runner.assert_terminal_semantics()?; + } +} diff --git a/ql-fsm/src/tests/session.rs b/ql-fsm/src/tests/session.rs new file mode 100644 index 00000000..c55e51c1 --- /dev/null +++ b/ql-fsm/src/tests/session.rs @@ -0,0 +1,532 @@ +use std::time::Duration; + +use bytes::Bytes; +use ql_wire::{RouteId, SessionClose, StreamId, VarInt}; + +use super::*; +use crate::{state::LinkState, CommitReadError, Event, NoSessionError, PeerStatus, StreamError}; + +fn stream_id(value: u32) -> StreamId { + StreamId(VarInt::from_u32(value)) +} + +fn route_id(value: u32) -> RouteId { + RouteId::from_u32(value) +} + +fn opened(stream_id: StreamId) -> Event { + Event::Opened { + stream_id, + route_id: route_id(1), + } +} + +fn open_stream_id(fsm: &mut QlFsm) -> StreamId { + fsm.open_stream(route_id(1)).unwrap().stream_id() +} + +fn write_stream_bytes( + fsm: &mut QlFsm, + stream_id: StreamId, + bytes: &[u8], +) -> Result { + let mut bytes = Bytes::copy_from_slice(bytes); + let mut stream = fsm.stream(stream_id)?; + let Some(mut writer) = stream.writer() else { + return Err(StreamError::NotWritable); + }; + Ok(writer.write(&mut bytes)) +} + +fn read_stream_all(fsm: &mut QlFsm, stream_id: StreamId) -> Vec { + let mut out = Vec::new(); + let Ok(mut stream) = fsm.stream(stream_id) else { + return out; + }; + loop { + let mut read = 0; + for chunk in stream.read() { + out.extend_from_slice(&chunk); + read += chunk.len(); + } + if read == 0 { + break; + } + stream.commit_read(read).unwrap(); + } + out +} + +#[test] +fn connected_fsms_deliver_stream_data() { + let mut harness = Harness::connected(QlFsmConfig::default()); + + let stream_id = open_stream_id(&mut harness.a.fsm); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"hello").unwrap(), + 5 + ); + harness + .a + .fsm + .stream(stream_id) + .unwrap() + .writer() + .unwrap() + .finish(); + + harness.pump(); + + assert_eq!(harness.take_event(Side::B), Some(opened(stream_id))); + assert_eq!( + harness.take_event(Side::B), + Some(Event::Readable(stream_id)) + ); + assert_eq!( + read_stream_all(&mut harness.b.fsm, stream_id), + b"hello".to_vec() + ); + assert_eq!( + harness.take_event(Side::B), + Some(Event::Finished(stream_id)) + ); + harness.advance(QlFsmConfig::default().session_record_ack_delay); + harness.on_timer(Side::B); + harness.pump(); + assert_eq!( + harness.take_event(Side::A), + Some(Event::OutboundFinished(stream_id)) + ); +} + +#[test] +fn session_retransmit_uses_new_record_seq() { + let config = QlFsmConfig::default(); + let mut harness = Harness::connected(config); + + let stream_id = open_stream_id(&mut harness.a.fsm); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"retry").unwrap(), + 5 + ); + + let first = harness.next_decoded_outbound(Side::A).unwrap(); + + harness.advance(config.session_record_retransmit_timeout + Duration::from_millis(1)); + harness.on_timer(Side::A); + + let retried = harness.next_decoded_outbound(Side::A).unwrap(); + + assert_ne!(retried.header.seq, first.header.seq); + assert_eq!(retried.frames, first.frames); + + harness.deliver(Side::B, retried.record); + harness.advance(config.session_record_ack_delay); + harness.on_timer(Side::A); + harness.on_timer(Side::B); + harness.pump(); + + assert_eq!(harness.take_event(Side::B), Some(opened(stream_id))); + assert_eq!( + harness.take_event(Side::B), + Some(Event::Readable(stream_id)) + ); + assert_eq!( + read_stream_all(&mut harness.b.fsm, stream_id), + b"retry".to_vec() + ); + + harness.advance(config.session_record_retransmit_timeout + Duration::from_millis(1)); + harness.on_timer(Side::A); + assert!(harness.next_outbound(Side::A).is_none()); +} + +#[test] +fn simultaneous_opens_use_even_and_odd_stream_ids() { + let mut harness = Harness::connected(QlFsmConfig::default()); + + let stream_id_a = open_stream_id(&mut harness.a.fsm); + let stream_id_b = open_stream_id(&mut harness.b.fsm); + + assert_ne!(stream_id_a, stream_id_b); + assert!( + StreamParity::for_local(harness.a.fsm.identity.qid, harness.b.fsm.identity.qid) + .matches(stream_id_a) + ); + assert!( + StreamParity::for_local(harness.b.fsm.identity.qid, harness.a.fsm.identity.qid) + .matches(stream_id_b) + ); + + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id_a, b"from-a").unwrap(), + 6 + ); + assert_eq!( + write_stream_bytes(&mut harness.b.fsm, stream_id_b, b"from-b").unwrap(), + 6 + ); + + harness.pump(); + + assert_eq!(harness.take_event(Side::A), Some(opened(stream_id_b))); + assert_eq!( + harness.take_event(Side::A), + Some(Event::Readable(stream_id_b)) + ); + assert_eq!( + read_stream_all(&mut harness.a.fsm, stream_id_b), + b"from-b".to_vec() + ); + assert_eq!(harness.take_event(Side::B), Some(opened(stream_id_a))); + assert_eq!( + harness.take_event(Side::B), + Some(Event::Readable(stream_id_a)) + ); + assert_eq!( + read_stream_all(&mut harness.b.fsm, stream_id_a), + b"from-a".to_vec() + ); +} + +#[test] +fn disconnected_stream_operations_fail_with_no_session() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + let missing = stream_id(0); + + assert!(matches!( + harness.a.fsm.open_stream(route_id(1)), + Err(NoSessionError) + )); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, missing, b"queued"), + Err(StreamError::NoSession) + ); + assert_eq!( + harness + .a + .fsm + .stream(missing) + .map(|mut stream| stream.writer().unwrap().finish()), + Err(StreamError::NoSession) + ); + assert_eq!( + harness.a.fsm.stream(missing).map(|mut stream| { + stream.close( + ql_wire::CloseTarget::Both, + ql_wire::StreamCloseCode::CANCELLED, + ); + }), + Err(StreamError::NoSession) + ); + assert_eq!(harness.a.fsm.queue_ping(), Err(NoSessionError)); + assert!(matches!( + harness.a.fsm.stream(missing), + Err(StreamError::NoSession) + )); +} + +#[test] +fn disconnected_stream_read_accessors_return_none() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + let missing = stream_id(0); + + assert!(matches!( + harness.a.fsm.stream(missing), + Err(StreamError::NoSession) + )); +} + +#[test] +fn commit_read_rejects_lengths_past_readable_prefix() { + let mut harness = Harness::connected(QlFsmConfig::default()); + + let stream_id = open_stream_id(&mut harness.a.fsm); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"hi").unwrap(), + 2 + ); + harness.pump(); + + let mut stream = harness.b.fsm.stream(stream_id).unwrap(); + assert_eq!(stream.commit_read(3), Err(CommitReadError)); +} + +#[test] +fn returned_session_write_is_reissued_with_new_record_seq() { + let mut harness = Harness::connected(QlFsmConfig::default()); + + let stream_id = open_stream_id(&mut harness.a.fsm); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"retry").unwrap(), + 5 + ); + + let first = harness.next_decoded_write(Side::A).unwrap(); + let id = first.write_id.expect("expected session write"); + + harness.reject_write(Side::A, id); + + let reissued = harness.next_decoded_write(Side::A).unwrap(); + let reissued_id = reissued.write_id.expect("expected reissued write"); + + assert_ne!(reissued_id, id); + assert_ne!(reissued.header.seq, first.header.seq); + assert_eq!(reissued.frames, first.frames); + + harness.confirm_write(Side::A, reissued_id); + harness.deliver(Side::B, reissued.record); + harness.pump(); + + assert_eq!(harness.take_event(Side::B), Some(opened(stream_id))); + assert_eq!( + harness.take_event(Side::B), + Some(Event::Readable(stream_id)) + ); + assert_eq!( + read_stream_all(&mut harness.b.fsm, stream_id), + b"retry".to_vec() + ); +} + +#[test] +fn unconfirmed_session_write_does_not_start_retransmit_timer() { + let config = QlFsmConfig::default(); + let mut harness = Harness::connected(config); + + let stream_id = open_stream_id(&mut harness.a.fsm); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"retry").unwrap(), + 5 + ); + + let first = harness.next_decoded_write(Side::A).unwrap(); + let id = first.write_id.expect("expected session write"); + + harness.advance(config.session_record_retransmit_timeout + Duration::from_millis(1)); + harness.on_timer(Side::A); + assert!(harness.next_write(Side::A).is_none()); + + harness.confirm_write(Side::A, id); + harness.advance(config.session_record_retransmit_timeout + Duration::from_millis(1)); + harness.on_timer(Side::A); + + let retried = harness.next_decoded_write(Side::A).unwrap(); + + assert_ne!(retried.header.seq, first.header.seq); + assert_eq!(retried.frames, first.frames); +} + +#[test] +fn ack_frame_releases_stream_capacity_and_emits_writable() { + let config = QlFsmConfig { + session_stream_send_buffer_size: 4, + ..QlFsmConfig::default() + }; + let mut harness = Harness::connected(config); + + let stream_id = open_stream_id(&mut harness.a.fsm); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"abcd").unwrap(), + 4 + ); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"z").unwrap(), + 0 + ); + + let record = harness.next_outbound(Side::A).unwrap(); + harness.deliver(Side::B, record); + harness.advance(config.session_record_ack_delay); + harness.on_timer(Side::A); + harness.on_timer(Side::B); + harness.pump(); + + assert_eq!( + harness.take_event(Side::A), + Some(Event::Writable(stream_id)) + ); +} + +#[test] +fn close_session_disconnects_locally() { + let mut harness = Harness::connected(QlFsmConfig::default()); + + harness + .a + .fsm + .close_session(ql_wire::SessionCloseCode::CANCELLED); + + assert!(matches!( + harness.take_event(Side::A), + Some(Event::SessionClosed(SessionClose { + code: ql_wire::SessionCloseCode::CANCELLED, + })) + )); + assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); + assert!(matches!( + harness.a.fsm.open_stream(route_id(1)), + Err(NoSessionError) + )); + assert_eq!(harness.a.fsm.queue_ping(), Err(NoSessionError)); + + let close = harness.next_decoded_outbound(Side::A).unwrap(); + assert!(matches!( + close.frames.as_slice(), + [ql_wire::SessionFrame::Close(_)] + )); + + assert!(matches!(harness.a.fsm.state.link, LinkState::Idle)); + assert_eq!( + harness.take_event(Side::A), + Some(Event::PeerStatusChanged(PeerStatus::Disconnected)) + ); +} + +#[test] +fn unpair_clears_bound_peer_and_emits_unpair_frame() { + let mut harness = Harness::connected(QlFsmConfig::default()); + + harness.a.fsm.unpair(); + + assert_eq!( + harness.take_event(Side::A), + Some(Event::PeerStatusChanged(PeerStatus::Unpaired)) + ); + assert!(harness.a.fsm.peer().is_none()); + assert!(matches!( + harness.a.fsm.open_stream(route_id(1)), + Err(NoSessionError) + )); + assert_eq!(harness.a.fsm.queue_ping(), Err(NoSessionError)); + + let unpair = harness.next_decoded_outbound(Side::A).unwrap(); + assert!(matches!( + unpair.frames.as_slice(), + [ql_wire::SessionFrame::Unpair] + )); + assert!(matches!(harness.a.fsm.state.link, LinkState::Idle)); +} + +#[test] +fn inbound_unpair_clears_remote_peer_binding() { + let mut harness = Harness::connected(QlFsmConfig::default()); + + harness.a.fsm.unpair(); + let unpair = harness.next_outbound(Side::A).unwrap(); + harness.deliver(Side::B, unpair); + + assert_eq!( + harness.take_event(Side::B), + Some(Event::PeerStatusChanged(PeerStatus::Unpaired)) + ); + assert!(harness.b.fsm.peer().is_none()); + assert!(matches!( + harness.b.fsm.open_stream(route_id(1)), + Err(NoSessionError) + )); + assert!(matches!(harness.connect_ik(Side::B), Err(NoPeerError))); + + let reply_key = harness.b.fsm.state.link.transport().unwrap().tx_key.clone(); + let reply = harness.next_outbound(Side::B).unwrap(); + let (_header, frames) = decrypt_record(&harness.b.crypto, &reply, &reply_key); + assert!(matches!(frames.as_slice(), [ql_wire::SessionFrame::Unpair])); + assert!(matches!(harness.b.fsm.state.link, LinkState::Idle)); +} + +#[test] +fn local_unpair_without_session_emits_unpaired_immediately() { + let mut harness = Harness::paired_known(QlFsmConfig::default()); + + harness.a.fsm.unpair(); + + assert_eq!( + harness.take_event(Side::A), + Some(Event::PeerStatusChanged(PeerStatus::Unpaired)) + ); + assert!(harness.a.fsm.peer().is_none()); + assert_eq!(harness.take_event(Side::A), None); +} + +#[test] +fn session_records_contain_ack_frames_after_delivery() { + let config = QlFsmConfig::default(); + let mut harness = Harness::connected(config); + + let stream_id = open_stream_id(&mut harness.a.fsm); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"x").unwrap(), + 1 + ); + + let data = harness.next_outbound(Side::A).unwrap(); + harness.deliver(Side::B, data); + harness.advance(config.session_record_ack_delay); + harness.on_timer(Side::B); + + let ack = harness.next_decoded_outbound(Side::B).unwrap(); + assert!(matches!( + ack.frames.as_slice(), + [ql_wire::SessionFrame::Ack(_)] + )); +} + +#[test] +fn first_stream_data_uses_negotiated_initial_peer_credit() { + let mut harness = Harness::paired_known_with_configs( + QlFsmConfig { + session_stream_receive_buffer_size: 8, + ..QlFsmConfig::default() + }, + QlFsmConfig { + session_stream_receive_buffer_size: 3, + ..QlFsmConfig::default() + }, + ); + + harness.connect_ik(Side::A).unwrap(); + let ik1 = harness.next_outbound(Side::A).unwrap(); + harness.deliver(Side::B, ik1); + let ik2 = harness.next_outbound(Side::B).unwrap(); + harness.deliver(Side::A, ik2); + + let stream_id = open_stream_id(&mut harness.a.fsm); + assert_eq!( + write_stream_bytes(&mut harness.a.fsm, stream_id, b"hello").unwrap(), + 5 + ); + + assert!(matches!( + harness.next_decoded_outbound(Side::A).unwrap().frames.as_slice(), + [ql_wire::SessionFrame::StreamData(frame)] if frame.stream_id == stream_id && frame.bytes.as_slice() == b"hel" + )); +} + +#[test] +fn session_timeout_emits_close_before_disconnect() { + let config = QlFsmConfig { + session_peer_timeout: Duration::from_millis(30), + ..QlFsmConfig::default() + }; + let mut harness = Harness::connected(config); + + harness.advance(config.session_peer_timeout); + harness.on_timer(Side::A); + + assert_eq!( + harness.drain_events(Side::A), + vec![Event::SessionClosed(SessionClose { + code: ql_wire::SessionCloseCode::TIMEOUT, + })] + ); + + let close = harness.next_decoded_outbound(Side::A).unwrap(); + assert!(matches!( + close.frames.as_slice(), + [ql_wire::SessionFrame::Close(_)] + )); + assert_eq!( + harness.take_event(Side::A), + Some(Event::PeerStatusChanged(PeerStatus::Disconnected)) + ); +} From ba41308051cb164ce14bca4849182fdaf800c152 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 4 Jun 2026 09:25:14 -0400 Subject: [PATCH 04/59] ql-rpc: add RPC modality layer --- Cargo.lock | 19 ++ Cargo.toml | 2 + ql-rpc/Cargo.toml | 10 + ql-rpc/src/chunk_queue.rs | 252 ++++++++++++++++++ ql-rpc/src/codec.rs | 83 ++++++ ql-rpc/src/error.rs | 112 ++++++++ ql-rpc/src/framed_value.rs | 127 +++++++++ ql-rpc/src/lib.rs | 44 ++++ ql-rpc/src/route_id.rs | 19 ++ ql-rpc/src/router/builder.rs | 354 ++++++++++++++++++++++++++ ql-rpc/src/router/config.rs | 12 + ql-rpc/src/router/mod.rs | 89 +++++++ ql-rpc/src/router/mode.rs | 21 ++ ql-rpc/src/rpc/download/client.rs | 246 ++++++++++++++++++ ql-rpc/src/rpc/download/mod.rs | 31 +++ ql-rpc/src/rpc/download/server.rs | 221 ++++++++++++++++ ql-rpc/src/rpc/duplex/client.rs | 164 ++++++++++++ ql-rpc/src/rpc/duplex/codec.rs | 86 +++++++ ql-rpc/src/rpc/duplex/mod.rs | 24 ++ ql-rpc/src/rpc/duplex/server.rs | 47 ++++ ql-rpc/src/rpc/mod.rs | 32 +++ ql-rpc/src/rpc/notification/client.rs | 10 + ql-rpc/src/rpc/notification/mod.rs | 19 ++ ql-rpc/src/rpc/notification/server.rs | 48 ++++ ql-rpc/src/rpc/parts.rs | 283 ++++++++++++++++++++ ql-rpc/src/rpc/progress/client.rs | 152 +++++++++++ ql-rpc/src/rpc/progress/codec.rs | 145 +++++++++++ ql-rpc/src/rpc/progress/mod.rs | 27 ++ ql-rpc/src/rpc/progress/server.rs | 106 ++++++++ ql-rpc/src/rpc/request/client.rs | 34 +++ ql-rpc/src/rpc/request/mod.rs | 23 ++ ql-rpc/src/rpc/request/server.rs | 96 +++++++ ql-rpc/src/rpc/subscription/client.rs | 99 +++++++ ql-rpc/src/rpc/subscription/codec.rs | 58 +++++ ql-rpc/src/rpc/subscription/mod.rs | 23 ++ ql-rpc/src/rpc/subscription/server.rs | 105 ++++++++ ql-rpc/src/rpc/upload/client.rs | 146 +++++++++++ ql-rpc/src/rpc/upload/mod.rs | 26 ++ ql-rpc/src/rpc/upload/server.rs | 243 ++++++++++++++++++ ql-rpc/src/rpc/utils.rs | 120 +++++++++ ql-rpc/src/stream.rs | 89 +++++++ 41 files changed, 3847 insertions(+) create mode 100644 ql-rpc/Cargo.toml create mode 100644 ql-rpc/src/chunk_queue.rs create mode 100644 ql-rpc/src/codec.rs create mode 100644 ql-rpc/src/error.rs create mode 100644 ql-rpc/src/framed_value.rs create mode 100644 ql-rpc/src/lib.rs create mode 100644 ql-rpc/src/route_id.rs create mode 100644 ql-rpc/src/router/builder.rs create mode 100644 ql-rpc/src/router/config.rs create mode 100644 ql-rpc/src/router/mod.rs create mode 100644 ql-rpc/src/router/mode.rs create mode 100644 ql-rpc/src/rpc/download/client.rs create mode 100644 ql-rpc/src/rpc/download/mod.rs create mode 100644 ql-rpc/src/rpc/download/server.rs create mode 100644 ql-rpc/src/rpc/duplex/client.rs create mode 100644 ql-rpc/src/rpc/duplex/codec.rs create mode 100644 ql-rpc/src/rpc/duplex/mod.rs create mode 100644 ql-rpc/src/rpc/duplex/server.rs create mode 100644 ql-rpc/src/rpc/mod.rs create mode 100644 ql-rpc/src/rpc/notification/client.rs create mode 100644 ql-rpc/src/rpc/notification/mod.rs create mode 100644 ql-rpc/src/rpc/notification/server.rs create mode 100644 ql-rpc/src/rpc/parts.rs create mode 100644 ql-rpc/src/rpc/progress/client.rs create mode 100644 ql-rpc/src/rpc/progress/codec.rs create mode 100644 ql-rpc/src/rpc/progress/mod.rs create mode 100644 ql-rpc/src/rpc/progress/server.rs create mode 100644 ql-rpc/src/rpc/request/client.rs create mode 100644 ql-rpc/src/rpc/request/mod.rs create mode 100644 ql-rpc/src/rpc/request/server.rs create mode 100644 ql-rpc/src/rpc/subscription/client.rs create mode 100644 ql-rpc/src/rpc/subscription/codec.rs create mode 100644 ql-rpc/src/rpc/subscription/mod.rs create mode 100644 ql-rpc/src/rpc/subscription/server.rs create mode 100644 ql-rpc/src/rpc/upload/client.rs create mode 100644 ql-rpc/src/rpc/upload/mod.rs create mode 100644 ql-rpc/src/rpc/upload/server.rs create mode 100644 ql-rpc/src/rpc/utils.rs create mode 100644 ql-rpc/src/stream.rs diff --git a/Cargo.lock b/Cargo.lock index 016c88c3..071d36b8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2099,6 +2099,14 @@ dependencies = [ "ql-wire", ] +[[package]] +name = "ql-rpc" +version = "0.1.0" +dependencies = [ + "bytes", + "trait-variant", +] + [[package]] name = "ql-wire" version = "0.1.0" @@ -2771,6 +2779,17 @@ dependencies = [ "slab", ] +[[package]] +name = "trait-variant" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "70977707304198400eb4835a78f6a9f928bf41bba420deb8fdb175cd965d77a7" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.106", +] + [[package]] name = "typenum" version = "1.18.0" diff --git a/Cargo.toml b/Cargo.toml index 84dde3e4..de8c80b7 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -5,6 +5,7 @@ members = [ "backup-shard", "btp", "ql-fsm", + "ql-rpc", "ql-wire", "quantum-link-macros", ] @@ -33,6 +34,7 @@ btp = { path = "btp" } foundation-api = { path = "api" } quantum-link-macros = { path = "quantum-link-macros" } ql-fsm = { path = "ql-fsm" } +ql-rpc = { path = "ql-rpc" } ql-wire = { path = "ql-wire" } [patch.crates-io] diff --git a/ql-rpc/Cargo.toml b/ql-rpc/Cargo.toml new file mode 100644 index 00000000..836e9d2b --- /dev/null +++ b/ql-rpc/Cargo.toml @@ -0,0 +1,10 @@ +[package] +name = "ql-rpc" +version = "0.1.0" +edition = "2021" +description = "QuantumLink RPC protocol" +license = "Proprietary" + +[dependencies] +bytes = { workspace = true } +trait-variant = { version = "0.1" } diff --git a/ql-rpc/src/chunk_queue.rs b/ql-rpc/src/chunk_queue.rs new file mode 100644 index 00000000..33f62998 --- /dev/null +++ b/ql-rpc/src/chunk_queue.rs @@ -0,0 +1,252 @@ +use std::collections::VecDeque; + +use bytes::{Buf, Bytes}; + +use crate::{CodecError, Error}; + +const LENGTH_SIZE: usize = 8; + +#[derive(Debug, Default)] +pub struct ChunkQueue { + chunks: VecDeque, + remaining: usize, +} + +impl ChunkQueue { + pub fn push(&mut self, chunk: Bytes) { + if chunk.is_empty() { + return; + } + self.remaining += chunk.len(); + self.chunks.push_back(chunk); + } + + pub fn remaining(&self) -> usize { + self.remaining + } + + pub fn expect_empty(&self) -> Result<(), CodecError> { + if self.remaining > 0 { + Err(CodecError::Rpc(Error::TrailingBytes)) + } else { + Ok(()) + } + } + + pub fn pop_front(&mut self, max_len: usize) -> Option { + let front = self.chunks.front_mut()?; + let chunk = if max_len >= front.len() { + self.chunks.pop_front().expect("buffered chunk is present") + } else { + front.split_to(max_len) + }; + self.remaining -= chunk.len(); + Some(chunk) + } + + pub fn pop_front_chunk(&mut self) -> Option { + self.pop_front(usize::MAX) + } + + pub fn try_take_part(&mut self) -> Result>, Error> { + let Some(len) = self.peek_next_part_len()? else { + return Ok(None); + }; + self.advance(LENGTH_SIZE); + Ok(Some(DrainBuf::new(self, len))) + } + + pub fn try_take_tagged_part(&mut self) -> Result)>, Error> { + let mut bytes = self.peek(); + let Ok(kind) = bytes.try_get_u8() else { + return Ok(None); + }; + let Some(len) = read_next_part_len(&mut bytes)? else { + return Ok(None); + }; + + self.advance(1 + LENGTH_SIZE); + Ok(Some((kind, DrainBuf::new(self, len)))) + } + + pub fn try_take_tagged_part_header(&mut self) -> Result, Error> { + let mut bytes = self.peek(); + let Ok(kind) = bytes.try_get_u8() else { + return Ok(None); + }; + let Some(len) = read_part_len_header(&mut bytes)? else { + return Ok(None); + }; + + self.advance(1 + LENGTH_SIZE); + Ok(Some((kind, len))) + } + + pub fn try_take_body(&mut self, len: usize) -> Option> { + if self.remaining < len { + return None; + } + + Some(DrainBuf::new(self, len)) + } + + fn peek_next_part_len(&self) -> Result, Error> { + let mut bytes = self.peek(); + read_next_part_len(&mut bytes) + } + + fn peek(&self) -> ChunkQueuePeek<'_> { + ChunkQueuePeek { + chunks: &self.chunks, + chunk_index: 0, + chunk_offset: 0, + remaining: self.remaining, + } + } + + fn front_chunk(&self, limit: usize) -> &[u8] { + let Some(chunk) = self.chunks.front() else { + return &[]; + }; + &chunk[..chunk.len().min(limit)] + } + + fn advance_inner(&mut self, mut cnt: usize) { + assert!(cnt <= self.remaining, "advanced past buffered data"); + self.remaining -= cnt; + while cnt > 0 { + let front = self.chunks.front_mut().expect("buffered data present"); + let consumed = cnt.min(front.len()); + front.advance(consumed); + cnt -= consumed; + if front.is_empty() { + self.chunks.pop_front(); + } + } + } +} + +struct ChunkQueuePeek<'a> { + chunks: &'a VecDeque, + chunk_index: usize, + chunk_offset: usize, + remaining: usize, +} + +impl Buf for ChunkQueuePeek<'_> { + fn remaining(&self) -> usize { + self.remaining + } + + fn chunk(&self) -> &[u8] { + if self.remaining == 0 { + return &[]; + } + + let Some(chunk) = self.chunks.get(self.chunk_index) else { + return &[]; + }; + &chunk[self.chunk_offset..] + } + + fn advance(&mut self, mut cnt: usize) { + assert!(cnt <= self.remaining, "advanced past buffered data"); + self.remaining -= cnt; + + while cnt > 0 { + let chunk = self + .chunks + .get(self.chunk_index) + .expect("buffered data present"); + let available = chunk.len() - self.chunk_offset; + let step = cnt.min(available); + self.chunk_offset += step; + cnt -= step; + if self.chunk_offset == chunk.len() { + self.chunk_index += 1; + self.chunk_offset = 0; + } + } + } +} + +impl Buf for ChunkQueue { + fn remaining(&self) -> usize { + self.remaining + } + + fn chunk(&self) -> &[u8] { + self.front_chunk(self.remaining) + } + + fn advance(&mut self, cnt: usize) { + assert!(cnt <= self.remaining, "advanced past buffered data"); + self.advance_inner(cnt); + } +} + +pub struct DrainBuf<'a> { + bytes: &'a mut ChunkQueue, + remaining: usize, +} + +impl<'a> DrainBuf<'a> { + pub fn new(bytes: &'a mut ChunkQueue, len: usize) -> Self { + debug_assert!(bytes.remaining() >= len); + Self { + bytes, + remaining: len, + } + } + + pub fn expect_empty(&self) -> Result<(), CodecError> { + if self.remaining > 0 { + Err(CodecError::Rpc(Error::TrailingBytes)) + } else { + Ok(()) + } + } +} + +impl Buf for DrainBuf<'_> { + fn remaining(&self) -> usize { + self.remaining + } + + fn chunk(&self) -> &[u8] { + self.bytes.front_chunk(self.remaining) + } + + fn advance(&mut self, cnt: usize) { + assert!(cnt <= self.remaining(), "advanced past payload boundary"); + self.bytes.advance_inner(cnt); + self.remaining -= cnt; + } +} + +impl Drop for DrainBuf<'_> { + fn drop(&mut self) { + if self.remaining > 0 { + self.bytes.advance_inner(self.remaining); + self.remaining = 0; + } + } +} + +fn read_next_part_len(bytes: &mut B) -> Result, Error> { + let Some(len) = read_part_len_header(bytes)? else { + return Ok(None); + }; + if bytes.remaining() < len { + return Ok(None); + } + Ok(Some(len)) +} + +fn read_part_len_header(bytes: &mut B) -> Result, Error> { + let Ok(len) = bytes.try_get_u64_le() else { + return Ok(None); + }; + let len: usize = len.try_into().map_err(|_| Error::LengthOverflow)?; + Ok(Some(len)) +} diff --git a/ql-rpc/src/codec.rs b/ql-rpc/src/codec.rs new file mode 100644 index 00000000..51da527b --- /dev/null +++ b/ql-rpc/src/codec.rs @@ -0,0 +1,83 @@ +use std::{convert::Infallible, str::Utf8Error}; + +use bytes::{Buf, BufMut, Bytes}; + +pub use crate::chunk_queue::ChunkQueue; + +pub trait RpcCodec: Sized { + type Error; + + fn encode_value(&self, out: &mut B); + fn decode_value(bytes: &mut B) -> Result; +} + +impl RpcCodec for String { + type Error = Utf8Error; + + fn encode_value(&self, out: &mut B) { + out.put_slice(self.as_bytes()); + } + + fn decode_value(bytes: &mut B) -> Result { + let len = bytes.remaining(); + if bytes.chunk().len() == len { + let s = std::str::from_utf8(bytes.chunk())?.to_owned(); + bytes.advance(len); + Ok(s) + } else { + let mut buf = vec![0; len]; + bytes.copy_to_slice(&mut buf); + String::from_utf8(buf).map_err(|err| err.utf8_error()) + } + } +} + +impl RpcCodec for Vec { + type Error = Infallible; + + fn encode_value(&self, out: &mut B) { + out.put_slice(self.as_slice()); + } + + fn decode_value(bytes: &mut B) -> Result { + let len = bytes.remaining(); + let mut buf = vec![0; len]; + bytes.copy_to_slice(&mut buf); + Ok(buf) + } +} + +impl RpcCodec for Bytes { + type Error = Infallible; + + fn encode_value(&self, out: &mut B) { + out.put_slice(self.as_ref()); + } + + fn decode_value(bytes: &mut B) -> Result { + Ok(bytes.copy_to_bytes(bytes.remaining())) + } +} + +const LENGTH_SIZE: usize = 8; + +pub fn encode_value_part>(value: &T, out: &mut B) { + let payload_start = reserve_length(out); + value.encode_value(out); + backpatch_length(out, payload_start); +} + +/// reads one length-delimited rpc value from buffered byte chunks +pub fn reserve_length>(out: &mut B) -> usize { + let start = out.as_mut().len(); + out.put_bytes(0, LENGTH_SIZE); + start +} + +pub fn backpatch_length + ?Sized>(out: &mut B, start: usize) { + let out = out.as_mut(); + let payload_start = start + LENGTH_SIZE; + let payload_len = out.len() - payload_start; + let payload_len = u64::try_from(payload_len).expect("rpc payload exceeds u64 length framing"); + out[start..payload_start].copy_from_slice(&payload_len.to_le_bytes()); +} diff --git a/ql-rpc/src/error.rs b/ql-rpc/src/error.rs new file mode 100644 index 00000000..7404a22e --- /dev/null +++ b/ql-rpc/src/error.rs @@ -0,0 +1,112 @@ +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Error { + Truncated, + LengthOverflow, + UnexpectedFrameKind(u8), + MissingResponse, + TrailingBytes, +} + +impl std::fmt::Display for Error { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Truncated => f.write_str("truncated rpc payload"), + Self::LengthOverflow => f.write_str("rpc payload length overflow"), + Self::UnexpectedFrameKind(kind) => write!(f, "unexpected rpc frame kind {kind}"), + Self::MissingResponse => f.write_str("missing terminal rpc response"), + Self::TrailingBytes => f.write_str("trailing rpc bytes"), + } + } +} + +impl std::error::Error for Error {} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum CodecError { + Rpc(Error), + Codec(E), +} + +impl std::error::Error for CodecError +where + E: std::error::Error + 'static, +{ + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + CodecError::Rpc(e) => Some(e), + CodecError::Codec(e) => Some(e), + } + } + + fn cause(&self) -> Option<&dyn std::error::Error> { + self.source() + } +} + +impl std::fmt::Display for CodecError +where + E: std::fmt::Display, +{ + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + CodecError::Rpc(e) => write!(f, "{e}"), + CodecError::Codec(e) => write!(f, "{e}"), + } + } +} + +impl From for CodecError { + fn from(error: Error) -> Self { + Self::Rpc(error) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum CallError { + Protocol(Error), + Codec(C), + Transport(T), +} + +impl std::fmt::Display for CallError +where + C: std::fmt::Display, + T: std::fmt::Display, +{ + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::Protocol(error) => write!(f, "{error}"), + Self::Codec(error) => write!(f, "{error}"), + Self::Transport(error) => write!(f, "{error}"), + } + } +} + +impl std::error::Error for CallError +where + C: std::error::Error + 'static, + T: std::error::Error + 'static, +{ + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + CallError::Protocol(error) => Some(error), + CallError::Codec(error) => Some(error), + CallError::Transport(error) => Some(error), + } + } +} + +impl From for CallError { + fn from(error: Error) -> Self { + Self::Protocol(error) + } +} + +impl From> for CallError { + fn from(error: CodecError) -> Self { + match error { + CodecError::Rpc(error) => Self::Protocol(error), + CodecError::Codec(error) => Self::Codec(error), + } + } +} diff --git a/ql-rpc/src/framed_value.rs b/ql-rpc/src/framed_value.rs new file mode 100644 index 00000000..600357da --- /dev/null +++ b/ql-rpc/src/framed_value.rs @@ -0,0 +1,127 @@ +use std::marker::PhantomData; + +use bytes::Bytes; + +use crate::{chunk_queue::ChunkQueue, CodecError, RpcCodec}; + +/// reads one length-delimited rpc value from buffered byte chunks +pub struct FramedReader { + bytes: ChunkQueue, + marker: PhantomData T>, +} + +pub enum FramedReadStep { + NeedMore(FramedReader), + Value(T), +} + +pub enum FramedPrefixStep { + NeedMore(FramedReader), + Value { value: T, bytes: ChunkQueue }, +} + +impl Default for FramedReader { + fn default() -> Self { + Self { + bytes: ChunkQueue::default(), + marker: PhantomData, + } + } +} + +impl FramedReader { + pub fn push(mut self, chunk: Bytes) -> Self { + self.bytes.push(chunk); + self + } + + pub fn advance(self) -> Result, CodecError> { + match self.advance_prefix()? { + FramedPrefixStep::NeedMore(next) => Ok(FramedReadStep::NeedMore(next)), + FramedPrefixStep::Value { value, bytes } => { + bytes.expect_empty()?; + Ok(FramedReadStep::Value(value)) + } + } + } + + pub fn advance_prefix(self) -> Result, CodecError> { + let mut this = self; + let Some(mut body) = this.bytes.try_take_part()? else { + return Ok(FramedPrefixStep::NeedMore(this)); + }; + + let value = T::decode_value(&mut body).map_err(CodecError::Codec)?; + drop(body); + Ok(FramedPrefixStep::Value { + value, + bytes: this.bytes, + }) + } +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use super::{FramedPrefixStep, FramedReadStep, FramedReader}; + use crate::codec::encode_value_part; + + #[test] + fn value_reader_round_trips_framed_values() { + let mut encoded = Vec::new(); + encode_value_part(&b"hello".to_vec(), &mut encoded); + + match FramedReader::>::default() + .push(Bytes::from(encoded)) + .advance() + .unwrap() + { + FramedReadStep::Value(value) => assert_eq!(value, b"hello".to_vec()), + _ => unreachable!(), + } + } + + #[test] + fn value_reader_waits_for_complete_frame() { + let mut encoded = Vec::new(); + encode_value_part(&b"hello".to_vec(), &mut encoded); + let encoded = Bytes::from(encoded); + + let reader = match FramedReader::>::default() + .push(encoded.slice(..4)) + .advance() + .unwrap() + { + FramedReadStep::NeedMore(next) => next, + _ => unreachable!(), + }; + + match reader.push(encoded.slice(4..)).advance().unwrap() { + FramedReadStep::Value(value) => assert_eq!(value, b"hello".to_vec()), + _ => unreachable!(), + } + } + + #[test] + fn value_reader_returns_prefix_remainder() { + let mut encoded = Vec::new(); + encode_value_part(&b"hello".to_vec(), &mut encoded); + encoded.extend_from_slice(b"tail"); + + match FramedReader::>::default() + .push(Bytes::from(encoded)) + .advance_prefix() + .unwrap() + { + FramedPrefixStep::Value { value, mut bytes } => { + assert_eq!(value, b"hello".to_vec()); + assert_eq!( + bytes.pop_front(usize::MAX), + Some(Bytes::from_static(b"tail")) + ); + } + _ => unreachable!(), + } + } +} diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs new file mode 100644 index 00000000..efea0250 --- /dev/null +++ b/ql-rpc/src/lib.rs @@ -0,0 +1,44 @@ +#![allow(clippy::type_complexity)] + +//! QuantumLink RPC protocol + +mod chunk_queue; +pub(crate) mod codec; +mod error; +mod framed_value; +mod route_id; +mod router; +mod rpc; +mod stream; + +pub use chunk_queue::ChunkQueue; +pub use codec::RpcCodec; +pub use error::*; +use framed_value::*; +pub use route_id::RouteId; +pub use router::*; +pub use rpc::*; +pub use stream::*; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct StreamCloseCode(pub u16); + +impl StreamCloseCode { + /// operation was cancelled + pub const CANCELLED: Self = Self(0); + /// local internal error + pub const INTERNAL: Self = Self(1); + /// request was refused + pub const REFUSED: Self = Self(2); + /// operation timed out + pub const TIMEOUT: Self = Self(3); + /// configured limit was exceeded + pub const LIMIT: Self = Self(4); + /// route identifier was unknown + pub const UNKNOWN_ROUTE: Self = Self(5); + + pub const fn into_inner(self) -> u16 { + self.0 + } +} diff --git a/ql-rpc/src/route_id.rs b/ql-rpc/src/route_id.rs new file mode 100644 index 00000000..1b054e74 --- /dev/null +++ b/ql-rpc/src/route_id.rs @@ -0,0 +1,19 @@ +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct RouteId(pub u32); + +impl RouteId { + pub const fn from_u32(value: u32) -> Self { + Self(value) + } + + pub const fn into_inner(self) -> u32 { + self.0 + } +} + +impl From for RouteId { + fn from(value: u32) -> Self { + Self::from_u32(value) + } +} diff --git a/ql-rpc/src/router/builder.rs b/ql-rpc/src/router/builder.rs new file mode 100644 index 00000000..b59a84e6 --- /dev/null +++ b/ql-rpc/src/router/builder.rs @@ -0,0 +1,354 @@ +use std::marker::PhantomData; + +use super::{ + LocalSpawner, RouteEntry, RouteFn, Router, RouterConfig, RpcStream, SendSpawner, Spawner, +}; +use crate::{ + download::{server::*, Download as DownloadRpc}, + duplex::{server::*, Duplex as DuplexRpc}, + notification::{server::*, Notification as NotificationRpc}, + progress::{server::*, Progress as ProgressRpc}, + request::{server::*, Request as RequestRpc}, + subscription::{server::*, Subscription as SubscriptionRpc}, + upload::{server::*, Upload as UploadRpc}, +}; + +pub struct LocalRoutes; +pub struct SendRoutes; + +pub struct RouterBuilder +where + Sp: Spawner, +{ + config: RouterConfig, + spawner: Sp, + routes: Vec>, + marker: PhantomData Mode>, +} + +impl RouterBuilder +where + Sp: Spawner, +{ + pub(crate) fn new(spawner: Sp) -> Self { + Self { + config: RouterConfig::default(), + spawner, + routes: Vec::new(), + marker: PhantomData, + } + } + + pub fn config(mut self, config: RouterConfig) -> Self { + self.config = config; + self + } + + pub fn max_request_bytes(mut self, max_request_bytes: usize) -> Self { + self.config.max_request_bytes = max_request_bytes; + self + } + + pub fn build(mut self, state: S) -> Router { + self.routes.sort_by_key(|entry| entry.route_id); + self.routes.shrink_to_fit(); + Router { + config: self.config, + state, + spawner: self.spawner, + routes: self.routes, + } + } + + fn add_route(mut self, route_id: crate::RouteId, route: RouteFn) -> Self { + if self.routes.iter().any(|entry| entry.route_id == route_id) { + panic!("duplicate rpc route {}", route_id.into_inner()); + } + self.routes.push(RouteEntry::new(route_id, route)); + self + } +} + +impl RouterBuilder +where + Sp: LocalSpawner, + St: RpcStream + 'static, +{ + pub fn request(self) -> Self + where + M: RequestRpc + 'static, + S: RequestHandlerLocal + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_request_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn notification(self) -> Self + where + M: NotificationRpc + 'static, + S: NotificationHandlerLocal + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_notification_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn duplex(self) -> Self + where + M: DuplexRpc + 'static, + S: DuplexHandlerLocal + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_duplex_inner::( + state, + config, + reader, + writer, + S::handle, + )) + }) + } + + pub fn download(self) -> Self + where + M: DownloadRpc + 'static, + S: DownloadHandlerLocal + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_download_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn subscription(self) -> Self + where + M: SubscriptionRpc + 'static, + S: SubscriptionHandlerLocal + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_subscription_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn progress(self) -> Self + where + M: ProgressRpc + 'static, + S: ProgressHandlerLocal + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_progress_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn upload(self) -> Self + where + M: UploadRpc + 'static, + S: UploadHandlerLocal + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_upload_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } +} + +impl RouterBuilder +where + Sp: SendSpawner + Send, + St: RpcStream + 'static, +{ + pub fn request(self) -> Self + where + M: RequestRpc + 'static, + M::Request: Send + 'static, + S: RequestHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_request_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn notification(self) -> Self + where + M: NotificationRpc + 'static, + M::Payload: Send + 'static, + S: NotificationHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_notification_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn duplex(self) -> Self + where + M: DuplexRpc + 'static, + M::InitiatorEvent: Send + 'static, + M::ResponderEvent: Send + 'static, + S: DuplexHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_duplex_inner::( + state, + config, + reader, + writer, + S::handle, + )) + }) + } + + pub fn download(self) -> Self + where + M: DownloadRpc + 'static, + M::Request: Send + 'static, + S: DownloadHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_download_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn subscription(self) -> Self + where + M: SubscriptionRpc + 'static, + M::Request: Send + 'static, + S: SubscriptionHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_subscription_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn progress(self) -> Self + where + M: ProgressRpc + 'static, + M::Request: Send + 'static, + S: ProgressHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_progress_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } + + pub fn upload(self) -> Self + where + M: UploadRpc + 'static, + M::Request: Send + 'static, + S: UploadHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + self.add_route(M::ROUTE, |spawner, state, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_upload_inner::( + state, + config, + reader, + writer, + S::handle, + S::handle_transport_error, + )) + }) + } +} diff --git a/ql-rpc/src/router/config.rs b/ql-rpc/src/router/config.rs new file mode 100644 index 00000000..d6fb048f --- /dev/null +++ b/ql-rpc/src/router/config.rs @@ -0,0 +1,12 @@ +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RouterConfig { + pub max_request_bytes: usize, +} + +impl Default for RouterConfig { + fn default() -> Self { + Self { + max_request_bytes: usize::MAX, + } + } +} diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs new file mode 100644 index 00000000..31e973ac --- /dev/null +++ b/ql-rpc/src/router/mod.rs @@ -0,0 +1,89 @@ +use crate::{RouteId, StreamCloseCode}; + +mod builder; +mod config; +mod mode; + +pub use self::{ + builder::{LocalRoutes, RouterBuilder, SendRoutes}, + config::RouterConfig, + mode::*, +}; +use crate::{close_stream, RpcStream}; +pub use crate::{ + download::{DownloadHandler, DownloadHandlerLocal, DownloadStart, DownloadWriter}, + duplex::{DuplexHandler, DuplexHandlerLocal, DuplexPeer}, + notification::{NotificationHandler, NotificationHandlerLocal}, + progress::{ProgressHandler, ProgressHandlerLocal, ProgressResponder}, + request::{RequestHandler, RequestHandlerLocal, Response}, + subscription::{SubscriptionHandler, SubscriptionHandlerLocal, SubscriptionResponder}, + upload::{UploadHandler, UploadHandlerLocal, UploadReader, UploadResponder}, +}; + +pub struct Router +where + Sp: Spawner, +{ + config: RouterConfig, + state: S, + spawner: Sp, + routes: Vec>, +} + +struct RouteEntry +where + Sp: Spawner, +{ + route_id: RouteId, + route: RouteFn, +} + +impl RouteEntry +where + Sp: Spawner, +{ + fn new(route_id: RouteId, route: RouteFn) -> Self { + Self { route_id, route } + } +} + +impl Router +where + S: Clone + 'static, + St: RpcStream, + Sp: Spawner, +{ + pub fn builder_local(spawner: Sp) -> RouterBuilder + where + Sp: LocalSpawner, + { + RouterBuilder::::new(spawner) + } + + pub fn builder_send(spawner: Sp) -> RouterBuilder + where + Sp: SendSpawner, + { + RouterBuilder::::new(spawner) + } + + pub fn handle(&self, stream: St) -> Option<(RouteId, Sp::Handle)> { + let route_id = stream.route_id()?; + let Ok(index) = self + .routes + .binary_search_by_key(&route_id, |entry| entry.route_id) + else { + close_stream(stream, StreamCloseCode::UNKNOWN_ROUTE); + return None; + }; + let route = self.routes[index].route; + Some(( + route_id, + route(&self.spawner, self.state.clone(), self.config, stream), + )) + } + + pub fn route_ids(&self) -> impl ExactSizeIterator + '_ { + self.routes.iter().map(|entry| entry.route_id) + } +} diff --git a/ql-rpc/src/router/mode.rs b/ql-rpc/src/router/mode.rs new file mode 100644 index 00000000..33b6c06a --- /dev/null +++ b/ql-rpc/src/router/mode.rs @@ -0,0 +1,21 @@ +use std::future::Future; + +use crate::RouterConfig; + +pub type RouteFn = fn(&Sp, S, RouterConfig, St) -> ::Handle; + +pub trait Spawner: Clone + 'static { + type Handle; +} + +pub trait LocalSpawner: Spawner { + fn spawn(&self, fut: F) -> Self::Handle + where + F: Future + 'static; +} + +pub trait SendSpawner: Spawner { + fn spawn(&self, fut: F) -> Self::Handle + where + F: Future + Send + 'static; +} diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs new file mode 100644 index 00000000..9a648181 --- /dev/null +++ b/ql-rpc/src/rpc/download/client.rs @@ -0,0 +1,246 @@ +use std::future::poll_fn; + +use bytes::{BufMut, Bytes}; + +use crate::{ + download::{Download, PartReadStep}, + rpc::parts::FrameKind, + CallError, FramedPrefixStep, FramedReader, RpcCodec, RpcRead, StreamCloseCode, +}; + +pub struct DownloadCall +where + M: Download, + R: RpcRead, +{ + stream: Option, + reader: Option>, +} + +pub struct DownloadPart<'a, M, R> +where + M: Download, + R: RpcRead, +{ + parent: &'a mut DownloadReader, + finished: bool, +} + +pub struct DownloadReader +where + M: Download, + R: RpcRead, +{ + stream: Option, + reader: crate::download::PartFrameReader, +} + +impl DownloadCall +where + M: Download, + R: RpcRead, +{ + pub fn new(stream: R) -> Self { + Self { + stream: Some(stream), + reader: Some(FramedReader::default()), + } + } + + pub async fn start( + mut self, + ) -> Result<(M::ResponseHeader, DownloadReader), CallError> { + loop { + let reader = self.reader.take().unwrap(); + let reader = match reader.advance_prefix() { + Ok(FramedPrefixStep::Value { value, bytes }) => { + let stream = self.stream.take().unwrap(); + return Ok(( + value, + DownloadReader { + stream: Some(stream), + reader: crate::download::PartFrameReader::::new(bytes), + }, + )); + } + Ok(FramedPrefixStep::NeedMore(next)) => next, + Err(error) => return Err(error.into()), + }; + + let stream = self.stream.as_mut().unwrap(); + match poll_fn(|cx| stream.poll_read(usize::MAX, cx)).await { + Ok(Some(chunk)) => { + self.reader = Some(reader.push(chunk)); + } + Ok(None) => return Err(crate::Error::Truncated.into()), + Err(error) => return Err(CallError::Transport(error)), + } + } + } + + pub fn close(mut self, code: StreamCloseCode) { + self.close_inner(code); + } + + fn close_inner(&mut self, code: StreamCloseCode) { + if let Some(stream) = self.stream.take() { + stream.close(code); + } + } +} + +impl Drop for DownloadCall +where + M: Download, + R: RpcRead, +{ + fn drop(&mut self) { + self.close_inner(StreamCloseCode::CANCELLED); + } +} + +impl DownloadReader +where + M: Download, + R: RpcRead, +{ + pub async fn next_part( + &mut self, + ) -> Result)>, CallError> + { + if self.stream.is_none() { + return Ok(None); + } + + match self.read_frame().await? { + PartReadStep::PartHeader(value) => Ok(Some(( + value, + DownloadPart { + parent: self, + finished: false, + }, + ))), + PartReadStep::Finish => { + self.stream.take(); + Ok(None) + } + PartReadStep::BodyBytes(_) => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::BodyChunk.tag()).into()) + } + PartReadStep::EndPart => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::EndPart.tag()).into()) + } + PartReadStep::NeedMore => unreachable!("read_frame waits for a complete frame"), + } + } + + pub async fn complete(mut self) -> Result<(), CallError> { + match self.read_frame().await? { + PartReadStep::Finish => { + self.stream.take(); + Ok(()) + } + PartReadStep::PartHeader(_) => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::PartHeader.tag()).into()) + } + PartReadStep::BodyBytes(_) => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::BodyChunk.tag()).into()) + } + PartReadStep::EndPart => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::EndPart.tag()).into()) + } + PartReadStep::NeedMore => unreachable!("read_frame waits for a complete frame"), + } + } + + pub fn close(mut self, code: StreamCloseCode) { + self.close_inner(code); + } + + async fn read_frame( + &mut self, + ) -> Result, CallError> { + loop { + match self.reader.advance() { + Ok(PartReadStep::NeedMore) => {} + Ok(step) => return Ok(step), + Err(error) => return Err(error.into()), + } + + let stream = self.stream.as_mut().unwrap(); + match poll_fn(|cx| stream.poll_read(usize::MAX, cx)).await { + Ok(Some(chunk)) => { + self.reader.push(chunk); + } + Ok(None) => return Err(crate::Error::Truncated.into()), + Err(error) => return Err(CallError::Transport(error)), + } + } + } + + fn close_inner(&mut self, code: StreamCloseCode) { + if let Some(stream) = self.stream.take() { + stream.close(code); + } + } +} + +impl Drop for DownloadReader +where + M: Download, + R: RpcRead, +{ + fn drop(&mut self) { + if self.stream.is_some() { + self.close_inner(StreamCloseCode::CANCELLED); + } + } +} + +impl DownloadPart<'_, M, R> +where + M: Download, + R: RpcRead, +{ + pub async fn read_chunk(&mut self) -> Result, CallError> { + if self.finished { + return Ok(None); + } + + match self.parent.read_frame().await? { + PartReadStep::BodyBytes(bytes) => Ok(Some(bytes)), + PartReadStep::EndPart => { + self.finished = true; + Ok(None) + } + PartReadStep::PartHeader(_) => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::PartHeader.tag()).into()) + } + PartReadStep::Finish => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::Finish.tag()).into()) + } + PartReadStep::NeedMore => unreachable!("read_frame waits for a complete frame"), + } + } + + pub fn close(mut self, code: StreamCloseCode) { + self.parent.close_inner(code); + self.finished = true; + } +} + +impl Drop for DownloadPart<'_, M, R> +where + M: Download, + R: RpcRead, +{ + fn drop(&mut self) { + if !self.finished { + self.parent.close_inner(StreamCloseCode::CANCELLED); + } + } +} + +pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { + request.encode_value(out) +} diff --git a/ql-rpc/src/rpc/download/mod.rs b/ql-rpc/src/rpc/download/mod.rs new file mode 100644 index 00000000..5ed34aed --- /dev/null +++ b/ql-rpc/src/rpc/download/mod.rs @@ -0,0 +1,31 @@ +use super::Route; +use crate::RpcCodec; + +pub(crate) mod client; +pub(crate) mod server; + +pub use client::{encode_request, DownloadCall, DownloadPart, DownloadReader}; +pub use server::{ + DownloadHandler, DownloadHandlerLocal, DownloadPartWriter, DownloadStart, DownloadWriter, +}; + +pub use crate::rpc::parts::{ + encode_body_chunk, encode_end_part, encode_finish, encode_part_header, PartFrameReader, + PartReadStep, +}; + +/// rpc where the responder returns metadata first and then zero or more byte parts +/// +/// the typed portion of the response ends at [`Self::ResponseHeader`] +/// after the header is decoded, the rest of the stream is exposed as typed +/// part headers followed by raw byte chunks through [`DownloadReader`] +pub trait Download: Route { + /// codec error shared by request and response header values + type Error; + /// typed input needed to start the download + type Request: RpcCodec; + /// typed metadata available before parts arrive + type ResponseHeader: RpcCodec; + /// typed metadata available before each byte part arrives + type PartHeader: RpcCodec; +} diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs new file mode 100644 index 00000000..fcdcb047 --- /dev/null +++ b/ql-rpc/src/rpc/download/server.rs @@ -0,0 +1,221 @@ +use std::{future::Future, marker::PhantomData}; + +use bytes::Bytes; + +use crate::{ + codec, + download::Download as DownloadRpc, + finish_bytes, + rpc::{ + parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, + read_eof_request, + }, + write_bytes, RouterConfig, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError, +}; + +#[trait_variant::make(DownloadHandler: Send)] +pub trait DownloadHandlerLocal +where + M: DownloadRpc, + St: RpcStream, +{ + async fn handle(self, message: M::Request, download: DownloadStart); + + fn handle_transport_error(&self, _error: &St::Error) {} +} + +pub struct DownloadStart +where + M: DownloadRpc, + W: RpcWrite, +{ + writer: Option, + marker: PhantomData M>, +} + +pub struct DownloadWriter +where + M: DownloadRpc, + W: RpcWrite, +{ + writer: Option, + marker: PhantomData M>, +} + +pub struct DownloadPartWriter<'a, M, W> +where + M: DownloadRpc, + W: RpcWrite, +{ + parent: &'a mut DownloadWriter, + finished: bool, +} + +impl DownloadStart +where + M: DownloadRpc, + W: RpcWrite, +{ + pub(crate) fn new(writer: W) -> Self { + Self { + writer: Some(writer), + marker: PhantomData, + } + } + + /// send the response header and begin streaming parts + pub async fn start( + mut self, + response_header: M::ResponseHeader, + ) -> Result, W::Error> { + let mut writer = self.writer.take().unwrap(); + let mut encoded = Vec::new(); + codec::encode_value_part(&response_header, &mut encoded); + write_bytes(&mut writer, Bytes::from(encoded)).await?; + Ok(DownloadWriter { + writer: Some(writer), + marker: PhantomData, + }) + } + + /// send a header-only response and finish the stream + pub async fn complete(mut self, response_header: M::ResponseHeader) -> Result<(), W::Error> { + let mut writer = self.writer.take().unwrap(); + let mut encoded = Vec::new(); + codec::encode_value_part(&response_header, &mut encoded); + encode_finish(&mut encoded); + write_bytes(&mut writer, Bytes::from(encoded)).await?; + finish_bytes(&mut writer).await + } + + /// close the stream with a transport code + pub fn close(mut self, code: StreamCloseCode) { + if let Some(writer) = self.writer.take() { + writer.close(code); + } + } +} + +impl Drop for DownloadStart +where + M: DownloadRpc, + W: RpcWrite, +{ + fn drop(&mut self) { + if let Some(writer) = self.writer.take() { + writer.close(StreamCloseCode::CANCELLED); + } + } +} + +impl DownloadWriter +where + M: DownloadRpc, + W: RpcWrite, +{ + pub async fn start_part( + &mut self, + part_header: M::PartHeader, + ) -> Result, W::Error> { + let writer = self.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + encode_part_header(&part_header, &mut encoded); + write_bytes(writer, Bytes::from(encoded)).await?; + Ok(DownloadPartWriter { + parent: self, + finished: false, + }) + } + + pub async fn finish(mut self) -> Result<(), W::Error> { + let mut writer = self.writer.take().unwrap(); + let mut encoded = Vec::new(); + encode_finish(&mut encoded); + write_bytes(&mut writer, Bytes::from(encoded)).await?; + finish_bytes(&mut writer).await + } + + pub fn close(mut self, code: StreamCloseCode) { + if let Some(writer) = self.writer.take() { + writer.close(code); + } + } +} + +impl Drop for DownloadWriter +where + M: DownloadRpc, + W: RpcWrite, +{ + fn drop(&mut self) { + if let Some(writer) = self.writer.take() { + writer.close(StreamCloseCode::CANCELLED); + } + } +} + +impl DownloadPartWriter<'_, M, W> +where + M: DownloadRpc, + W: RpcWrite, +{ + pub async fn send(&mut self, bytes: Bytes) -> Result<(), W::Error> { + let writer = self.parent.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + encode_body_chunk(&bytes, &mut encoded); + write_bytes(writer, Bytes::from(encoded)).await + } + + pub async fn finish(mut self) -> Result<(), W::Error> { + let writer = self.parent.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + encode_end_part(&mut encoded); + write_bytes(writer, Bytes::from(encoded)).await?; + self.finished = true; + Ok(()) + } +} + +impl Drop for DownloadPartWriter<'_, M, W> +where + M: DownloadRpc, + W: RpcWrite, +{ + fn drop(&mut self) { + if !self.finished { + if let Some(writer) = self.parent.writer.take() { + writer.close(StreamCloseCode::CANCELLED); + } + } + } +} + +pub(crate) async fn handle_download_inner( + state: S, + config: RouterConfig, + mut reader: St::Reader, + writer: St::Writer, + handle: H, + handle_transport_error: E, +) where + M: DownloadRpc + 'static, + St: RpcStream + 'static, + H: FnOnce(S, M::Request, DownloadStart) -> HF, + HF: Future, + E: FnOnce(&S, &St::Error), +{ + let request = match read_eof_request::(&mut reader, config).await { + Ok(request) => request, + Err(error) => { + let code = error.close_code(); + handle_transport_error(&state, &error); + if let Some(code) = code { + reader.close(code); + writer.close(code); + } + return; + } + }; + + handle(state, request, DownloadStart::new(writer)).await; +} diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs new file mode 100644 index 00000000..e76050a6 --- /dev/null +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -0,0 +1,164 @@ +use std::{ + future::poll_fn, + marker::PhantomData, + task::{Context, Poll}, +}; + +use bytes::Bytes; + +use crate::{ + duplex::{codec, Duplex, EventReader, ReadStep}, + finish_bytes, write_bytes, CallError, RpcCodec, RpcRead, RpcWrite, StreamCloseCode, +}; + +pub struct DuplexCall +where + M: Duplex, + W: RpcWrite, + R: RpcRead, +{ + pub sender: DuplexSender, + pub receiver: DuplexReceiver, +} + +pub struct DuplexSender +where + T: RpcCodec, + W: RpcWrite, +{ + writer: Option, + marker: PhantomData T>, +} + +pub struct DuplexReceiver +where + T: RpcCodec, + R: RpcRead, +{ + stream: Option, + reader: EventReader, +} + +impl DuplexSender +where + T: RpcCodec, + W: RpcWrite, +{ + pub fn new(writer: W) -> Self { + Self { + writer: Some(writer), + marker: PhantomData, + } + } + + pub async fn send(&mut self, event: &T) -> Result<(), W::Error> { + let writer = self.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + codec::encode_event(event, &mut encoded); + write_bytes(writer, Bytes::from(encoded)).await + } + + pub async fn finish(mut self) -> Result<(), W::Error> { + let mut writer = self.writer.take().unwrap(); + finish_bytes(&mut writer).await + } + + pub fn close(mut self, code: StreamCloseCode) { + if let Some(writer) = self.writer.take() { + writer.close(code); + } + } +} + +impl Drop for DuplexSender +where + T: RpcCodec, + W: RpcWrite, +{ + fn drop(&mut self) { + if let Some(writer) = self.writer.take() { + writer.close(StreamCloseCode::CANCELLED); + } + } +} + +impl DuplexReceiver +where + T: RpcCodec, + R: RpcRead, +{ + pub fn new(stream: R) -> Self { + Self { + stream: Some(stream), + reader: EventReader::default(), + } + } + + pub async fn next_event(&mut self) -> Option>> { + poll_fn(|cx| self.poll_next_event(cx)).await + } + + pub fn poll_next_event( + &mut self, + cx: &mut Context<'_>, + ) -> Poll>>> { + if self.stream.is_none() { + return Poll::Ready(None); + } + + loop { + match self.reader.advance() { + Ok(ReadStep::Event(value)) => return Poll::Ready(Some(Ok(value))), + Ok(ReadStep::NeedMore) => {} + Err(error) => { + self.stream.take(); + return Poll::Ready(Some(Err(error.into()))); + } + } + + let stream = self.stream.as_mut().unwrap(); + match stream.poll_read(usize::MAX, cx) { + Poll::Ready(Ok(Some(chunk))) => { + self.reader.push(chunk); + } + Poll::Ready(Ok(None)) => { + if self.reader.is_empty() { + self.stream.take(); + return Poll::Ready(None); + } + self.stream.take(); + return Poll::Ready(Some(Err(crate::Error::Truncated.into()))); + } + Poll::Ready(Err(error)) => { + self.stream.take(); + return Poll::Ready(Some(Err(CallError::Transport(error)))); + } + Poll::Pending => { + return Poll::Pending; + } + } + } + } + + pub fn close(mut self, code: StreamCloseCode) { + self.close_inner(code); + } + + fn close_inner(&mut self, code: StreamCloseCode) { + if let Some(stream) = self.stream.take() { + stream.close(code); + } + } +} + +impl Drop for DuplexReceiver +where + T: RpcCodec, + R: RpcRead, +{ + fn drop(&mut self) { + if self.stream.is_some() { + self.close_inner(StreamCloseCode::CANCELLED); + } + } +} diff --git a/ql-rpc/src/rpc/duplex/codec.rs b/ql-rpc/src/rpc/duplex/codec.rs new file mode 100644 index 00000000..68bc87c7 --- /dev/null +++ b/ql-rpc/src/rpc/duplex/codec.rs @@ -0,0 +1,86 @@ +use std::marker::PhantomData; + +use bytes::{BufMut, Bytes}; + +use crate::{codec, CodecError, RpcCodec}; + +pub fn encode_event(event: &T, out: &mut (impl BufMut + AsMut<[u8]>)) +where + T: RpcCodec, +{ + codec::encode_value_part(event, out) +} + +pub enum ReadStep { + NeedMore, + Event(T), +} + +pub struct EventReader { + bytes: codec::ChunkQueue, + marker: PhantomData T>, +} + +impl Default for EventReader { + fn default() -> Self { + Self { + bytes: codec::ChunkQueue::default(), + marker: PhantomData, + } + } +} + +impl EventReader { + pub fn push(&mut self, chunk: Bytes) { + self.bytes.push(chunk); + } + + pub fn is_empty(&self) -> bool { + self.bytes.remaining() == 0 + } + + pub fn advance(&mut self) -> Result, CodecError> { + let Some(mut body) = self.bytes.try_take_part()? else { + return Ok(ReadStep::NeedMore); + }; + + let value = { + let value = T::decode_value(&mut body).map_err(CodecError::Codec)?; + drop(body); + value + }; + Ok(ReadStep::Event(value)) + } +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use super::{encode_event, EventReader, ReadStep}; + + #[test] + fn event_reader_emits_multiple_events() { + let mut encoded = Vec::new(); + encode_event(&b"one".to_vec(), &mut encoded); + encode_event(&b"two".to_vec(), &mut encoded); + + let mut reader = EventReader::>::default(); + reader.push(Bytes::from(encoded)); + + match reader.advance().unwrap() { + ReadStep::Event(value) => { + assert_eq!(value, b"one".to_vec()); + } + _ => unreachable!(), + }; + + match reader.advance().unwrap() { + ReadStep::Event(value) => { + assert_eq!(value, b"two".to_vec()); + assert!(reader.is_empty()); + } + _ => unreachable!(), + } + } +} diff --git a/ql-rpc/src/rpc/duplex/mod.rs b/ql-rpc/src/rpc/duplex/mod.rs new file mode 100644 index 00000000..a9622029 --- /dev/null +++ b/ql-rpc/src/rpc/duplex/mod.rs @@ -0,0 +1,24 @@ +use super::Route; +use crate::RpcCodec; + +pub(crate) mod client; +pub(crate) mod codec; +pub(crate) mod server; + +pub use client::{DuplexCall, DuplexReceiver, DuplexSender}; +pub use codec::{encode_event, EventReader, ReadStep}; +pub use server::{DuplexHandler, DuplexHandlerLocal, DuplexPeer}; + +/// rpc where both sides exchange typed events on the same stream +/// +/// The initiator opens the routed stream. After that, either side may send any +/// number of events of its directional event type until it finishes or closes +/// its write side. +pub trait Duplex: Route { + /// codec error shared by both directional event values + type Error; + /// typed event sent by the side that opened the stream + type InitiatorEvent: RpcCodec; + /// typed event sent by the side handling the route + type ResponderEvent: RpcCodec; +} diff --git a/ql-rpc/src/rpc/duplex/server.rs b/ql-rpc/src/rpc/duplex/server.rs new file mode 100644 index 00000000..bf024335 --- /dev/null +++ b/ql-rpc/src/rpc/duplex/server.rs @@ -0,0 +1,47 @@ +use std::future::Future; + +use crate::{ + duplex::{Duplex, DuplexReceiver, DuplexSender}, + RpcRead, RpcStream, RpcWrite, +}; + +#[trait_variant::make(DuplexHandler: Send)] +pub trait DuplexHandlerLocal +where + M: Duplex, + St: RpcStream, +{ + async fn handle(self, peer: DuplexPeer); +} + +pub struct DuplexPeer +where + M: Duplex, + W: RpcWrite, + R: RpcRead, +{ + pub sender: DuplexSender, + pub receiver: DuplexReceiver, +} + +pub(crate) async fn handle_duplex_inner( + state: S, + _config: crate::RouterConfig, + reader: St::Reader, + writer: St::Writer, + handle: H, +) where + M: Duplex + 'static, + St: RpcStream + 'static, + H: FnOnce(S, DuplexPeer) -> HF, + HF: Future, +{ + handle( + state, + DuplexPeer { + sender: DuplexSender::new(writer), + receiver: DuplexReceiver::new(reader), + }, + ) + .await; +} diff --git a/ql-rpc/src/rpc/mod.rs b/ql-rpc/src/rpc/mod.rs new file mode 100644 index 00000000..2d84f050 --- /dev/null +++ b/ql-rpc/src/rpc/mod.rs @@ -0,0 +1,32 @@ +//! rpc protocol families built on top of one stream per call +//! +//! each trait in this module names one rpc shape and the typed values that +//! travel on that stream +//! route dispatch uses [`crate::RouteId`] and the submodules provide the matching +//! client and server helpers for encoding, decoding, and handler glue + +use crate::RouteId; + +pub mod download; +pub mod duplex; +pub mod notification; +pub(crate) mod parts; +pub mod progress; +pub mod request; +pub mod subscription; +pub mod upload; +mod utils; + +pub trait Route { + /// route used to dispatch this rpc family + const ROUTE: RouteId; +} + +pub use download::Download; +pub use duplex::Duplex; +pub use notification::Notification; +pub use progress::Progress; +pub use request::Request; +pub use subscription::Subscription; +pub use upload::Upload; +use utils::*; diff --git a/ql-rpc/src/rpc/notification/client.rs b/ql-rpc/src/rpc/notification/client.rs new file mode 100644 index 00000000..72b6900a --- /dev/null +++ b/ql-rpc/src/rpc/notification/client.rs @@ -0,0 +1,10 @@ +use bytes::BufMut; + +use crate::{notification::Notification, RpcCodec}; + +pub fn encode_notification( + payload: &M::Payload, + out: &mut (impl BufMut + AsMut<[u8]>), +) { + payload.encode_value(out) +} diff --git a/ql-rpc/src/rpc/notification/mod.rs b/ql-rpc/src/rpc/notification/mod.rs new file mode 100644 index 00000000..4740a64f --- /dev/null +++ b/ql-rpc/src/rpc/notification/mod.rs @@ -0,0 +1,19 @@ +use super::Route; +use crate::RpcCodec; + +pub(crate) mod client; +pub(crate) mod server; + +pub use client::encode_notification; +pub use server::{NotificationHandler, NotificationHandlerLocal}; + +/// one-way rpc that carries a single typed payload and no typed response +/// +/// the server reads [`Self::Payload`] to eof and then closes the response side +/// of the stream +pub trait Notification: Route { + /// codec error for the notification payload + type Error; + /// typed payload emitted by the caller + type Payload: RpcCodec; +} diff --git a/ql-rpc/src/rpc/notification/server.rs b/ql-rpc/src/rpc/notification/server.rs new file mode 100644 index 00000000..c9a4fdba --- /dev/null +++ b/ql-rpc/src/rpc/notification/server.rs @@ -0,0 +1,48 @@ +use std::future::Future; + +use crate::{ + notification::Notification as NotificationRpc, rpc::read_eof_request, RouterConfig, RpcRead, + RpcStream, RpcWrite, StreamCloseCode, StreamError, +}; + +#[trait_variant::make(NotificationHandler: Send)] +pub trait NotificationHandlerLocal +where + M: NotificationRpc, + St: RpcStream, +{ + async fn handle(self, message: M::Payload); + + fn handle_transport_error(&self, _error: &St::Error) {} +} + +pub(crate) async fn handle_notification_inner( + state: S, + config: RouterConfig, + mut reader: St::Reader, + writer: St::Writer, + handle: H, + handle_transport_error: E, +) where + M: NotificationRpc + 'static, + St: RpcStream + 'static, + H: FnOnce(S, M::Payload) -> HF, + HF: Future, + E: FnOnce(&S, &St::Error), +{ + let notification = match read_eof_request::(&mut reader, config).await { + Ok(notification) => notification, + Err(error) => { + let code = error.close_code(); + handle_transport_error(&state, &error); + if let Some(code) = code { + reader.close(code); + writer.close(code); + } + return; + } + }; + + writer.close(StreamCloseCode::CANCELLED); + handle(state, notification).await; +} diff --git a/ql-rpc/src/rpc/parts.rs b/ql-rpc/src/rpc/parts.rs new file mode 100644 index 00000000..47ff1e87 --- /dev/null +++ b/ql-rpc/src/rpc/parts.rs @@ -0,0 +1,283 @@ +use std::marker::PhantomData; + +use bytes::{BufMut, Bytes}; + +use crate::{codec, ChunkQueue, CodecError, RpcCodec}; + +pub enum PartReadStep { + NeedMore, + PartHeader(H), + BodyBytes(Bytes), + EndPart, + Finish, +} + +pub struct PartFrameReader { + bytes: codec::ChunkQueue, + pending_frame: PendingFrame, + marker: PhantomData H>, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum PendingFrame { + None, + Control { kind: FrameKind, len: usize }, + Body { remaining: usize }, +} + +impl PendingFrame { + fn take(&mut self) -> Self { + std::mem::replace(self, Self::None) + } +} + +impl PartFrameReader { + pub fn new(bytes: ChunkQueue) -> Self { + Self { + bytes, + pending_frame: PendingFrame::None, + marker: PhantomData, + } + } + + pub fn push(&mut self, chunk: Bytes) { + self.bytes.push(chunk); + } + + pub fn advance(&mut self) -> Result, CodecError> { + loop { + match self.pending_frame.take() { + PendingFrame::Body { remaining } => { + if remaining == 0 { + continue; + } + + let Some(bytes) = self.bytes.pop_front(remaining) else { + self.pending_frame = PendingFrame::Body { remaining }; + return Ok(PartReadStep::NeedMore); + }; + + let remaining = remaining - bytes.len(); + self.pending_frame = if remaining == 0 { + PendingFrame::None + } else { + PendingFrame::Body { remaining } + }; + return Ok(PartReadStep::BodyBytes(bytes)); + } + PendingFrame::Control { kind, len } => { + let Some(mut body) = self.bytes.try_take_body(len) else { + self.pending_frame = PendingFrame::Control { kind, len }; + return Ok(PartReadStep::NeedMore); + }; + + match kind { + FrameKind::PartHeader => { + let value = H::decode_value(&mut body).map_err(CodecError::Codec)?; + return Ok(PartReadStep::PartHeader(value)); + } + FrameKind::BodyChunk => unreachable!("body chunk is not a control frame"), + FrameKind::EndPart => { + body.expect_empty()?; + return Ok(PartReadStep::EndPart); + } + FrameKind::Finish => { + body.expect_empty()?; + drop(body); + self.bytes.expect_empty()?; + return Ok(PartReadStep::Finish); + } + } + } + PendingFrame::None => { + let Some((kind, len)) = self + .bytes + .try_take_tagged_part_header() + .map_err(CodecError::Rpc)? + else { + return Ok(PartReadStep::NeedMore); + }; + + let kind = FrameKind::try_from(kind).map_err(CodecError::Rpc)?; + self.pending_frame = if kind == FrameKind::BodyChunk { + PendingFrame::Body { remaining: len } + } else { + PendingFrame::Control { kind, len } + }; + } + } + } + } +} + +pub fn encode_part_header(part_header: &H, out: &mut (impl BufMut + AsMut<[u8]>)) { + encode_tagged_value_part(FrameKind::PartHeader, part_header, out) +} + +pub fn encode_body_chunk(bytes: &Bytes, out: &mut (impl BufMut + AsMut<[u8]>)) { + encode_tagged_value_part(FrameKind::BodyChunk, bytes, out) +} + +pub fn encode_end_part(out: &mut (impl BufMut + AsMut<[u8]>)) { + encode_tagged_empty_part(FrameKind::EndPart, out) +} + +pub fn encode_finish(out: &mut (impl BufMut + AsMut<[u8]>)) { + encode_tagged_empty_part(FrameKind::Finish, out) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +pub(super) enum FrameKind { + PartHeader = 1, + BodyChunk = 2, + EndPart = 3, + Finish = 4, +} + +impl FrameKind { + pub fn tag(self) -> u8 { + self as u8 + } +} + +impl TryFrom for FrameKind { + type Error = crate::Error; + + fn try_from(value: u8) -> Result { + match value { + x if x == Self::PartHeader.tag() => Ok(Self::PartHeader), + x if x == Self::BodyChunk.tag() => Ok(Self::BodyChunk), + x if x == Self::EndPart.tag() => Ok(Self::EndPart), + x if x == Self::Finish.tag() => Ok(Self::Finish), + other => Err(crate::Error::UnexpectedFrameKind(other)), + } + } +} + +fn encode_tagged_value_part>( + kind: FrameKind, + value: &T, + out: &mut B, +) { + out.put_u8(kind.tag()); + let payload_start = codec::reserve_length(out); + value.encode_value(out); + codec::backpatch_length(out, payload_start); +} + +fn encode_tagged_empty_part>(kind: FrameKind, out: &mut B) { + out.put_u8(kind.tag()); + let payload_start = codec::reserve_length(out); + codec::backpatch_length(out, payload_start); +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use super::{ + encode_body_chunk, encode_end_part, encode_finish, encode_part_header, PartFrameReader, + PartReadStep, + }; + + #[test] + fn part_reader_emits_multipart_sequence() { + let mut encoded = Vec::new(); + encode_part_header(&b"a.txt".to_vec(), &mut encoded); + encode_body_chunk(&Bytes::from_static(b"hel"), &mut encoded); + encode_body_chunk(&Bytes::from_static(b"lo"), &mut encoded); + encode_end_part(&mut encoded); + encode_part_header(&b"b.txt".to_vec(), &mut encoded); + encode_end_part(&mut encoded); + encode_finish(&mut encoded); + + let mut reader = PartFrameReader::>::new(Default::default()); + reader.push(Bytes::from(encoded)); + + match reader.advance().unwrap() { + PartReadStep::PartHeader(value) => { + assert_eq!(value, b"a.txt".to_vec()); + } + _ => unreachable!(), + }; + + match reader.advance().unwrap() { + PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"hel")), + _ => unreachable!(), + }; + + match reader.advance().unwrap() { + PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"lo")), + _ => unreachable!(), + }; + + match reader.advance().unwrap() { + PartReadStep::EndPart => {} + _ => unreachable!(), + }; + + match reader.advance().unwrap() { + PartReadStep::PartHeader(value) => { + assert_eq!(value, b"b.txt".to_vec()); + } + _ => unreachable!(), + }; + + match reader.advance().unwrap() { + PartReadStep::EndPart => {} + _ => unreachable!(), + }; + + match reader.advance().unwrap() { + PartReadStep::Finish => {} + _ => unreachable!(), + } + } + + #[test] + fn part_reader_waits_for_complete_header_frame() { + let mut encoded = Vec::new(); + encode_part_header(&b"a.txt".to_vec(), &mut encoded); + let encoded = Bytes::from(encoded); + + let mut reader = PartFrameReader::>::new(Default::default()); + reader.push(encoded.slice(..4)); + match reader.advance().unwrap() { + PartReadStep::NeedMore => {} + _ => unreachable!(), + }; + + reader.push(encoded.slice(4..)); + match reader.advance().unwrap() { + PartReadStep::PartHeader(value) => assert_eq!(value, b"a.txt".to_vec()), + _ => unreachable!(), + } + } + + #[test] + fn body_chunk_frame_streams_after_header() { + let mut encoded = Vec::new(); + encode_body_chunk(&Bytes::from_static(b"hello"), &mut encoded); + let encoded = Bytes::from(encoded); + + let mut reader = PartFrameReader::>::new(Default::default()); + reader.push(encoded.slice(..9)); + match reader.advance().unwrap() { + PartReadStep::NeedMore => {} + _ => unreachable!(), + }; + + reader.push(encoded.slice(9..11)); + match reader.advance().unwrap() { + PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"he")), + _ => unreachable!(), + }; + + reader.push(encoded.slice(11..)); + match reader.advance().unwrap() { + PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"llo")), + _ => unreachable!(), + }; + } +} diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs new file mode 100644 index 00000000..c2218c97 --- /dev/null +++ b/ql-rpc/src/rpc/progress/client.rs @@ -0,0 +1,152 @@ +use std::{ + future::{poll_fn, Future}, + pin::Pin, + task::{Context, Poll}, +}; + +use crate::{ + progress::{Progress, ReadStep, ResponseReader}, + CallError, Error, RpcRead, StreamCloseCode, +}; + +pub struct ProgressCall +where + M: Progress, + R: RpcRead, +{ + stream: Option, + state: State, +} + +enum State +where + M: Progress, +{ + Invalid, + Reading(ResponseReader), + Terminal(Result>), + Done, +} + +impl Unpin for ProgressCall +where + M: Progress, + R: RpcRead, +{ +} + +impl ProgressCall +where + M: Progress, + R: RpcRead, +{ + pub fn new(stream: R) -> Self { + Self { + stream: Some(stream), + state: State::Reading(ResponseReader::default()), + } + } + + pub async fn next_progress(&mut self) -> Option { + poll_fn(|cx| self.poll_next_progress(cx)).await + } + + fn poll_step(&mut self, cx: &mut Context<'_>) -> Poll> { + loop { + let reader = match &mut self.state { + State::Reading(reader) => reader, + State::Terminal(_) | State::Done => return Poll::Ready(None), + State::Invalid => panic!("invalid state"), + }; + + match reader.advance() { + Ok(ReadStep::Progress(value)) => return Poll::Ready(Some(value)), + Ok(ReadStep::Response(response)) => { + self.state = State::Terminal(Ok(response)); + return Poll::Ready(None); + } + Ok(ReadStep::NeedMore) => {} + Err(error) => { + self.state = State::Terminal(Err(error.into())); + return Poll::Ready(None); + } + } + + let stream = self.stream.as_mut().unwrap(); + match stream.poll_read(usize::MAX, cx) { + Poll::Ready(Ok(Some(chunk))) => { + let State::Reading(reader) = &mut self.state else { + panic!("invalid state"); + }; + reader.push(chunk); + } + Poll::Ready(Ok(None)) => { + self.state = State::Terminal(Err(Error::MissingResponse.into())); + return Poll::Ready(None); + } + Poll::Ready(Err(error)) => { + self.state = State::Terminal(Err(CallError::Transport(error))); + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } + + pub fn poll_next_progress(&mut self, cx: &mut Context<'_>) -> Poll> { + self.poll_step(cx) + } + + pub fn close(mut self, code: StreamCloseCode) { + self.close_inner(code); + } + + fn close_inner(&mut self, code: StreamCloseCode) { + self.state = State::Done; + if let Some(stream) = self.stream.take() { + stream.close(code); + } + } +} + +impl Drop for ProgressCall +where + M: Progress, + R: RpcRead, +{ + fn drop(&mut self) { + if matches!(self.state, State::Reading(_)) { + self.close_inner(StreamCloseCode::CANCELLED); + } + } +} + +impl Future for ProgressCall +where + M: Progress, + R: RpcRead, +{ + type Output = Result>; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + let this = self.get_mut(); + + loop { + match this.poll_step(cx) { + Poll::Ready(Some(_)) => {} + Poll::Ready(None) => match std::mem::replace(&mut this.state, State::Invalid) { + State::Terminal(result) => { + this.state = State::Done; + return Poll::Ready(result); + } + State::Done => panic!("polled after completion"), + State::Invalid => panic!("polled during state transition"), + State::Reading(_) => { + panic!("progress call reached terminal step without result") + } + }, + Poll::Pending => return Poll::Pending, + } + } + } +} diff --git a/ql-rpc/src/rpc/progress/codec.rs b/ql-rpc/src/rpc/progress/codec.rs new file mode 100644 index 00000000..a0dc1b8c --- /dev/null +++ b/ql-rpc/src/rpc/progress/codec.rs @@ -0,0 +1,145 @@ +use std::marker::PhantomData; + +use bytes::{BufMut, Bytes}; + +use crate::{codec, progress::Progress, CodecError, Error, RpcCodec}; + +pub enum ReadStep { + NeedMore, + Progress(M::Progress), + Response(M::Response), +} + +pub struct ResponseReader { + bytes: codec::ChunkQueue, + marker: PhantomData M>, +} + +impl Default for ResponseReader { + fn default() -> Self { + Self { + bytes: codec::ChunkQueue::default(), + marker: PhantomData, + } + } +} + +impl ResponseReader { + pub fn push(&mut self, chunk: Bytes) { + self.bytes.push(chunk); + } + + pub fn advance(&mut self) -> Result, CodecError> { + let Some((kind, mut body)) = self.bytes.try_take_tagged_part().map_err(CodecError::Rpc)? + else { + return Ok(ReadStep::NeedMore); + }; + + match kind { + x if x == FrameKind::Progress as u8 => { + let value = { + let value = M::Progress::decode_value(&mut body).map_err(CodecError::Codec)?; + drop(body); + value + }; + Ok(ReadStep::Progress(value)) + } + x if x == FrameKind::Response as u8 => { + let response = M::Response::decode_value(&mut body).map_err(CodecError::Codec)?; + drop(body); + if self.bytes.remaining() > 0 { + Err(CodecError::Rpc(Error::TrailingBytes)) + } else { + Ok(ReadStep::Response(response)) + } + } + other => Err(CodecError::Rpc(Error::UnexpectedFrameKind(other))), + } + } +} + +pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { + codec::encode_value_part(request, out) +} + +pub fn encode_progress(progress: &M::Progress, out: &mut (impl BufMut + AsMut<[u8]>)) { + encode_tagged_value_part(FrameKind::Progress, progress, out) +} + +pub fn encode_response(response: &M::Response, out: &mut (impl BufMut + AsMut<[u8]>)) { + encode_tagged_value_part(FrameKind::Response, response, out) +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +enum FrameKind { + Progress = 1, + Response = 2, +} + +fn encode_tagged_value_part>( + kind: FrameKind, + value: &T, + out: &mut B, +) { + out.put_u8(kind as u8); + let payload_start = codec::reserve_length(out); + value.encode_value(out); + codec::backpatch_length(out, payload_start); +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use super::{encode_progress, encode_response, ReadStep, ResponseReader}; + use crate::{progress::Progress, Route, RouteId}; + + struct Watch; + + impl Route for Watch { + const ROUTE: RouteId = RouteId::from_u32(11); + } + + impl Progress for Watch { + type Error = core::convert::Infallible; + type Request = Vec; + type Progress = Vec; + type Response = Vec; + } + + #[test] + fn response_reader_emits_progress_then_response() { + let mut encoded = Vec::new(); + encode_progress::(&b"10%".to_vec(), &mut encoded); + encode_response::(&b"done".to_vec(), &mut encoded); + + let mut reader = ResponseReader::::default(); + reader.push(Bytes::from(encoded)); + + match reader.advance().unwrap() { + ReadStep::Progress(value) => { + assert_eq!(value, b"10%".to_vec()); + } + _ => unreachable!(), + }; + match reader.advance().unwrap() { + ReadStep::Response(value) => assert_eq!(value, b"done".to_vec()), + _ => unreachable!(), + } + } + + #[test] + fn response_reader_handles_response_only() { + let mut encoded = Vec::new(); + encode_response::(&b"done".to_vec(), &mut encoded); + + let mut reader = ResponseReader::::default(); + reader.push(Bytes::from(encoded)); + + match reader.advance().unwrap() { + ReadStep::Response(value) => assert_eq!(value, b"done".to_vec()), + _ => unreachable!(), + } + } +} diff --git a/ql-rpc/src/rpc/progress/mod.rs b/ql-rpc/src/rpc/progress/mod.rs new file mode 100644 index 00000000..b21c826d --- /dev/null +++ b/ql-rpc/src/rpc/progress/mod.rs @@ -0,0 +1,27 @@ +use super::Route; +use crate::RpcCodec; + +pub(crate) mod client; +pub(crate) mod codec; +pub(crate) mod server; + +pub use client::ProgressCall; +pub use codec::{encode_progress, encode_request, encode_response, ReadStep, ResponseReader}; +pub use server::{ProgressHandler, ProgressHandlerLocal, ProgressResponder}; + +/// rpc where the responder streams progress values before a final response +/// +/// the request is length-delimited +/// response frames are tagged so the client can distinguish +/// [`Self::Progress`] items from the final [`Self::Response`] +/// reaching eof before the final response is an error +pub trait Progress: Route { + /// codec error shared by request, progress, and response values + type Error; + /// typed input sent by the caller + type Request: RpcCodec; + /// typed progress item emitted before completion + type Progress: RpcCodec; + /// typed terminal response that completes the call + type Response: RpcCodec; +} diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs new file mode 100644 index 00000000..b94421cf --- /dev/null +++ b/ql-rpc/src/rpc/progress/server.rs @@ -0,0 +1,106 @@ +use std::{future::Future, marker::PhantomData}; + +use bytes::Bytes; + +use crate::{ + finish_bytes, + progress::{encode_progress, encode_response, Progress}, + rpc::read_framed_request, + write_bytes, RouterConfig, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError, +}; + +#[trait_variant::make(ProgressHandler: Send)] +pub trait ProgressHandlerLocal +where + M: Progress, + St: RpcStream, +{ + async fn handle(self, request: M::Request, responder: ProgressResponder); + + fn handle_transport_error(&self, _error: &St::Error) {} +} + +pub struct ProgressResponder +where + M: Progress, + W: RpcWrite, +{ + writer: Option, + marker: PhantomData M>, +} + +impl ProgressResponder +where + M: Progress, + W: RpcWrite, +{ + pub(crate) fn new(writer: W) -> Self { + Self { + writer: Some(writer), + marker: PhantomData, + } + } + + pub async fn send(&mut self, progress: M::Progress) -> Result<(), W::Error> { + let writer = self.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + encode_progress::(&progress, &mut encoded); + write_bytes(writer, Bytes::from(encoded)).await + } + + pub async fn finish(mut self, response: M::Response) -> Result<(), W::Error> { + let mut writer = self.writer.take().unwrap(); + let mut encoded = Vec::new(); + encode_response::(&response, &mut encoded); + write_bytes(&mut writer, Bytes::from(encoded)).await?; + finish_bytes(&mut writer).await + } + + pub fn close(mut self, code: StreamCloseCode) { + if let Some(writer) = self.writer.take() { + writer.close(code); + } + } +} + +impl Drop for ProgressResponder +where + M: Progress, + W: RpcWrite, +{ + fn drop(&mut self) { + if let Some(writer) = self.writer.take() { + writer.close(StreamCloseCode::CANCELLED); + } + } +} + +pub(crate) async fn handle_progress_inner( + state: S, + config: RouterConfig, + mut reader: St::Reader, + writer: St::Writer, + handle: H, + handle_transport_error: E, +) where + M: Progress + 'static, + St: RpcStream + 'static, + H: FnOnce(S, M::Request, ProgressResponder) -> HF, + HF: Future, + E: FnOnce(&S, &St::Error), +{ + let request = match read_framed_request::(&mut reader, config).await { + Ok(request) => request, + Err(error) => { + let code = error.close_code(); + handle_transport_error(&state, &error); + if let Some(code) = code { + reader.close(code); + writer.close(code); + } + return; + } + }; + + handle(state, request, ProgressResponder::new(writer)).await; +} diff --git a/ql-rpc/src/rpc/request/client.rs b/ql-rpc/src/rpc/request/client.rs new file mode 100644 index 00000000..e7ffb845 --- /dev/null +++ b/ql-rpc/src/rpc/request/client.rs @@ -0,0 +1,34 @@ +use bytes::BufMut; + +use crate::{read_bytes, request::Request, CallError, ChunkQueue, RpcCodec, RpcRead}; + +pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { + request.encode_value(out) +} + +pub fn encode_response(response: &M::Response, out: &mut (impl BufMut + AsMut<[u8]>)) { + response.encode_value(out) +} + +pub async fn read_response( + mut reader: R, +) -> Result> +where + M: Request, + R: RpcRead, +{ + let mut bytes = ChunkQueue::default(); + + while let Some(chunk) = read_bytes(&mut reader, usize::MAX) + .await + .map_err(CallError::Transport)? + { + bytes.push(chunk); + } + + let value = M::Response::decode_value(&mut bytes).map_err(CallError::Codec)?; + if bytes.remaining() > 0 { + return Err(crate::Error::TrailingBytes.into()); + } + Ok(value) +} diff --git a/ql-rpc/src/rpc/request/mod.rs b/ql-rpc/src/rpc/request/mod.rs new file mode 100644 index 00000000..adf32597 --- /dev/null +++ b/ql-rpc/src/rpc/request/mod.rs @@ -0,0 +1,23 @@ +use super::Route; +use crate::RpcCodec; + +pub(crate) mod client; +pub(crate) mod server; + +pub use client::{encode_request, encode_response, read_response}; +pub use server::{RequestHandler, RequestHandlerLocal, Response}; + +/// request-response rpc with exactly one typed value in each direction +/// +/// the request is read to eof on the server side, so callers must finish the +/// request stream after encoding [`Self::Request`] +/// the response is also read to eof and rejects trailing bytes after +/// [`Self::Response`] +pub trait Request: Route { + /// codec error shared by request and response values + type Error; + /// typed input sent by the caller + type Request: RpcCodec; + /// typed output returned by the responder + type Response: RpcCodec; +} diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs new file mode 100644 index 00000000..5211cce2 --- /dev/null +++ b/ql-rpc/src/rpc/request/server.rs @@ -0,0 +1,96 @@ +use std::{future::Future, marker::PhantomData}; + +use bytes::Bytes; + +use crate::{ + finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, RouterConfig, + RpcCodec, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError, +}; + +#[trait_variant::make(RequestHandler: Send)] +pub trait RequestHandlerLocal +where + M: RequestRpc, + St: RpcStream, +{ + async fn handle(self, message: M::Request, responder: Response); + + fn handle_transport_error(&self, _error: &St::Error) {} +} + +pub struct Response +where + W: RpcWrite, +{ + writer: Option, + marker: PhantomData T>, +} + +impl Response +where + T: RpcCodec, + W: RpcWrite, +{ + pub(crate) fn new(writer: W) -> Self { + Self { + writer: Some(writer), + marker: PhantomData, + } + } + + pub async fn respond(mut self, response: T) -> Result<(), W::Error> { + let mut writer = self.writer.take().unwrap(); + let mut encoded = Vec::new(); + response.encode_value(&mut encoded); + write_bytes(&mut writer, Bytes::from(encoded)).await?; + finish_bytes(&mut writer).await?; + Ok(()) + } + + pub fn close(mut self, code: StreamCloseCode) { + if let Some(writer) = self.writer.take() { + writer.close(code); + } + } +} + +impl Drop for Response +where + W: RpcWrite, +{ + fn drop(&mut self) { + if let Some(writer) = self.writer.take() { + writer.close(StreamCloseCode::CANCELLED); + } + } +} + +pub(crate) async fn handle_request_inner( + state: S, + config: RouterConfig, + mut reader: St::Reader, + writer: St::Writer, + handle: H, + handle_transport_error: E, +) where + M: RequestRpc + 'static, + St: RpcStream + 'static, + H: FnOnce(S, M::Request, Response) -> HF, + HF: Future, + E: FnOnce(&S, &St::Error), +{ + let request = match read_eof_request::(&mut reader, config).await { + Ok(request) => request, + Err(error) => { + let code = error.close_code(); + handle_transport_error(&state, &error); + if let Some(code) = code { + reader.close(code); + writer.close(code); + } + return; + } + }; + + handle(state, request, Response::new(writer)).await; +} diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs new file mode 100644 index 00000000..fe6aa5b1 --- /dev/null +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -0,0 +1,99 @@ +use std::{ + future::poll_fn, + task::{Context, Poll}, +}; + +use crate::{ + subscription::{ReadStep, ResponseReader, Subscription}, + CallError, RpcRead, StreamCloseCode, +}; + +pub struct SubscriptionCall +where + M: Subscription, + R: RpcRead, +{ + stream: Option, + reader: ResponseReader, +} + +impl SubscriptionCall +where + M: Subscription, + R: RpcRead, +{ + pub fn new(stream: R) -> Self { + Self { + stream: Some(stream), + reader: ResponseReader::default(), + } + } + + pub async fn next_event(&mut self) -> Option>> { + poll_fn(|cx| self.poll_next_event(cx)).await + } + + pub fn poll_next_event( + &mut self, + cx: &mut Context<'_>, + ) -> Poll>>> { + if self.stream.is_none() { + return Poll::Ready(None); + } + + loop { + match self.reader.advance() { + Ok(ReadStep::Item(value)) => return Poll::Ready(Some(Ok(value))), + Ok(ReadStep::NeedMore) => {} + Err(error) => { + self.stream.take(); + return Poll::Ready(Some(Err(error.into()))); + } + } + + let stream = self.stream.as_mut().unwrap(); + match stream.poll_read(usize::MAX, cx) { + Poll::Ready(Ok(Some(chunk))) => { + self.reader.push(chunk); + } + Poll::Ready(Ok(None)) => { + if self.reader.is_empty() { + self.stream.take(); + return Poll::Ready(None); + } + self.stream.take(); + return Poll::Ready(Some(Err(crate::Error::Truncated.into()))); + } + Poll::Ready(Err(error)) => { + self.stream.take(); + return Poll::Ready(Some(Err(CallError::Transport(error)))); + } + Poll::Pending => { + return Poll::Pending; + } + } + } + } + + pub fn close(mut self, code: StreamCloseCode) { + self.close_inner(code); + } + + fn close_inner(&mut self, code: StreamCloseCode) { + if let Some(stream) = self.stream.take() { + stream.close(code); + } + } +} + +impl Drop for SubscriptionCall +where + M: Subscription, + R: RpcRead, +{ + fn drop(&mut self) { + if self.stream.is_some() { + self.close_inner(StreamCloseCode::CANCELLED); + } + } +} diff --git a/ql-rpc/src/rpc/subscription/codec.rs b/ql-rpc/src/rpc/subscription/codec.rs new file mode 100644 index 00000000..bdd16209 --- /dev/null +++ b/ql-rpc/src/rpc/subscription/codec.rs @@ -0,0 +1,58 @@ +use std::marker::PhantomData; + +use bytes::{BufMut, Bytes}; + +use crate::{codec, subscription::Subscription, CodecError, RpcCodec}; + +pub fn encode_request( + request: &M::Request, + out: &mut (impl BufMut + AsMut<[u8]>), +) { + request.encode_value(out) +} + +pub fn encode_item(item: &M::Event, out: &mut (impl BufMut + AsMut<[u8]>)) { + codec::encode_value_part(item, out) +} + +pub enum ReadStep { + NeedMore, + Item(M::Event), +} + +pub struct ResponseReader { + bytes: codec::ChunkQueue, + marker: PhantomData M>, +} + +impl Default for ResponseReader { + fn default() -> Self { + Self { + bytes: codec::ChunkQueue::default(), + marker: PhantomData, + } + } +} + +impl ResponseReader { + pub fn push(&mut self, chunk: Bytes) { + self.bytes.push(chunk); + } + + pub fn is_empty(&self) -> bool { + self.bytes.remaining() == 0 + } + + pub fn advance(&mut self) -> Result, CodecError> { + let Some(mut body) = self.bytes.try_take_part()? else { + return Ok(ReadStep::NeedMore); + }; + + let item = { + let item = M::Event::decode_value(&mut body).map_err(CodecError::Codec)?; + drop(body); + item + }; + Ok(ReadStep::Item(item)) + } +} diff --git a/ql-rpc/src/rpc/subscription/mod.rs b/ql-rpc/src/rpc/subscription/mod.rs new file mode 100644 index 00000000..672eb9bc --- /dev/null +++ b/ql-rpc/src/rpc/subscription/mod.rs @@ -0,0 +1,23 @@ +use super::Route; +use crate::RpcCodec; + +pub(crate) mod client; +pub(crate) mod codec; +pub(crate) mod server; + +pub use client::SubscriptionCall; +pub use codec::{encode_item, encode_request, ReadStep, ResponseReader}; +pub use server::{SubscriptionHandler, SubscriptionHandlerLocal, SubscriptionResponder}; + +/// rpc where one request opens a stream of typed events +/// +/// event frames are length-delimited and the stream ends cleanly at eof +/// any partial trailing frame is reported as truncation on the client side +pub trait Subscription: Route { + /// codec error shared by request and event values + type Error; + /// typed input that starts the subscription + type Request: RpcCodec; + /// typed event yielded by the responder + type Event: RpcCodec; +} diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs new file mode 100644 index 00000000..6dfdd4b0 --- /dev/null +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -0,0 +1,105 @@ +use std::{future::Future, marker::PhantomData}; + +use bytes::Bytes; + +use crate::{ + codec, finish_bytes, rpc::read_eof_request, subscription::Subscription as SubscriptionRpc, + write_bytes, RouterConfig, RpcCodec, RpcRead, RpcStream, RpcWrite, StreamCloseCode, + StreamError, +}; + +#[trait_variant::make(SubscriptionHandler: Send)] +pub trait SubscriptionHandlerLocal +where + M: SubscriptionRpc, + St: RpcStream, +{ + async fn handle( + self, + message: M::Request, + responder: SubscriptionResponder, + ); + + fn handle_transport_error(&self, _error: &St::Error) {} +} + +pub struct SubscriptionResponder +where + W: RpcWrite, +{ + writer: Option, + marker: PhantomData T>, +} + +impl SubscriptionResponder +where + T: RpcCodec, + W: RpcWrite, +{ + pub(crate) fn new(writer: W) -> Self { + Self { + writer: Some(writer), + marker: PhantomData, + } + } + + pub async fn send(&mut self, event: T) -> Result<(), W::Error> { + let writer = self.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + codec::encode_value_part(&event, &mut encoded); + write_bytes(writer, Bytes::from(encoded)).await?; + Ok(()) + } + + pub async fn finish(mut self) -> Result<(), W::Error> { + let mut writer = self.writer.take().unwrap(); + finish_bytes(&mut writer).await + } + + pub fn close(mut self, code: StreamCloseCode) { + if let Some(writer) = self.writer.take() { + writer.close(code); + } + } +} + +impl Drop for SubscriptionResponder +where + W: RpcWrite, +{ + fn drop(&mut self) { + if let Some(writer) = self.writer.take() { + writer.close(StreamCloseCode::CANCELLED); + } + } +} + +pub(crate) async fn handle_subscription_inner( + state: S, + config: RouterConfig, + mut reader: St::Reader, + writer: St::Writer, + handle: H, + handle_transport_error: E, +) where + M: SubscriptionRpc + 'static, + St: RpcStream + 'static, + H: FnOnce(S, M::Request, SubscriptionResponder) -> HF, + HF: Future, + E: FnOnce(&S, &St::Error), +{ + let request = match read_eof_request::(&mut reader, config).await { + Ok(request) => request, + Err(error) => { + let code = error.close_code(); + handle_transport_error(&state, &error); + if let Some(code) = code { + reader.close(code); + writer.close(code); + } + return; + } + }; + + handle(state, request, SubscriptionResponder::new(writer)).await; +} diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs new file mode 100644 index 00000000..b41dedcd --- /dev/null +++ b/ql-rpc/src/rpc/upload/client.rs @@ -0,0 +1,146 @@ +use bytes::{BufMut, Bytes}; + +use crate::{ + finish_bytes, read_bytes, + rpc::parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, + upload::Upload, + write_bytes, CallError, ChunkQueue, RpcCodec, RpcRead, RpcWrite, StreamCloseCode, +}; + +pub struct UploadCall +where + M: Upload, + W: RpcWrite, + R: RpcRead, +{ + writer: Option, + reader: Option, + marker: std::marker::PhantomData M>, +} + +pub struct UploadPartWriter<'a, M, W, R> +where + M: Upload, + W: RpcWrite, + R: RpcRead, +{ + parent: &'a mut UploadCall, + finished: bool, +} + +impl UploadCall +where + M: Upload, + W: RpcWrite, + R: RpcRead, +{ + pub fn new(writer: W, reader: R) -> Self { + Self { + writer: Some(writer), + reader: Some(reader), + marker: std::marker::PhantomData, + } + } + + pub async fn start_part( + &mut self, + part_header: M::PartHeader, + ) -> Result, W::Error> { + let writer = self.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + encode_part_header(&part_header, &mut encoded); + write_bytes(writer, Bytes::from(encoded)).await?; + Ok(UploadPartWriter { + parent: self, + finished: false, + }) + } + + pub async fn finish(mut self) -> Result> { + let mut writer = self.writer.take().unwrap(); + let mut encoded = Vec::new(); + encode_finish(&mut encoded); + write_bytes(&mut writer, Bytes::from(encoded)) + .await + .map_err(CallError::Transport)?; + finish_bytes(&mut writer) + .await + .map_err(CallError::Transport)?; + + let mut reader = self.reader.take().unwrap(); + let mut bytes = ChunkQueue::default(); + + while let Some(chunk) = read_bytes(&mut reader, usize::MAX) + .await + .map_err(CallError::Transport)? + { + bytes.push(chunk); + } + + let value = M::Response::decode_value(&mut bytes).map_err(CallError::Codec)?; + if bytes.remaining() > 0 { + return Err(crate::Error::TrailingBytes.into()); + } + Ok(value) + } + + fn close(&mut self, code: StreamCloseCode) { + if let Some(reader) = self.reader.take() { + reader.close(code); + } + if let Some(writer) = self.writer.take() { + writer.close(code); + } + } +} + +impl Drop for UploadCall +where + M: Upload, + W: RpcWrite, + R: RpcRead, +{ + fn drop(&mut self) { + self.close(StreamCloseCode::CANCELLED); + } +} + +impl UploadPartWriter<'_, M, W, R> +where + M: Upload, + W: RpcWrite, + R: RpcRead, +{ + pub async fn send(&mut self, bytes: Bytes) -> Result<(), W::Error> { + let writer = self.parent.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + encode_body_chunk(&bytes, &mut encoded); + write_bytes(writer, Bytes::from(encoded)).await + } + + pub async fn finish(mut self) -> Result<(), W::Error> { + let writer = self.parent.writer.as_mut().unwrap(); + let mut encoded = Vec::new(); + encode_end_part(&mut encoded); + write_bytes(writer, Bytes::from(encoded)).await?; + self.finished = true; + Ok(()) + } +} + +impl Drop for UploadPartWriter<'_, M, W, R> +where + M: Upload, + W: RpcWrite, + R: RpcRead, +{ + fn drop(&mut self) { + if !self.finished { + self.parent.close(StreamCloseCode::CANCELLED); + } + } +} + +pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { + crate::codec::encode_value_part(request, out) +} diff --git a/ql-rpc/src/rpc/upload/mod.rs b/ql-rpc/src/rpc/upload/mod.rs new file mode 100644 index 00000000..9f96a824 --- /dev/null +++ b/ql-rpc/src/rpc/upload/mod.rs @@ -0,0 +1,26 @@ +use super::Route; +use crate::RpcCodec; + +pub(crate) mod client; +pub(crate) mod server; + +pub use client::{encode_request, UploadCall, UploadPartWriter}; +pub use server::{UploadHandler, UploadHandlerLocal, UploadPart, UploadReader, UploadResponder}; + +/// rpc where the caller uploads zero or more byte parts after a typed request +/// +/// the typed request usually describes how the responder should interpret the +/// following parts +/// the request is length-delimited so raw upload bytes can follow immediately +/// once the upload reaches eof, the responder returns one typed +/// [`Self::Response`] +pub trait Upload: Route { + /// codec error shared by request and response values + type Error; + /// typed input needed before request body bytes arrive + type Request: RpcCodec; + /// typed metadata available before each byte part arrives + type PartHeader: RpcCodec; + /// typed terminal result after the upload body is fully read + type Response: RpcCodec; +} diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs new file mode 100644 index 00000000..d2e6765b --- /dev/null +++ b/ql-rpc/src/rpc/upload/server.rs @@ -0,0 +1,243 @@ +use std::future::{poll_fn, Future}; + +use bytes::Bytes; + +use crate::{ + request::Response, + rpc::{ + parts::{FrameKind, PartFrameReader, PartReadStep}, + read_framed_request_prefix, + }, + RouterConfig, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError, Upload, +}; + +#[trait_variant::make(UploadHandler: Send)] +pub trait UploadHandlerLocal +where + M: Upload, + St: RpcStream, +{ + async fn handle( + self, + request: M::Request, + upload: UploadReader, + responder: UploadResponder, + ); + + fn handle_transport_error(&self, _error: &St::Error) {} +} + +pub struct UploadReader +where + M: Upload, + R: RpcRead, +{ + stream: Option, + reader: PartFrameReader, +} + +pub struct UploadPart<'a, M, R> +where + M: Upload, + R: RpcRead, +{ + parent: &'a mut UploadReader, + finished: bool, +} + +pub struct UploadResponder +where + W: RpcWrite, +{ + inner: Response, +} + +impl UploadReader +where + M: Upload, + R: RpcRead, +{ + pub async fn next_part( + &mut self, + ) -> Result)>, crate::CallError> + { + if self.stream.is_none() { + return Ok(None); + } + + match self.read_frame().await? { + PartReadStep::PartHeader(value) => Ok(Some(( + value, + UploadPart { + parent: self, + finished: false, + }, + ))), + PartReadStep::Finish => { + self.stream.take(); + Ok(None) + } + PartReadStep::BodyBytes(_) => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::BodyChunk.tag()).into()) + } + PartReadStep::EndPart => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::EndPart.tag()).into()) + } + PartReadStep::NeedMore => unreachable!("read_frame waits for a complete frame"), + } + } + + async fn read_frame( + &mut self, + ) -> Result, crate::CallError> { + loop { + match self.reader.advance() { + Ok(PartReadStep::NeedMore) => {} + Ok(step) => return Ok(step), + Err(error) => return Err(error.into()), + } + + let stream = self.stream.as_mut().unwrap(); + match poll_fn(|cx| stream.poll_read(usize::MAX, cx)).await { + Ok(Some(chunk)) => { + self.reader.push(chunk); + } + Ok(None) => return Err(crate::Error::Truncated.into()), + Err(error) => return Err(crate::CallError::Transport(error)), + } + } + } + + pub fn close(mut self, code: StreamCloseCode) { + self.close_inner(code); + } + + fn close_inner(&mut self, code: StreamCloseCode) { + if let Some(stream) = self.stream.take() { + stream.close(code); + } + } +} + +impl Drop for UploadReader +where + M: Upload, + R: RpcRead, +{ + fn drop(&mut self) { + if self.stream.is_some() { + self.close_inner(StreamCloseCode::CANCELLED); + } + } +} + +impl UploadPart<'_, M, R> +where + M: Upload, + R: RpcRead, +{ + pub async fn read_chunk( + &mut self, + ) -> Result, crate::CallError> { + if self.finished { + return Ok(None); + } + + match self.parent.read_frame().await? { + PartReadStep::BodyBytes(bytes) => Ok(Some(bytes)), + PartReadStep::EndPart => { + self.finished = true; + Ok(None) + } + PartReadStep::PartHeader(_) => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::PartHeader.tag()).into()) + } + PartReadStep::Finish => { + Err(crate::Error::UnexpectedFrameKind(FrameKind::Finish.tag()).into()) + } + PartReadStep::NeedMore => unreachable!("read_frame waits for a complete frame"), + } + } + + pub fn close(mut self, code: StreamCloseCode) { + self.parent.close_inner(code); + self.finished = true; + } +} + +impl Drop for UploadPart<'_, M, R> +where + M: Upload, + R: RpcRead, +{ + fn drop(&mut self) { + if !self.finished { + self.parent.close_inner(StreamCloseCode::CANCELLED); + } + } +} + +impl UploadResponder +where + T: crate::RpcCodec, + W: RpcWrite, +{ + pub(crate) fn new(writer: W) -> Self { + Self { + inner: Response::new(writer), + } + } + + pub async fn respond(self, response: T) -> Result<(), W::Error> { + self.inner.respond(response).await + } + + pub fn close(self, code: StreamCloseCode) { + self.inner.close(code); + } +} + +pub(crate) async fn handle_upload_inner( + state: S, + config: RouterConfig, + mut reader: St::Reader, + writer: St::Writer, + handle: H, + handle_transport_error: E, +) where + M: Upload + 'static, + St: RpcStream + 'static, + H: FnOnce( + S, + M::Request, + UploadReader, + UploadResponder, + ) -> HF, + HF: Future, + E: FnOnce(&S, &St::Error), +{ + let (request, buffered) = + match read_framed_request_prefix::(&mut reader, config).await { + Ok(value) => value, + Err(error) => { + let code = error.close_code(); + handle_transport_error(&state, &error); + if let Some(code) = code { + reader.close(code); + writer.close(code); + } + return; + } + }; + + handle( + state, + request, + UploadReader { + stream: Some(reader), + reader: PartFrameReader::new(buffered), + }, + UploadResponder::new(writer), + ) + .await; +} diff --git a/ql-rpc/src/rpc/utils.rs b/ql-rpc/src/rpc/utils.rs new file mode 100644 index 00000000..bf5f49ea --- /dev/null +++ b/ql-rpc/src/rpc/utils.rs @@ -0,0 +1,120 @@ +use crate::{ + read_bytes, ChunkQueue, CodecError, FramedPrefixStep, FramedReadStep, FramedReader, + RouterConfig, RpcCodec, RpcRead, StreamCloseCode, +}; + +/// reads one length-delimited value and rejects trailing bytes +pub(crate) async fn read_framed_request( + reader: &mut R, + config: RouterConfig, +) -> Result +where + T: RpcCodec, + R: RpcRead, +{ + let mut value_reader = FramedReader::::default(); + let mut total_read = 0usize; + + let value = loop { + match value_reader.advance() { + Ok(FramedReadStep::Value(value)) => break value, + Ok(FramedReadStep::NeedMore(next)) => value_reader = next, + Err(CodecError::Rpc(_error)) => return Err(StreamCloseCode::REFUSED.into()), + Err(CodecError::Codec(_error)) => return Err(StreamCloseCode::REFUSED.into()), + } + + let remaining = config.max_request_bytes.saturating_sub(total_read); + if remaining == 0 { + return Err(StreamCloseCode::LIMIT.into()); + } + + match read_bytes(reader, remaining).await { + Ok(Some(chunk)) => { + total_read += chunk.len(); + value_reader = value_reader.push(chunk); + } + Ok(None) => return Err(StreamCloseCode::REFUSED.into()), + Err(error) => return Err(error), + } + }; + + let remaining = config.max_request_bytes.saturating_sub(total_read); + let probe = remaining.max(1); + match read_bytes(reader, probe).await { + Ok(None) => Ok(value), + Ok(Some(_)) if remaining == 0 => Err(StreamCloseCode::LIMIT.into()), + Ok(Some(_)) => Err(StreamCloseCode::REFUSED.into()), + Err(error) => Err(error), + } +} + +/// reads one length-delimited value and returns any bytes already buffered +pub(crate) async fn read_framed_request_prefix( + reader: &mut R, + config: RouterConfig, +) -> Result<(T, ChunkQueue), R::Error> +where + T: RpcCodec, + R: RpcRead, +{ + let mut value_reader = FramedReader::::default(); + let mut total_read = 0usize; + + loop { + match value_reader.advance_prefix() { + Ok(FramedPrefixStep::Value { value, bytes }) => return Ok((value, bytes)), + Ok(FramedPrefixStep::NeedMore(next)) => value_reader = next, + Err(CodecError::Rpc(_error)) => return Err(StreamCloseCode::REFUSED.into()), + Err(CodecError::Codec(_error)) => return Err(StreamCloseCode::REFUSED.into()), + } + + let remaining = config.max_request_bytes.saturating_sub(total_read); + if remaining == 0 { + return Err(StreamCloseCode::LIMIT.into()); + } + + match read_bytes(reader, remaining).await { + Ok(Some(chunk)) => { + total_read += chunk.len(); + value_reader = value_reader.push(chunk); + } + Ok(None) => return Err(StreamCloseCode::REFUSED.into()), + Err(error) => return Err(error), + } + } +} + +/// reads one eof-delimited value up to the configured request limit +pub(crate) async fn read_eof_request( + reader: &mut R, + config: RouterConfig, +) -> Result +where + T: RpcCodec, + R: RpcRead, +{ + let mut bytes = ChunkQueue::default(); + let mut total_read = 0usize; + + loop { + let remaining = config.max_request_bytes.saturating_sub(total_read); + let probe = remaining.max(1); + match read_bytes(reader, probe).await { + Ok(Some(chunk)) => { + if chunk.len() > remaining { + return Err(StreamCloseCode::LIMIT.into()); + } + total_read += chunk.len(); + bytes.push(chunk); + } + Ok(None) => break, + Err(error) => return Err(error), + } + } + + let value = T::decode_value(&mut bytes).map_err(|_error| StreamCloseCode::REFUSED)?; + if bytes.remaining() > 0 { + return Err(StreamCloseCode::REFUSED.into()); + } + Ok(value) +} diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs new file mode 100644 index 00000000..f6174efd --- /dev/null +++ b/ql-rpc/src/stream.rs @@ -0,0 +1,89 @@ +use std::{ + future::poll_fn, + task::{Context, Poll}, +}; + +use bytes::Bytes; + +use crate::{RouteId, StreamCloseCode}; + +pub trait RpcStream { + type Error: StreamError; + type Reader: RpcRead; + type Writer: RpcWrite; + + fn route_id(&self) -> Option; + fn split(self) -> (Self::Reader, Self::Writer); +} + +pub trait RpcRead { + type Error: StreamError; + + /// reads inbound bytes until eof or error + fn poll_read( + &mut self, + max_len: usize, + cx: &mut Context<'_>, + ) -> Poll, Self::Error>>; + + /// aborts the read side + fn close(self, code: StreamCloseCode); +} + +pub trait RpcWrite { + type Error: StreamError; + + /// writes outbound bytes before finish or close + fn poll_write( + &mut self, + bytes: &mut Bytes, + cx: &mut Context<'_>, + ) -> Poll>; + + /// completes the write side and must be polled until ready without further write or close calls + fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll>; + + /// aborts the write side before finish + fn close(self, code: StreamCloseCode); +} + +pub trait StreamError: From { + fn close_code(&self) -> Option; +} + +impl StreamError for StreamCloseCode { + fn close_code(&self) -> Option { + Some(*self) + } +} + +pub async fn read_bytes(reader: &mut R, max_len: usize) -> Result, R::Error> +where + R: RpcRead, +{ + poll_fn(|cx| reader.poll_read(max_len, cx)).await +} + +pub async fn write_bytes(writer: &mut W, bytes: Bytes) -> Result<(), W::Error> +where + W: RpcWrite, +{ + let mut bytes = bytes; + poll_fn(|cx| writer.poll_write(&mut bytes, cx)).await +} + +pub async fn finish_bytes(writer: &mut W) -> Result<(), W::Error> +where + W: RpcWrite, +{ + poll_fn(|cx| writer.poll_finish(cx)).await +} + +pub fn close_stream(stream: St, code: StreamCloseCode) +where + St: RpcStream, +{ + let (reader, writer) = stream.split(); + reader.close(code); + writer.close(code); +} From 154a4d9ccb4a5e74bbbbbfd5667e70bdcccf33a0 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 4 Jun 2026 09:26:25 -0400 Subject: [PATCH 05/59] ql-runtime: add async runtime --- Cargo.lock | 447 ++++++++++++++++-- Cargo.toml | 1 + ql-runtime/Cargo.toml | 35 ++ ql-runtime/src/command.rs | 57 +++ ql-runtime/src/driver/mod.rs | 609 +++++++++++++++++++++++++ ql-runtime/src/driver/state.rs | 165 +++++++ ql-runtime/src/driver/test.rs | 207 +++++++++ ql-runtime/src/error.rs | 18 + ql-runtime/src/handle/mod.rs | 96 ++++ ql-runtime/src/io/inner.rs | 643 ++++++++++++++++++++++++++ ql-runtime/src/io/mod.rs | 59 +++ ql-runtime/src/io/reader.rs | 236 ++++++++++ ql-runtime/src/io/slot.rs | 175 +++++++ ql-runtime/src/io/sync.rs | 89 ++++ ql-runtime/src/io/writer.rs | 294 ++++++++++++ ql-runtime/src/lib.rs | 63 +++ ql-runtime/src/log.rs | 54 +++ ql-runtime/src/platform.rs | 43 ++ ql-runtime/src/rpc/adapter.rs | 83 ++++ ql-runtime/src/rpc/download.rs | 67 +++ ql-runtime/src/rpc/duplex.rs | 59 +++ ql-runtime/src/rpc/error.rs | 79 ++++ ql-runtime/src/rpc/mod.rs | 154 +++++++ ql-runtime/src/rpc/progress.rs | 50 ++ ql-runtime/src/rpc/subscription.rs | 43 ++ ql-runtime/src/rpc/upload.rs | 44 ++ ql-runtime/src/tests/handshake.rs | 178 ++++++++ ql-runtime/src/tests/mod.rs | 710 +++++++++++++++++++++++++++++ ql-runtime/src/tests/rpc.rs | 677 +++++++++++++++++++++++++++ ql-runtime/src/tests/session.rs | 213 +++++++++ ql-runtime/src/tests/stream.rs | 673 +++++++++++++++++++++++++++ 31 files changed, 6295 insertions(+), 26 deletions(-) create mode 100644 ql-runtime/Cargo.toml create mode 100644 ql-runtime/src/command.rs create mode 100644 ql-runtime/src/driver/mod.rs create mode 100644 ql-runtime/src/driver/state.rs create mode 100644 ql-runtime/src/driver/test.rs create mode 100644 ql-runtime/src/error.rs create mode 100644 ql-runtime/src/handle/mod.rs create mode 100644 ql-runtime/src/io/inner.rs create mode 100644 ql-runtime/src/io/mod.rs create mode 100644 ql-runtime/src/io/reader.rs create mode 100644 ql-runtime/src/io/slot.rs create mode 100644 ql-runtime/src/io/sync.rs create mode 100644 ql-runtime/src/io/writer.rs create mode 100644 ql-runtime/src/lib.rs create mode 100644 ql-runtime/src/log.rs create mode 100644 ql-runtime/src/platform.rs create mode 100644 ql-runtime/src/rpc/adapter.rs create mode 100644 ql-runtime/src/rpc/download.rs create mode 100644 ql-runtime/src/rpc/duplex.rs create mode 100644 ql-runtime/src/rpc/error.rs create mode 100644 ql-runtime/src/rpc/mod.rs create mode 100644 ql-runtime/src/rpc/progress.rs create mode 100644 ql-runtime/src/rpc/subscription.rs create mode 100644 ql-runtime/src/rpc/upload.rs create mode 100644 ql-runtime/src/tests/handshake.rs create mode 100644 ql-runtime/src/tests/mod.rs create mode 100644 ql-runtime/src/tests/rpc.rs create mode 100644 ql-runtime/src/tests/session.rs create mode 100644 ql-runtime/src/tests/stream.rs diff --git a/Cargo.lock b/Cargo.lock index 071d36b8..123d0e59 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -72,7 +72,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "dbb4e440d04be07da1f1bf44fb4495ebd58669372fe0cffa6e48595ac5bd88a3" dependencies = [ "android_log-sys", - "env_filter", + "env_filter 0.1.3", "log", ] @@ -85,6 +85,56 @@ dependencies = [ "libc", ] +[[package]] +name = "anstream" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" +dependencies = [ + "anstyle", + "anstyle-parse", + "anstyle-query", + "anstyle-wincon", + "colorchoice", + "is_terminal_polyfill", + "utf8parse", +] + +[[package]] +name = "anstyle" +version = "1.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" + +[[package]] +name = "anstyle-parse" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" +dependencies = [ + "utf8parse", +] + +[[package]] +name = "anstyle-query" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "anstyle-wincon" +version = "3.0.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" +dependencies = [ + "anstyle", + "once_cell_polyfill", + "windows-sys 0.61.2", +] + [[package]] name = "anyhow" version = "1.0.99" @@ -109,6 +159,18 @@ version = "0.7.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" +[[package]] +name = "async-channel" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "924ed96dd52d1b75e9c1a3e6275715fd320f5f9439fb5a4a11fa51f4221158d2" +dependencies = [ + "concurrent-queue", + "event-listener-strategy", + "futures-core", + "pin-project-lite", +] + [[package]] name = "atomic" version = "0.5.3" @@ -341,9 +403,9 @@ dependencies = [ [[package]] name = "bitflags" -version = "2.12.1" +version = "2.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "84d7ced0ae9557296835c32bf1b1e02b44c746701f898460fb000d7eaa84f00a" +checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" [[package]] name = "blake2" @@ -494,7 +556,7 @@ dependencies = [ "num-traits", "serde", "wasm-bindgen", - "windows-link", + "windows-link 0.1.3", ] [[package]] @@ -508,6 +570,22 @@ dependencies = [ "zeroize", ] +[[package]] +name = "colorchoice" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" + +[[package]] +name = "concurrent-queue" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ca0197aee26d1ae37445ee532fefce43251d24cc7c166799f4d46817f1d3973" +dependencies = [ + "crossbeam-utils", + "loom", +] + [[package]] name = "console" version = "0.15.11" @@ -517,7 +595,7 @@ dependencies = [ "encode_unicode", "libc", "once_cell", - "windows-sys", + "windows-sys 0.59.0", ] [[package]] @@ -591,6 +669,12 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "crossbeam-utils" +version = "0.8.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" + [[package]] name = "crunchy" version = "0.2.4" @@ -704,6 +788,12 @@ dependencies = [ "zeroize", ] +[[package]] +name = "diatomic-waker" +version = "0.2.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab03c107fafeb3ee9f5925686dbb7a73bc76e3932abb0d2b365cb64b169cf04c" + [[package]] name = "digest" version = "0.10.7" @@ -829,6 +919,29 @@ dependencies = [ "regex", ] +[[package]] +name = "env_filter" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32e90c2accc4b07a8456ea0debdc2e7587bdd890680d71173a15d4ae604f6eef" +dependencies = [ + "log", + "regex", +] + +[[package]] +name = "env_logger" +version = "0.11.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0621c04f2196ac3f488dd583365b9c09be011a4ab8b9f37248ffcc8f6198b56a" +dependencies = [ + "anstream", + "anstyle", + "env_filter 1.0.1", + "jiff", + "log", +] + [[package]] name = "equivalent" version = "1.0.2" @@ -842,14 +955,36 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" dependencies = [ "libc", - "windows-sys", + "windows-sys 0.61.2", +] + +[[package]] +name = "event-listener" +version = "5.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e13b66accf52311f30a0db42147dadea9850cb48cd070028831ae5f5d4b856ab" +dependencies = [ + "concurrent-queue", + "loom", + "parking", + "pin-project-lite", +] + +[[package]] +name = "event-listener-strategy" +version = "0.5.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8be9f3dfaaffdae2972880079a491a1a8bb7cbed0b8dd7a347f668b4150a3b93" +dependencies = [ + "event-listener", + "pin-project-lite", ] [[package]] name = "fastrand" -version = "2.4.1" +version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9f1f227452a390804cdb637b74a86990f2a7d7ba4b7d5693aac9b4dd6defd8d6" +checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be" [[package]] name = "ff" @@ -989,6 +1124,19 @@ version = "0.3.31" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9e5c1b78ca4aae1ac06c48a526a655760685149f0d465d21f37abfe57ce075c6" +[[package]] +name = "futures-lite" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f78e10609fe0e0b3f4157ffab1876319b5b0db102a2c60dc4626306dc46b44ad" +dependencies = [ + "fastrand", + "futures-core", + "futures-io", + "parking", + "pin-project-lite", +] + [[package]] name = "futures-macro" version = "0.3.31" @@ -1030,6 +1178,21 @@ dependencies = [ "slab", ] +[[package]] +name = "generator" +version = "0.8.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52f04ae4152da20c76fe800fa48659201d5cf627c5149ca0b707b69d7eef6cf9" +dependencies = [ + "cc", + "cfg-if", + "libc", + "log", + "rustversion", + "windows-link 0.2.1", + "windows-result", +] + [[package]] name = "generic-array" version = "0.14.7" @@ -1380,6 +1543,12 @@ dependencies = [ "libc", ] +[[package]] +name = "is_terminal_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" + [[package]] name = "itertools" version = "0.11.0" @@ -1395,6 +1564,30 @@ version = "1.0.15" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4a5f13b858c8d314ee3e8f639011f7ccefe71f97f96e50151fb991f267928e2c" +[[package]] +name = "jiff" +version = "0.2.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1a3546dc96b6d42c5f24902af9e2538e82e39ad350b0c766eb3fbf2d8f3d8359" +dependencies = [ + "jiff-static", + "log", + "portable-atomic", + "portable-atomic-util", + "serde_core", +] + +[[package]] +name = "jiff-static" +version = "0.2.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2a8c8b344124222efd714b73bb41f8b5120b27a7cc1c75593a6ff768d9d05aa4" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.106", +] + [[package]] name = "jobserver" version = "0.1.33" @@ -1549,9 +1742,31 @@ dependencies = [ [[package]] name = "log" -version = "0.4.27" +version = "0.4.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" + +[[package]] +name = "loom" +version = "0.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "13dc2df351e3202783a1fe0d44375f7295ffb4049267b0f3018346dc122a1d94" +checksum = "419e0dc8046cb947daa77eb95ae174acfbddb7673b4151f56d1eed8e93fbfaca" +dependencies = [ + "cfg-if", + "generator", + "scoped-tls", + "tracing", + "tracing-subscriber", +] + +[[package]] +name = "matchers" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d1525a2a28c7f4fa0fc98bb91ae755d1e2d1505079e05539e35bc876b5d65ae9" +dependencies = [ + "regex-automata", +] [[package]] name = "md-5" @@ -1615,7 +1830,7 @@ checksum = "78bed444cc8a2160f01cbcf811ef18cac863ad68ae8ca62092e8db51d51c761c" dependencies = [ "libc", "wasi 0.11.1+wasi-snapshot-preview1", - "windows-sys", + "windows-sys 0.59.0", ] [[package]] @@ -1638,6 +1853,15 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "nu-ansi-term" +version = "0.50.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "num-bigint" version = "0.4.6" @@ -1719,6 +1943,18 @@ version = "1.21.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" +[[package]] +name = "once_cell_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" + +[[package]] +name = "oneshot" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4ce411919553d3f9fa53a0880544cda985a112117a0444d5ff1e870a893d6ea" + [[package]] name = "opaque-debug" version = "0.3.1" @@ -1774,6 +2010,15 @@ dependencies = [ "sha2", ] +[[package]] +name = "parking" +version = "2.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f38d5652c16fde515bb1ecef450ab0f6a219d619a7274976324d5e377f7dceba" +dependencies = [ + "loom", +] + [[package]] name = "parking_lot_core" version = "0.9.11" @@ -1806,9 +2051,9 @@ checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" [[package]] name = "pastey" -version = "0.2.3" +version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4" +checksum = "b867cad97c0791bbd3aaa6472142568c6c9e8f71937e98379f584cfb0cf35bec" [[package]] name = "pbkdf2" @@ -1927,6 +2172,15 @@ version = "1.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f84267b20a16ea918e43c6a88433c2d54fa145c92a811b5b047ccbe153674483" +[[package]] +name = "portable-atomic-util" +version = "0.2.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a106d1259c23fac8e543272398ae0e3c0b8d33c88ed73d0cc71b0f1d902618" +dependencies = [ + "portable-atomic", +] + [[package]] name = "potential_utf" version = "0.1.2" @@ -2107,6 +2361,25 @@ dependencies = [ "trait-variant", ] +[[package]] +name = "ql-runtime" +version = "0.1.0" +dependencies = [ + "async-channel", + "bytes", + "diatomic-waker", + "env_logger", + "event-listener", + "futures-lite", + "log", + "loom", + "oneshot", + "ql-fsm", + "ql-rpc", + "ql-wire", + "tokio", +] + [[package]] name = "ql-wire" version = "0.1.0" @@ -2245,9 +2518,9 @@ dependencies = [ [[package]] name = "regex" -version = "1.11.1" +version = "1.12.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b544ef1b4eac5dc2db33ea63606ae9ffcfac26c1416a2806ae0bf5f56b201191" +checksum = "e10754a14b9137dd7b1e3e5b0493cc9171fdd105e0ab477f51b72e7f3ac0e276" dependencies = [ "aho-corasick", "memchr", @@ -2257,9 +2530,9 @@ dependencies = [ [[package]] name = "regex-automata" -version = "0.4.9" +version = "0.4.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "809e8dc61f6de73b46c85f4c96486310fe304c434cfa43669d7b40f711150908" +checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" dependencies = [ "aho-corasick", "memchr", @@ -2367,7 +2640,7 @@ dependencies = [ "errno", "libc", "linux-raw-sys", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -2403,6 +2676,12 @@ dependencies = [ "cipher", ] +[[package]] +name = "scoped-tls" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e1cf6437eb19a8f4a6cc0f7dca544973b0b78843adbfeb3683d1a94a0024a294" + [[package]] name = "scopeguard" version = "1.2.0" @@ -2462,18 +2741,28 @@ checksum = "56e6fa9c48d24d85fb3de5ad847117517440f6beceb7798af16b4a87d616b8d0" [[package]] name = "serde" -version = "1.0.219" +version = "1.0.228" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5f0e2c6ed6606019b4e29e69dbaba95b11854410e5347d525002456dbbb786b6" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" dependencies = [ "serde_derive", ] [[package]] name = "serde_derive" -version = "1.0.219" +version = "1.0.228" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5b0276cf7f2c73365f7157c8123c21cd9a50fbbd844757af28ca1f5925fc2a00" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" dependencies = [ "proc-macro2", "quote", @@ -2514,6 +2803,15 @@ dependencies = [ "digest", ] +[[package]] +name = "sharded-slab" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6" +dependencies = [ + "lazy_static", +] + [[package]] name = "shlex" version = "1.3.0" @@ -2687,7 +2985,7 @@ dependencies = [ "getrandom 0.3.3", "once_cell", "rustix", - "windows-sys", + "windows-sys 0.61.2", ] [[package]] @@ -2710,6 +3008,15 @@ dependencies = [ "syn 2.0.106", ] +[[package]] +name = "thread_local" +version = "1.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f60246a4944f24f6e018aa17cdeffb7818b76356965d03b07d6a9886e8962185" +dependencies = [ + "cfg-if", +] + [[package]] name = "threadpool" version = "1.8.1" @@ -2777,6 +3084,67 @@ dependencies = [ "mio", "pin-project-lite", "slab", + "tokio-macros", +] + +[[package]] +name = "tokio-macros" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e06d43f1345a3bcd39f6a56dbb7dcab2ba47e68e8ac134855e7e2bdbaf8cab8" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.106", +] + +[[package]] +name = "tracing" +version = "0.1.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" +dependencies = [ + "pin-project-lite", + "tracing-core", +] + +[[package]] +name = "tracing-core" +version = "0.1.36" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" +dependencies = [ + "once_cell", + "valuable", +] + +[[package]] +name = "tracing-log" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ee855f1f400bd0e5c02d150ae5de3840039a3f54b025156404e34c23c03f47c3" +dependencies = [ + "log", + "once_cell", + "tracing-core", +] + +[[package]] +name = "tracing-subscriber" +version = "0.3.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319" +dependencies = [ + "matchers", + "nu-ansi-term", + "once_cell", + "regex-automata", + "sharded-slab", + "smallvec", + "thread_local", + "tracing", + "tracing-core", + "tracing-log", ] [[package]] @@ -2857,6 +3225,12 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +[[package]] +name = "utf8parse" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" + [[package]] name = "uuid" version = "1.18.1" @@ -2868,6 +3242,12 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "valuable" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" + [[package]] name = "version_check" version = "0.9.5" @@ -2987,7 +3367,7 @@ checksum = "c0fdd3ddb90610c7638aa2b3a3ab2904fb9e5cdbecc643ddb3647212781c4ae3" dependencies = [ "windows-implement", "windows-interface", - "windows-link", + "windows-link 0.1.3", "windows-result", "windows-strings", ] @@ -3020,13 +3400,19 @@ version = "0.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e6ad25900d524eaabdbbb96d20b4311e1e7ae1699af4fb28c17ae66c80d798a" +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + [[package]] name = "windows-result" version = "0.3.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "56f42bd332cc6c8eac5af113fc0c1fd6a8fd2aa08a0119358686e5160d0586c6" dependencies = [ - "windows-link", + "windows-link 0.1.3", ] [[package]] @@ -3035,7 +3421,7 @@ version = "0.4.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "56e6c93f3a0c3b36176cb1327a4958a0353d5d166c2a35cb268ace15e91d3b57" dependencies = [ - "windows-link", + "windows-link 0.1.3", ] [[package]] @@ -3047,6 +3433,15 @@ dependencies = [ "windows-targets", ] +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link 0.2.1", +] + [[package]] name = "windows-targets" version = "0.52.6" diff --git a/Cargo.toml b/Cargo.toml index de8c80b7..b2492c48 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -6,6 +6,7 @@ members = [ "btp", "ql-fsm", "ql-rpc", + "ql-runtime", "ql-wire", "quantum-link-macros", ] diff --git a/ql-runtime/Cargo.toml b/ql-runtime/Cargo.toml new file mode 100644 index 00000000..564cee1c --- /dev/null +++ b/ql-runtime/Cargo.toml @@ -0,0 +1,35 @@ +[package] +name = "ql-runtime" +version = "0.1.0" +edition = "2021" +description = "QuantumLink async runtime" +license = "Proprietary" + +[features] +default = [] +log = ["dep:log"] +rpc = ["dep:ql-rpc"] + +[dependencies] +async-channel = { version = "2.5" } +bytes = { workspace = true } +diatomic-waker = { version = "0.2.3", default-features = false } +futures-lite = { version = "2.5" } +log = { version = "0.4", optional = true } +oneshot = { version = "0.1.11" } +ql-fsm = { workspace = true } +ql-rpc = { workspace = true, optional = true } +ql-wire = { workspace = true } + +[dev-dependencies] +env_logger = "0.11" +log = "0.4" +ql-wire = { workspace = true, features = ["test-utils"] } +tokio = { version = "1.44", features = ["macros", "rt", "time", "sync"] } + +[target.'cfg(loom)'.dev-dependencies] +event-listener = { version = "5.4", features = ["loom"] } +loom = "0.7" + +[lints.rust] +unexpected_cfgs = { level = "warn", check-cfg = ['cfg(loom)'] } diff --git a/ql-runtime/src/command.rs b/ql-runtime/src/command.rs new file mode 100644 index 00000000..4a47a45e --- /dev/null +++ b/ql-runtime/src/command.rs @@ -0,0 +1,57 @@ +use ql_fsm::{NoSessionError, PairingInvite}; +use ql_wire::{ + CloseTarget, PairingToken, PeerBundle, RouteId, SessionCloseCode, StreamCloseCode, StreamId, +}; + +use crate::{StreamReader, StreamWriter}; + +pub enum Command { + BindPeer { + peer: PeerBundle, + }, + Connect, + ArmPairing { + token: PairingToken, + }, + DisarmPairing, + StartPairing { + invite: PairingInvite, + }, + OpenStream { + route_id: RouteId, + start: oneshot::Sender>, + }, + PollInbound { + stream_id: StreamId, + }, + PollStream { + stream_id: StreamId, + }, + CloseSession { + code: SessionCloseCode, + }, + Unpair, + CloseStream { + stream_id: StreamId, + target: CloseTarget, + code: StreamCloseCode, + }, +} + +impl Command { + pub fn kind(&self) -> &'static str { + match self { + Self::BindPeer { .. } => "BindPeer", + Self::Connect => "Connect", + Self::ArmPairing { .. } => "ArmPairing", + Self::DisarmPairing => "DisarmPairing", + Self::StartPairing { .. } => "StartPairing", + Self::OpenStream { .. } => "OpenStream", + Self::PollInbound { .. } => "PollInbound", + Self::PollStream { .. } => "PollStream", + Self::CloseSession { .. } => "CloseSession", + Self::Unpair => "Unpair", + Self::CloseStream { .. } => "CloseStream", + } + } +} diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs new file mode 100644 index 00000000..35de1bf0 --- /dev/null +++ b/ql-runtime/src/driver/mod.rs @@ -0,0 +1,609 @@ +mod state; +#[cfg(test)] +mod test; + +use std::{ + collections::{ + hash_map::{Entry, OccupiedEntry}, + HashMap, + }, + future::Future, + pin::{pin, Pin}, + task::{Context, Poll}, + time::Instant, +}; + +use async_channel::Recv; +use futures_lite::future::{poll_fn, yield_now}; +use ql_fsm::{Event, QlFsm, WriteId}; +use ql_wire::{CloseTarget, StreamCloseCode, StreamId}; + +use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; +use crate::{ + command::Command, + handle::QlStream, + io, log, + platform::{QlInbound, QlPlatform, QlTimer}, + QlStreamError, Runtime, RuntimeHandle, +}; + +impl Runtime

{ + #[allow(clippy::future_not_send)] + pub async fn run(self) { + let Self { + identity, + mut platform, + config, + rx, + tx, + } = self; + + let mut fsm = QlFsm::new(config.fsm, identity, Instant::now()); + + let mut state = DriverState { + streams: HashMap::new(), + runtime_tx: tx, + max_concurrent_message_writes: config.max_concurrent_message_writes, + }; + + let mut in_flight = Vec::new(); + let timer = platform.timer(); + let mut timer = pin!(timer); + let inbound = platform.inbound(); + let mut inbound = pin!(inbound); + let recv_future = rx.recv(); + let mut recv_future = Some(pin!(recv_future)); + let mut poll_cursor = 0usize; + + loop { + state.drain_fsm_events(&mut fsm, &platform); + if state.fill_write_slots(&mut fsm, &platform, &mut in_flight) { + state.drain_fsm_events(&mut fsm, &platform); + } + timer.as_mut().set_deadline(fsm.next_deadline()); + + let step = poll_fn(|cx| { + next_step( + cx, + recv_future.as_mut().map(|future| future.as_mut()), + inbound.as_mut(), + timer.as_mut(), + &mut in_flight, + poll_cursor, + ) + }) + .await; + poll_cursor = (poll_cursor + 1) % STEP_COUNT; + + match step { + DriverStep::Command(command) => { + log::trace!("processing command: kind={}", command.kind()); + state.drive_command(&mut fsm, command, &platform); + } + DriverStep::Inbound(bytes) => { + log::trace!("received transport frame: len={}", bytes.len()); + if let Err(e) = fsm.receive(Instant::now(), bytes, &platform) { + log::info!("receive rejected frame: error={e:?}"); + platform.handle_recv_error(e); + } + } + DriverStep::WriteCompleted { index, success } => { + let write = in_flight.swap_remove(index); + let write_id = write.write_id; + log::trace!( + "write completed: success={success} index={index} write_id={write_id:?}", + ); + DriverState::drive_write_completed(&mut fsm, write_id, success); + yield_now().await; + } + DriverStep::TimerExpired => { + log::trace!("timer expired"); + fsm.on_timer(Instant::now()); + } + DriverStep::Closed => { + log::debug!( + "command channel closed: in_flight_writes={}", + in_flight.len() + ); + recv_future = None; + if in_flight.is_empty() && !fsm.has_shutdown_work() { + break; + } + } + } + } + log::info!("runtime stopped"); + } +} + +struct InFlightWrite { + write_id: Option, + future: F, +} + +enum DriverStep { + Command(Command), + Inbound(Vec), + WriteCompleted { index: usize, success: bool }, + TimerExpired, + Closed, +} + +const STEP_COUNT: usize = 4; + +fn next_step( + cx: &mut Context<'_>, + mut recv_future: Option>>, + mut inbound: Pin<&mut I>, + mut timer: Pin<&mut T>, + in_flight: &mut [InFlightWrite], + start: usize, +) -> Poll +where + T: QlTimer, + F: Future + Unpin, + I: QlInbound, +{ + for offset in 0..STEP_COUNT { + let step = (start + offset) % STEP_COUNT; + let poll = match step { + 0 => recv_future.as_mut().map_or(Poll::Pending, |recv_future| { + recv_future + .as_mut() + .poll(cx) + .map(|res| res.map_or(DriverStep::Closed, DriverStep::Command)) + }), + 1 => inbound.as_mut().poll_recv(cx).map(DriverStep::Inbound), + 2 => { + for (index, write) in in_flight.iter_mut().enumerate() { + if let Poll::Ready(success) = Pin::new(&mut write.future).poll(cx) { + return Poll::Ready(DriverStep::WriteCompleted { index, success }); + } + } + Poll::Pending + } + 3 => timer + .as_mut() + .poll_wait(cx) + .map(|()| DriverStep::TimerExpired), + _ => unreachable!(), + }; + if poll.is_ready() { + return poll; + } + } + + Poll::Pending +} + +impl DriverState { + #[allow(clippy::too_many_lines)] + fn drive_command(&mut self, fsm: &mut QlFsm, command: Command, platform: &P) { + match command { + Command::BindPeer { peer } => { + log::info!("binding peer"); + fsm.bind_peer(peer); + } + Command::Connect => { + log::info!("starting IK connect"); + if fsm.connect_ik(Instant::now(), platform).is_err() { + log::warn!("IK connect ignored: no bound peer"); + } + } + Command::ArmPairing { token } => { + log::info!("arming inbound pairing"); + fsm.arm_pairing(token); + } + Command::DisarmPairing => { + log::info!("disarming inbound pairing"); + fsm.disarm_pairing(); + } + Command::StartPairing { invite } => { + log::info!(" starting XX pairing"); + fsm.connect_xx(Instant::now(), invite, platform); + } + Command::CloseSession { code } => { + log::info!("closing session: code={code:?}"); + fsm.close_session(code); + } + Command::Unpair => { + log::info!("unpairing peer"); + fsm.unpair(); + } + Command::OpenStream { route_id, start } => { + log::info!("open stream requested: route_id={route_id}"); + let Some(runtime_tx) = self.runtime_tx.upgrade() else { + log::warn!("open stream aborted: runtime channel unavailable"); + let _ = start.send(Err(ql_fsm::NoSessionError)); + return; + }; + + let mut stream_ops = match fsm.open_stream(route_id) { + Ok(stream_ops) => stream_ops, + Err(error) => { + log::warn!("open stream failed: route_id={route_id}"); + let _ = start.send(Err(error)); + return; + } + }; + let stream_id = stream_ops.stream_id(); + log::info!("open stream allocated: route_id={route_id} stream_id={stream_id}"); + let (reader, writer, reader_io, writer_io) = io::new_stream( + stream_id, + CloseTarget::Return, + CloseTarget::Origin, + RuntimeHandle::new(runtime_tx), + ); + self.streams.insert( + stream_id, + DriverStreamIo::new( + true, + Some(OutboundIo::new(writer_io)), + Some(InboundIo::new(reader_io)), + ), + ); + if start.send(Ok((stream_id, reader, writer))).is_err() { + log::warn!("open stream cancelled before delivery: stream_id={stream_id}"); + if let Some(stream) = self.streams.get_mut(&stream_id) { + stream.inbound_close(); + stream.outbound_close(); + } + stream_ops.close(CloseTarget::Both, StreamCloseCode::CANCELLED); + drop(stream_ops); + return; + } + drop(stream_ops); + self.poll_stream(fsm, stream_id); + } + Command::PollInbound { stream_id } => { + log::trace!("poll inbound requested: stream_id={stream_id}"); + self.handle_inbound_readable(fsm, stream_id); + } + Command::PollStream { stream_id } => { + log::trace!("poll stream requested: stream_id={stream_id}"); + self.poll_stream(fsm, stream_id); + } + Command::CloseStream { + stream_id, + target, + code, + } => { + log::debug!( + "close stream command: stream_id={stream_id} target={target:?} code={code:?}" + ); + if let Entry::Occupied(mut entry) = self.streams.entry(stream_id) { + let stream = entry.get_mut(); + if target == CloseTarget::Both || target == stream.inbound_target() { + stream.inbound_close(); + } + if target == CloseTarget::Both || target == stream.outbound_target() { + stream.outbound_close(); + } + Self::try_reap_stream(entry); + } + if let Ok(mut stream) = fsm.stream(stream_id) { + stream.close(target, code); + } + } + } + } + + fn drive_write_completed(fsm: &mut QlFsm, session_write_id: Option, success: bool) { + if let Some(write_id) = session_write_id { + fsm.complete_write(Instant::now(), write_id, success); + } + } + + fn drain_fsm_events(&mut self, fsm: &mut QlFsm, platform: &P) { + while let Some(event) = fsm.poll_event() { + log::trace!("polled FSM event: event={event:?}"); + match event { + Event::NewPeer => { + log::info!("new ql peer"); + if let Some(peer) = fsm.peer().cloned() { + platform.persist_peer(peer); + } + } + Event::PeerStatusChanged(status) => { + let peer = fsm.peer().map(|peer| peer.qid); + log::info!("peer status changed: peer={peer:?} status={status:?}"); + if status == ql_fsm::PeerStatus::Unpaired { + for (_, mut stream) in self.streams.drain() { + stream.fail_all(); + } + } + platform.handle_peer_status(peer, status); + } + Event::Opened { + stream_id, + route_id, + } => { + log::info!("inbound stream opened: stream_id={stream_id} route_id={route_id}"); + self.handle_opened_stream(fsm, platform, stream_id, route_id); + } + Event::Readable(stream_id) => { + log::trace!("stream readable: stream_id={stream_id}"); + self.handle_inbound_readable(fsm, stream_id); + } + Event::Writable(stream_id) => { + log::trace!("stream writable: stream_id={stream_id}"); + self.poll_stream(fsm, stream_id); + } + Event::Finished(stream_id) => { + log::info!("peer finished stream writes: stream_id={stream_id}"); + self.handle_inbound_finished(stream_id); + } + Event::OutboundFinished(stream_id) => { + log::info!("outbound finish acknowledged: stream_id={stream_id}"); + self.handle_outbound_finished(stream_id); + } + Event::Closed(frame) => { + self.handle_closed_stream(&frame); + } + Event::WritableClosed(frame) => { + self.handle_writable_closed(&frame); + } + Event::SessionClosed(close) => { + log::info!("session closed: frame={close:?}"); + for (_, mut stream) in self.streams.drain() { + stream.fail_all(); + } + } + } + } + } + + fn handle_opened_stream( + &mut self, + fsm: &mut QlFsm, + platform: &P, + stream_id: StreamId, + route_id: ql_wire::RouteId, + ) { + let Some(runtime_tx) = self.runtime_tx.upgrade() else { + log::warn!( + "dropping inbound stream because handle channel is unavailable: stream_id={stream_id}" + ); + if let Ok(mut stream) = fsm.stream(stream_id) { + stream.close(CloseTarget::Both, StreamCloseCode::CANCELLED); + } + return; + }; + + let (reader, writer, reader_io, writer_io) = io::new_stream( + stream_id, + CloseTarget::Origin, + CloseTarget::Return, + RuntimeHandle::new(runtime_tx), + ); + + self.streams.insert( + stream_id, + DriverStreamIo::new( + false, + Some(OutboundIo::new(writer_io)), + Some(InboundIo::new(reader_io)), + ), + ); + + log::info!( + "delivering inbound stream to platform: stream_id={stream_id} route_id={route_id}" + ); + platform.handle_inbound(QlStream { + stream_id, + route_id, + writer, + reader, + }); + } + + fn handle_inbound_readable(&mut self, fsm: &mut QlFsm, stream_id: StreamId) { + let Ok(mut stream_ops) = fsm.stream(stream_id) else { + log::info!("inbound readable for unknown stream: stream_id={stream_id}"); + return; + }; + let readable = stream_ops.readable_bytes(); + if readable == 0 { + return; + } + log::trace!("draining inbound bytes: stream_id={stream_id} readable={readable}"); + let mut accepted = 0usize; + let mut peer_closed = false; + let target; + { + let Some(stream) = self.streams.get_mut(&stream_id) else { + return; + }; + target = stream.inbound_target(); + for chunk in stream_ops.read() { + if chunk.is_empty() { + continue; + } + match stream.inbound_try_write(chunk) { + InboundWriteResult::Accepted(n) => { + accepted += n; + } + InboundWriteResult::Full => { + log::debug!( + "inbound backpressure: stream_id={stream_id} accepted={accepted}" + ); + break; + } + InboundWriteResult::Closed => { + log::warn!( + "inbound consumer closed; sending CANCELLED: stream_id={stream_id} target={target:?}" + ); + peer_closed = true; + break; + } + } + } + } + + if accepted > 0 { + log::trace!("committed inbound bytes: stream_id={stream_id:?} accepted={accepted}"); + stream_ops.commit_read(accepted).unwrap(); + } + if peer_closed { + stream_ops.close(target, StreamCloseCode::CANCELLED); + if let Entry::Occupied(entry) = self.streams.entry(stream_id) { + Self::try_reap_stream(entry); + } + } + + drop(stream_ops); + } + + fn handle_inbound_finished(&mut self, stream_id: StreamId) { + log::info!("inbound finished event: stream_id={stream_id}"); + let Entry::Occupied(mut entry) = self.streams.entry(stream_id) else { + return; + }; + log::info!("delivering clean inbound finish: stream_id={stream_id}"); + entry.get_mut().inbound_finish(); + Self::try_reap_stream(entry); + } + + fn handle_closed_stream(&mut self, frame: &ql_wire::StreamClose) { + log::info!( + "inbound close frame: stream_id={} target={:?} code={}", + frame.stream_id, + frame.target, + frame.code + ); + let Entry::Occupied(mut entry) = self.streams.entry(frame.stream_id) else { + return; + }; + let stream = entry.get_mut(); + + if frame.target == CloseTarget::Both || frame.target == stream.inbound_target() { + stream.inbound_fail(QlStreamError::StreamClosed { code: frame.code }); + } + if frame.target == CloseTarget::Both || frame.target == stream.outbound_target() { + stream.outbound_fail(QlStreamError::StreamClosed { code: frame.code }); + } + Self::try_reap_stream(entry); + } + + fn handle_writable_closed(&mut self, frame: &ql_wire::StreamClose) { + log::info!( + "writable close frame: stream_id={} target={:?} code={}", + frame.stream_id, + frame.target, + frame.code + ); + let Entry::Occupied(mut entry) = self.streams.entry(frame.stream_id) else { + return; + }; + let stream = entry.get_mut(); + stream.outbound_fail(QlStreamError::StreamClosed { code: frame.code }); + Self::try_reap_stream(entry); + } + + fn handle_outbound_finished(&mut self, stream_id: StreamId) { + log::info!("outbound finish acknowledged: stream_id={stream_id}"); + let Entry::Occupied(mut entry) = self.streams.entry(stream_id) else { + return; + }; + let stream = entry.get_mut(); + if !stream.outbound_finish_pending() { + return; + } + stream.outbound_finish(); + Self::try_reap_stream(entry); + } + + fn fill_write_slots<'a, P: QlPlatform + 'a>( + &self, + fsm: &mut QlFsm, + platform: &'a P, + in_flight: &mut Vec>>, + ) -> bool { + let mut filled = false; + while in_flight.len() < self.max_concurrent_message_writes { + let Some(write) = fsm.take_next_write(Instant::now(), platform) else { + break; + }; + filled = true; + log::trace!( + "queueing transport write: bytes={} write_id={:?}", + write.record.len(), + write.write_id + ); + in_flight.push(InFlightWrite { + write_id: write.write_id, + future: platform.write_message(write.record), + }); + } + filled + } + + fn poll_stream(&mut self, fsm: &mut QlFsm, stream_id: StreamId) { + let Entry::Occupied(mut entry) = self.streams.entry(stream_id) else { + return; + }; + let stream = entry.get_mut(); + let Some(writer_io) = stream.outbound_writer_mut() else { + log::trace!("poll stream skipped without outbound writer: stream_id={stream_id}"); + return; + }; + + if writer_io.is_finished() { + log::info!("observed outbound writer finished before write: stream_id={stream_id}"); + if let Ok(mut stream_ops) = fsm.stream(stream_id) { + if let Some(writer) = stream_ops.writer() { + writer.finish(); + } + } + stream.outbound_queue_finish(); + if stream.is_closed() { + entry.remove(); + } + return; + } + + let Ok(mut stream_ops) = fsm.stream(stream_id) else { + return; + }; + let Some(mut writer) = stream_ops.writer() else { + log::trace!("poll stream skipped without session writer: stream_id={stream_id}"); + return; + }; + + loop { + let capacity = writer.capacity(); + log::trace!("stream write capacity: stream_id={stream_id} capacity={capacity}"); + if capacity == 0 { + break; + } + + let Ok(mut bytes) = writer_io.try_read(capacity) else { + break; + }; + if bytes.is_empty() { + break; + } + + log::trace!( + "writing stream bytes: stream_id={stream_id} len={}", + bytes.len() + ); + let _ = writer.write(&mut bytes); + } + + if writer_io.is_finished() { + log::info!("observed outbound writer finished after write: stream_id={stream_id}"); + writer.finish(); + stream.outbound_queue_finish(); + if stream.is_closed() { + entry.remove(); + } + } + } + + fn try_reap_stream(entry: OccupiedEntry<'_, StreamId, DriverStreamIo>) { + if entry.get().is_closed() { + entry.remove(); + } + } +} diff --git a/ql-runtime/src/driver/state.rs b/ql-runtime/src/driver/state.rs new file mode 100644 index 00000000..0ff8eca8 --- /dev/null +++ b/ql-runtime/src/driver/state.rs @@ -0,0 +1,165 @@ +use std::collections::HashMap; + +use bytes::Bytes; +use ql_wire::{CloseTarget, StreamId}; + +use crate::{ + command::Command, + io::{PushError, Rx, Tx}, + QlStreamError, +}; + +pub struct DriverState { + pub streams: HashMap, + pub runtime_tx: async_channel::WeakSender, + pub max_concurrent_message_writes: usize, +} + +pub struct DriverStreamIo { + is_initiator: bool, + outbound: Option, + inbound: Option, +} + +impl DriverStreamIo { + pub fn new( + is_initiator: bool, + outbound: Option, + inbound: Option, + ) -> Self { + Self { + is_initiator, + outbound, + inbound, + } + } + + pub fn inbound_target(&self) -> CloseTarget { + if self.is_initiator { + CloseTarget::Return + } else { + CloseTarget::Origin + } + } + + pub fn outbound_target(&self) -> CloseTarget { + if self.is_initiator { + CloseTarget::Origin + } else { + CloseTarget::Return + } + } + + pub fn fail_all(&mut self) { + self.inbound_fail(QlStreamError::NoSession); + self.outbound_fail(QlStreamError::NoSession); + } + + pub fn is_closed(&self) -> bool { + self.outbound.is_none() && self.inbound.is_none() + } + + pub fn outbound_close(&mut self) { + self.outbound = None; + } + + pub fn outbound_finish(&mut self) { + if let Some(outbound) = self.outbound.take() { + outbound.tx.finish(); + } + } + + pub fn outbound_fail(&mut self, error: QlStreamError) { + if let Some(outbound) = self.outbound.take() { + let _ = outbound.tx.fail(error); + } + } + + pub fn outbound_writer_mut(&mut self) -> Option<&mut OutboundIo> { + self.outbound.as_mut() + } + + pub fn outbound_queue_finish(&mut self) { + if let Some(outbound) = self.outbound.as_mut() { + outbound.finish_pending = true; + } + } + + pub fn outbound_finish_pending(&self) -> bool { + self.outbound + .as_ref() + .is_some_and(|outbound| outbound.finish_pending) + } + + pub fn inbound_close(&mut self) { + self.inbound = None; + } + + pub fn inbound_try_write(&mut self, bytes: Bytes) -> InboundWriteResult { + let Some(inbound) = self.inbound.as_mut() else { + return InboundWriteResult::Closed; + }; + + let len = bytes.len(); + match inbound.rx.try_write(bytes) { + Ok(()) => InboundWriteResult::Accepted(len), + Err(PushError::Full(_)) => InboundWriteResult::Full, + Err(PushError::Closed(_)) => { + self.inbound = None; + InboundWriteResult::Closed + } + } + } + + pub fn inbound_finish(&mut self) { + if let Some(inbound) = self.inbound.take() { + inbound.rx.finish(); + } + } + + pub fn inbound_fail(&mut self, error: QlStreamError) { + if let Some(inbound) = self.inbound.take() { + inbound.rx.fail(error); + } + } +} + +pub struct OutboundIo { + tx: Tx, + pending: Bytes, + finish_pending: bool, +} + +impl OutboundIo { + pub fn new(tx: Tx) -> Self { + Self { + tx, + pending: Bytes::new(), + finish_pending: false, + } + } + + pub fn is_finished(&self) -> bool { + self.pending.is_empty() && self.tx.is_finished() + } + + pub fn try_read(&mut self, max_len: usize) -> Result { + self.tx.try_read(&mut self.pending, max_len) + } +} + +pub struct InboundIo { + rx: Rx, +} + +pub enum InboundWriteResult { + Accepted(usize), + Full, + Closed, +} + +impl InboundIo { + pub fn new(rx: Rx) -> Self { + Self { rx } + } +} diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs new file mode 100644 index 00000000..af4ab63a --- /dev/null +++ b/ql-runtime/src/driver/test.rs @@ -0,0 +1,207 @@ +use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto, StreamClose, QID}; + +use super::*; +use crate::{ + driver::state::{InboundIo, OutboundIo}, + io, + platform::QlInbound, +}; + +pub struct NoopTimer; +pub struct NoopInbound; + +impl crate::platform::QlTimer for NoopTimer { + fn set_deadline(self: Pin<&mut Self>, _deadline: Option) {} + + fn poll_wait(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<()> { + Poll::Pending + } +} + +impl QlPlatform for NoopCrypto { + type Timer = NoopTimer; + type WriteMessageFut<'a> = std::future::Ready; + type Inbound = NoopInbound; + + fn write_message(&self, _message: Vec) -> Self::WriteMessageFut<'_> { + std::future::ready(true) + } + + fn inbound(&mut self) -> Self::Inbound { + NoopInbound + } + + fn timer(&self) -> Self::Timer { + NoopTimer + } + + fn persist_peer(&self, _peer: PeerBundle) {} + + fn handle_peer_status(&self, _peer: Option, _status: ql_fsm::PeerStatus) {} + + fn handle_inbound(&self, _event: QlStream) {} +} + +impl QlInbound for NoopInbound { + fn poll_recv(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll> { + Poll::Pending + } +} + +fn new_driver_state() -> (DriverState, QlFsm) { + let (runtime_tx, _runtime_rx) = async_channel::unbounded(); + ( + DriverState { + streams: HashMap::new(), + runtime_tx: runtime_tx.downgrade(), + max_concurrent_message_writes: 1, + }, + QlFsm::new( + ql_fsm::QlFsmConfig::default(), + generate_identity(&SoftwareCrypto, "driver").unwrap(), + Instant::now(), + ), + ) +} + +fn new_inbound_io(capacity: usize) -> InboundIo { + let _ = capacity; + let (runtime_tx, _runtime_rx) = async_channel::unbounded(); + let stream = io::new_stream( + StreamId(99u32.into()), + CloseTarget::Origin, + CloseTarget::Return, + RuntimeHandle::new(runtime_tx), + ); + let (_, _, reader_io, _) = stream; + InboundIo::new(reader_io) +} + +fn new_outbound_io() -> OutboundIo { + let (runtime_tx, _runtime_rx) = async_channel::unbounded(); + let stream = io::new_stream( + StreamId(100u32.into()), + CloseTarget::Return, + CloseTarget::Origin, + RuntimeHandle::new(runtime_tx), + ); + let (_, _, _, writer_io) = stream; + OutboundIo::new(writer_io) +} + +#[test] +fn handle_inbound_finished_reaps_closed_initiator_stream() { + let (mut state, _fsm) = new_driver_state(); + let stream_id = StreamId(1u32.into()); + + state.streams.insert( + stream_id, + DriverStreamIo::new(true, None, Some(new_inbound_io(1))), + ); + + state.handle_inbound_finished(stream_id); + + assert!(!state.streams.contains_key(&stream_id)); +} + +#[test] +fn handle_closed_stream_reaps_when_both_halves_close() { + let (mut state, _fsm) = new_driver_state(); + let stream_id = StreamId(1u32.into()); + + state.streams.insert( + stream_id, + DriverStreamIo::new(false, Some(new_outbound_io()), Some(new_inbound_io(1))), + ); + + state.handle_closed_stream(&StreamClose { + stream_id, + target: CloseTarget::Both, + code: StreamCloseCode::CANCELLED, + }); + + assert!(!state.streams.contains_key(&stream_id)); +} + +#[test] +fn poll_stream_keeps_outbound_pending_after_local_finish_when_inbound_is_closed() { + let (mut state, mut fsm) = new_driver_state(); + let stream_id = StreamId(1u32.into()); + let (runtime_tx, _runtime_rx) = async_channel::unbounded(); + let (_, mut writer, _, writer_io) = io::new_stream( + stream_id, + CloseTarget::Return, + CloseTarget::Origin, + RuntimeHandle::new(runtime_tx), + ); + writer.queue_finish(); + state.streams.insert( + stream_id, + DriverStreamIo::new(true, Some(OutboundIo::new(writer_io)), None), + ); + + state.poll_stream(&mut fsm, stream_id); + + let stream = state.streams.get(&stream_id).unwrap(); + assert!(stream.outbound_finish_pending()); + assert!(!stream.is_closed()); +} + +#[test] +fn local_close_command_reaps_when_other_half_is_already_closed() { + let (mut state, mut fsm) = new_driver_state(); + let stream_id = StreamId(1u32.into()); + let (runtime_tx, _runtime_rx) = async_channel::unbounded(); + let (_, _, _, writer_io) = io::new_stream( + stream_id, + CloseTarget::Return, + CloseTarget::Origin, + RuntimeHandle::new(runtime_tx), + ); + + state.streams.insert( + stream_id, + DriverStreamIo::new(true, Some(OutboundIo::new(writer_io)), None), + ); + + state.drive_command( + &mut fsm, + Command::CloseStream { + stream_id, + target: CloseTarget::Origin, + code: StreamCloseCode::CANCELLED, + }, + &NoopCrypto, + ); + + assert!(!state.streams.contains_key(&stream_id)); +} + +#[test] +fn unpaired_status_fails_and_reaps_all_streams() { + let (mut state, mut fsm) = new_driver_state(); + let peer = generate_identity(&SoftwareCrypto, "peer").unwrap().bundle(); + let stream_id = StreamId(1u32.into()); + let (runtime_tx, _runtime_rx) = async_channel::unbounded(); + let (_, _, reader_io, writer_io) = io::new_stream( + stream_id, + CloseTarget::Origin, + CloseTarget::Return, + RuntimeHandle::new(runtime_tx), + ); + + state.streams.insert( + stream_id, + DriverStreamIo::new( + false, + Some(OutboundIo::new(writer_io)), + Some(InboundIo::new(reader_io)), + ), + ); + fsm.bind_peer(peer); + fsm.unpair(); + + state.drain_fsm_events(&mut fsm, &NoopCrypto); + + assert!(state.streams.is_empty()); +} diff --git a/ql-runtime/src/error.rs b/ql-runtime/src/error.rs new file mode 100644 index 00000000..5b74bcf8 --- /dev/null +++ b/ql-runtime/src/error.rs @@ -0,0 +1,18 @@ +use ql_wire::StreamCloseCode; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum QlStreamError { + StreamClosed { code: StreamCloseCode }, + NoSession, +} + +impl std::fmt::Display for QlStreamError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::StreamClosed { code } => write!(f, "stream closed {code:?}"), + Self::NoSession => f.write_str("no session"), + } + } +} + +impl std::error::Error for QlStreamError {} diff --git a/ql-runtime/src/handle/mod.rs b/ql-runtime/src/handle/mod.rs new file mode 100644 index 00000000..1782c17a --- /dev/null +++ b/ql-runtime/src/handle/mod.rs @@ -0,0 +1,96 @@ +use ql_fsm::{NoSessionError, PairingInvite}; +use ql_wire::{PairingToken, PeerBundle, RouteId, SessionCloseCode, StreamId}; + +use crate::command::Command; +pub use crate::io::{StreamReader, StreamWriter}; + +#[derive(Debug)] +pub struct QlStream { + pub stream_id: StreamId, + pub route_id: RouteId, + pub writer: StreamWriter, + pub reader: StreamReader, +} + +#[derive(Clone)] +pub struct RuntimeHandle { + tx: async_channel::Sender, +} + +impl RuntimeHandle { + /// binds the remote peer + pub fn bind_peer(&self, peer: PeerBundle) { + self.send(Command::BindPeer { peer }); + } + + /// starts an IK handshake with the bound peer + pub fn connect(&self) { + self.send(Command::Connect); + } + + /// arms acceptance of inbound xx pairings for a single token + pub fn arm_pairing(&self, token: PairingToken) { + self.send(Command::ArmPairing { token }); + } + + /// disarms inbound xx pairing + pub fn disarm_pairing(&self) { + self.send(Command::DisarmPairing); + } + + /// starts an outbound xx handshake using an out-of-band pairing invite + pub fn start_pairing(&self, invite: PairingInvite) { + self.send(Command::StartPairing { invite }); + } + + /// closes the current encrypted session + pub fn close_session(&self, code: SessionCloseCode) { + self.send(Command::CloseSession { code }); + } + + /// forgets the currently bound peer and initiates session unpairing if connected + pub fn unpair(&self) { + self.send(Command::Unpair); + } + + /// opens a new stream on the active encrypted session + pub async fn open_stream(&self, route_id: RouteId) -> Result { + let (start_tx, start_rx) = oneshot::channel(); + + self.send(Command::OpenStream { + route_id, + start: start_tx, + }); + + // runtime cannot be shutdown while we have a handle + let (stream_id, reader, writer) = start_rx.await.unwrap()?; + + Ok(QlStream { + stream_id, + route_id, + writer, + reader, + }) + } + + #[cfg(feature = "rpc")] + pub fn rpc(&self) -> crate::rpc::RpcHandle { + crate::rpc::RpcHandle::new(self.clone()) + } +} + +impl RuntimeHandle { + pub(crate) fn new(tx: async_channel::Sender) -> Self { + Self { tx } + } + + #[inline] + #[track_caller] + pub(crate) fn send(&self, cmd: Command) { + self.tx.try_send(cmd).expect("runtime is alive"); + } + + pub(crate) fn try_send(&self, cmd: Command) -> bool { + self.tx.try_send(cmd).is_ok() + } +} diff --git a/ql-runtime/src/io/inner.rs b/ql-runtime/src/io/inner.rs new file mode 100644 index 00000000..64df6ced --- /dev/null +++ b/ql-runtime/src/io/inner.rs @@ -0,0 +1,643 @@ +//! per-stream shared io state +//! each lane has one slot and one waker +//! the low slot bits belong to `slot.rs` and the higher bits here carry lane-specific flags + +use std::task::Waker; + +use bytes::Bytes; +use diatomic_waker::DiatomicWaker; +use ql_wire::StreamId; + +use super::{ + slot::{PopError, PushError, Slot}, + sync::Arc, +}; +use crate::QlStreamError; + +pub(super) fn new(stream_id: StreamId) -> Arc { + Arc::new(Inner { + stream_id, + rx: RxInner::new(), + tx: TxInner::new(), + }) +} + +pub(super) struct Inner { + pub(super) stream_id: StreamId, + pub(super) rx: RxInner, + pub(super) tx: TxInner, +} + +pub enum Item { + Chunk(Bytes), + Error(QlStreamError), +} + +#[derive(Debug, PartialEq, Eq)] +pub struct ForcePushError(pub T); + +/// reader-lane shared state +pub struct RxInner { + slot: Slot, + changed: DiatomicWaker, +} + +impl RxInner { + const FINISHED: usize = 1 << 2; + + fn new() -> Self { + Self { + slot: Slot::new(), + changed: DiatomicWaker::new(), + } + } + + pub fn try_write(&self, bytes: Bytes) -> Result<(), PushError> { + try_write_chunk(&self.slot, &self.changed, bytes, Self::FINISHED) + } + + /// marks clean reader eof + pub fn finish(&self) { + if self.slot.fetch_or(Self::FINISHED) & Self::FINISHED == 0 { + self.changed.notify(); + } + } + + /// stores a terminal reader error + pub fn fail(&self, error: QlStreamError) -> Option { + let displaced = self.slot.force_push(Item::Error(error)); + self.changed.notify(); + displaced_bytes(displaced) + } + + pub fn load_state(&self) -> usize { + self.slot.load_state() + } + + pub fn is_finished(state: usize) -> bool { + state & Self::FINISHED != 0 + } + + pub fn pop(&self) -> Result { + pop_item(&self.slot, &self.changed) + } + + /// registers the sole reader-lane waiter + pub fn register_waiter(&self, waker: &Waker) { + // Safety: StreamReader is the only reader-lane registrar for this + // shared state, so register/unregister never run concurrently. + unsafe { self.changed.register(waker) }; + } + + /// unregisters the sole reader-lane waiter + pub fn unregister_waiter(&self) { + // Safety: StreamReader is the only reader-lane registrar for this + // shared state, so register/unregister never run concurrently. + unsafe { self.changed.unregister() }; + } +} + +/// writer-lane shared state +/// +/// finish and fail race to establish the terminal result +/// terminal errors are stored in the slot +pub struct TxInner { + slot: Slot, + changed: DiatomicWaker, +} + +impl TxInner { + const FINISH_REQUESTED: usize = 1 << 2; + const TERMINAL_READY: usize = 1 << 3; + const TERMINAL_OK: usize = 1 << 4; + + fn new() -> Self { + Self { + slot: Slot::new(), + changed: DiatomicWaker::new(), + } + } + + pub fn load_state(&self) -> usize { + self.slot.load_state() + } + + pub fn finish_requested(state: usize) -> bool { + state & Self::FINISH_REQUESTED != 0 + } + + pub fn terminal_ready(state: usize) -> bool { + state & Self::TERMINAL_READY != 0 + } + + pub fn terminal_ok(state: usize) -> bool { + state & Self::TERMINAL_OK != 0 + } + + pub fn try_write(&self, bytes: Bytes) -> Result<(), PushError> { + try_write_chunk( + &self.slot, + &self.changed, + bytes, + Self::FINISH_REQUESTED | Self::TERMINAL_READY, + ) + } + + /// prevents future chunk writes once observed + pub fn request_finish(&self) { + if self.slot.fetch_or(Self::FINISH_REQUESTED) & Self::FINISH_REQUESTED == 0 { + self.changed.notify(); + } + } + + /// commits a clean writer eof + pub fn finish(&self) { + let mut state = self.slot.load_state(); + loop { + if Self::terminal_ready(state) { + return; + } + + let new_state = state | Self::TERMINAL_READY | Self::TERMINAL_OK; + match self.slot.compare_exchange(state, new_state) { + Ok(()) => { + self.changed.notify(); + return; + } + Err(actual) => state = actual, + } + } + } + + /// stores a terminal writer error + /// futures calls will have no effect + pub fn fail( + &self, + error: QlStreamError, + ) -> Result, ForcePushError> { + let mut state = self.slot.load_state(); + loop { + if Self::terminal_ready(state) { + return Err(ForcePushError(error)); + } + + let new_state = state | Self::TERMINAL_READY; + match self.slot.compare_exchange(state, new_state) { + Ok(()) => break, + Err(actual) => state = actual, + } + } + + let displaced = self.slot.force_push(Item::Error(error)); + self.changed.notify(); + Ok(displaced_bytes(displaced)) + } + + pub fn pop(&self) -> Result { + pop_item(&self.slot, &self.changed) + } + + /// registers the sole writer-lane waiter + pub fn register_waiter(&self, waker: &Waker) { + // Safety: StreamWriter is the only writer-lane registrar for this + // shared state, so register/unregister never run concurrently. + unsafe { self.changed.register(waker) }; + } + + /// unregisters the sole writer-lane waiter + pub fn unregister_waiter(&self) { + // Safety: StreamWriter is the only writer-lane registrar for this + // shared state, so register/unregister never run concurrently. + unsafe { self.changed.unregister() }; + } + + /// returns true once finish was requested and buffered data is drained + pub fn is_finished(&self) -> bool { + let state = self.load_state(); + Self::finish_requested(state) && Slot::::is_empty_state(state) + } + + pub fn try_read(&self, pending: &mut Bytes, max_len: usize) -> Result { + if !pending.is_empty() { + return Ok(if pending.len() <= max_len { + std::mem::take(pending) + } else { + pending.split_to(max_len) + }); + } + + let state = self.load_state(); + if Self::terminal_ready(state) { + return Err(()); + } + + match self.pop() { + Ok(Item::Chunk(mut bytes)) => { + if bytes.len() <= max_len { + Ok(bytes) + } else { + let head = bytes.split_to(max_len); + *pending = bytes; + Ok(head) + } + } + Ok(Item::Error(_)) => Err(()), + Err(PopError) => Ok(Bytes::new()), + } + } +} + +#[inline] +fn try_write_chunk( + slot: &Slot, + changed: &DiatomicWaker, + bytes: Bytes, + closed_mask: usize, +) -> Result<(), PushError> { + match slot.try_push(Item::Chunk(bytes), closed_mask) { + Ok(()) => { + changed.notify(); + Ok(()) + } + Err(PushError::Closed(Item::Chunk(bytes))) => Err(PushError::Closed(bytes)), + Err(PushError::Full(Item::Chunk(bytes))) => Err(PushError::Full(bytes)), + Err(PushError::Closed(Item::Error(_)) | PushError::Full(Item::Error(_))) => { + unreachable!("chunk write cannot recover an error payload") + } + } +} + +#[inline] +fn displaced_bytes(displaced: Option) -> Option { + match displaced { + Some(Item::Chunk(bytes)) => Some(bytes), + Some(Item::Error(_)) | None => None, + } +} + +#[inline] +fn pop_item(slot: &Slot, changed: &DiatomicWaker) -> Result { + match slot.pop() { + item @ Ok(Item::Chunk(_)) => { + changed.notify(); + item + } + item @ (Ok(Item::Error(_)) | Err(_)) => item, + } +} + +#[cfg(all(test, loom))] +mod loom_tests { + use std::task::Waker; + + use bytes::Bytes; + use loom::thread; + use ql_wire::StreamCloseCode; + + use super::*; + use crate::{ + io::{sync::loom::*, Tx}, + QlStreamError, + }; + + #[test] + fn reader_waiter_registration_survives_finish() { + check_model(|| { + let shared = shared(); + shared.rx.register_waiter(Waker::noop()); + + let finisher = { + let shared = shared.clone(); + thread::spawn(move || { + shared.rx.finish(); + }) + }; + + finisher.join().unwrap(); + assert!(RxInner::is_finished(shared.rx.load_state())); + + shared.rx.unregister_waiter(); + }); + } + + #[test] + fn reader_chunk_remains_available_after_finish() { + check_model(|| { + let shared = shared(); + + let producer = { + let shared = shared.clone(); + thread::spawn(move || { + shared.rx.try_write(Bytes::from_static(b"abc")).unwrap(); + shared.rx.finish(); + }) + }; + + producer.join().unwrap(); + + match shared.rx.pop() { + Ok(Item::Chunk(bytes)) => assert_eq!(bytes, Bytes::from_static(b"abc")), + _ => panic!("expected buffered reader chunk"), + } + assert!(RxInner::is_finished(shared.rx.load_state())); + assert!(matches!(shared.rx.pop(), Err(PopError))); + }); + } + + #[test] + fn reader_rejects_write_after_finish() { + check_model(|| { + let shared = shared(); + + shared.rx.finish(); + + assert_eq!( + shared.rx.try_write(Bytes::from_static(b"abc")), + Err(PushError::Closed(Bytes::from_static(b"abc"))) + ); + assert!(RxInner::is_finished(shared.rx.load_state())); + assert!(matches!(shared.rx.pop(), Err(PopError))); + }); + } + + #[test] + fn reader_write_races_with_finish_has_coherent_outcome() { + check_model(|| { + let shared = shared(); + + let writer = { + let shared = shared.clone(); + thread::spawn(move || shared.rx.try_write(Bytes::from_static(b"abc"))) + }; + let finisher = { + let shared = shared.clone(); + thread::spawn(move || shared.rx.finish()) + }; + + let write_result = writer.join().unwrap(); + finisher.join().unwrap(); + + assert!(RxInner::is_finished(shared.rx.load_state())); + match write_result { + Ok(()) => match shared.rx.pop() { + Ok(Item::Chunk(bytes)) => assert_eq!(bytes, Bytes::from_static(b"abc")), + _ => panic!("expected buffered reader chunk"), + }, + Err(PushError::Closed(bytes)) => { + assert_eq!(bytes, Bytes::from_static(b"abc")); + assert!(matches!(shared.rx.pop(), Err(PopError))); + return; + } + Err(PushError::Full(_)) => panic!("empty reader slot must not report full"), + } + assert!(matches!(shared.rx.pop(), Err(PopError))); + }); + } + + #[test] + fn reader_fail_racing_with_pop_preserves_terminal_outcome() { + check_model(|| { + let shared = shared(); + shared.rx.try_write(Bytes::from_static(b"abc")).unwrap(); + + let popper = { + let shared = shared.clone(); + thread::spawn(move || shared.rx.pop()) + }; + let failer = { + let shared = shared.clone(); + thread::spawn(move || { + shared.rx.fail(QlStreamError::StreamClosed { + code: StreamCloseCode::CANCELLED, + }) + }) + }; + + let pop_result = popper.join().unwrap(); + let fail_result = failer.join().unwrap(); + + match (pop_result, fail_result) { + (Ok(Item::Chunk(bytes)), None) => { + assert_eq!(bytes, Bytes::from_static(b"abc")); + match shared.rx.pop() { + Ok(Item::Error(QlStreamError::StreamClosed { code })) => { + assert_eq!(code, StreamCloseCode::CANCELLED); + } + _ => panic!("expected terminal reader error"), + } + } + (Ok(Item::Error(QlStreamError::StreamClosed { code })), Some(bytes)) => { + assert_eq!(code, StreamCloseCode::CANCELLED); + assert_eq!(bytes, Bytes::from_static(b"abc")); + assert!(matches!(shared.rx.pop(), Err(PopError))); + } + _ => panic!("unexpected reader fail/pop race outcome"), + } + }); + } + + #[test] + fn writer_is_finished_only_after_drain() { + check_model(|| { + let shared = shared(); + let tx = Tx(shared.clone()); + let mut pending = Bytes::new(); + + shared.tx.try_write(Bytes::from_static(b"abc")).unwrap(); + shared.tx.request_finish(); + + assert!(!(pending.is_empty() && tx.is_finished())); + assert_eq!(tx.try_read(&mut pending, 2), Ok(Bytes::from_static(b"ab"))); + assert!(!(pending.is_empty() && tx.is_finished())); + assert_eq!(tx.try_read(&mut pending, 8), Ok(Bytes::from_static(b"c"))); + assert!(pending.is_empty() && tx.is_finished()); + }); + } + + #[test] + fn writer_write_races_with_request_finish() { + check_model(|| { + let shared = shared(); + let tx = Tx(shared.clone()); + let mut pending = Bytes::new(); + + let writer = { + let shared = shared.clone(); + thread::spawn(move || shared.tx.try_write(Bytes::from_static(b"abc"))) + }; + let finisher = { + let shared = shared.clone(); + thread::spawn(move || shared.tx.request_finish()) + }; + + let write_result = writer.join().unwrap(); + finisher.join().unwrap(); + + assert!(TxInner::finish_requested(shared.tx.load_state())); + match write_result { + Ok(()) => { + assert_eq!(tx.try_read(&mut pending, 8), Ok(Bytes::from_static(b"abc"))); + assert!(pending.is_empty() && tx.is_finished()); + } + Err(PushError::Closed(bytes)) => { + assert_eq!(bytes, Bytes::from_static(b"abc")); + assert!(pending.is_empty() && tx.is_finished()); + } + Err(PushError::Full(_)) => panic!("empty writer slot must not report full"), + } + }); + } + + #[test] + fn writer_fail_overwrites_buffered_chunk_and_keeps_terminal_state_observable() { + check_model(|| { + let shared = shared(); + shared.tx.try_write(Bytes::from_static(b"abc")).unwrap(); + shared.tx.register_waiter(Waker::noop()); + + let failer = { + let shared = shared.clone(); + thread::spawn(move || { + let displaced = shared.tx.fail(QlStreamError::StreamClosed { + code: StreamCloseCode::CANCELLED, + }); + assert_eq!(displaced.unwrap(), Some(Bytes::from_static(b"abc"))); + }) + }; + + failer.join().unwrap(); + + assert!(TxInner::terminal_ready(shared.tx.load_state())); + shared.tx.unregister_waiter(); + match shared.tx.pop() { + Ok(Item::Error(QlStreamError::StreamClosed { code })) => { + assert_eq!(code, StreamCloseCode::CANCELLED); + } + _ => panic!("expected terminal writer error"), + } + }); + } + + #[test] + fn reader_waiter_registration_can_be_reused_after_notification() { + check_model(|| { + let shared = shared(); + + shared.rx.register_waiter(Waker::noop()); + shared.rx.try_write(Bytes::from_static(b"abc")).unwrap(); + match shared.rx.pop() { + Ok(Item::Chunk(bytes)) => assert_eq!(bytes, Bytes::from_static(b"abc")), + _ => panic!("expected buffered reader chunk"), + } + + shared.rx.register_waiter(Waker::noop()); + shared.rx.finish(); + assert!(RxInner::is_finished(shared.rx.load_state())); + shared.rx.unregister_waiter(); + }); + } + + #[test] + fn writer_waiter_registration_can_be_reused_after_notification() { + check_model(|| { + let shared = shared(); + + shared.tx.register_waiter(Waker::noop()); + shared.tx.try_write(Bytes::from_static(b"abc")).unwrap(); + match shared.tx.pop() { + Ok(Item::Chunk(bytes)) => assert_eq!(bytes, Bytes::from_static(b"abc")), + _ => panic!("expected buffered writer chunk"), + } + + shared.tx.register_waiter(Waker::noop()); + shared.tx.finish(); + assert!(TxInner::terminal_ready(shared.tx.load_state())); + shared.tx.unregister_waiter(); + }); + } + + #[test] + fn writer_write_races_with_fail() { + check_model(|| { + let shared = shared(); + + let writer = { + let shared = shared.clone(); + thread::spawn(move || shared.tx.try_write(Bytes::from_static(b"abc"))) + }; + let failer = { + let shared = shared.clone(); + thread::spawn(move || { + shared.tx.fail(QlStreamError::StreamClosed { + code: StreamCloseCode::CANCELLED, + }) + }) + }; + + let write_result = writer.join().unwrap(); + let fail_result = failer.join().unwrap(); + + assert!(TxInner::terminal_ready(shared.tx.load_state())); + match (&write_result, &fail_result) { + (Ok(()), Ok(Some(bytes))) => { + assert_eq!(Bytes::from_static(b"abc"), bytes.clone()); + } + (Err(PushError::Closed(bytes)), Ok(None)) => { + assert_eq!(Bytes::from_static(b"abc"), bytes.clone()); + } + (Err(PushError::Full(bytes)), Ok(None)) => { + assert_eq!(Bytes::from_static(b"abc"), bytes.clone()); + } + _ => panic!( + "unexpected writer fail/write race outcome: write={write_result:?} fail={fail_result:?}" + ), + } + + match shared.tx.pop() { + Ok(Item::Error(QlStreamError::StreamClosed { code })) => { + assert_eq!(code, StreamCloseCode::CANCELLED); + } + _ => panic!("expected terminal writer error"), + } + }); + } + + #[test] + fn writer_finish_races_with_fail_without_masking_error() { + check_model(|| { + let shared = shared(); + + let finisher = { + let shared = shared.clone(); + thread::spawn(move || shared.tx.finish()) + }; + let failer = { + let shared = shared.clone(); + thread::spawn(move || { + shared.tx.fail(QlStreamError::StreamClosed { + code: StreamCloseCode::CANCELLED, + }) + }) + }; + + finisher.join().unwrap(); + let fail_result = failer.join().unwrap(); + + assert!(TxInner::terminal_ready(shared.tx.load_state())); + match fail_result { + Err(_) => { + assert!(TxInner::terminal_ok(shared.tx.load_state())); + } + Ok(_) => { + assert!(!TxInner::terminal_ok(shared.tx.load_state())); + match shared.tx.pop() { + Ok(Item::Error(QlStreamError::StreamClosed { code })) => { + assert_eq!(code, StreamCloseCode::CANCELLED); + } + _ => panic!("expected terminal writer error"), + } + } + } + }); + } +} diff --git a/ql-runtime/src/io/mod.rs b/ql-runtime/src/io/mod.rs new file mode 100644 index 00000000..2eb7f0f0 --- /dev/null +++ b/ql-runtime/src/io/mod.rs @@ -0,0 +1,59 @@ +mod inner; +mod reader; +mod slot; +mod sync; +mod writer; + +use std::ops::Deref; + +use ql_wire::{CloseTarget, StreamId}; + +pub use self::{reader::StreamReader, slot::PushError, writer::StreamWriter}; +use crate::RuntimeHandle; + +pub struct Rx(sync::Arc); + +impl Deref for Rx { + type Target = inner::RxInner; + + fn deref(&self) -> &Self::Target { + &self.0.rx + } +} + +impl Rx { + pub fn stream_id(&self) -> StreamId { + self.0.stream_id + } +} + +pub struct Tx(sync::Arc); + +impl Deref for Tx { + type Target = inner::TxInner; + + fn deref(&self) -> &Self::Target { + &self.0.tx + } +} + +impl Tx { + pub fn stream_id(&self) -> StreamId { + self.0.stream_id + } +} + +pub fn new_stream( + stream_id: StreamId, + reader_target: CloseTarget, + writer_target: CloseTarget, + handle: RuntimeHandle, +) -> (StreamReader, StreamWriter, Rx, Tx) { + let shared = inner::new(stream_id); + ( + StreamReader::new(Rx(shared.clone()), reader_target, handle.clone()), + StreamWriter::new(Tx(shared.clone()), writer_target, handle), + Rx(shared.clone()), + Tx(shared), + ) +} diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs new file mode 100644 index 00000000..8c40ccd3 --- /dev/null +++ b/ql-runtime/src/io/reader.rs @@ -0,0 +1,236 @@ +use std::{ + future::poll_fn, + task::{Context, Poll}, +}; + +use bytes::Bytes; +use ql_wire::{CloseTarget, StreamCloseCode}; + +use super::{ + inner::{Item, RxInner}, + slot::PopError, + Rx, +}; +use crate::{command::Command, log, QlStreamError, RuntimeHandle}; + +pub struct StreamReader { + rx: Rx, + target: CloseTarget, + pending: Bytes, + terminal: ReaderTerminalState, + handle: RuntimeHandle, +} + +enum ReaderTerminalState { + Open, + Delivered, +} + +unsafe impl Sync for StreamReader {} + +impl std::fmt::Debug for StreamReader { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("StreamReader") + .field("stream_id", &self.rx.stream_id()) + .field("target", &self.target) + .field( + "terminal", + &matches!(self.terminal, ReaderTerminalState::Delivered), + ) + .finish_non_exhaustive() + } +} + +impl StreamReader { + pub(crate) fn new(shared: Rx, target: CloseTarget, handle: RuntimeHandle) -> Self { + Self { + rx: shared, + target, + pending: Bytes::new(), + terminal: ReaderTerminalState::Open, + handle, + } + } + + pub fn poll_read( + &mut self, + max_len: usize, + cx: &mut Context<'_>, + ) -> Poll, QlStreamError>> { + if matches!(self.terminal, ReaderTerminalState::Delivered) { + return Poll::Ready(Ok(None)); + } + + match self.try_read_ready(max_len) { + Poll::Ready(result) => return Poll::Ready(result), + Poll::Pending => {} + } + + self.rx.register_waiter(cx.waker()); + + match self.try_read_ready(max_len) { + Poll::Ready(result) => { + self.rx.unregister_waiter(); + Poll::Ready(result) + } + Poll::Pending => Poll::Pending, + } + } + + fn try_read_ready(&mut self, max_len: usize) -> Poll, QlStreamError>> { + if !self.pending.is_empty() { + let pending = &mut self.pending; + let bytes = if pending.len() <= max_len { + std::mem::take(pending) + } else { + pending.split_to(max_len) + }; + self.handle.try_send(Command::PollInbound { + stream_id: self.rx.stream_id(), + }); + return Poll::Ready(Ok(Some(bytes))); + } + + match self.rx.pop() { + Ok(Item::Chunk(mut bytes)) => { + log::trace!( + "byte reader received chunk: stream_id={} target={:?} len={}", + self.rx.stream_id(), + self.target, + bytes.len() + ); + self.handle.try_send(Command::PollInbound { + stream_id: self.rx.stream_id(), + }); + if bytes.len() <= max_len { + return Poll::Ready(Ok(Some(bytes))); + } + let head = bytes.split_to(max_len); + self.pending = bytes; + Poll::Ready(Ok(Some(head))) + } + Ok(Item::Error(error)) => { + log::debug!( + "byte reader delivered terminal error: stream_id={} target={:?} error={:?}", + self.rx.stream_id(), + self.target, + error + ); + self.terminal = ReaderTerminalState::Delivered; + Poll::Ready(Err(error)) + } + Err(PopError) => { + if RxInner::is_finished(self.rx.load_state()) { + log::debug!( + "byte reader delivered clean eof: stream_id={} target={:?}", + self.rx.stream_id(), + self.target + ); + self.terminal = ReaderTerminalState::Delivered; + return Poll::Ready(Ok(None)); + } + Poll::Pending + } + } + } + + pub fn poll_read_chunk( + &mut self, + cx: &mut Context<'_>, + ) -> Poll, QlStreamError>> { + self.poll_read(usize::MAX, cx) + } + + pub async fn read(&mut self, max_len: usize) -> Result, QlStreamError> { + poll_fn(|cx| self.poll_read(max_len, cx)).await + } + + pub async fn read_chunk(&mut self) -> Result, QlStreamError> { + self.read(usize::MAX).await + } + + pub fn close(mut self, code: StreamCloseCode) { + self.close_inner(code); + } + + fn close_inner(&mut self, code: StreamCloseCode) { + if matches!(self.terminal, ReaderTerminalState::Delivered) { + return; + } + log::debug!( + "byte reader explicit close: stream_id={:?} target={:?} code={:?}", + self.rx.stream_id(), + self.target, + code + ); + self.terminal = ReaderTerminalState::Delivered; + self.handle.try_send(Command::CloseStream { + stream_id: self.rx.stream_id(), + target: self.target, + code, + }); + } +} + +impl Drop for StreamReader { + fn drop(&mut self) { + if matches!(self.terminal, ReaderTerminalState::Delivered) { + return; + } + log::debug!( + "byte reader drop close: stream_id={:?} target={:?} code={:?}", + self.rx.stream_id(), + self.target, + StreamCloseCode::CANCELLED + ); + self.handle.try_send(Command::CloseStream { + stream_id: self.rx.stream_id(), + target: self.target, + code: StreamCloseCode::CANCELLED, + }); + } +} + +#[cfg(all(test, loom))] +mod loom_tests { + use std::task::{Context, Poll, Waker}; + + use bytes::Bytes; + use loom::thread; + use ql_wire::CloseTarget; + + use super::*; + use crate::io::sync::loom::*; + + #[test] + fn poll_read_observes_chunk_racing_with_registration() { + check_model(|| { + let inner = shared(); + let mut reader = StreamReader::new(Rx(inner.clone()), CloseTarget::Origin, handle()); + let mut cx = Context::from_waker(Waker::noop()); + + let producer = { + let inner = inner.clone(); + thread::spawn(move || { + inner.rx.try_write(Bytes::from_static(b"abc")).unwrap(); + }) + }; + + let first = reader.poll_read(usize::MAX, &mut cx); + producer.join().unwrap(); + + match first { + Poll::Ready(Ok(Some(bytes))) => { + assert_eq!(bytes, Bytes::from_static(b"abc")); + } + Poll::Pending => { + assert_eq!( + reader.poll_read(usize::MAX, &mut cx), + Poll::Ready(Ok(Some(Bytes::from_static(b"abc")))) + ); + } + other => panic!("unexpected first poll result: {other:?}"), + } + }); + } +} diff --git a/ql-runtime/src/io/slot.rs b/ql-runtime/src/io/slot.rs new file mode 100644 index 00000000..f71f1b0c --- /dev/null +++ b/ql-runtime/src/io/slot.rs @@ -0,0 +1,175 @@ +//! local single-slot queue for stream io +//! copied from `concurrent_queue::single::Single` in `concurrent-queue` + +use core::mem::MaybeUninit; + +#[allow(clippy::wildcard_imports)] +use super::sync::*; + +const LOCKED: usize = 1 << 0; +const PUSHED: usize = 1 << 1; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct PopError; + +#[derive(Debug, PartialEq, Eq)] +pub enum PushError { + Full(T), + Closed(T), +} + +/// A single-element queue. +pub struct Slot { + state: AtomicUsize, + value: UnsafeCell>, +} + +unsafe impl Send for Slot {} +unsafe impl Sync for Slot {} + +impl Slot { + /// Creates a new single-element queue. + pub fn new() -> Self { + Self { + state: AtomicUsize::new(0), + value: UnsafeCell::new(MaybeUninit::uninit()), + } + } + + #[inline] + pub fn load_state(&self) -> usize { + self.state.load(Ordering::Acquire) + } + + #[inline] + pub fn fetch_or(&self, bits: usize) -> usize { + self.state.fetch_or(bits, Ordering::Release) + } + + #[inline] + pub fn compare_exchange(&self, current: usize, new: usize) -> Result<(), usize> { + self.state + .compare_exchange(current, new, Ordering::AcqRel, Ordering::Acquire) + .map(|_| ()) + } + + /// Attempts to push an item into the queue. + pub fn try_push(&self, value: T, closed_mask: usize) -> Result<(), PushError> { + let mut state = self.load_state(); + loop { + if state & closed_mask != 0 { + return Err(PushError::Closed(value)); + } + if state & LOCKED != 0 { + busy_wait(); + state = self.load_state(); + continue; + } + if state & PUSHED != 0 { + return Err(PushError::Full(value)); + } + + // Lock and fill the slot. + let new_state = state | LOCKED | PUSHED; + match self.compare_exchange(state, new_state) { + Ok(()) => { + // Write the value and unlock. + self.value.with_mut(|slot| unsafe { + slot.write(MaybeUninit::new(value)); + }); + self.state.fetch_and(!LOCKED, Ordering::Release); + return Ok(()); + } + Err(actual) => state = actual, + } + } + } + + /// Attempts to push an item into the queue, displacing another if necessary. + pub fn force_push(&self, value: T) -> Option { + // Attempt to lock the slot. + let mut state = self.load_state(); + + loop { + if state & LOCKED != 0 { + busy_wait(); + state = self.load_state(); + continue; + } + + // Lock the slot. + let new_state = state | LOCKED | PUSHED; + match self.compare_exchange(state, new_state) { + Ok(()) => { + // If the value was pushed, swap out the value. + let displaced = if state & PUSHED == 0 { + // SAFETY: write is safe because we have locked the state. + self.value.with_mut(|slot| unsafe { + slot.write(MaybeUninit::new(value)); + }); + None + } else { + // SAFETY: replace is safe because we have locked the state, and + // assume_init is safe because we have checked that the value was pushed. + self.value.with_mut(move |slot| unsafe { + Some(std::ptr::replace(slot, MaybeUninit::new(value)).assume_init()) + }) + }; + + // We can unlock the slot now. + self.state.fetch_and(!LOCKED, Ordering::Release); + return displaced; + } + Err(actual) => state = actual, + } + } + } + + /// Attempts to pop an item from the queue. + pub fn pop(&self) -> Result { + let mut state = PUSHED; + loop { + if state & LOCKED != 0 { + busy_wait(); + state = self.load_state(); + continue; + } + if state & PUSHED == 0 { + return Err(PopError); + } + + // Lock and empty the slot. + let new_state = (state | LOCKED) & !PUSHED; + match self.compare_exchange(state, new_state) { + Ok(()) => { + // Read the value and unlock. + let value = self + .value + .with_mut(|slot| unsafe { slot.read().assume_init() }); + self.state.fetch_and(!LOCKED, Ordering::Release); + return Ok(value); + } + Err(actual) => state = actual, + } + } + } + + #[inline] + pub fn is_empty_state(state: usize) -> bool { + state & PUSHED == 0 + } +} + +impl Drop for Slot { + fn drop(&mut self) { + // Drop the value in the slot. + self.state.with_mut(|state| { + if *state & PUSHED != 0 { + self.value.with_mut(|slot| unsafe { + let value = &mut *slot; + value.as_mut_ptr().drop_in_place(); + }); + } + }); + } +} diff --git a/ql-runtime/src/io/sync.rs b/ql-runtime/src/io/sync.rs new file mode 100644 index 00000000..c5034076 --- /dev/null +++ b/ql-runtime/src/io/sync.rs @@ -0,0 +1,89 @@ +#[cfg(not(all(test, loom)))] +mod inner { + pub use std::{ + cell::UnsafeCell, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }, + }; + + pub fn busy_wait() { + std::thread::yield_now(); + } + + pub trait UnsafeCellExt { + type Value; + + fn with_mut(&self, f: F) -> R + where + F: FnOnce(*mut Self::Value) -> R; + } + + impl UnsafeCellExt for UnsafeCell { + type Value = T; + + fn with_mut(&self, f: F) -> R + where + F: FnOnce(*mut Self::Value) -> R, + { + f(self.get()) + } + } + + pub trait AtomicExt { + type Value; + + fn with_mut(&mut self, f: F) -> R + where + F: FnOnce(&mut Self::Value) -> R; + } + + impl AtomicExt for AtomicUsize { + type Value = usize; + + fn with_mut(&mut self, f: F) -> R + where + F: FnOnce(&mut Self::Value) -> R, + { + f(self.get_mut()) + } + } +} + +#[cfg(all(test, loom))] +mod inner { + pub use loom::{ + cell::UnsafeCell, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, + }, + thread::yield_now as busy_wait, + }; +} + +pub use inner::*; + +#[cfg(all(test, loom))] +pub(crate) mod loom { + use loom::model; + use ql_wire::StreamId; + + use super::Arc; + use crate::{io::inner::Inner, RuntimeHandle}; + + pub(crate) fn check_model(f: impl Fn() + Sync + Send + 'static) { + let builder = model::Builder::new(); + builder.check(f); + } + + pub(crate) fn shared() -> Arc { + crate::io::inner::new(StreamId(1u32.into())) + } + + pub(crate) fn handle() -> RuntimeHandle { + let (tx, _rx) = async_channel::unbounded(); + RuntimeHandle::new(tx) + } +} diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs new file mode 100644 index 00000000..cfad3196 --- /dev/null +++ b/ql-runtime/src/io/writer.rs @@ -0,0 +1,294 @@ +use std::{ + future::poll_fn, + task::{Context, Poll}, +}; + +use bytes::Bytes; +use ql_wire::{CloseTarget, StreamCloseCode}; + +use super::{ + inner::{Item, TxInner}, + slot::PopError, + PushError, Tx, +}; +use crate::{command::Command, log, QlStreamError, RuntimeHandle}; + +pub struct StreamWriter { + tx: Tx, + target: CloseTarget, + open: bool, + terminal: WriterTerminalState, + handle: RuntimeHandle, +} + +enum WriterTerminalState { + Pending, + Terminal(Result<(), QlStreamError>), +} + +unsafe impl Sync for StreamWriter {} + +impl std::fmt::Debug for StreamWriter { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("StreamWriter") + .field("stream_id", &self.tx.stream_id()) + .field("target", &self.target) + .field("closed", &!self.open) + .finish_non_exhaustive() + } +} + +impl StreamWriter { + pub(crate) fn new(shared: Tx, target: CloseTarget, handle: RuntimeHandle) -> Self { + Self { + tx: shared, + target, + open: true, + terminal: WriterTerminalState::Pending, + handle, + } + } + + pub fn poll_write( + &mut self, + bytes: &mut Bytes, + cx: &mut Context<'_>, + ) -> Poll> { + if bytes.is_empty() { + return Poll::Ready(Ok(())); + } + + if !self.open { + return self.poll_terminal(cx); + } + + match self.tx.try_write(std::mem::take(bytes)) { + Ok(()) => { + log::trace!( + "byte writer accepted chunk: stream_id={} target={:?}", + self.tx.stream_id(), + self.target + ); + self.poll_runtime(); + return Poll::Ready(Ok(())); + } + Err(PushError::Closed(chunk)) => { + *bytes = chunk; + self.open = false; + return self.poll_terminal(cx); + } + Err(PushError::Full(chunk)) => { + *bytes = chunk; + } + } + + self.tx.register_waiter(cx.waker()); + + match self.tx.try_write(std::mem::take(bytes)) { + Ok(()) => { + self.tx.unregister_waiter(); + log::trace!( + "byte writer accepted chunk: stream_id={} target={:?}", + self.tx.stream_id(), + self.target + ); + self.poll_runtime(); + Poll::Ready(Ok(())) + } + Err(PushError::Closed(chunk)) => { + self.tx.unregister_waiter(); + *bytes = chunk; + self.open = false; + self.poll_terminal(cx) + } + Err(PushError::Full(chunk)) => { + *bytes = chunk; + Poll::Pending + } + } + } + + pub async fn write(&mut self, bytes: Bytes) -> Result<(), QlStreamError> { + let mut bytes = bytes; + poll_fn(|cx| self.poll_write(&mut bytes, cx)).await + } + + pub fn queue_finish(&mut self) { + if !self.open { + return; + } + log::debug!( + "byte writer finish: stream_id={} target={:?}", + self.tx.stream_id(), + self.target + ); + self.open = false; + self.tx.request_finish(); + self.poll_runtime(); + } + + pub async fn finish(mut self) -> Result<(), QlStreamError> { + self.queue_finish(); + poll_fn(|cx| self.poll_terminal(cx)).await + } + + pub fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll> { + if self.open { + self.queue_finish(); + } + self.poll_terminal(cx) + } + + pub fn close(mut self, code: StreamCloseCode) { + self.close_inner(code); + } + + fn poll_runtime(&self) { + self.handle.try_send(Command::PollStream { + stream_id: self.tx.stream_id(), + }); + } + + fn poll_terminal(&mut self, cx: &Context<'_>) -> Poll> { + match &self.terminal { + WriterTerminalState::Terminal(result) => return Poll::Ready(result.clone()), + WriterTerminalState::Pending => {} + } + + match self.try_poll_terminal_ready() { + Poll::Ready(result) => return Poll::Ready(result), + Poll::Pending => {} + } + + self.tx.register_waiter(cx.waker()); + + match self.try_poll_terminal_ready() { + Poll::Ready(result) => { + self.tx.unregister_waiter(); + Poll::Ready(result) + } + Poll::Pending => Poll::Pending, + } + } + + fn try_poll_terminal_ready(&mut self) -> Poll> { + let state = self.tx.load_state(); + if TxInner::terminal_ready(state) { + if TxInner::terminal_ok(state) { + self.terminal = WriterTerminalState::Terminal(Ok(())); + return Poll::Ready(Ok(())); + } + + match self.tx.pop() { + Ok(Item::Error(error)) => { + self.terminal = WriterTerminalState::Terminal(Err(error.clone())); + return Poll::Ready(Err(error)); + } + Ok(Item::Chunk(_)) => { + panic!("writer terminal phase contained chunk data") + } + Err(PopError) => {} + } + } + + Poll::Pending + } + + fn close_inner(&mut self, code: StreamCloseCode) { + if !self.open { + return; + } + self.open = false; + log::debug!( + "byte writer close: stream_id={:?} target={:?} code={:?}", + self.tx.stream_id(), + self.target, + code + ); + self.handle.try_send(Command::CloseStream { + stream_id: self.tx.stream_id(), + target: self.target, + code, + }); + } +} + +impl Drop for StreamWriter { + fn drop(&mut self) { + self.close_inner(StreamCloseCode::CANCELLED); + } +} + +#[cfg(all(test, loom))] +mod loom_tests { + use std::task::{Context, Poll, Waker}; + + use bytes::Bytes; + use loom::thread; + use ql_wire::CloseTarget; + + use super::*; + use crate::io::sync::loom::*; + + #[test] + fn poll_write_observes_capacity_racing_with_registration() { + check_model(|| { + let inner = shared(); + inner.tx.try_write(Bytes::from_static(b"abc")).unwrap(); + + let mut writer = StreamWriter::new(Tx(inner.clone()), CloseTarget::Origin, handle()); + let mut bytes = Bytes::from_static(b"xyz"); + let mut cx = Context::from_waker(Waker::noop()); + + let drainer = { + let inner = inner.clone(); + thread::spawn(move || { + assert!(matches!(inner.tx.pop(), Ok(Item::Chunk(_)))); + }) + }; + + let first = writer.poll_write(&mut bytes, &mut cx); + drainer.join().unwrap(); + + match first { + Poll::Ready(Ok(())) => { + assert!(bytes.is_empty()); + } + Poll::Pending => { + assert_eq!(writer.poll_write(&mut bytes, &mut cx), Poll::Ready(Ok(()))); + assert!(bytes.is_empty()); + } + other => panic!("unexpected first poll result: {other:?}"), + } + }); + } + + #[test] + fn poll_finish_observes_terminal_racing_with_registration() { + check_model(|| { + let inner = shared(); + let mut writer = StreamWriter::new(Tx(inner.clone()), CloseTarget::Origin, handle()); + let mut cx = Context::from_waker(Waker::noop()); + + writer.queue_finish(); + + let finisher = { + let inner = inner.clone(); + thread::spawn(move || { + inner.tx.finish(); + }) + }; + + let first = writer.poll_finish(&mut cx); + finisher.join().unwrap(); + + match first { + Poll::Ready(Ok(())) => {} + Poll::Pending => { + assert_eq!(writer.poll_finish(&mut cx), Poll::Ready(Ok(()))); + } + other => panic!("unexpected first poll result: {other:?}"), + } + }); + } +} diff --git a/ql-runtime/src/lib.rs b/ql-runtime/src/lib.rs new file mode 100644 index 00000000..33783456 --- /dev/null +++ b/ql-runtime/src/lib.rs @@ -0,0 +1,63 @@ +pub use ql_fsm::{NoSessionError, PairingInvite}; + +pub use self::{error::QlStreamError, handle::*, platform::*}; + +pub(crate) mod command; +pub(crate) mod driver; +mod error; +pub mod handle; +pub(crate) mod io; +pub mod log; +pub mod platform; +#[cfg(feature = "rpc")] +pub mod rpc; + +#[cfg(test)] +mod tests; + +use ql_fsm::QlFsmConfig; +use ql_wire::QlIdentity; + +#[derive(Debug, Clone, Copy)] +pub struct RuntimeConfig { + pub fsm: QlFsmConfig, + pub max_concurrent_message_writes: usize, +} + +impl Default for RuntimeConfig { + fn default() -> Self { + Self { + fsm: QlFsmConfig::default(), + max_concurrent_message_writes: 4, + } + } +} + +pub struct Runtime

{ + identity: QlIdentity, + platform: P, + config: RuntimeConfig, + rx: async_channel::Receiver, + tx: async_channel::WeakSender, +} + +pub fn new_runtime

( + identity: QlIdentity, + platform: P, + config: RuntimeConfig, +) -> (Runtime

, RuntimeHandle) +where + P: QlPlatform, +{ + let (tx, rx) = async_channel::unbounded(); + ( + Runtime { + identity, + platform, + config, + rx, + tx: tx.downgrade(), + }, + RuntimeHandle::new(tx), + ) +} diff --git a/ql-runtime/src/log.rs b/ql-runtime/src/log.rs new file mode 100644 index 00000000..a0908f79 --- /dev/null +++ b/ql-runtime/src/log.rs @@ -0,0 +1,54 @@ +#![allow(unused_imports, unused_macros)] + +#[cfg(any(feature = "log", test))] +macro_rules! log { + ($level:ident, $($arg:tt)*) => { + ::log::log!(::log::Level::$level, $($arg)*) + }; +} + +#[cfg(not(any(feature = "log", test)))] +macro_rules! log { + ($level:ident, $($arg:tt)*) => { + if false { + let _ = format_args!($($arg)*); + } + }; +} + +macro_rules! trace { + ($($arg:tt)*) => { + $crate::log::log!(Trace, $($arg)*) + }; +} + +macro_rules! debug { + ($($arg:tt)*) => { + $crate::log::log!(Debug, $($arg)*) + }; +} + +macro_rules! info { + ($($arg:tt)*) => { + $crate::log::log!(Info, $($arg)*) + }; +} + +macro_rules! warn_ { + ($($arg:tt)*) => { + $crate::log::log!(Warn, $($arg)*) + }; +} + +macro_rules! error { + ($($arg:tt)*) => { + $crate::log::log!(Error, $($arg)*) + }; +} + +pub(crate) use debug; +pub(crate) use error; +pub(crate) use info; +pub(crate) use log; +pub(crate) use trace; +pub(crate) use warn_ as warn; diff --git a/ql-runtime/src/platform.rs b/ql-runtime/src/platform.rs new file mode 100644 index 00000000..331bfe7a --- /dev/null +++ b/ql-runtime/src/platform.rs @@ -0,0 +1,43 @@ +use std::{ + future::Future, + pin::Pin, + task::{Context, Poll}, + time::Instant, +}; + +use ql_fsm::{PeerStatus, ReceiveError}; +use ql_wire::{PeerBundle, QlCrypto, QID}; + +use crate::QlStream; + +pub trait QlTimer { + fn set_deadline(self: Pin<&mut Self>, deadline: Option); + fn poll_wait(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()>; +} + +pub trait QlInbound { + fn poll_recv(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll>; +} + +pub trait QlPlatform: QlCrypto { + type Timer: QlTimer; + type WriteMessageFut<'a>: Future + Unpin + 'a + where + Self: 'a; + type Inbound: QlInbound; + + fn write_message(&self, message: Vec) -> Self::WriteMessageFut<'_>; + /// Returns the platform's inbound transport poller. + /// + /// The runtime calls this once while starting the driver loop and retains the returned + /// poller for the lifetime of the runtime. Platform implementations may panic if this is + /// called more than once. + fn inbound(&mut self) -> Self::Inbound; + fn timer(&self) -> Self::Timer; + + fn persist_peer(&self, peer: PeerBundle); + + fn handle_peer_status(&self, peer: Option, status: PeerStatus); + fn handle_inbound(&self, event: QlStream); + fn handle_recv_error(&self, _error: ReceiveError) {} +} diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs new file mode 100644 index 00000000..a7347602 --- /dev/null +++ b/ql-runtime/src/rpc/adapter.rs @@ -0,0 +1,83 @@ +use std::task::{Context, Poll}; + +use bytes::Bytes; +use ql_rpc::{RouteId, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError}; +use ql_wire::{RouteId as WireRouteId, StreamCloseCode as WireStreamCloseCode}; + +use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; + +impl RpcStream for QlStream { + type Error = QlStreamError; + type Reader = StreamReader; + type Writer = StreamWriter; + + fn route_id(&self) -> Option { + let route_id = u32::try_from(self.route_id.into_inner()).ok()?; + Some(RouteId::from_u32(route_id)) + } + + fn split(self) -> (Self::Reader, Self::Writer) { + (self.reader, self.writer) + } +} + +impl RpcRead for StreamReader { + type Error = QlStreamError; + + fn poll_read( + &mut self, + max_len: usize, + cx: &mut Context<'_>, + ) -> Poll, QlStreamError>> { + StreamReader::poll_read(self, max_len, cx) + } + + fn close(self, code: StreamCloseCode) { + StreamReader::close(self, to_wire_close_code(code)); + } +} + +impl RpcWrite for StreamWriter { + type Error = QlStreamError; + + fn poll_write( + &mut self, + bytes: &mut Bytes, + cx: &mut Context<'_>, + ) -> Poll> { + StreamWriter::poll_write(self, bytes, cx) + } + + fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll> { + StreamWriter::poll_finish(self, cx) + } + + fn close(self, code: StreamCloseCode) { + StreamWriter::close(self, to_wire_close_code(code)); + } +} + +pub(super) fn to_wire_route_id(route_id: RouteId) -> WireRouteId { + WireRouteId::from_u32(route_id.into_inner()) +} + +pub(super) fn to_wire_close_code(code: StreamCloseCode) -> WireStreamCloseCode { + WireStreamCloseCode(code.into_inner()) +} + +impl From for QlStreamError { + fn from(code: StreamCloseCode) -> Self { + Self::StreamClosed { + code: WireStreamCloseCode(code.into_inner()), + } + } +} + +impl StreamError for QlStreamError { + fn close_code(&self) -> Option { + match self { + QlStreamError::StreamClosed { code } => Some(StreamCloseCode(code.0)), + QlStreamError::NoSession => None, + } + } +} diff --git a/ql-runtime/src/rpc/download.rs b/ql-runtime/src/rpc/download.rs new file mode 100644 index 00000000..d3b63585 --- /dev/null +++ b/ql-runtime/src/rpc/download.rs @@ -0,0 +1,67 @@ +use bytes::Bytes; +use ql_rpc::download::Download as DownloadRpc; + +use super::RpcError; +use crate::StreamReader; + +pub struct DownloadCall { + pub(super) inner: ql_rpc::download::DownloadCall, +} + +pub struct DownloadReader { + pub(super) inner: ql_rpc::download::DownloadReader, +} + +pub struct DownloadPart<'a, M: DownloadRpc> { + inner: ql_rpc::download::DownloadPart<'a, M, StreamReader>, +} + +impl DownloadCall +where + M: DownloadRpc, +{ + pub async fn start(self) -> Result<(M::ResponseHeader, DownloadReader), RpcError> { + let (header, inner) = self.inner.start().await?; + Ok((header, DownloadReader { inner })) + } + + pub fn close(self, code: ql_wire::StreamCloseCode) { + self.inner.close(ql_rpc::StreamCloseCode(code.0)); + } +} + +impl DownloadReader +where + M: DownloadRpc, +{ + pub async fn next_part( + &mut self, + ) -> Result)>, RpcError> { + Ok(self + .inner + .next_part() + .await? + .map(|(header, inner)| (header, DownloadPart { inner }))) + } + + pub async fn complete(self) -> Result<(), RpcError> { + self.inner.complete().await.map_err(RpcError::from) + } + + pub fn close(self, code: ql_wire::StreamCloseCode) { + self.inner.close(ql_rpc::StreamCloseCode(code.0)); + } +} + +impl DownloadPart<'_, M> +where + M: DownloadRpc, +{ + pub async fn read_chunk(&mut self) -> Result, RpcError> { + Ok(self.inner.read_chunk().await?) + } + + pub fn close(self, code: ql_wire::StreamCloseCode) { + self.inner.close(ql_rpc::StreamCloseCode(code.0)); + } +} diff --git a/ql-runtime/src/rpc/duplex.rs b/ql-runtime/src/rpc/duplex.rs new file mode 100644 index 00000000..cdad6670 --- /dev/null +++ b/ql-runtime/src/rpc/duplex.rs @@ -0,0 +1,59 @@ +use futures_lite::future::poll_fn; +use ql_rpc::duplex::Duplex as DuplexRpc; + +use super::RpcError; +use crate::{QlStreamError, StreamReader, StreamWriter}; + +pub struct DuplexCall { + pub sender: DuplexSender, + pub receiver: DuplexReceiver, +} + +pub struct DuplexSender +where + T: ql_rpc::RpcCodec, +{ + pub(super) inner: ql_rpc::duplex::DuplexSender, +} + +pub struct DuplexReceiver +where + T: ql_rpc::RpcCodec, +{ + pub(super) inner: ql_rpc::duplex::DuplexReceiver, +} + +impl DuplexSender +where + T: ql_rpc::RpcCodec, +{ + pub async fn send(&mut self, event: &T) -> Result<(), QlStreamError> { + self.inner.send(event).await + } + + pub async fn finish(self) -> Result<(), QlStreamError> { + self.inner.finish().await + } + + pub fn close(self, code: ql_wire::StreamCloseCode) { + self.inner.close(ql_rpc::StreamCloseCode(code.0)); + } +} + +impl DuplexReceiver +where + T: ql_rpc::RpcCodec, +{ + pub async fn next_event(&mut self) -> Option>> { + poll_fn(|cx| { + self.inner + .poll_next_event(cx) + .map(|item| item.map(|result| Ok(result?))) + }) + .await + } + + pub fn close(self, code: ql_wire::StreamCloseCode) { + self.inner.close(ql_rpc::StreamCloseCode(code.0)); + } +} diff --git a/ql-runtime/src/rpc/error.rs b/ql-runtime/src/rpc/error.rs new file mode 100644 index 00000000..4cc9e176 --- /dev/null +++ b/ql-runtime/src/rpc/error.rs @@ -0,0 +1,79 @@ +use ql_fsm::NoSessionError; + +use crate::QlStreamError; + +#[derive(Debug)] +pub enum RpcError { + NoSession, + Closed(ql_rpc::StreamCloseCode), + Protocol(ql_rpc::Error), + Codec(E), +} + +impl From for RpcError { + fn from(_: NoSessionError) -> Self { + Self::NoSession + } +} + +impl From for RpcError { + fn from(error: QlStreamError) -> Self { + match error { + QlStreamError::StreamClosed { code } => Self::Closed(ql_rpc::StreamCloseCode(code.0)), + QlStreamError::NoSession => Self::NoSession, + } + } +} + +impl From for RpcError { + fn from(error: ql_rpc::Error) -> Self { + Self::Protocol(error) + } +} + +impl From> for RpcError { + fn from(error: ql_rpc::CodecError) -> Self { + match error { + ql_rpc::CodecError::Rpc(error) => Self::Protocol(error), + ql_rpc::CodecError::Codec(error) => Self::Codec(error), + } + } +} + +impl From> for RpcError { + fn from(error: ql_rpc::CallError) -> Self { + match error { + ql_rpc::CallError::Protocol(error) => Self::Protocol(error), + ql_rpc::CallError::Codec(error) => Self::Codec(error), + ql_rpc::CallError::Transport(error) => error.into(), + } + } +} + +impl std::fmt::Display for RpcError +where + E: std::fmt::Display, +{ + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::NoSession => write!(f, "no session"), + Self::Closed(code) => write!(f, "stream closed {code:?}"), + Self::Protocol(error) => write!(f, "{error}"), + Self::Codec(error) => write!(f, "{error}"), + } + } +} + +impl std::error::Error for RpcError +where + E: std::error::Error + 'static, +{ + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + match self { + Self::Protocol(error) => Some(error), + Self::Codec(error) => Some(error), + RpcError::NoSession => None, + RpcError::Closed(_) => None, + } + } +} diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs new file mode 100644 index 00000000..d8be02c5 --- /dev/null +++ b/ql-runtime/src/rpc/mod.rs @@ -0,0 +1,154 @@ +pub use self::{download::*, duplex::*, error::*, progress::*, subscription::*, upload::*}; + +mod adapter; +mod download; +mod duplex; +mod error; +mod progress; +mod subscription; +mod upload; + +use bytes::Bytes; +use ql_rpc::{ + download::{self as rpc_download, Download as DownloadRpc}, + duplex::{self as rpc_duplex, Duplex as DuplexRpc}, + notification::{self, Notification}, + progress::{self as rpc_progress, Progress}, + request::{self, Request as RequestRpc}, + subscription::{self as rpc_subscription, Subscription as SubscriptionRpc}, + upload::{self as rpc_upload, Upload as UploadRpc}, +}; + +use crate::{RuntimeHandle, StreamReader}; + +#[derive(Clone)] +pub struct RpcHandle { + inner: RuntimeHandle, +} + +impl RpcHandle { + pub async fn notification(&self, event: &M::Payload) -> Result<(), RpcError> + where + M: Notification, + { + let mut payload = Vec::new(); + notification::encode_notification::(event, &mut payload); + let mut stream = self + .inner + .open_stream(adapter::to_wire_route_id(M::ROUTE)) + .await?; + stream.reader.close(ql_wire::StreamCloseCode::CANCELLED); + stream.writer.write(Bytes::from(payload)).await?; + stream.writer.finish().await?; + Ok(()) + } + + pub async fn request(&self, request: &M::Request) -> Result> + where + M: RequestRpc, + { + let mut payload = Vec::new(); + request::encode_request::(request, &mut payload); + let response = self.start_request(M::ROUTE, payload).await?; + Ok(request::read_response::(response).await?) + } + + pub async fn subscribe( + &self, + request: &M::Request, + ) -> Result, RpcError> + where + M: SubscriptionRpc, + { + let mut payload = Vec::new(); + rpc_subscription::encode_request::(request, &mut payload); + let response = self.start_request(M::ROUTE, payload).await?; + Ok(Subscription { + inner: rpc_subscription::SubscriptionCall::new(response), + }) + } + + pub async fn download( + &self, + request: &M::Request, + ) -> Result, RpcError> + where + M: DownloadRpc, + { + let mut payload = Vec::new(); + rpc_download::encode_request::(request, &mut payload); + let response = self.start_request(M::ROUTE, payload).await?; + Ok(DownloadCall { + inner: rpc_download::DownloadCall::new(response), + }) + } + + pub async fn progress( + &self, + request: &M::Request, + ) -> Result, RpcError> + where + M: Progress, + { + let mut payload = Vec::new(); + rpc_progress::encode_request::(request, &mut payload); + let response = self.start_request(M::ROUTE, payload).await?; + Ok(ProgressCall { + inner: rpc_progress::ProgressCall::new(response), + }) + } + + pub async fn upload(&self, request: &M::Request) -> Result, RpcError> + where + M: UploadRpc, + { + let mut payload = Vec::new(); + rpc_upload::encode_request::(request, &mut payload); + let mut stream = self + .inner + .open_stream(adapter::to_wire_route_id(M::ROUTE)) + .await?; + stream.writer.write(Bytes::from(payload)).await?; + Ok(UploadCall { + inner: rpc_upload::UploadCall::new(stream.writer, stream.reader), + }) + } + + pub async fn duplex(&self) -> Result, RpcError> + where + M: DuplexRpc, + { + let stream = self + .inner + .open_stream(adapter::to_wire_route_id(M::ROUTE)) + .await?; + Ok(DuplexCall { + sender: DuplexSender { + inner: rpc_duplex::DuplexSender::new(stream.writer), + }, + receiver: DuplexReceiver { + inner: rpc_duplex::DuplexReceiver::new(stream.reader), + }, + }) + } +} + +impl RpcHandle { + pub(super) fn new(inner: RuntimeHandle) -> Self { + Self { inner } + } + + async fn start_request( + &self, + route_id: ql_rpc::RouteId, + payload: Vec, + ) -> Result> { + let mut stream = self + .inner + .open_stream(adapter::to_wire_route_id(route_id)) + .await?; + stream.writer.write(Bytes::from(payload)).await?; + stream.writer.finish().await?; + Ok(stream.reader) + } +} diff --git a/ql-runtime/src/rpc/progress.rs b/ql-runtime/src/rpc/progress.rs new file mode 100644 index 00000000..a22da20f --- /dev/null +++ b/ql-runtime/src/rpc/progress.rs @@ -0,0 +1,50 @@ +use std::{ + future::Future, + pin::Pin, + task::{Context, Poll}, +}; + +use futures_lite::Stream; +use ql_rpc::progress::Progress; + +use super::RpcError; +use crate::StreamReader; + +pub struct ProgressCall { + pub(super) inner: ql_rpc::progress::ProgressCall, +} + +impl Unpin for ProgressCall where M: Progress {} + +impl ProgressCall +where + M: Progress, +{ + pub fn close(self, code: ql_wire::StreamCloseCode) { + self.inner.close(ql_rpc::StreamCloseCode(code.0)); + } +} + +impl Stream for ProgressCall +where + M: Progress, +{ + type Item = M::Progress; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.get_mut().inner.poll_next_progress(cx) + } +} + +impl Future for ProgressCall +where + M: Progress, +{ + type Output = Result>; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + Pin::new(&mut self.get_mut().inner) + .poll(cx) + .map(|result| result.map_err(RpcError::from)) + } +} diff --git a/ql-runtime/src/rpc/subscription.rs b/ql-runtime/src/rpc/subscription.rs new file mode 100644 index 00000000..45a08a6b --- /dev/null +++ b/ql-runtime/src/rpc/subscription.rs @@ -0,0 +1,43 @@ +use std::{ + pin::Pin, + task::{Context, Poll}, +}; + +use futures_lite::{future::poll_fn, Stream}; +use ql_rpc::subscription::Subscription as SubscriptionRpc; + +use super::RpcError; +use crate::StreamReader; + +pub struct Subscription { + pub(super) inner: ql_rpc::subscription::SubscriptionCall, +} + +impl Unpin for Subscription where M: SubscriptionRpc {} + +impl Subscription +where + M: SubscriptionRpc, +{ + pub async fn next_event(&mut self) -> Option>> { + poll_fn(|cx| Pin::new(&mut *self).poll_next(cx)).await + } + + pub fn close(self, code: ql_wire::StreamCloseCode) { + self.inner.close(ql_rpc::StreamCloseCode(code.0)); + } +} + +impl Stream for Subscription +where + M: SubscriptionRpc, +{ + type Item = Result>; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.get_mut() + .inner + .poll_next_event(cx) + .map(|item| item.map(|result| Ok(result?))) + } +} diff --git a/ql-runtime/src/rpc/upload.rs b/ql-runtime/src/rpc/upload.rs new file mode 100644 index 00000000..33ee3665 --- /dev/null +++ b/ql-runtime/src/rpc/upload.rs @@ -0,0 +1,44 @@ +use bytes::Bytes; +use ql_rpc::upload::Upload as UploadRpc; + +use super::RpcError; +use crate::QlStreamError; + +pub struct UploadCall { + pub(super) inner: ql_rpc::upload::UploadCall, +} + +pub struct UploadPartWriter<'a, M: UploadRpc> { + inner: ql_rpc::upload::UploadPartWriter<'a, M, crate::StreamWriter, crate::StreamReader>, +} + +impl UploadCall +where + M: UploadRpc, +{ + pub async fn start_part( + &mut self, + part_header: M::PartHeader, + ) -> Result, QlStreamError> { + Ok(UploadPartWriter { + inner: self.inner.start_part(part_header).await?, + }) + } + + pub async fn finish(self) -> Result> { + self.inner.finish().await.map_err(RpcError::from) + } +} + +impl UploadPartWriter<'_, M> +where + M: UploadRpc, +{ + pub async fn send(&mut self, bytes: Bytes) -> Result<(), QlStreamError> { + self.inner.send(bytes).await + } + + pub async fn finish(self) -> Result<(), QlStreamError> { + self.inner.finish().await + } +} diff --git a/ql-runtime/src/tests/handshake.rs b/ql-runtime/src/tests/handshake.rs new file mode 100644 index 00000000..65731bbc --- /dev/null +++ b/ql-runtime/src/tests/handshake.rs @@ -0,0 +1,178 @@ +use std::time::Duration; + +use bytes::Bytes; + +use super::*; + +#[tokio::test(flavor = "current_thread")] +async fn connect_round_trip_changes_peer_status() { + run_local_test(async { + let pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn opening_stream_requires_connection() { + run_local_test(async { + let pair = TestPair::new(default_runtime_config()); + assert!(matches!( + pair.side(Side::A).handle.open_stream(test_route_id()).await, + Err(NoSessionError) + )); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn handshake_timeout_disconnects() { + run_local_test(async { + let config = RuntimeConfig { + fsm: QlFsmConfig { + handshake_timeout: Duration::from_millis(60), + ..default_runtime_config().fsm + }, + ..default_runtime_config() + }; + let (platform_a, _outbound_a, _inbound_a, status_a) = TestPlatform::new(); + let (platform_b, _outbound_b, _inbound_b, _status_b) = TestPlatform::new(); + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + + let (runtime_a, handle_a) = new_runtime(identity_a.clone(), platform_a, config); + let (runtime_b, handle_b) = new_runtime(identity_b.clone(), platform_b, config); + + tokio::task::spawn_local(async move { runtime_a.run().await }); + tokio::task::spawn_local(async move { runtime_b.run().await }); + + register_peers(&handle_a, &handle_b, &identity_a, &identity_b); + handle_a.connect(); + + await_status(&status_a, Some(identity_b.qid), PeerStatus::Disconnected).await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn rejected_session_write_is_reissued() { + run_local_test(async { + let config = default_runtime_config(); + let (platform_a, outbound_a, inbound_a_tx, status_a) = + TestPlatform::new_with_session_write_failure(1); + let (platform_b, outbound_b, inbound_b_tx, status_b, inbound_b) = + TestPlatform::new_with_inbound(); + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + + let (runtime_a, handle_a) = new_runtime(identity_a.clone(), platform_a, config); + let (runtime_b, handle_b) = new_runtime(identity_b.clone(), platform_b, config); + + tokio::task::spawn_local(async move { runtime_a.run().await }); + tokio::task::spawn_local(async move { runtime_b.run().await }); + + spawn_forwarder(outbound_a, inbound_b_tx); + spawn_forwarder(outbound_b, inbound_a_tx); + + register_peers(&handle_a, &handle_b, &identity_a, &identity_b); + handle_a.connect(); + + await_status(&status_a, Some(identity_b.qid), PeerStatus::Connected).await; + await_status(&status_b, Some(identity_a.qid), PeerStatus::Connected).await; + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let request = read_all(stream.reader).await.unwrap(); + stream.writer.finish().await.unwrap(); + request + }); + + let mut stream = handle_a.open_stream(test_route_id()).await.unwrap(); + stream + .writer + .write(Bytes::from_static(b"retry")) + .await + .unwrap(); + stream.writer.finish().await.unwrap(); + assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), None); + + assert_eq!( + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(), + b"retry".to_vec() + ); + + assert_no_status_for( + &status_a, + Some(identity_b.qid), + PeerStatus::Disconnected, + Duration::from_millis(150), + ) + .await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn start_pairing_round_trip_connects_when_armed() { + run_local_test(async { + let config = default_runtime_config(); + let (platform_a, outbound_a, inbound_a_tx, status_a) = TestPlatform::new(); + let (platform_b, outbound_b, inbound_b_tx, status_b) = TestPlatform::new(); + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + let token = pairing_token(7); + + let (runtime_a, handle_a) = new_runtime(identity_a.clone(), platform_a, config); + let (runtime_b, handle_b) = new_runtime(identity_b.clone(), platform_b, config); + + tokio::task::spawn_local(async move { runtime_a.run().await }); + tokio::task::spawn_local(async move { runtime_b.run().await }); + + spawn_forwarder(outbound_a, inbound_b_tx); + spawn_forwarder(outbound_b, inbound_a_tx); + + handle_b.arm_pairing(token); + handle_a.start_pairing(PairingInvite { + qid: identity_b.qid, + token, + }); + + await_status(&status_a, Some(identity_b.qid), PeerStatus::Connected).await; + await_status(&status_b, Some(identity_a.qid), PeerStatus::Connected).await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn start_pairing_does_not_connect_when_unarmed() { + run_local_test(async { + let config = default_runtime_config(); + let (platform_a, outbound_a, inbound_a_tx, status_a) = TestPlatform::new(); + let (platform_b, outbound_b, inbound_b_tx, _status_b) = TestPlatform::new(); + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + let token = pairing_token(8); + + let (runtime_a, handle_a) = new_runtime(identity_a.clone(), platform_a, config); + let (runtime_b, _handle_b) = new_runtime(identity_b.clone(), platform_b, config); + + tokio::task::spawn_local(async move { runtime_a.run().await }); + tokio::task::spawn_local(async move { runtime_b.run().await }); + + spawn_forwarder(outbound_a, inbound_b_tx); + spawn_forwarder(outbound_b, inbound_a_tx); + + handle_a.start_pairing(PairingInvite { + qid: identity_b.qid, + token, + }); + + assert_no_status_for( + &status_a, + Some(identity_b.qid), + PeerStatus::Connected, + Duration::from_millis(150), + ) + .await; + }) + .await; +} diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs new file mode 100644 index 00000000..af368738 --- /dev/null +++ b/ql-runtime/src/tests/mod.rs @@ -0,0 +1,710 @@ +use std::{ + future::Future, + pin::Pin, + sync::{ + atomic::{AtomicUsize, Ordering}, + Arc, Mutex, Once, + }, + task::{Context, Poll}, + time::Duration, +}; + +use async_channel::{Receiver, Sender}; +use futures_lite::Stream; +use ql_fsm::PeerStatus; +use ql_wire::{ + generate_identity, test_identities, MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, + MlKemPublicKey, Nonce, PairingToken, PeerBundle, QlAead, QlHash, QlIdentity, QlKem, QlRandom, + RecordHeader, RecordType, RouteId, SessionKey, SoftwareCrypto, WireDecode, QID, +}; +use tokio::{task::LocalSet, time::Sleep}; + +use crate::{ + new_runtime, platform::QlTimer, NoSessionError, PairingInvite, QlFsmConfig, QlStream, + QlStreamError, RuntimeConfig, RuntimeHandle, +}; + +mod handshake; +#[cfg(feature = "rpc")] +mod rpc; +mod session; +mod stream; + +fn init_test_logger() { + static INIT: Once = Once::new(); + + INIT.call_once(|| { + let env = env_logger::Env::default().default_filter_or("ql_runtime=info"); + let mut builder = env_logger::Builder::from_env(env); + builder.is_test(true); + let _ = builder.try_init(); + }); +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +struct StatusEvent { + peer: Option, + status: PeerStatus, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Side { + A, + B, +} + +impl Side { + fn opposite(self) -> Self { + match self { + Self::A => Self::B, + Self::B => Self::A, + } + } +} + +fn test_route_id() -> RouteId { + RouteId::from_u32(1) +} + +#[derive(Debug, Clone)] +struct WriteStats { + active: Arc, + max_active: Arc, +} + +impl WriteStats { + fn new() -> Self { + Self { + active: Arc::new(AtomicUsize::new(0)), + max_active: Arc::new(AtomicUsize::new(0)), + } + } + + fn max_active(&self) -> usize { + self.max_active.load(Ordering::Relaxed) + } +} + +struct TestPlatform { + outbound: Sender>, + _inbound_messages_tx: Sender>, + inbound_messages: Option>>, + status: Sender, + inbound: Option>, + crypto: SoftwareCrypto, + encrypted_write_counter: AtomicUsize, + fail_encrypted_write_at: Option, + write_delay: Duration, + write_stats: Option, +} + +struct TestInbound { + receiver: Receiver>, +} + +type TestPlatformParts = ( + TestPlatform, + Receiver>, + Sender>, + Receiver, +); + +type TestPlatformPartsWithInbound = ( + TestPlatform, + Receiver>, + Sender>, + Receiver, + Receiver, +); + +impl TestPlatform { + fn new() -> TestPlatformParts { + Self::new_inner(None, None, Duration::ZERO, None) + } + + fn new_with_inbound() -> TestPlatformPartsWithInbound { + let (inbound_tx, inbound_rx) = async_channel::unbounded(); + let (platform, outbound_rx, inbound_messages_tx, status_rx) = + Self::new_inner(Some(inbound_tx), None, Duration::ZERO, None); + ( + platform, + outbound_rx, + inbound_messages_tx, + status_rx, + inbound_rx, + ) + } + + fn new_with_session_write_failure(fail_encrypted_write_at: usize) -> TestPlatformParts { + Self::new_inner(None, Some(fail_encrypted_write_at), Duration::ZERO, None) + } + + fn new_with_delayed_writes(delay: Duration, write_stats: WriteStats) -> TestPlatformParts { + Self::new_inner(None, None, delay, Some(write_stats)) + } + + fn new_inner( + inbound: Option>, + fail_encrypted_write_at: Option, + write_delay: Duration, + write_stats: Option, + ) -> TestPlatformParts { + let (outbound, outbound_rx) = async_channel::unbounded(); + let (inbound_messages_tx, inbound_messages_rx) = async_channel::unbounded(); + let (status, status_rx) = async_channel::unbounded(); + ( + Self { + outbound, + _inbound_messages_tx: inbound_messages_tx.clone(), + inbound_messages: Some(inbound_messages_rx), + status, + inbound, + crypto: SoftwareCrypto, + encrypted_write_counter: AtomicUsize::new(0), + fail_encrypted_write_at, + write_delay, + write_stats, + }, + outbound_rx, + inbound_messages_tx, + status_rx, + ) + } +} + +struct TestSide { + handle: RuntimeHandle, + status: Receiver, + peer: QID, + inbound: Receiver, +} + +struct TestPair { + a: TestSide, + b: TestSide, +} + +#[derive(Debug, Clone, Copy, Default)] +struct LinkBehavior { + base_delay: Duration, + drop_encrypted_every: Option, + duplicate_encrypted_every: Option, + delay_encrypted_every: Option<(usize, Duration)>, +} + +#[derive(Clone, Default)] +struct LinkController { + behavior: Arc>, +} + +impl LinkController { + fn new(behavior: LinkBehavior) -> Self { + Self { + behavior: Arc::new(Mutex::new(behavior)), + } + } + + fn load(&self) -> LinkBehavior { + *self.behavior.lock().unwrap() + } + + fn store(&self, behavior: LinkBehavior) { + *self.behavior.lock().unwrap() = behavior; + } +} + +#[derive(Clone)] +struct ControlledLinks { + a_to_b: LinkController, + b_to_a: LinkController, +} + +impl TestPair { + fn new(config: RuntimeConfig) -> Self { + Self::new_with_links(config, LinkBehavior::default(), LinkBehavior::default()) + } + + fn new_with_links(config: RuntimeConfig, a_to_b: LinkBehavior, b_to_a: LinkBehavior) -> Self { + let (pair, _links) = Self::new_with_controlled_links(config, a_to_b, b_to_a); + pair + } + + fn new_with_controlled_links( + config: RuntimeConfig, + a_to_b: LinkBehavior, + b_to_a: LinkBehavior, + ) -> (Self, ControlledLinks) { + let (platform_a, outbound_a, inbound_a_tx, status_a, inbound_a) = + TestPlatform::new_with_inbound(); + let (platform_b, outbound_b, inbound_b_tx, status_b, inbound_b) = + TestPlatform::new_with_inbound(); + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + let links = ControlledLinks { + a_to_b: LinkController::new(a_to_b), + b_to_a: LinkController::new(b_to_a), + }; + + let (runtime_a, handle_a) = new_runtime(identity_a.clone(), platform_a, config); + let (runtime_b, handle_b) = new_runtime(identity_b.clone(), platform_b, config); + + tokio::task::spawn_local(async move { runtime_a.run().await }); + tokio::task::spawn_local(async move { runtime_b.run().await }); + + spawn_simulated_forwarder(outbound_a, inbound_b_tx, links.a_to_b.clone()); + spawn_simulated_forwarder(outbound_b, inbound_a_tx, links.b_to_a.clone()); + register_peers(&handle_a, &handle_b, &identity_a, &identity_b); + + ( + Self { + a: TestSide { + handle: handle_a, + status: status_a, + peer: identity_a.qid, + inbound: inbound_a, + }, + b: TestSide { + handle: handle_b, + status: status_b, + peer: identity_b.qid, + inbound: inbound_b, + }, + }, + links, + ) + } + + fn side(&self, side: Side) -> &TestSide { + match side { + Side::A => &self.a, + Side::B => &self.b, + } + } + + fn side_mut(&mut self, side: Side) -> &mut TestSide { + match side { + Side::A => &mut self.a, + Side::B => &mut self.b, + } + } + + async fn connect_and_wait(&self, initiator: Side) { + self.side(initiator).handle.connect(); + await_status( + &self.side(initiator).status, + Some(self.side(initiator.opposite()).peer), + PeerStatus::Connected, + ) + .await; + await_status( + &self.side(initiator.opposite()).status, + Some(self.side(initiator).peer), + PeerStatus::Connected, + ) + .await; + } + + fn take_inbound(&mut self, side: Side) -> Receiver { + let replacement = async_channel::unbounded().1; + std::mem::replace(&mut self.side_mut(side).inbound, replacement) + } +} + +struct TokioTimer { + sleep: Pin>, +} + +impl TokioTimer { + fn new() -> Self { + Self { + sleep: Box::pin(tokio::time::sleep_until(parked_deadline())), + } + } +} + +impl QlTimer for TokioTimer { + fn set_deadline(mut self: Pin<&mut Self>, deadline: Option) { + let deadline = deadline.map_or_else(parked_deadline, tokio::time::Instant::from_std); + self.as_mut().get_mut().sleep.as_mut().reset(deadline); + } + + fn poll_wait(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<()> { + self.as_mut().get_mut().sleep.as_mut().poll(cx) + } +} + +impl QlRandom for TestPlatform { + fn fill_random_bytes(&self, data: &mut [u8]) { + self.crypto.fill_random_bytes(data); + } +} + +impl QlHash for TestPlatform { + fn sha256(&self, parts: &[&[u8]]) -> [u8; 32] { + self.crypto.sha256(parts) + } +} + +impl QlAead for TestPlatform { + fn aes256_gcm_encrypt( + &self, + key: &SessionKey, + nonce: &Nonce, + aad: &[u8], + buffer: &mut [u8], + ) -> [u8; ql_wire::ENCRYPTED_MESSAGE_AUTH_SIZE] { + self.crypto.aes256_gcm_encrypt(key, nonce, aad, buffer) + } + + fn aes256_gcm_decrypt( + &self, + key: &SessionKey, + nonce: &Nonce, + aad: &[u8], + buffer: &mut [u8], + auth_tag: &[u8; ql_wire::ENCRYPTED_MESSAGE_AUTH_SIZE], + ) -> bool { + self.crypto + .aes256_gcm_decrypt(key, nonce, aad, buffer, auth_tag) + } +} + +impl QlKem for TestPlatform { + fn mlkem_generate_keypair(&self) -> MlKemKeyPair { + self.crypto.mlkem_generate_keypair() + } + + fn mlkem_encapsulate(&self, public_key: &MlKemPublicKey) -> (MlKemCiphertext, SessionKey) { + self.crypto.mlkem_encapsulate(public_key) + } + + fn mlkem_decapsulate(&self, pk: &MlKemPrivateKey, cipher: &MlKemCiphertext) -> SessionKey { + self.crypto.mlkem_decapsulate(pk, cipher) + } +} + +impl crate::platform::QlPlatform for TestPlatform { + type Timer = TokioTimer; + type WriteMessageFut<'a> = Pin + Send + 'a>>; + type Inbound = TestInbound; + + fn write_message(&self, message: Vec) -> Self::WriteMessageFut<'_> { + let outbound = self.outbound.clone(); + let write_delay = self.write_delay; + let fail_encrypted_write_at = self.fail_encrypted_write_at; + let write_stats = self.write_stats.clone(); + + Box::pin(async move { + if let Some(stats) = write_stats.as_ref() { + let active = stats.active.fetch_add(1, Ordering::Relaxed) + 1; + stats.max_active.fetch_max(active, Ordering::Relaxed); + } + + if !write_delay.is_zero() { + tokio::time::sleep(write_delay).await; + } + + let should_fail = if is_encrypted_payload(&message) { + let count = self.encrypted_write_counter.fetch_add(1, Ordering::Relaxed) + 1; + fail_encrypted_write_at == Some(count) + } else { + false + }; + + let success = if should_fail { + false + } else { + outbound.send(message).await.is_ok() + }; + + if let Some(stats) = write_stats.as_ref() { + stats.active.fetch_sub(1, Ordering::Relaxed); + } + + success + }) + } + + fn inbound(&mut self) -> Self::Inbound { + TestInbound { + receiver: self + .inbound_messages + .take() + .expect("TestPlatform::inbound may only be called once"), + } + } + + fn timer(&self) -> Self::Timer { + TokioTimer::new() + } + + fn persist_peer(&self, _peer: PeerBundle) {} + + fn handle_peer_status(&self, peer: Option, status: PeerStatus) { + let _ = self.status.try_send(StatusEvent { peer, status }); + } + + fn handle_inbound(&self, event: QlStream) { + if let Some(tx) = &self.inbound { + let _ = tx.try_send(event); + } + } +} + +impl crate::platform::QlInbound for TestInbound { + fn poll_recv(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + match unsafe { self.as_mut().map_unchecked_mut(|this| &mut this.receiver) }.poll_next(cx) { + Poll::Ready(Some(bytes)) => Poll::Ready(bytes), + Poll::Ready(None) => panic!("TestInbound channel closed"), + Poll::Pending => Poll::Pending, + } + } +} + +fn parked_deadline() -> tokio::time::Instant { + tokio::time::Instant::now() + Duration::from_secs(60 * 60 * 24 * 365 * 100) +} + +fn is_encrypted_payload(bytes: &[u8]) -> bool { + RecordHeader::decode_bytes(bytes) + .ok() + .is_some_and(|header| header.record_type == RecordType::Session) +} + +fn pairing_token(byte: u8) -> PairingToken { + PairingToken([byte; PairingToken::SIZE]) +} + +fn register_peers( + handle_a: &RuntimeHandle, + handle_b: &RuntimeHandle, + id_a: &QlIdentity, + id_b: &QlIdentity, +) { + handle_a.bind_peer(id_b.bundle()); + handle_b.bind_peer(id_a.bundle()); +} + +fn spawn_forwarder(outbound: Receiver>, inbound: Sender>) { + spawn_simulated_forwarder( + outbound, + inbound, + LinkController::new(LinkBehavior::default()), + ); +} + +fn spawn_simulated_forwarder( + outbound: Receiver>, + inbound: Sender>, + controller: LinkController, +) { + tokio::task::spawn_local(async move { + let mut encrypted_count = 0usize; + while let Ok(bytes) = outbound.recv().await { + let behavior = controller.load(); + let encrypted = is_encrypted_payload(&bytes); + let ordinal = if encrypted { + encrypted_count = encrypted_count.saturating_add(1); + Some(encrypted_count) + } else { + None + }; + + if ordinal.is_some_and(|count| { + behavior + .drop_encrypted_every + .is_some_and(|nth| nth != 0 && count % nth == 0) + }) { + continue; + } + + let mut delay = behavior.base_delay; + if let Some(count) = ordinal { + if let Some((nth, extra_delay)) = behavior.delay_encrypted_every { + if nth != 0 && count % nth == 0 { + delay += extra_delay; + } + } + } + + let primary = bytes.clone(); + let primary_inbound = inbound.clone(); + tokio::task::spawn_local(async move { + if !delay.is_zero() { + tokio::time::sleep(delay).await; + } + let _ = primary_inbound.try_send(primary); + }); + + if ordinal.is_some_and(|count| { + behavior + .duplicate_encrypted_every + .is_some_and(|nth| nth != 0 && count % nth == 0) + }) { + let duplicate_inbound = inbound.clone(); + tokio::task::spawn_local(async move { + let duplicate_delay = delay + Duration::from_millis(1); + if !duplicate_delay.is_zero() { + tokio::time::sleep(duplicate_delay).await; + } + let _ = duplicate_inbound.try_send(bytes); + }); + } + } + }); +} + +fn spawn_drop_every_nth_encrypted_forwarder( + outbound: Receiver>, + inbound: Sender>, + nth: usize, +) { + tokio::task::spawn_local(async move { + let mut encrypted_count = 0usize; + while let Ok(bytes) = outbound.recv().await { + if nth > 0 && is_encrypted_payload(&bytes) { + encrypted_count = encrypted_count.saturating_add(1); + if encrypted_count % nth == 0 { + continue; + } + } + let _ = inbound.try_send(bytes); + } + }); +} + +fn spawn_gated_forwarder( + outbound: Receiver>, + inbound: Sender>, + drop_flag: Arc, +) { + tokio::task::spawn_local(async move { + while let Ok(bytes) = outbound.recv().await { + if drop_flag.load(Ordering::Relaxed) { + continue; + } + let _ = inbound.try_send(bytes); + } + }); +} + +#[allow(clippy::future_not_send)] +async fn run_local_test(future: F) +where + F: Future, +{ + run_local_test_timeout(Duration::from_secs(5), future).await; +} + +#[allow(clippy::future_not_send)] +async fn run_local_test_timeout(duration: Duration, future: F) +where + F: Future, +{ + init_test_logger(); + let local = LocalSet::new(); + let future = local.run_until(future); + tokio::time::timeout(duration, future) + .await + .unwrap_or_else(|_| panic!("local runtime test exceeded {duration:?}")); +} + +async fn await_status(receiver: &Receiver, peer: Option, stage: PeerStatus) { + tokio::time::timeout(Duration::from_secs(2), async { + loop { + if let Ok(event) = receiver.recv().await { + if event.peer == peer && event.status == stage { + return; + } + } + } + }) + .await + .unwrap(); +} + +async fn assert_no_status_for( + receiver: &Receiver, + peer: Option, + status: PeerStatus, + window: Duration, +) { + let res = tokio::time::timeout(window, async { + loop { + let event = receiver.recv().await.unwrap(); + if event.peer == peer && event.status == status { + return; + } + } + }) + .await; + assert!(res.is_err(), "unexpected status event: {status:?}"); +} + +async fn read_all(mut stream: crate::StreamReader) -> Result, QlStreamError> { + let mut data = Vec::new(); + while let Some(chunk) = next_chunk(&mut stream).await? { + data.extend_from_slice(&chunk); + } + Ok(data) +} + +async fn next_chunk_max( + stream: &mut crate::StreamReader, + max_len: usize, +) -> Result>, crate::QlStreamError> { + stream + .read(max_len) + .await + .map(|chunk| chunk.map(|bytes| bytes.to_vec())) +} + +async fn next_chunk(stream: &mut crate::StreamReader) -> Result>, QlStreamError> { + next_chunk_max(stream, usize::MAX).await +} + +fn default_runtime_config() -> RuntimeConfig { + RuntimeConfig { + fsm: QlFsmConfig { + handshake_timeout: Duration::from_millis(300), + session_record_retransmit_timeout: Duration::from_millis(30), + session_keepalive_interval: Duration::ZERO, + session_peer_timeout: Duration::ZERO, + ..Default::default() + }, + ..Default::default() + } +} + +// runtime is send, if platform is send +#[test] +fn runtime_is_send() { + let config = default_runtime_config(); + let identity = generate_identity(&SoftwareCrypto, "runtime").unwrap(); + let (platform, _, _, _) = TestPlatform::new(); + let (runtime, _handle) = new_runtime(identity, platform, config); + let _run: Box + Send> = Box::new(runtime.run()); +} + +#[test] +fn runtime_exits_when_last_handle_drops() { + let config = default_runtime_config(); + let identity = generate_identity(&SoftwareCrypto, "runtime").unwrap(); + let (platform, _, _, _) = TestPlatform::new(); + let (runtime, handle) = new_runtime(identity, platform, config); + let (done_tx, done_rx) = oneshot::channel(); + + std::thread::spawn(move || { + tokio::runtime::Builder::new_current_thread() + .enable_time() + .build() + .unwrap() + .block_on(runtime.run()); + done_tx.send(()).unwrap(); + }); + + drop(handle); + + done_rx + .recv_timeout(Duration::from_secs(1)) + .expect("runtime should stop once the last sender is dropped"); +} diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs new file mode 100644 index 00000000..3244c587 --- /dev/null +++ b/ql-runtime/src/tests/rpc.rs @@ -0,0 +1,677 @@ +use std::{ + cell::RefCell, + future::Future, + rc::Rc, + str::Utf8Error, + sync::{Arc, Mutex}, + time::Duration, +}; + +use bytes::Bytes; +use futures_lite::StreamExt; +use ql_rpc::{ + DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, + NotificationHandlerLocal, ProgressHandlerLocal, ProgressResponder, RequestHandler, + RequestHandlerLocal, Response, RouteId, SendSpawner, Spawner, StreamCloseCode, + SubscriptionHandlerLocal, SubscriptionResponder, UploadHandlerLocal, UploadReader, + UploadResponder, +}; + +use super::*; +use crate::{rpc::RpcError, QlStream, StreamWriter}; + +#[derive(Debug, Clone, Copy)] +struct TokioLocalSpawner; + +impl Spawner for TokioLocalSpawner { + type Handle = tokio::task::JoinHandle<()>; +} + +impl LocalSpawner for TokioLocalSpawner { + fn spawn(&self, fut: F) -> Self::Handle + where + F: Future + 'static, + { + tokio::task::spawn_local(fut) + } +} + +#[derive(Debug, Clone, Copy)] +struct TokioSendSpawner; + +impl Spawner for TokioSendSpawner { + type Handle = tokio::task::JoinHandle<()>; +} + +impl SendSpawner for TokioSendSpawner { + fn spawn(&self, fut: F) -> Self::Handle + where + F: Future + Send + 'static, + { + tokio::task::spawn(fut) + } +} + +struct Echo; + +impl ql_rpc::Route for Echo { + const ROUTE: RouteId = RouteId::from_u32(51); +} + +impl ql_rpc::request::Request for Echo { + type Error = Utf8Error; + + type Request = String; + type Response = String; +} + +struct Feed; + +impl ql_rpc::Route for Feed { + const ROUTE: RouteId = RouteId::from_u32(52); +} + +impl ql_rpc::subscription::Subscription for Feed { + type Error = core::convert::Infallible; + type Request = Vec; + type Event = Vec; +} + +struct Notice; + +impl ql_rpc::Route for Notice { + const ROUTE: RouteId = RouteId::from_u32(521); +} + +impl ql_rpc::notification::Notification for Notice { + type Error = core::convert::Infallible; + type Payload = Vec; +} + +struct Download; + +impl ql_rpc::Route for Download { + const ROUTE: RouteId = RouteId::from_u32(53); +} + +impl ql_rpc::progress::Progress for Download { + type Error = core::convert::Infallible; + type Request = Vec; + type Progress = Vec; + type Response = Vec; +} + +struct BlobDownload; + +impl ql_rpc::Route for BlobDownload { + const ROUTE: RouteId = RouteId::from_u32(54); +} + +impl ql_rpc::download::Download for BlobDownload { + type Error = core::convert::Infallible; + type Request = Vec; + type ResponseHeader = Vec; + type PartHeader = Vec; +} + +struct BlobUpload; + +impl ql_rpc::Route for BlobUpload { + const ROUTE: RouteId = RouteId::from_u32(55); +} + +impl ql_rpc::upload::Upload for BlobUpload { + type Error = core::convert::Infallible; + type Request = Vec; + type PartHeader = Vec; + type Response = Vec; +} + +struct Chat; + +impl ql_rpc::Route for Chat { + const ROUTE: RouteId = RouteId::from_u32(56); +} + +impl ql_rpc::duplex::Duplex for Chat { + type Error = core::convert::Infallible; + type InitiatorEvent = Vec; + type ResponderEvent = Vec; +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_request() { + #[derive(Clone)] + struct RouterState { + seen: Arc>>, + } + + impl RequestHandler for RouterState { + async fn handle(self, request: String, response: Response) { + let seen = self.seen.clone(); + seen.lock().unwrap().push(request); + let _ = response.respond("world".into()).await; + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let seen = Arc::new(Mutex::new(Vec::new())); + + let router = + ql_rpc::Router::<_, QlStream, TokioSendSpawner>::builder_send(TokioSendSpawner) + .request::() + .build(RouterState { seen: seen.clone() }); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + let fut = assert_send(fut); + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let response = rpc.request::(&"hello".into()).await.unwrap(); + assert_eq!(response, "world"); + assert_eq!(&*seen.lock().unwrap(), &["hello".to_string()]); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +fn assert_send(value: T) -> T { + value +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_notification() { + #[derive(Clone)] + struct RouterState { + seen: Rc>>>, + } + + impl NotificationHandlerLocal for RouterState { + async fn handle(self, payload: Vec) { + self.seen.borrow_mut().push(payload); + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let seen = Rc::new(RefCell::new(Vec::new())); + + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .notification::() + .build(RouterState { seen: seen.clone() }); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + rpc.notification::(&b"hello".to_vec()) + .await + .unwrap(); + assert_eq!(seen.borrow().as_slice(), &[b"hello".to_vec()]); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_subscrption() { + #[derive(Clone)] + struct RouterState { + seen: Rc>>>, + } + + impl SubscriptionHandlerLocal for RouterState { + async fn handle( + self, + request: Vec, + mut response: SubscriptionResponder, StreamWriter>, + ) { + let seen = self.seen.clone(); + seen.borrow_mut().push(request); + let _ = response.send(b"one".to_vec()).await; + let _ = response.send(b"two".to_vec()).await; + let _ = response.finish().await; + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + + let seen = Rc::new(RefCell::new(Vec::new())); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .subscription::() + .build(RouterState { seen: seen.clone() }); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let mut subscription = rpc.subscribe::(&b"watch".to_vec()).await.unwrap(); + assert_eq!(subscription.next().await.unwrap().unwrap(), b"one".to_vec()); + assert_eq!(subscription.next().await.unwrap().unwrap(), b"two".to_vec()); + assert!(subscription.next().await.is_none()); + assert_eq!(seen.borrow().as_slice(), &[b"watch".to_vec()]); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_router_enforces_max_request_bytes() { + #[derive(Clone)] + struct LimitedState; + + impl RequestHandlerLocal for LimitedState { + async fn handle(self, request: String, response: Response) { + let _ = response.respond(request).await; + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .max_request_bytes(4) + .request::() + .build(LimitedState); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let response = rpc.request::(&"hello".to_string()).await; + assert!(matches!( + response, + Err(RpcError::Closed(code)) if code == StreamCloseCode::LIMIT + )); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_progress() { + #[derive(Clone)] + struct RouterState { + seen: Rc>>>, + } + + impl ProgressHandlerLocal for RouterState { + async fn handle( + self, + request: Vec, + mut responder: ProgressResponder, + ) { + let seen = self.seen.clone(); + seen.borrow_mut().push(request); + responder.send(b"10".to_vec()).await.unwrap(); + responder.send(b"90".to_vec()).await.unwrap(); + responder.finish(b"done".to_vec()).await.unwrap(); + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let seen = Rc::new(RefCell::new(Vec::new())); + + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .progress::() + .build(RouterState { seen: seen.clone() }); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let mut download = rpc.progress::(&b"logo".to_vec()).await.unwrap(); + + assert_eq!(download.next().await, Some(b"10".to_vec())); + assert_eq!(download.next().await, Some(b"90".to_vec())); + assert_eq!(download.next().await, None); + assert_eq!(download.await.unwrap(), b"done".to_vec()); + assert_eq!(seen.borrow().as_slice(), &[b"logo".to_vec()]); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_download() { + #[derive(Clone)] + struct RouterState { + seen: Rc>>>, + } + + impl DownloadHandlerLocal for RouterState { + async fn handle( + self, + request: Vec, + download: DownloadStart, + ) { + let seen = self.seen.clone(); + seen.borrow_mut().push(request); + let mut writer = download.start(b"image/png".to_vec()).await.unwrap(); + let mut part = writer.start_part(b"icon".to_vec()).await.unwrap(); + part.send(Bytes::from_static(b"abc")).await.unwrap(); + part.send(Bytes::from_static(b"def")).await.unwrap(); + part.finish().await.unwrap(); + let mut part = writer.start_part(b"manifest".to_vec()).await.unwrap(); + part.send(Bytes::from_static(b"{}")).await.unwrap(); + part.finish().await.unwrap(); + writer.finish().await.unwrap(); + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let seen = Rc::new(RefCell::new(Vec::new())); + + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .download::() + .build(RouterState { seen: seen.clone() }); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let download = rpc + .download::(&b"logo".to_vec()) + .await + .unwrap(); + let (header, mut reader) = download.start().await.unwrap(); + assert_eq!(header, b"image/png".to_vec()); + { + let (part_header, mut part) = reader.next_part().await.unwrap().unwrap(); + assert_eq!(part_header, b"icon".to_vec()); + assert_eq!( + part.read_chunk().await.unwrap(), + Some(Bytes::from_static(b"abc")) + ); + assert_eq!( + part.read_chunk().await.unwrap(), + Some(Bytes::from_static(b"def")) + ); + assert_eq!(part.read_chunk().await.unwrap(), None); + } + { + let (part_header, mut part) = reader.next_part().await.unwrap().unwrap(); + assert_eq!(part_header, b"manifest".to_vec()); + assert_eq!( + part.read_chunk().await.unwrap(), + Some(Bytes::from_static(b"{}")) + ); + assert_eq!(part.read_chunk().await.unwrap(), None); + } + assert!(reader.next_part().await.unwrap().is_none()); + assert_eq!(seen.borrow().as_slice(), &[b"logo".to_vec()]); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_download_complete() { + #[derive(Clone)] + struct RouterState { + seen: Rc>>>, + } + + impl DownloadHandlerLocal for RouterState { + async fn handle( + self, + request: Vec, + download: DownloadStart, + ) { + self.seen.borrow_mut().push(request); + download.complete(b"not found".to_vec()).await.unwrap(); + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let seen = Rc::new(RefCell::new(Vec::new())); + + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .download::() + .build(RouterState { seen: seen.clone() }); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let download = rpc + .download::(&b"logo".to_vec()) + .await + .unwrap(); + let (header, reader) = download.start().await.unwrap(); + assert_eq!(header, b"not found".to_vec()); + reader.complete().await.unwrap(); + assert_eq!(seen.borrow().as_slice(), &[b"logo".to_vec()]); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_upload() { + #[derive(Clone)] + struct RouterState { + requests: Rc>>>, + uploads: Rc>>>, + } + + impl UploadHandlerLocal for RouterState { + async fn handle( + self, + request: Vec, + mut upload: UploadReader, + responder: UploadResponder, StreamWriter>, + ) { + let requests = self.requests.clone(); + let uploads = self.uploads.clone(); + requests.borrow_mut().push(request); + + let mut body = Vec::new(); + while let Some((part_header, mut part)) = upload.next_part().await.unwrap() { + body.extend_from_slice(&part_header); + body.push(b':'); + while let Some(chunk) = part.read_chunk().await.unwrap() { + body.extend_from_slice(&chunk); + } + body.push(b';'); + } + uploads.borrow_mut().push(body.clone()); + + responder.respond(body).await.unwrap(); + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let requests = Rc::new(RefCell::new(Vec::new())); + let uploads = Rc::new(RefCell::new(Vec::new())); + + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .upload::() + .build(RouterState { + requests: requests.clone(), + uploads: uploads.clone(), + }); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let mut upload = rpc.upload::(&b"logo".to_vec()).await.unwrap(); + let mut part = upload.start_part(b"icon".to_vec()).await.unwrap(); + part.send(Bytes::from_static(b"abc")).await.unwrap(); + part.send(Bytes::from_static(b"def")).await.unwrap(); + part.finish().await.unwrap(); + let mut part = upload.start_part(b"manifest".to_vec()).await.unwrap(); + part.send(Bytes::from_static(b"{}")).await.unwrap(); + part.finish().await.unwrap(); + let response = upload.finish().await.unwrap(); + + assert_eq!(response, b"icon:abcdef;manifest:{};".to_vec()); + assert_eq!(requests.borrow().as_slice(), &[b"logo".to_vec()]); + assert_eq!( + uploads.borrow().as_slice(), + &[b"icon:abcdef;manifest:{};".to_vec()] + ); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn rpc_duplex() { + #[derive(Clone)] + struct RouterState { + seen: Rc>>>, + } + + impl DuplexHandlerLocal for RouterState { + async fn handle(self, mut peer: DuplexPeer) { + let seen = self.seen.clone(); + let first = peer.receiver.next_event().await.unwrap().unwrap(); + seen.borrow_mut().push(first); + + peer.sender + .send(&b"challenge-response".to_vec()) + .await + .unwrap(); + + let second = peer.receiver.next_event().await.unwrap().unwrap(); + seen.borrow_mut().push(second); + + peer.sender.finish().await.unwrap(); + } + } + + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let seen = Rc::new(RefCell::new(Vec::new())); + + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .duplex::() + .build(RouterState { seen: seen.clone() }); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(inbound) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let mut chat = rpc.duplex::().await.unwrap(); + chat.sender.send(&b"challenge".to_vec()).await.unwrap(); + assert_eq!( + chat.receiver.next_event().await.unwrap().unwrap(), + b"challenge-response".to_vec() + ); + chat.sender.send(&b"verification".to_vec()).await.unwrap(); + chat.sender.finish().await.unwrap(); + assert!(chat.receiver.next_event().await.is_none()); + + assert_eq!( + seen.borrow().as_slice(), + &[b"challenge".to_vec(), b"verification".to_vec()] + ); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} diff --git a/ql-runtime/src/tests/session.rs b/ql-runtime/src/tests/session.rs new file mode 100644 index 00000000..ec351e35 --- /dev/null +++ b/ql-runtime/src/tests/session.rs @@ -0,0 +1,213 @@ +use std::{ + sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }, + time::Duration, +}; + +use bytes::Bytes; +use ql_wire::SessionCloseCode; + +use super::*; +use crate::QlStreamError; + +#[tokio::test(flavor = "current_thread")] +async fn close_session_aborts_active_streams_and_allows_reconnect() { + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + let inbound_b = pair.take_inbound(Side::B); + let (received_tx, received_rx) = async_channel::bounded(1); + pair.connect_and_wait(Side::A).await; + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let mut reader = stream.reader; + + assert_eq!( + next_chunk(&mut reader).await.unwrap(), + Some(vec![1, 2, 3, 4]) + ); + received_tx.send(()).await.unwrap(); + + let err = next_chunk(&mut reader).await.unwrap_err(); + assert_eq!(err, QlStreamError::NoSession); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + stream + .writer + .write(Bytes::from_static(&[1, 2, 3, 4])) + .await + .unwrap(); + received_rx.recv().await.unwrap(); + + pair.side(Side::A) + .handle + .close_session(SessionCloseCode::CANCELLED); + + let err = stream.writer.finish().await.unwrap_err(); + assert_eq!(err, QlStreamError::NoSession); + + await_status( + &pair.side(Side::A).status, + Some(pair.side(Side::B).peer), + PeerStatus::Disconnected, + ) + .await; + await_status( + &pair.side(Side::B).status, + Some(pair.side(Side::A).peer), + PeerStatus::Disconnected, + ) + .await; + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + + pair.connect_and_wait(Side::A).await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn unpair_aborts_active_streams_and_prevents_reconnect() { + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + let inbound_b = pair.take_inbound(Side::B); + let (received_tx, received_rx) = async_channel::bounded(1); + pair.connect_and_wait(Side::A).await; + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let mut reader = stream.reader; + + assert_eq!( + next_chunk(&mut reader).await.unwrap(), + Some(vec![5, 6, 7, 8]) + ); + received_tx.send(()).await.unwrap(); + + let err = next_chunk(&mut reader).await.unwrap_err(); + assert_eq!(err, QlStreamError::NoSession); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + stream + .writer + .write(Bytes::from_static(&[5, 6, 7, 8])) + .await + .unwrap(); + received_rx.recv().await.unwrap(); + + pair.side(Side::A).handle.unpair(); + + let err = stream.writer.finish().await.unwrap_err(); + assert_eq!(err, QlStreamError::NoSession); + + await_status(&pair.side(Side::A).status, None, PeerStatus::Unpaired).await; + await_status(&pair.side(Side::B).status, None, PeerStatus::Unpaired).await; + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + + assert!(matches!( + pair.side(Side::A).handle.open_stream(test_route_id()).await, + Err(NoSessionError) + )); + assert!(matches!( + pair.side(Side::B).handle.open_stream(test_route_id()).await, + Err(NoSessionError) + )); + + pair.side(Side::B).handle.connect(); + assert_no_status_for( + &pair.side(Side::B).status, + None, + PeerStatus::Initiator, + Duration::from_millis(150), + ) + .await; + assert_no_status_for( + &pair.side(Side::B).status, + None, + PeerStatus::Connected, + Duration::from_millis(150), + ) + .await; + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn session_timeout_disconnects_and_fails_pending_open() { + run_local_test(async { + let config_a = RuntimeConfig { + fsm: QlFsmConfig { + session_keepalive_interval: Duration::from_millis(40), + session_peer_timeout: Duration::from_millis(60), + ..default_runtime_config().fsm + }, + ..default_runtime_config() + }; + let config_b = default_runtime_config(); + let (platform_a, outbound_a, inbound_a_tx, status_a) = TestPlatform::new(); + let (platform_b, outbound_b, inbound_b_tx, status_b, inbound_b) = + TestPlatform::new_with_inbound(); + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + + let (runtime_a, handle_a) = new_runtime(identity_a.clone(), platform_a, config_a); + let (runtime_b, handle_b) = new_runtime(identity_b.clone(), platform_b, config_b); + + tokio::task::spawn_local(async move { runtime_a.run().await }); + tokio::task::spawn_local(async move { runtime_b.run().await }); + + let drop_flag = Arc::new(AtomicBool::new(false)); + spawn_forwarder(outbound_a, inbound_b_tx); + spawn_gated_forwarder(outbound_b, inbound_a_tx, drop_flag.clone()); + + register_peers(&handle_a, &handle_b, &identity_a, &identity_b); + handle_a.connect(); + + await_status(&status_a, Some(identity_b.qid), PeerStatus::Connected).await; + await_status(&status_b, Some(identity_a.qid), PeerStatus::Connected).await; + + let responder_task = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let _ = read_all(stream.reader).await; + let err = stream.writer.finish().await.unwrap_err(); + assert!(matches!(err, QlStreamError::NoSession)); + }); + + drop_flag.store(true, Ordering::Relaxed); + + let mut pending = handle_a.open_stream(test_route_id()).await.unwrap(); + let err = pending.writer.finish().await.unwrap_err(); + assert!(matches!(err, QlStreamError::NoSession)); + + await_status(&status_a, Some(identity_b.qid), PeerStatus::Disconnected).await; + + let result = + tokio::time::timeout(Duration::from_millis(300), next_chunk(&mut pending.reader)) + .await + .unwrap(); + assert!(matches!(result, Err(QlStreamError::NoSession))); + + responder_task.abort(); + }) + .await; +} diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs new file mode 100644 index 00000000..176711c8 --- /dev/null +++ b/ql-runtime/src/tests/stream.rs @@ -0,0 +1,673 @@ +use std::time::Duration; + +use bytes::Bytes; +use ql_wire::StreamCloseCode; + +use super::*; +use crate::QlStreamError; + +#[tokio::test(flavor = "current_thread")] +async fn open_stream_duplex_happy_path() { + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + + let mut writer = inbound.writer; + let mut reader = inbound.reader; + + assert_eq!(next_chunk(&mut reader).await.unwrap(), Some(vec![1, 2])); + writer.write(Bytes::from_static(&[9])).await.unwrap(); + assert_eq!(next_chunk(&mut reader).await.unwrap(), Some(vec![3, 4])); + writer.write(Bytes::from_static(&[8, 7])).await.unwrap(); + assert_eq!(next_chunk(&mut reader).await.unwrap(), None); + writer.finish().await.unwrap(); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + stream + .writer + .write(Bytes::from_static(&[1, 2])) + .await + .unwrap(); + assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), Some(vec![9])); + stream + .writer + .write(Bytes::from_static(&[3, 4])) + .await + .unwrap(); + stream.writer.finish().await.unwrap(); + assert_eq!( + next_chunk(&mut stream.reader).await.unwrap(), + Some(vec![8, 7]) + ); + assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), None); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn reader_respects_max_len() { + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + let mut reader = inbound.reader; + + assert_eq!( + next_chunk_max(&mut reader, 2).await.unwrap(), + Some(vec![1, 2]) + ); + assert_eq!( + next_chunk_max(&mut reader, 2).await.unwrap(), + Some(vec![3, 4]) + ); + assert_eq!( + next_chunk_max(&mut reader, 2).await.unwrap(), + Some(vec![5, 6]) + ); + assert_eq!(next_chunk(&mut reader).await.unwrap(), None); + + inbound.writer.finish().await.unwrap(); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + stream + .writer + .write(Bytes::from_static(&[1, 2, 3, 4, 5, 6])) + .await + .unwrap(); + stream.writer.finish().await.unwrap(); + assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), None); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn large_stream_payload_round_trips() { + run_local_test(async { + let payload: Vec = (0..40).collect(); + let mut pair = TestPair::new(default_runtime_config()); + let (done_tx, done_rx) = async_channel::bounded(1); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let request_data = read_all(stream.reader).await.unwrap(); + stream.writer.finish().await.unwrap(); + done_tx.send(request_data).await.unwrap(); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + stream + .writer + .write(Bytes::from(payload.clone())) + .await + .unwrap(); + stream.writer.finish().await.unwrap(); + assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), None); + + let received = tokio::time::timeout(Duration::from_secs(2), done_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(received, payload); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn dropping_responder_closes_initiator_response() { + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + drop(stream.reader); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + let err = stream.writer.finish().await.unwrap_err(); + assert!(matches!( + err, + QlStreamError::StreamClosed { code } if code == StreamCloseCode::CANCELLED + )); + + let err = next_chunk(&mut stream.reader).await.unwrap_err(); + assert!(matches!( + err, + QlStreamError::StreamClosed { code } if code == StreamCloseCode::CANCELLED + )); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn dropping_inbound_reader_cancels_remote_writer() { + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + let inbound_b = pair.take_inbound(Side::B); + let (go_tx, go_rx) = async_channel::bounded(1); + pair.connect_and_wait(Side::A).await; + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let mut writer = stream.writer; + let mut reader = stream.reader; + assert_eq!(next_chunk(&mut reader).await.unwrap(), None); + writer + .write(Bytes::from_static(&[1, 2, 3, 4])) + .await + .unwrap(); + go_rx.recv().await.unwrap(); + let _ = writer.write(Bytes::from(vec![5; 64])).await; + let err = writer.finish().await.unwrap_err(); + assert!(matches!( + err, + QlStreamError::StreamClosed { code } if code == StreamCloseCode::CANCELLED + )); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + stream.writer.finish().await.unwrap(); + assert_eq!( + next_chunk(&mut stream.reader).await.unwrap(), + Some(vec![1, 2, 3, 4]) + ); + drop(stream.reader); + go_tx.send(()).await.unwrap(); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn closing_initiator_reader_preserves_initiator_writer() { + run_local_test(async { + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + let (done_tx, done_rx) = async_channel::bounded(1); + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let request = read_all(stream.reader).await.unwrap(); + done_tx.send(request).await.unwrap(); + }); + + let stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + let mut writer = stream.writer; + stream.reader.close(StreamCloseCode::CANCELLED); + + writer.write(Bytes::from_static(&[1, 2])).await.unwrap(); + writer.write(Bytes::from_static(&[3, 4])).await.unwrap(); + writer.finish().await.unwrap(); + + let request = tokio::time::timeout(Duration::from_secs(2), done_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(request, vec![1, 2, 3, 4]); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn max_concurrent_message_writes_is_respected() { + run_local_test(async { + let stats = WriteStats::new(); + let config = RuntimeConfig { + max_concurrent_message_writes: 2, + ..default_runtime_config() + }; + let (platform_a, outbound_a, inbound_a_tx, status_a) = + TestPlatform::new_with_delayed_writes(Duration::from_millis(40), stats.clone()); + let (platform_b, outbound_b, inbound_b_tx, status_b, inbound_b) = + TestPlatform::new_with_inbound(); + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + + let (runtime_a, handle_a) = new_runtime(identity_a.clone(), platform_a, config); + let (runtime_b, handle_b) = new_runtime(identity_b.clone(), platform_b, config); + + tokio::task::spawn_local(async move { runtime_a.run().await }); + tokio::task::spawn_local(async move { runtime_b.run().await }); + + spawn_forwarder(outbound_a, inbound_b_tx); + spawn_forwarder(outbound_b, inbound_a_tx); + + register_peers(&handle_a, &handle_b, &identity_a, &identity_b); + handle_a.connect(); + + await_status(&status_a, Some(identity_b.qid), PeerStatus::Connected).await; + await_status(&status_b, Some(identity_a.qid), PeerStatus::Connected).await; + + let responder = tokio::task::spawn_local(async move { + for _ in 0..4 { + let stream = inbound_b.recv().await.unwrap(); + let _ = read_all(stream.reader).await; + let mut writer = stream.writer; + writer.queue_finish(); + } + }); + + let mut tasks = Vec::new(); + for i in 0..4u8 { + let handle = handle_a.clone(); + tasks.push(tokio::task::spawn_local(async move { + let mut stream = handle.open_stream(test_route_id()).await.unwrap(); + stream.writer.write(Bytes::from(vec![i; 8])).await.unwrap(); + stream.writer.finish().await.unwrap(); + assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), None); + })); + } + + for task in tasks { + tokio::time::timeout(Duration::from_secs(4), task) + .await + .unwrap() + .unwrap(); + } + + tokio::time::timeout(Duration::from_secs(4), responder) + .await + .unwrap() + .unwrap(); + + assert!( + stats.max_active() <= 2, + "max active writes exceeded: {}", + stats.max_active() + ); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn stream_round_trip_survives_encrypted_packet_drops() { + run_local_test(async { + let config = RuntimeConfig { + fsm: QlFsmConfig { + session_record_retransmit_timeout: Duration::from_millis(20), + ..default_runtime_config().fsm + }, + ..default_runtime_config() + }; + let (platform_a, outbound_a, inbound_a_tx, status_a) = TestPlatform::new(); + let (platform_b, outbound_b, inbound_b_tx, status_b, inbound_b) = + TestPlatform::new_with_inbound(); + let (identity_a, identity_b) = test_identities(&SoftwareCrypto); + + let request_payload: Vec = (0..32).collect(); + let response_payload: Vec = (100..132).collect(); + let expected_response = response_payload.clone(); + + let (runtime_a, handle_a) = new_runtime(identity_a.clone(), platform_a, config); + let (runtime_b, handle_b) = new_runtime(identity_b.clone(), platform_b, config); + + tokio::task::spawn_local(async move { runtime_a.run().await }); + tokio::task::spawn_local(async move { runtime_b.run().await }); + + spawn_drop_every_nth_encrypted_forwarder(outbound_a, inbound_b_tx, 3); + spawn_drop_every_nth_encrypted_forwarder(outbound_b, inbound_a_tx, 3); + + register_peers(&handle_a, &handle_b, &identity_a, &identity_b); + handle_a.connect(); + + await_status(&status_a, Some(identity_b.qid), PeerStatus::Connected).await; + await_status(&status_b, Some(identity_a.qid), PeerStatus::Connected).await; + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let received_request = read_all(stream.reader).await.unwrap(); + let mut writer = stream.writer; + writer + .write(Bytes::from(response_payload.clone())) + .await + .unwrap(); + writer.finish().await.unwrap(); + received_request + }); + + let mut stream = handle_a.open_stream(test_route_id()).await.unwrap(); + stream + .writer + .write(Bytes::from(request_payload.clone())) + .await + .unwrap(); + stream.writer.finish().await.unwrap(); + + let mut received_response = Vec::new(); + while let Some(chunk) = next_chunk(&mut stream.reader).await.unwrap() { + received_response.extend_from_slice(&chunk); + } + assert_eq!(received_response, expected_response); + + let received_request = tokio::time::timeout(Duration::from_secs(4), responder) + .await + .unwrap() + .unwrap(); + assert_eq!(received_request, request_payload); + }) + .await; +} + +#[allow(clippy::too_many_lines)] +#[tokio::test(flavor = "current_thread")] +async fn multi_megabyte_stream_survives_asymmetric_loss_and_delay() { + run_local_test_timeout(Duration::from_secs(10), async { + let payload_len = 2 * 1024 * 1024; + let chunk_len = 16 * 1024; + let payload: Vec = (0..payload_len) + .map(|i| u8::try_from(i % 251).unwrap()) + .collect(); + let expected = payload.clone(); + let config = RuntimeConfig { + fsm: QlFsmConfig { + session_record_max_size: 16 * 1024, + session_record_ack_delay: Duration::from_millis(2), + session_record_retransmit_timeout: Duration::from_millis(25), + session_stream_send_buffer_size: 4 * 1024 * 1024, + session_stream_receive_buffer_size: 4 * 1024 * 1024, + session_accepted_record_window: 16 * 1024, + session_pending_ack_range_limit: 4 * 1024, + ..default_runtime_config().fsm + }, + ..default_runtime_config() + }; + let (mut pair, links) = TestPair::new_with_controlled_links( + config, + LinkBehavior { + base_delay: Duration::from_millis(1), + drop_encrypted_every: Some(41), + delay_encrypted_every: Some((13, Duration::from_millis(12))), + ..LinkBehavior::default() + }, + LinkBehavior { + base_delay: Duration::from_millis(1), + ..LinkBehavior::default() + }, + ); + pair.connect_and_wait(Side::A).await; + links.b_to_a.store(LinkBehavior { + base_delay: Duration::from_millis(3), + drop_encrypted_every: Some(7), + duplicate_encrypted_every: Some(19), + delay_encrypted_every: Some((3, Duration::from_millis(25))), + }); + let inbound_b = pair.take_inbound(Side::B); + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + eprintln!("responder accepted inbound stream"); + let mut reader = stream.reader; + let mut received = Vec::new(); + while let Some(chunk) = next_chunk(&mut reader).await.unwrap() { + if received.len() >= 36 * chunk_len { + eprintln!("responder received chunk of {} bytes", chunk.len()); + } + received.extend_from_slice(&chunk); + if received.len() % (256 * 1024) == 0 { + eprintln!("responder received {} bytes", received.len()); + } + } + stream.writer.finish().await.unwrap(); + received + }); + + let recovery_links = links.clone(); + let recovery = tokio::task::spawn_local(async move { + tokio::time::sleep(Duration::from_millis(300)).await; + eprintln!("restoring reverse path"); + recovery_links.b_to_a.store(LinkBehavior { + base_delay: Duration::from_millis(1), + delay_encrypted_every: Some((17, Duration::from_millis(8))), + ..LinkBehavior::default() + }); + }); + + let writer = tokio::task::spawn_local(async move { + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + for (index, chunk) in payload.chunks(chunk_len).enumerate() { + if index + 1 >= 40 { + eprintln!("writer attempting chunk {}", index + 1); + } + stream + .writer + .write(Bytes::copy_from_slice(chunk)) + .await + .unwrap(); + if index + 1 >= 40 { + eprintln!("writer queued chunk {}", index + 1); + } + if index % 16 == 15 { + eprintln!("writer queued {} chunks", index + 1); + } + } + eprintln!("writer finished queueing"); + stream.writer.finish().await.unwrap(); + eprintln!("writer waiting for eof"); + assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), None); + eprintln!("writer observed eof"); + }); + + tokio::time::timeout(Duration::from_secs(30), writer) + .await + .unwrap() + .unwrap(); + tokio::time::timeout(Duration::from_secs(2), recovery) + .await + .unwrap() + .unwrap(); + let received = tokio::time::timeout(Duration::from_secs(30), responder) + .await + .unwrap() + .unwrap(); + assert_eq!(received, expected); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn reproducer_writer_stalls_after_reverse_path_impairment() { + run_local_test_timeout(Duration::from_secs(10), async { + let payload_len = 2 * 1024 * 1024; + let chunk_len = 16 * 1024; + let payload: Vec = (0..payload_len) + .map(|i| u8::try_from(i % 251).unwrap()) + .collect(); + let config = RuntimeConfig { + fsm: QlFsmConfig { + session_record_max_size: 16 * 1024, + session_record_ack_delay: Duration::from_millis(2), + session_record_retransmit_timeout: Duration::from_millis(25), + session_stream_send_buffer_size: 4 * 1024 * 1024, + session_stream_receive_buffer_size: 4 * 1024 * 1024, + session_accepted_record_window: 16 * 1024, + session_pending_ack_range_limit: 4 * 1024, + ..default_runtime_config().fsm + }, + ..default_runtime_config() + }; + let (mut pair, links) = TestPair::new_with_controlled_links( + config, + LinkBehavior { + base_delay: Duration::from_millis(1), + drop_encrypted_every: Some(41), + delay_encrypted_every: Some((13, Duration::from_millis(12))), + ..LinkBehavior::default() + }, + LinkBehavior { + base_delay: Duration::from_millis(1), + ..LinkBehavior::default() + }, + ); + pair.connect_and_wait(Side::A).await; + links.b_to_a.store(LinkBehavior { + base_delay: Duration::from_millis(3), + drop_encrypted_every: Some(7), + duplicate_encrypted_every: Some(19), + delay_encrypted_every: Some((3, Duration::from_millis(25))), + }); + let inbound_b = pair.take_inbound(Side::B); + + let responder = tokio::task::spawn_local(async move { + let stream = inbound_b.recv().await.unwrap(); + let mut reader = stream.reader; + while next_chunk(&mut reader).await.unwrap().is_some() {} + }); + + let recovery_links = links.clone(); + let recovery = tokio::task::spawn_local(async move { + tokio::time::sleep(Duration::from_millis(300)).await; + recovery_links.b_to_a.store(LinkBehavior { + base_delay: Duration::from_millis(1), + delay_encrypted_every: Some((17, Duration::from_millis(8))), + ..LinkBehavior::default() + }); + }); + + let writer = tokio::task::spawn_local(async move { + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + for chunk in payload.chunks(chunk_len) { + stream + .writer + .write(Bytes::copy_from_slice(chunk)) + .await + .unwrap(); + } + stream.writer.queue_finish(); + let _ = next_chunk(&mut stream.reader).await; + }); + + let _ = tokio::time::timeout(Duration::from_secs(15), writer).await; + recovery.abort(); + responder.abort(); + }) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn responder_drains_multiple_local_chunks_per_writable_wake() { + run_local_test(async { + let chunk_len = 4104usize; + let chunk_count = 5usize; + let expected = vec![0x5a; chunk_len * chunk_count]; + let mut pair = TestPair::new(default_runtime_config()); + pair.connect_and_wait(Side::A).await; + let inbound_b = pair.take_inbound(Side::B); + + let responder = tokio::task::spawn_local(async move { + let inbound = inbound_b.recv().await.unwrap(); + let _ = read_all(inbound.reader).await.unwrap(); + + let mut writer = inbound.writer; + for _ in 0..chunk_count { + writer + .write(Bytes::from(vec![0x5a; chunk_len])) + .await + .unwrap(); + } + writer.finish().await.unwrap(); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_route_id()) + .await + .unwrap(); + stream + .writer + .write(Bytes::from_static(b"request")) + .await + .unwrap(); + stream.writer.finish().await.unwrap(); + + let received = read_all(stream.reader).await.unwrap(); + assert_eq!(received, expected); + + tokio::time::timeout(Duration::from_secs(2), responder) + .await + .unwrap() + .unwrap(); + }) + .await; +} From 6c0bb2f7e89d181adc5b911a8d213e710897fbc4 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 4 Jun 2026 09:38:37 -0400 Subject: [PATCH 06/59] workspace: remove legacy QLV1 crates --- Cargo.lock | 1845 +---------------- Cargo.toml | 17 - Justfile | 12 - README.md | 11 +- api/.gitignore | 2 - api/Cargo.toml | 31 - api/src/api/backup.rs | 264 --- api/src/api/bitcoin.rs | 32 - api/src/api/firmware.rs | 116 -- api/src/api/fx.rs | 34 - api/src/api/message.rs | 149 -- api/src/api/mod.rs | 13 - api/src/api/onboarding.rs | 47 - api/src/api/pairing.rs | 39 - api/src/api/passport.rs | 25 - api/src/api/quantum_link.rs | 431 ---- api/src/api/scv.rs | 52 - api/src/api/status.rs | 35 - api/src/api/tests.rs | 414 ---- api/src/lib.rs | 11 - api/tests/golden_tests.rs | 550 ----- .../golden_tests__golden_account_update.snap | 5 - ...n_tests__golden_apply_passphrase_none.snap | 5 - ...n_tests__golden_apply_passphrase_some.snap | 5 - ...en_tests__golden_backup_shard_request.snap | 5 - ...s__golden_backup_shard_response_error.snap | 5 - ..._golden_backup_shard_response_success.snap | 5 - ...n_tests__golden_broadcast_transaction.snap | 5 - ...olden_create_magic_backup_event_chunk.snap | 5 - ...olden_create_magic_backup_event_start.snap | 5 - ...lden_create_magic_backup_result_error.snap | 5 - ...en_create_magic_backup_result_success.snap | 5 - .../golden_tests__golden_device_status.snap | 5 - ..._tests__golden_device_status_updating.snap | 5 - ...en_envoy_magic_backup_enabled_request.snap | 5 - ...n_envoy_magic_backup_enabled_response.snap | 5 - .../golden_tests__golden_envoy_status.snap | 5 - .../golden_tests__golden_exchange_rate.snap | 5 - ...n_tests__golden_exchange_rate_history.snap | 5 - ...ts__golden_firmware_fetch_event_chunk.snap | 5 - ...lden_firmware_fetch_event_downloading.snap | 5 - ...ts__golden_firmware_fetch_event_error.snap | 5 - ...en_firmware_fetch_event_not_available.snap | 5 - ..._golden_firmware_fetch_event_starting.snap | 5 - ..._tests__golden_firmware_fetch_request.snap | 5 - ..._golden_firmware_update_check_request.snap | 5 - ...mware_update_check_response_available.snap | 5 - ...e_update_check_response_not_available.snap | 5 - ...__golden_firmware_update_result_error.snap | 5 - ..._firmware_update_result_error_install.snap | 5 - ...n_firmware_update_result_error_verify.snap | 5 - ...den_firmware_update_result_installing.snap | 5 - ...lden_firmware_update_result_rebooting.snap | 5 - ...golden_firmware_update_result_success.snap | 5 - ...irmware_update_result_update_verified.snap | 5 - .../golden_tests__golden_heartbeat.snap | 5 - ...ts__golden_onboarding_state_completed.snap | 5 - ...boarding_state_firmware_update_screen.snap | 5 - .../golden_tests__golden_pairing_request.snap | 5 - ...golden_tests__golden_pairing_response.snap | 6 - ...ts__golden_prime_magic_backup_enabled.snap | 5 - ...den_prime_magic_backup_status_request.snap | 5 - ...en_prime_magic_backup_status_response.snap | 5 - .../golden_tests__golden_raw_data.snap | 5 - ...lden_restore_magic_backup_event_chunk.snap | 5 - ...lden_restore_magic_backup_event_error.snap | 5 - ..._restore_magic_backup_event_no_backup.snap | 5 - ...n_restore_magic_backup_event_starting.snap | 5 - ...__golden_restore_magic_backup_request.snap | 5 - ...den_restore_magic_backup_result_error.snap | 5 - ...n_restore_magic_backup_result_success.snap | 5 - ...n_tests__golden_restore_shard_request.snap | 5 - ...__golden_restore_shard_response_error.snap | 5 - ...lden_restore_shard_response_not_found.snap | 5 - ...golden_restore_shard_response_success.snap | 5 - ...lden_security_check_challenge_request.snap | 5 - ...curity_check_challenge_response_error.snap | 5 - ...rity_check_challenge_response_success.snap | 5 - ...den_security_check_verification_error.snap | 5 - ...n_security_check_verification_success.snap | 5 - .../golden_tests__golden_sign_psbt.snap | 5 - quantum-link-macros/Cargo.toml | 14 - quantum-link-macros/src/lib.rs | 632 ------ 83 files changed, 45 insertions(+), 5032 deletions(-) delete mode 100644 api/.gitignore delete mode 100644 api/Cargo.toml delete mode 100644 api/src/api/backup.rs delete mode 100644 api/src/api/bitcoin.rs delete mode 100644 api/src/api/firmware.rs delete mode 100644 api/src/api/fx.rs delete mode 100644 api/src/api/message.rs delete mode 100644 api/src/api/mod.rs delete mode 100644 api/src/api/onboarding.rs delete mode 100644 api/src/api/pairing.rs delete mode 100644 api/src/api/passport.rs delete mode 100644 api/src/api/quantum_link.rs delete mode 100644 api/src/api/scv.rs delete mode 100644 api/src/api/status.rs delete mode 100644 api/src/api/tests.rs delete mode 100644 api/src/lib.rs delete mode 100644 api/tests/golden_tests.rs delete mode 100644 api/tests/snapshots/golden_tests__golden_account_update.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_apply_passphrase_none.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_apply_passphrase_some.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_backup_shard_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_backup_shard_response_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_backup_shard_response_success.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_broadcast_transaction.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_create_magic_backup_event_chunk.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_create_magic_backup_event_start.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_create_magic_backup_result_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_create_magic_backup_result_success.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_device_status.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_device_status_updating.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_envoy_magic_backup_enabled_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_envoy_magic_backup_enabled_response.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_envoy_status.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_exchange_rate.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_exchange_rate_history.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_fetch_event_chunk.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_fetch_event_downloading.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_fetch_event_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_fetch_event_not_available.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_fetch_event_starting.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_fetch_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_check_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_check_response_available.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_check_response_not_available.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_result_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_result_error_install.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_result_error_verify.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_result_installing.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_result_rebooting.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_result_success.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_firmware_update_result_update_verified.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_heartbeat.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_onboarding_state_completed.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_onboarding_state_firmware_update_screen.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_pairing_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_pairing_response.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_prime_magic_backup_enabled.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_prime_magic_backup_status_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_prime_magic_backup_status_response.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_raw_data.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_chunk.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_no_backup.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_starting.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_magic_backup_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_magic_backup_result_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_magic_backup_result_success.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_shard_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_shard_response_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_shard_response_not_found.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_restore_shard_response_success.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_security_check_challenge_request.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_security_check_challenge_response_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_security_check_challenge_response_success.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_security_check_verification_error.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_security_check_verification_success.snap delete mode 100644 api/tests/snapshots/golden_tests__golden_sign_psbt.snap delete mode 100644 quantum-link-macros/Cargo.toml delete mode 100644 quantum-link-macros/src/lib.rs diff --git a/Cargo.lock b/Cargo.lock index 123d0e59..c1e30d37 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -11,28 +11,12 @@ dependencies = [ "gimli", ] -[[package]] -name = "adler" -version = "1.0.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f26201604c87b1e01bd3d98f8d5d9a8fcbb815e8cedb41ffccbeb4bf593a35fe" - [[package]] name = "adler2" version = "2.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" -[[package]] -name = "aead" -version = "0.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d122413f284cf2d62fb1b7db97e02edb8cda96d769b16e443a4f6195e35662b0" -dependencies = [ - "crypto-common", - "generic-array", -] - [[package]] name = "aho-corasick" version = "1.1.3" @@ -42,40 +26,12 @@ dependencies = [ "memchr", ] -[[package]] -name = "allo-isolate" -version = "0.1.27" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "449e356a4864c017286dbbec0e12767ea07efba29e3b7d984194c2a7ff3c4550" -dependencies = [ - "anyhow", - "atomic", - "backtrace", -] - [[package]] name = "android-tzdata" version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e999941b234f3131b00bc13c22d06e8c5ff726d1b6318ac7eb276997bbb4fef0" -[[package]] -name = "android_log-sys" -version = "0.3.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "84521a3cf562bc62942e294181d9eef17eb38ceb8c68677bc49f144e4c3d4f8d" - -[[package]] -name = "android_logger" -version = "0.15.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dbb4e440d04be07da1f1bf44fb4495ebd58669372fe0cffa6e48595ac5bd88a3" -dependencies = [ - "android_log-sys", - "env_filter 0.1.3", - "log", -] - [[package]] name = "android_system_properties" version = "0.1.5" @@ -135,30 +91,6 @@ dependencies = [ "windows-sys 0.61.2", ] -[[package]] -name = "anyhow" -version = "1.0.99" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b0674a1ddeecb70197781e945de4b3b8ffb61fa939a5597bcf48503737663100" - -[[package]] -name = "argon2" -version = "0.5.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c3610892ee6e0cbce8ae2700349fcf8f98adb0dbfbee85aec3c9179d29cc072" -dependencies = [ - "base64ct", - "blake2", - "cpufeatures", - "password-hash", -] - -[[package]] -name = "arrayvec" -version = "0.7.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" - [[package]] name = "async-channel" version = "2.5.0" @@ -171,12 +103,6 @@ dependencies = [ "pin-project-lite", ] -[[package]] -name = "atomic" -version = "0.5.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c59bdb34bc650a32731b31bd8f0829cc15d24a708ee31559e0bb34f2bc320cba" - [[package]] name = "autocfg" version = "1.5.0" @@ -192,7 +118,7 @@ dependencies = [ "addr2line", "cfg-if", "libc", - "miniz_oxide 0.8.9", + "miniz_oxide", "object", "rustc-demangle", "windows-targets", @@ -208,153 +134,6 @@ dependencies = [ "zeroize", ] -[[package]] -name = "base16ct" -version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4c7f02d4ea65f2c1853089ffd8d2787bdbc63de2f0d29dedbcf8ccdfa0ccd4cf" - -[[package]] -name = "base64" -version = "0.22.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" - -[[package]] -name = "base64ct" -version = "1.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "55248b47b0caf0546f7988906588779981c43bb1bc9d0c44087278f80cdb44ba" - -[[package]] -name = "bc-components" -version = "0.28.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "64fc6326f9838e1332cb767fba7ce6a31fa8f14912f13dd42427125263722b9f" -dependencies = [ - "bc-crypto", - "bc-rand", - "bc-tags", - "bc-ur", - "dcbor", - "hex", - "miniz_oxide 0.7.4", - "pqcrypto-mldsa", - "pqcrypto-mlkem", - "pqcrypto-traits", - "rand_core 0.6.4", - "ssh-key", - "sskr", - "thiserror", - "url", - "zeroize", -] - -[[package]] -name = "bc-crypto" -version = "0.13.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9644245d48f4ab1bfa8c7eebfbd20d2bea7895e220de766f66876c5a71b14712" -dependencies = [ - "argon2", - "bc-rand", - "chacha20poly1305", - "crc32fast", - "ed25519-dalek", - "hex", - "hkdf", - "hmac", - "pbkdf2", - "rand 0.8.5", - "scrypt", - "secp256k1", - "sha2", - "thiserror", - "x25519-dalek", -] - -[[package]] -name = "bc-envelope" -version = "0.37.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "515acbccbbbc35f5ac024b890fdeec084607c73f4f39c0fb231a356823ef272c" -dependencies = [ - "bc-components", - "bc-crypto", - "bc-rand", - "bc-ur", - "bytes", - "dcbor", - "hex", - "itertools", - "known-values", - "paste", - "ssh-key", - "thiserror", -] - -[[package]] -name = "bc-rand" -version = "0.4.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fdade83e92b8dfb9acbccd68e09f9e9555dbaf64c7bc2e5fbb894fcc9b53b413" -dependencies = [ - "getrandom 0.2.16", - "lazy_static", - "num-traits", - "rand 0.8.5", - "rand_core 0.6.4", - "rand_xoshiro", -] - -[[package]] -name = "bc-shamir" -version = "0.12.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0cc9e00fbb348a889d951b0a57b04cb609ebd5b123231a6a7c18b4d057825823" -dependencies = [ - "bc-crypto", - "bc-rand", - "thiserror", -] - -[[package]] -name = "bc-tags" -version = "0.9.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "947dee941635701788b56a6557cf9f89b9750bcc6aae76cf10863759b9964c4a" -dependencies = [ - "dcbor", - "paste", -] - -[[package]] -name = "bc-ur" -version = "0.16.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ac0af650d34ec93be355e81f22df87b108a419d7bc775a5ee1fb00f5daeb9376" -dependencies = [ - "dcbor", - "thiserror", - "ur", -] - -[[package]] -name = "bc-xid" -version = "0.16.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2206b65d39a6057ad75301cb691a6ea0e2f09556551936640b436c6f269f260e" -dependencies = [ - "bc-components", - "bc-envelope", - "bc-rand", - "bc-ur", - "dcbor", - "hex", - "provenance-mark", - "thiserror", -] - [[package]] name = "bit-set" version = "0.8.0" @@ -370,52 +149,12 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" -[[package]] -name = "bitcoin-io" -version = "0.1.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0b47c4ab7a93edb0c7198c5535ed9b52b63095f4e9b45279c6736cec4b856baf" - -[[package]] -name = "bitcoin-private" -version = "0.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "73290177011694f38ec25e165d0387ab7ea749a4b81cd4c80dae5988229f7a57" - -[[package]] -name = "bitcoin_hashes" -version = "0.12.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5d7066118b13d4b20b23645932dfb3a81ce7e29f95726c2036fa33cd7b092501" -dependencies = [ - "bitcoin-private", -] - -[[package]] -name = "bitcoin_hashes" -version = "0.14.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bb18c03d0db0247e147a21a6faafd5a7eb851c743db062de72018b6b7e8e4d16" -dependencies = [ - "bitcoin-io", - "hex-conservative", -] - [[package]] name = "bitflags" version = "2.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" -[[package]] -name = "blake2" -version = "0.10.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "46502ad458c9a52b69d4d4d32775c788b7a1b85e8bc9d482d92250fc0e3f8efe" -dependencies = [ - "digest", -] - [[package]] name = "block-buffer" version = "0.10.4" @@ -432,16 +171,10 @@ dependencies = [ "bytemuck", "consts", "getrandom 0.2.16", - "rand 0.9.2", + "rand", "thiserror", ] -[[package]] -name = "build-target" -version = "0.4.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "832133bbabbbaa9fbdba793456a2827627a7d2b8fb96032fa1e7666d7895832b" - [[package]] name = "bumpalo" version = "3.19.0" @@ -468,7 +201,7 @@ checksum = "89385e82b5d1821d2219e0b095efa2cc1f246cbf99080f3be46a1a85c0d392d9" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -488,15 +221,9 @@ checksum = "4f154e572231cb6ba2bd1176980827e3d5dc04cc183a75dea38109fbdd672d29" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] -[[package]] -name = "byteorder" -version = "1.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" - [[package]] name = "bytes" version = "1.10.1" @@ -509,8 +236,6 @@ version = "1.2.34" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42bc4aea80032b7bf409b0bc7ccad88853858911b7713a8062fdc0623867bedc" dependencies = [ - "jobserver", - "libc", "shlex", ] @@ -520,30 +245,6 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2fd1289c04a9ea8cb22300a459a72a385d7c73d3259e2ed7dcb2af674838cfa9" -[[package]] -name = "chacha20" -version = "0.9.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c3613f74bd2eac03dad61bd53dbe620703d4371614fe0bc3b9f04dd36fe4e818" -dependencies = [ - "cfg-if", - "cipher", - "cpufeatures", -] - -[[package]] -name = "chacha20poly1305" -version = "0.10.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "10cd79432192d1c0f4e1a0fef9527696cc039165d729fb41b3f4f4f354c2dc35" -dependencies = [ - "aead", - "chacha20", - "cipher", - "poly1305", - "zeroize", -] - [[package]] name = "chrono" version = "0.4.41" @@ -554,22 +255,10 @@ dependencies = [ "iana-time-zone", "js-sys", "num-traits", - "serde", "wasm-bindgen", "windows-link 0.1.3", ] -[[package]] -name = "cipher" -version = "0.4.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" -dependencies = [ - "crypto-common", - "inout", - "zeroize", -] - [[package]] name = "colorchoice" version = "1.0.5" @@ -598,22 +287,6 @@ dependencies = [ "windows-sys 0.59.0", ] -[[package]] -name = "console_error_panic_hook" -version = "0.1.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a06aeb73f470f66dcdbf7223caeebb85984942f22f1adb2a088cf9668146bbbc" -dependencies = [ - "cfg-if", - "wasm-bindgen", -] - -[[package]] -name = "const-oid" -version = "0.9.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c2459377285ad874054d797f3ccebf984978aa39129f6eafde5cdc8315b612f8" - [[package]] name = "consts" version = "1.0.0" @@ -633,7 +306,7 @@ checksum = "657f625ff361906f779745d08375ae3cc9fef87a35fba5f22874cf773010daf4" dependencies = [ "hax-lib", "pastey", - "rand 0.9.2", + "rand", ] [[package]] @@ -645,30 +318,6 @@ dependencies = [ "libc", ] -[[package]] -name = "crc" -version = "3.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9710d3b3739c2e349eb44fe848ad0b7c8cb1e42bd87ee49371df2f7acaf3e675" -dependencies = [ - "crc-catalog", -] - -[[package]] -name = "crc-catalog" -version = "2.4.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "19d374276b40fb8bbdee95aef7c7fa6b5316ec764510eb64b8dd0e2ed0d7e7f5" - -[[package]] -name = "crc32fast" -version = "1.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9481c1c90cbf2ac953f07c8d4a58aa3945c425b7185c9154d67a65e4230da511" -dependencies = [ - "cfg-if", -] - [[package]] name = "crossbeam-utils" version = "0.8.21" @@ -681,18 +330,6 @@ version = "0.2.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" -[[package]] -name = "crypto-bigint" -version = "0.5.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0dc92fb57ca44df6db8059111ab3af99a63d5d0f8375d9972e319a379c6bab76" -dependencies = [ - "generic-array", - "rand_core 0.6.4", - "subtle", - "zeroize", -] - [[package]] name = "crypto-common" version = "0.1.6" @@ -700,59 +337,9 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1bfb12502f3fc46cca1bb51ac28df9d618d813cdc3d2f25b9fe775a34af26bb3" dependencies = [ "generic-array", - "rand_core 0.6.4", "typenum", ] -[[package]] -name = "curve25519-dalek" -version = "4.1.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "97fb8b7c4503de7d6ae7b42ab72a5a59857b4c937ec27a3d4539dba95b5ab2be" -dependencies = [ - "cfg-if", - "cpufeatures", - "curve25519-dalek-derive", - "digest", - "fiat-crypto", - "rustc_version", - "subtle", - "zeroize", -] - -[[package]] -name = "curve25519-dalek-derive" -version = "0.1.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f46882e17999c6cc590af592290432be3bce0428cb0d5f8b6715e4dc7b383eb3" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - -[[package]] -name = "dart-sys" -version = "4.1.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "57967e4b200d767d091b961d6ab42cc7d0cc14fe9e052e75d0d3cf9eb732d895" -dependencies = [ - "cc", -] - -[[package]] -name = "dashmap" -version = "5.5.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "978747c1d849a7d2ee5e8adc0159961c48fb7e5db2f06af6723b80123bb53856" -dependencies = [ - "cfg-if", - "hashbrown 0.14.5", - "lock_api", - "once_cell", - "parking_lot_core", -] - [[package]] name = "dcbor" version = "0.23.3" @@ -767,27 +354,6 @@ dependencies = [ "unicode-normalization", ] -[[package]] -name = "delegate-attr" -version = "0.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "51aac4c99b2e6775164b412ea33ae8441b2fde2dbf05a20bc0052a63d08c475b" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - -[[package]] -name = "der" -version = "0.7.10" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e7c1832837b905bbfb5101e07cc24c8deddf52f93225eee6ead5f4d63d53ddcb" -dependencies = [ - "const-oid", - "zeroize", -] - [[package]] name = "diatomic-waker" version = "0.2.3" @@ -801,124 +367,15 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ "block-buffer", - "const-oid", "crypto-common", - "subtle", ] [[package]] -name = "displaydoc" -version = "0.2.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "97369cbbc041bc366949bc74d34658d6cda5621039731c6310521892a3a20ae0" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - -[[package]] -name = "dsa" -version = "0.6.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "48bc224a9084ad760195584ce5abb3c2c34a225fa312a128ad245a6b412b7689" -dependencies = [ - "digest", - "num-bigint-dig", - "num-traits", - "pkcs8", - "rfc6979", - "sha2", - "signature", - "zeroize", -] - -[[package]] -name = "dunce" -version = "1.0.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813" - -[[package]] -name = "ecdsa" -version = "0.16.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ee27f32b5c5292967d2d4a9d7f1e0b0aed2c15daded5a60300e4abb9d8020bca" -dependencies = [ - "der", - "digest", - "elliptic-curve", - "rfc6979", - "signature", - "spki", -] - -[[package]] -name = "ed25519" -version = "2.2.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "115531babc129696a58c64a4fef0a8bf9e9698629fb97e9e40767d235cfbcd53" -dependencies = [ - "pkcs8", - "signature", -] - -[[package]] -name = "ed25519-dalek" -version = "2.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "70e796c081cee67dc755e1a36a0a172b897fab85fc3f6bc48307991f64e4eca9" -dependencies = [ - "curve25519-dalek", - "ed25519", - "rand_core 0.6.4", - "serde", - "sha2", - "subtle", - "zeroize", -] - -[[package]] -name = "either" -version = "1.15.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" - -[[package]] -name = "elliptic-curve" -version = "0.13.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b5e6043086bf7973472e0c7dff2142ea0b680d30e18d9cc40f267efbf222bd47" -dependencies = [ - "base16ct", - "crypto-bigint", - "digest", - "ff", - "generic-array", - "group", - "pkcs8", - "rand_core 0.6.4", - "sec1", - "subtle", - "zeroize", -] - -[[package]] -name = "encode_unicode" -version = "1.0.0" +name = "encode_unicode" +version = "1.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "34aa73646ffb006b8f5147f3dc182bd4bcb190227ce861fc4a4844bf8e3cb2c0" -[[package]] -name = "env_filter" -version = "0.1.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "186e05a59d4c50738528153b83b0b0194d3a29507dfec16eccd4b342903397d0" -dependencies = [ - "log", - "regex", -] - [[package]] name = "env_filter" version = "1.0.1" @@ -937,7 +394,7 @@ checksum = "0621c04f2196ac3f488dd583365b9c09be011a4ab8b9f37248ffcc8f6198b56a" dependencies = [ "anstream", "anstyle", - "env_filter 1.0.1", + "env_filter", "jiff", "log", ] @@ -986,138 +443,18 @@ version = "2.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be" -[[package]] -name = "ff" -version = "0.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c0b50bfb653653f9ca9095b427bed08ab8d75a137839d9ad64eb11810d5b6393" -dependencies = [ - "rand_core 0.6.4", - "subtle", -] - -[[package]] -name = "fiat-crypto" -version = "0.2.9" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d" - -[[package]] -name = "flutter_rust_bridge" -version = "2.11.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dde126295b2acc5f0a712e265e91b6fdc0ed38767496483e592ae7134db83725" -dependencies = [ - "allo-isolate", - "android_logger", - "anyhow", - "build-target", - "bytemuck", - "byteorder", - "console_error_panic_hook", - "dart-sys", - "delegate-attr", - "flutter_rust_bridge_macros", - "futures", - "js-sys", - "lazy_static", - "log", - "oslog", - "portable-atomic", - "threadpool", - "tokio", - "wasm-bindgen", - "wasm-bindgen-futures", - "web-sys", -] - -[[package]] -name = "flutter_rust_bridge_macros" -version = "2.11.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d5f0420326b13675321b194928bb7830043b68cf8b810e1c651285c747abb080" -dependencies = [ - "hex", - "md-5", - "proc-macro2", - "quote", - "syn 2.0.106", -] - [[package]] name = "fnv" version = "1.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" -[[package]] -name = "form_urlencoded" -version = "1.2.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" -dependencies = [ - "percent-encoding", -] - -[[package]] -name = "foundation-api" -version = "2.0.0" -dependencies = [ - "bc-components", - "bc-envelope", - "bc-xid", - "chrono", - "dcbor", - "flutter_rust_bridge", - "gstp", - "insta", - "quantum-link-macros", - "rkyv", - "thiserror", -] - -[[package]] -name = "futures" -version = "0.3.31" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "65bc07b1a8bc7c85c5f2e110c476c7389b4554ba72af57d8445ea63a576b0876" -dependencies = [ - "futures-channel", - "futures-core", - "futures-executor", - "futures-io", - "futures-sink", - "futures-task", - "futures-util", -] - -[[package]] -name = "futures-channel" -version = "0.3.31" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2dff15bf788c671c1934e366d07e30c1814a8ef514e1af724a602e8a2fbe1b10" -dependencies = [ - "futures-core", - "futures-sink", -] - [[package]] name = "futures-core" version = "0.3.31" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "05f29059c0c2090612e8d742178b0580d2dc940c837851ad723096f87af6663e" -[[package]] -name = "futures-executor" -version = "0.3.31" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1e28d1d997f585e54aebc3f97d39e72338912123a67330d723fdbb564d646c9f" -dependencies = [ - "futures-core", - "futures-task", - "futures-util", -] - [[package]] name = "futures-io" version = "0.3.31" @@ -1137,47 +474,6 @@ dependencies = [ "pin-project-lite", ] -[[package]] -name = "futures-macro" -version = "0.3.31" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "162ee34ebcb7c64a8abebc059ce0fee27c2262618d7b60ed8faf72fef13c3650" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - -[[package]] -name = "futures-sink" -version = "0.3.31" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e575fab7d1e0dcb8d0c7bcf9a63ee213816ab51902e6d244a95819acacf1d4f7" - -[[package]] -name = "futures-task" -version = "0.3.31" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f90f7dce0722e95104fcb095585910c0977252f286e354b5e3bd38902cd99988" - -[[package]] -name = "futures-util" -version = "0.3.31" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9fa08315bb612088cc391249efdc3bc77536f16c91f6cf495e6fbe85b20a4a81" -dependencies = [ - "futures-channel", - "futures-core", - "futures-io", - "futures-macro", - "futures-sink", - "futures-task", - "memchr", - "pin-project-lite", - "pin-utils", - "slab", -] - [[package]] name = "generator" version = "0.8.8" @@ -1201,7 +497,6 @@ checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" dependencies = [ "typenum", "version_check", - "zeroize", ] [[package]] @@ -1233,37 +528,6 @@ version = "0.31.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "07e28edb80900c19c28f1072f2e8aeca7fa06b23cd4169cefe1af5aa3260783f" -[[package]] -name = "glob" -version = "0.3.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280" - -[[package]] -name = "group" -version = "0.13.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f0f9ef7462f7c099f518d754361858f86d8a07af53ba9af0fe635bbccb151a63" -dependencies = [ - "ff", - "rand_core 0.6.4", - "subtle", -] - -[[package]] -name = "gstp" -version = "0.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9dd8214e6a70abd783f45565cba634b58e8afca35dd374a155517a7f4c437774" -dependencies = [ - "bc-components", - "bc-envelope", - "bc-rand", - "bc-xid", - "dcbor", - "thiserror", -] - [[package]] name = "half" version = "2.6.0" @@ -1274,12 +538,6 @@ dependencies = [ "crunchy", ] -[[package]] -name = "hashbrown" -version = "0.14.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e5274423e17b7c9fc20b6e7e208532f9b19825d82dfd615708b70edd83df41f1" - [[package]] name = "hashbrown" version = "0.15.5" @@ -1313,7 +571,7 @@ dependencies = [ "proc-macro-error2", "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -1329,47 +587,11 @@ dependencies = [ "uuid", ] -[[package]] -name = "hermit-abi" -version = "0.5.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" - [[package]] name = "hex" version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" -dependencies = [ - "serde", -] - -[[package]] -name = "hex-conservative" -version = "0.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5313b072ce3c597065a808dbf612c4c8e8590bdbf8b579508bf7a762c5eae6cd" -dependencies = [ - "arrayvec", -] - -[[package]] -name = "hkdf" -version = "0.12.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7b5f8eb2ad728638ea2c7d47a21db23b7b58a72ed6a38256b8a1849f15fbbdf7" -dependencies = [ - "hmac", -] - -[[package]] -name = "hmac" -version = "0.12.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" -dependencies = [ - "digest", -] [[package]] name = "iana-time-zone" @@ -1395,113 +617,6 @@ dependencies = [ "cc", ] -[[package]] -name = "icu_collections" -version = "2.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "200072f5d0e3614556f94a9930d5dc3e0662a652823904c3a75dc3b0af7fee47" -dependencies = [ - "displaydoc", - "potential_utf", - "yoke", - "zerofrom", - "zerovec", -] - -[[package]] -name = "icu_locale_core" -version = "2.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0cde2700ccaed3872079a65fb1a78f6c0a36c91570f28755dda67bc8f7d9f00a" -dependencies = [ - "displaydoc", - "litemap", - "tinystr", - "writeable", - "zerovec", -] - -[[package]] -name = "icu_normalizer" -version = "2.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "436880e8e18df4d7bbc06d58432329d6458cc84531f7ac5f024e93deadb37979" -dependencies = [ - "displaydoc", - "icu_collections", - "icu_normalizer_data", - "icu_properties", - "icu_provider", - "smallvec", - "zerovec", -] - -[[package]] -name = "icu_normalizer_data" -version = "2.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "00210d6893afc98edb752b664b8890f0ef174c8adbb8d0be9710fa66fbbf72d3" - -[[package]] -name = "icu_properties" -version = "2.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "016c619c1eeb94efb86809b015c58f479963de65bdb6253345c1a1276f22e32b" -dependencies = [ - "displaydoc", - "icu_collections", - "icu_locale_core", - "icu_properties_data", - "icu_provider", - "potential_utf", - "zerotrie", - "zerovec", -] - -[[package]] -name = "icu_properties_data" -version = "2.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "298459143998310acd25ffe6810ed544932242d3f07083eee1084d83a71bd632" - -[[package]] -name = "icu_provider" -version = "2.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "03c80da27b5f4187909049ee2d72f276f0d9f99a42c306bd0131ecfe04d8e5af" -dependencies = [ - "displaydoc", - "icu_locale_core", - "stable_deref_trait", - "tinystr", - "writeable", - "yoke", - "zerofrom", - "zerotrie", - "zerovec", -] - -[[package]] -name = "idna" -version = "1.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" -dependencies = [ - "idna_adapter", - "smallvec", - "utf8_iter", -] - -[[package]] -name = "idna_adapter" -version = "1.2.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3acae9609540aa318d1bc588455225fb2085b9ed0c4f6bd0d9d5bcd86f1a0344" -dependencies = [ - "icu_normalizer", - "icu_properties", -] - [[package]] name = "indexmap" version = "2.12.1" @@ -1512,15 +627,6 @@ dependencies = [ "hashbrown 0.16.1", ] -[[package]] -name = "inout" -version = "0.1.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "879f10e63c20629ecabbb64a8010319738c66a5cd0c29b02d63d272b03751d01" -dependencies = [ - "generic-array", -] - [[package]] name = "insta" version = "1.44.1" @@ -1549,15 +655,6 @@ version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" -[[package]] -name = "itertools" -version = "0.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b1c173a5686ce8bfa551b3563d0c2170bf24ca44da99c7ca4bfdab5418c3fe57" -dependencies = [ - "either", -] - [[package]] name = "itoa" version = "1.0.15" @@ -1583,19 +680,9 @@ version = "0.2.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2a8c8b344124222efd714b73bb41f8b5120b27a7cc1c75593a6ff768d9d05aa4" dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - -[[package]] -name = "jobserver" -version = "0.1.33" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "38f262f097c174adebe41eb73d66ae9c06b2844fb0da69969647bbddd9b0538a" -dependencies = [ - "getrandom 0.3.3", - "libc", + "proc-macro2", + "quote", + "syn", ] [[package]] @@ -1608,25 +695,11 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "known-values" -version = "0.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "efadaa833480ac053954ea1bf019eee3b3ed0123417434158d1d9ca46162dfed" -dependencies = [ - "bc-components", - "dcbor", - "paste", -] - [[package]] name = "lazy_static" version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" -dependencies = [ - "spin", -] [[package]] name = "libc" @@ -1668,7 +741,7 @@ dependencies = [ "libcrux-secrets", "libcrux-sha3", "libcrux-traits", - "rand 0.9.2", + "rand", "tls_codec", ] @@ -1709,37 +782,15 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "812e4fa89f3f5e34b47f928b22b1b78395a0d4ec23b1f583db635f128159d65f" dependencies = [ "libcrux-secrets", - "rand 0.9.2", + "rand", ] -[[package]] -name = "libm" -version = "0.2.15" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f9fbbcab51052fe104eb5e5d351cf728d30a5be1fe14d9be8a3b097481fb97de" - [[package]] name = "linux-raw-sys" version = "0.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "df1d3c3b53da64cf5760482273a98e575c651a67eec7f77df96b5b642de8f039" -[[package]] -name = "litemap" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "241eaef5fd12c88705a01fc1066c48c4b36e0dd4377dcdc7ec3942cea7a69956" - -[[package]] -name = "lock_api" -version = "0.4.13" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "96936507f153605bddfcda068dd804796c84324ed2510809e5b2a624c81da765" -dependencies = [ - "autocfg", - "scopeguard", -] - [[package]] name = "log" version = "0.4.29" @@ -1768,51 +819,12 @@ dependencies = [ "regex-automata", ] -[[package]] -name = "md-5" -version = "0.10.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" -dependencies = [ - "cfg-if", - "digest", -] - [[package]] name = "memchr" version = "2.7.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32a282da65faaf38286cf3be983213fcf1d2e2a58700e808f83f4ea9a4804bc0" -[[package]] -name = "minicbor" -version = "0.19.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d7005aaf257a59ff4de471a9d5538ec868a21586534fff7f85dd97d4043a6139" -dependencies = [ - "minicbor-derive", -] - -[[package]] -name = "minicbor-derive" -version = "0.13.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1154809406efdb7982841adb6311b3d095b46f78342dd646736122fe6b19e267" -dependencies = [ - "proc-macro2", - "quote", - "syn 1.0.109", -] - -[[package]] -name = "miniz_oxide" -version = "0.7.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b8a240ddb74feaf34a79a7add65a741f3167852fba007066dcac1ca548d89c08" -dependencies = [ - "adler", -] - [[package]] name = "miniz_oxide" version = "0.8.9" @@ -1850,7 +862,7 @@ checksum = "4568f25ccbd45ab5d5603dc34318c1ec56b117531781260002151b8530a9f931" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -1872,22 +884,6 @@ dependencies = [ "num-traits", ] -[[package]] -name = "num-bigint-dig" -version = "0.8.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e661dda6640fad38e827a6d4a310ff4763082116fe217f279885c97f511bb0b7" -dependencies = [ - "lazy_static", - "libm", - "num-integer", - "num-iter", - "num-traits", - "rand 0.8.5", - "smallvec", - "zeroize", -] - [[package]] name = "num-integer" version = "0.1.46" @@ -1897,17 +893,6 @@ dependencies = [ "num-traits", ] -[[package]] -name = "num-iter" -version = "0.1.45" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1429034a0490724d0075ebb2bc9e875d6503c3cf69e235a8941aa757d83ef5bf" -dependencies = [ - "autocfg", - "num-integer", - "num-traits", -] - [[package]] name = "num-traits" version = "0.2.19" @@ -1915,17 +900,6 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" dependencies = [ "autocfg", - "libm", -] - -[[package]] -name = "num_cpus" -version = "1.17.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" -dependencies = [ - "hermit-abi", - "libc", ] [[package]] @@ -1955,61 +929,6 @@ version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b4ce411919553d3f9fa53a0880544cda985a112117a0444d5ff1e870a893d6ea" -[[package]] -name = "opaque-debug" -version = "0.3.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" - -[[package]] -name = "oslog" -version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "80d2043d1f61d77cb2f4b1f7b7b2295f40507f5f8e9d1c8bf10a1ca5f97a3969" -dependencies = [ - "cc", - "dashmap", - "log", -] - -[[package]] -name = "p256" -version = "0.13.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c9863ad85fa8f4460f9c48cb909d38a0d689dba1f6f6988a5e3e0d31071bcd4b" -dependencies = [ - "ecdsa", - "elliptic-curve", - "primeorder", - "sha2", -] - -[[package]] -name = "p384" -version = "0.13.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fe42f1670a52a47d448f14b6a5c61dd78fce51856e68edaa38f7ae3a46b8d6b6" -dependencies = [ - "ecdsa", - "elliptic-curve", - "primeorder", - "sha2", -] - -[[package]] -name = "p521" -version = "0.13.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0fc9e2161f1f215afdfce23677034ae137bbd45016a880c2eb3ba8eb95f085b2" -dependencies = [ - "base16ct", - "ecdsa", - "elliptic-curve", - "primeorder", - "rand_core 0.6.4", - "sha2", -] - [[package]] name = "parking" version = "2.2.1" @@ -2019,30 +938,6 @@ dependencies = [ "loom", ] -[[package]] -name = "parking_lot_core" -version = "0.9.11" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bc838d2a56b5b1a6c25f55575dfc605fabb63bb2365f6c2353ef9159aa69e4a5" -dependencies = [ - "cfg-if", - "libc", - "redox_syscall", - "smallvec", - "windows-targets", -] - -[[package]] -name = "password-hash" -version = "0.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "346f04948ba92c43e8469c1ee6736c7563d71012b17d40745260fe106aac2166" -dependencies = [ - "base64ct", - "rand_core 0.6.4", - "subtle", -] - [[package]] name = "paste" version = "1.0.15" @@ -2055,117 +950,12 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b867cad97c0791bbd3aaa6472142568c6c9e8f71937e98379f584cfb0cf35bec" -[[package]] -name = "pbkdf2" -version = "0.12.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f8ed6a7761f76e3b9f92dfb0a60a6a6477c61024b775147ff0973a02653abaf2" -dependencies = [ - "digest", - "hmac", -] - -[[package]] -name = "pem-rfc7468" -version = "0.7.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "88b39c9bfcfc231068454382784bb460aae594343fb030d46e9f50a645418412" -dependencies = [ - "base64ct", -] - -[[package]] -name = "percent-encoding" -version = "2.3.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" - -[[package]] -name = "phf" -version = "0.11.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1fd6780a80ae0c52cc120a26a1a42c1ae51b247a253e4e06113d23d2c2edd078" -dependencies = [ - "phf_macros", - "phf_shared", -] - -[[package]] -name = "phf_generator" -version = "0.11.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c80231409c20246a13fddb31776fb942c38553c51e871f8cbd687a4cfb5843d" -dependencies = [ - "phf_shared", - "rand 0.8.5", -] - -[[package]] -name = "phf_macros" -version = "0.11.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f84ac04429c13a7ff43785d75ad27569f2951ce0ffd30a3321230db2fc727216" -dependencies = [ - "phf_generator", - "phf_shared", - "proc-macro2", - "quote", - "syn 2.0.106", -] - -[[package]] -name = "phf_shared" -version = "0.11.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "67eabc2ef2a60eb7faa00097bd1ffdb5bd28e62bf39990626a582201b7a754e5" -dependencies = [ - "siphasher", -] - [[package]] name = "pin-project-lite" version = "0.2.16" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3b3cff922bd51709b605d9ead9aa71031d81447142d828eb4a6eba76fe619f9b" -[[package]] -name = "pin-utils" -version = "0.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8b870d8c151b6f2fb93e84a13146138f05d02ed11c7e7c54f8826aaaf7c9f184" - -[[package]] -name = "pkcs1" -version = "0.7.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c8ffb9f10fa047879315e6625af03c164b16962a5368d724ed16323b68ace47f" -dependencies = [ - "der", - "pkcs8", - "spki", -] - -[[package]] -name = "pkcs8" -version = "0.10.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f950b2377845cebe5cf8b5165cb3cc1a5e0fa5cfa3e1f7f55707d8fd82e0a7b7" -dependencies = [ - "der", - "spki", -] - -[[package]] -name = "poly1305" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf" -dependencies = [ - "cpufeatures", - "opaque-debug", - "universal-hash", -] - [[package]] name = "portable-atomic" version = "1.11.1" @@ -2181,15 +971,6 @@ dependencies = [ "portable-atomic", ] -[[package]] -name = "potential_utf" -version = "0.1.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e5a7c30837279ca13e7c867e9e40053bc68740f988cb07f7ca6df43cc734b585" -dependencies = [ - "zerovec", -] - [[package]] name = "ppv-lite86" version = "0.2.21" @@ -2199,56 +980,6 @@ dependencies = [ "zerocopy", ] -[[package]] -name = "pqcrypto-internals" -version = "0.2.10" -source = "git+https://github.com/Foundation-Devices/pqcrypto?rev=ebadf71214f67cb970242fa1053b4acb65767737#ebadf71214f67cb970242fa1053b4acb65767737" -dependencies = [ - "cc", - "dunce", - "getrandom 0.2.16", - "libc", -] - -[[package]] -name = "pqcrypto-mldsa" -version = "0.1.1" -source = "git+https://github.com/Foundation-Devices/pqcrypto?rev=ebadf71214f67cb970242fa1053b4acb65767737#ebadf71214f67cb970242fa1053b4acb65767737" -dependencies = [ - "cc", - "glob", - "libc", - "paste", - "pqcrypto-internals", - "pqcrypto-traits", -] - -[[package]] -name = "pqcrypto-mlkem" -version = "0.1.0" -source = "git+https://github.com/Foundation-Devices/pqcrypto?rev=ebadf71214f67cb970242fa1053b4acb65767737#ebadf71214f67cb970242fa1053b4acb65767737" -dependencies = [ - "cc", - "glob", - "libc", - "pqcrypto-internals", - "pqcrypto-traits", -] - -[[package]] -name = "pqcrypto-traits" -version = "0.3.5" -source = "git+https://github.com/Foundation-Devices/pqcrypto?rev=ebadf71214f67cb970242fa1053b4acb65767737#ebadf71214f67cb970242fa1053b4acb65767737" - -[[package]] -name = "primeorder" -version = "0.13.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "353e1ca18966c16d9deb1c69278edbc5f194139612772bd9537af60ac231e1e6" -dependencies = [ - "elliptic-curve", -] - [[package]] name = "proc-macro-error-attr2" version = "2.0.0" @@ -2268,7 +999,7 @@ dependencies = [ "proc-macro-error-attr2", "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -2290,8 +1021,8 @@ dependencies = [ "bit-vec", "bitflags", "num-traits", - "rand 0.9.2", - "rand_chacha 0.9.0", + "rand", + "rand_chacha", "rand_xorshift", "regex-syntax", "rusty-fork", @@ -2299,30 +1030,6 @@ dependencies = [ "unarray", ] -[[package]] -name = "provenance-mark" -version = "0.16.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e8a2078ef3c515d873099557bdbcb4b1ad8e67af1d2b398b5a7aa224baca970d" -dependencies = [ - "base64", - "bc-envelope", - "bc-rand", - "bc-tags", - "bc-ur", - "chacha20", - "chrono", - "dcbor", - "hex", - "hkdf", - "rand_core 0.6.4", - "serde", - "serde_json", - "sha2", - "thiserror", - "url", -] - [[package]] name = "ptr_meta" version = "0.3.1" @@ -2340,7 +1047,7 @@ checksum = "7347867d0a7e1208d93b46767be83e2b8f978c3dad35f775ac8d8847551d6fe1" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -2391,15 +1098,6 @@ dependencies = [ "sha2", ] -[[package]] -name = "quantum-link-macros" -version = "0.1.0" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - [[package]] name = "quick-error" version = "1.2.3" @@ -2430,54 +1128,24 @@ dependencies = [ "ptr_meta", ] -[[package]] -name = "rand" -version = "0.8.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "34af8d1a0e25924bc5b7c43c079c942339d8f0a8b57c39049bef581b46327404" -dependencies = [ - "libc", - "rand_chacha 0.3.1", - "rand_core 0.6.4", -] - [[package]] name = "rand" version = "0.9.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6db2770f06117d490610c7488547d543617b21bfa07796d7a12f6f1bd53850d1" dependencies = [ - "rand_chacha 0.9.0", - "rand_core 0.9.3", -] - -[[package]] -name = "rand_chacha" -version = "0.3.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e6c10a63a0fa32252be49d21e7709d4d4baf8d231c2dbce1eaa8141b9b127d88" -dependencies = [ - "ppv-lite86", - "rand_core 0.6.4", -] - -[[package]] -name = "rand_chacha" -version = "0.9.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" -dependencies = [ - "ppv-lite86", - "rand_core 0.9.3", + "rand_chacha", + "rand_core", ] [[package]] -name = "rand_core" -version = "0.6.4" +name = "rand_chacha" +version = "0.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" dependencies = [ - "getrandom 0.2.16", + "ppv-lite86", + "rand_core", ] [[package]] @@ -2495,25 +1163,7 @@ version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "513962919efc330f829edb2535844d1b912b0fbe2ca165d613e4e8788bb05a5a" dependencies = [ - "rand_core 0.9.3", -] - -[[package]] -name = "rand_xoshiro" -version = "0.6.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6f97cdb2a36ed4183de61b2f824cc45c9f1037f28afe0a322e9fff4c108b5aaa" -dependencies = [ - "rand_core 0.6.4", -] - -[[package]] -name = "redox_syscall" -version = "0.5.17" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5407465600fb0548f1442edf71dd20683c6ed326200ace4b1ef0763521bb3b77" -dependencies = [ - "bitflags", + "rand_core", ] [[package]] @@ -2554,16 +1204,6 @@ dependencies = [ "bytecheck", ] -[[package]] -name = "rfc6979" -version = "0.4.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f8dd2a808d456c4a54e300a23e9f5a67e122c3024119acbfd73e3bf664491cb2" -dependencies = [ - "hmac", - "subtle", -] - [[package]] name = "rkyv" version = "0.8.12" @@ -2591,28 +1231,7 @@ checksum = "bd83f5f173ff41e00337d97f6572e416d022ef8a19f371817259ae960324c482" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", -] - -[[package]] -name = "rsa" -version = "0.9.10" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b8573f03f5883dcaebdfcf4725caa1ecb9c15b2ef50c43a07b816e06799bb12d" -dependencies = [ - "const-oid", - "digest", - "num-bigint-dig", - "num-integer", - "num-traits", - "pkcs1", - "pkcs8", - "rand_core 0.6.4", - "sha2", - "signature", - "spki", - "subtle", - "zeroize", + "syn", ] [[package]] @@ -2621,15 +1240,6 @@ version = "0.1.26" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "56f7d92ca342cea22a06f2121d944b4fd82af56988c270852495420f961d4ace" -[[package]] -name = "rustc_version" -version = "0.4.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92" -dependencies = [ - "semver", -] - [[package]] name = "rustix" version = "1.1.2" @@ -2667,78 +1277,12 @@ version = "1.0.20" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "28d3b2b1366ec20994f1fd18c3c594f05c5dd4bc44d8bb0c1c632c8d6829481f" -[[package]] -name = "salsa20" -version = "0.10.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "97a22f5af31f73a954c10289c93e8a50cc23d971e80ee446f1f6f7137a088213" -dependencies = [ - "cipher", -] - [[package]] name = "scoped-tls" version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e1cf6437eb19a8f4a6cc0f7dca544973b0b78843adbfeb3683d1a94a0024a294" -[[package]] -name = "scopeguard" -version = "1.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" - -[[package]] -name = "scrypt" -version = "0.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0516a385866c09368f0b5bcd1caff3366aace790fcd46e2bb032697bb172fd1f" -dependencies = [ - "pbkdf2", - "salsa20", - "sha2", -] - -[[package]] -name = "sec1" -version = "0.7.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d3e97a565f76233a6003f9f5c54be1d9c5bdfa3eccfb189469f11ec4901c47dc" -dependencies = [ - "base16ct", - "der", - "generic-array", - "pkcs8", - "subtle", - "zeroize", -] - -[[package]] -name = "secp256k1" -version = "0.30.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b50c5943d326858130af85e049f2661ba3c78b26589b8ab98e65e80ae44a1252" -dependencies = [ - "bitcoin_hashes 0.14.0", - "rand 0.8.5", - "secp256k1-sys", -] - -[[package]] -name = "secp256k1-sys" -version = "0.10.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d4387882333d3aa8cb20530a17c69a3752e97837832f34f6dccc760e715001d9" -dependencies = [ - "cc", -] - -[[package]] -name = "semver" -version = "1.0.26" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "56e6fa9c48d24d85fb3de5ad847117517440f6beceb7798af16b4a87d616b8d0" - [[package]] name = "serde" version = "1.0.228" @@ -2766,7 +1310,7 @@ checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -2781,17 +1325,6 @@ dependencies = [ "serde", ] -[[package]] -name = "sha1" -version = "0.10.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" -dependencies = [ - "cfg-if", - "cpufeatures", - "digest", -] - [[package]] name = "sha2" version = "0.10.9" @@ -2818,16 +1351,6 @@ version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0fda2ff0d084019ba4d7c6f371c95d8fd75ce3524c3cb8fb653a3023f6323e64" -[[package]] -name = "signature" -version = "2.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de" -dependencies = [ - "digest", - "rand_core 0.6.4", -] - [[package]] name = "simdutf8" version = "0.1.5" @@ -2840,12 +1363,6 @@ version = "2.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbbb5d9659141646ae647b42fe094daf6c6192d1620870b449d9557f748b2daa" -[[package]] -name = "siphasher" -version = "1.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "56199f7ddabf13fe5074ce809e7d3f42b42ae711800501b5b16ea82ad029c39d" - [[package]] name = "slab" version = "0.4.11" @@ -2858,101 +1375,6 @@ version = "1.15.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "67b1b7a3b5fe4f1376887184045fcf45c69e92af734b7aaddc05fb777b6fbd03" -[[package]] -name = "spin" -version = "0.9.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6980e8d7511241f8acf4aebddbb1ff938df5eebe98691418c4468d0b72a96a67" - -[[package]] -name = "spki" -version = "0.7.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d91ed6c858b01f942cd56b37a94b3e0a1798290327d1236e4d9cf4eaca44d29d" -dependencies = [ - "base64ct", - "der", -] - -[[package]] -name = "ssh-cipher" -version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "caac132742f0d33c3af65bfcde7f6aa8f62f0e991d80db99149eb9d44708784f" -dependencies = [ - "cipher", - "ssh-encoding", -] - -[[package]] -name = "ssh-encoding" -version = "0.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "eb9242b9ef4108a78e8cd1a2c98e193ef372437f8c22be363075233321dd4a15" -dependencies = [ - "base64ct", - "pem-rfc7468", - "sha2", -] - -[[package]] -name = "ssh-key" -version = "0.6.7" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3b86f5297f0f04d08cabaa0f6bff7cb6aec4d9c3b49d87990d63da9d9156a8c3" -dependencies = [ - "dsa", - "ed25519-dalek", - "num-bigint-dig", - "p256", - "p384", - "p521", - "rand_core 0.6.4", - "rsa", - "sec1", - "sha1", - "sha2", - "signature", - "ssh-cipher", - "ssh-encoding", - "subtle", - "zeroize", -] - -[[package]] -name = "sskr" -version = "0.11.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7228e0234fae61785706c7f2b2bc5e47b6b34397e8c1052e1cfcba8030536234" -dependencies = [ - "bc-rand", - "bc-shamir", - "thiserror", -] - -[[package]] -name = "stable_deref_trait" -version = "1.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a8f112729512f8e442d81f95a8a7ddf2b7c6b8a1a6f509a95864142b30cab2d3" - -[[package]] -name = "subtle" -version = "2.6.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" - -[[package]] -name = "syn" -version = "1.0.109" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "72b64191b275b66ffe2469e8af2c1cfe3bafa67b529ead792a6d0160888b4237" -dependencies = [ - "proc-macro2", - "quote", - "unicode-ident", -] - [[package]] name = "syn" version = "2.0.106" @@ -2964,17 +1386,6 @@ dependencies = [ "unicode-ident", ] -[[package]] -name = "synstructure" -version = "0.13.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - [[package]] name = "tempfile" version = "3.23.0" @@ -3005,7 +1416,7 @@ checksum = "3ff15c8ecd7de3849db632e14d18d2571fa09dfc5ed93479bc4485c7a517c913" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -3017,25 +1428,6 @@ dependencies = [ "cfg-if", ] -[[package]] -name = "threadpool" -version = "1.8.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d050e60b33d41c19108b32cea32164033a9013fe3b46cbd4457559bfbf77afaa" -dependencies = [ - "num_cpus", -] - -[[package]] -name = "tinystr" -version = "0.8.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5d4f6d1145dcb577acf783d4e601bc1d76a13337bb54e6233add580b07344c8b" -dependencies = [ - "displaydoc", - "zerovec", -] - [[package]] name = "tinyvec" version = "1.10.0" @@ -3069,7 +1461,7 @@ checksum = "2d2e76690929402faae40aebdda620a2c0e25dd6d3b9afe48867dfd95991f4bd" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -3095,7 +1487,7 @@ checksum = "6e06d43f1345a3bcd39f6a56dbb7dcab2ba47e68e8ac134855e7e2bdbaf8cab8" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -3155,7 +1547,7 @@ checksum = "70977707304198400eb4835a78f6a9f928bf41bba420deb8fdb175cd965d77a7" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -3185,46 +1577,6 @@ dependencies = [ "tinyvec", ] -[[package]] -name = "universal-hash" -version = "0.5.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fc1de2c688dc15305988b563c3854064043356019f97a4b46276fe734c4f07ea" -dependencies = [ - "crypto-common", - "subtle", -] - -[[package]] -name = "ur" -version = "0.4.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "010f24a953db5d22d0010969ca3bbf40b3857b89f47c0f7be0da4c2d7ded0760" -dependencies = [ - "bitcoin_hashes 0.12.0", - "crc", - "minicbor", - "phf", - "rand_xoshiro", -] - -[[package]] -name = "url" -version = "2.5.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "137a3c834eaf7139b73688502f3f1141a0337c5d8e4d9b536f9b8c796e26a7c4" -dependencies = [ - "form_urlencoded", - "idna", - "percent-encoding", -] - -[[package]] -name = "utf8_iter" -version = "1.0.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" - [[package]] name = "utf8parse" version = "0.2.2" @@ -3300,23 +1652,10 @@ dependencies = [ "log", "proc-macro2", "quote", - "syn 2.0.106", + "syn", "wasm-bindgen-shared", ] -[[package]] -name = "wasm-bindgen-futures" -version = "0.4.50" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "555d470ec0bc3bb57890405e5d4322cc9ea83cebb085523ced7be4144dac1e61" -dependencies = [ - "cfg-if", - "js-sys", - "once_cell", - "wasm-bindgen", - "web-sys", -] - [[package]] name = "wasm-bindgen-macro" version = "0.2.100" @@ -3335,7 +1674,7 @@ checksum = "8ae87ea40c9f689fc23f209965b6fb8a99ad69aeeb0231408be24920604395de" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", "wasm-bindgen-backend", "wasm-bindgen-shared", ] @@ -3349,16 +1688,6 @@ dependencies = [ "unicode-ident", ] -[[package]] -name = "web-sys" -version = "0.3.77" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "33b6dd2ef9186f1f2072e409e99cd22a975331a6b3591b12c764e0e55c60d5d2" -dependencies = [ - "js-sys", - "wasm-bindgen", -] - [[package]] name = "windows-core" version = "0.61.2" @@ -3380,7 +1709,7 @@ checksum = "a47fddd13af08290e67f4acabf4b459f647552718f683a7b415d290ac744a836" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -3391,7 +1720,7 @@ checksum = "bd9211b69f8dcdfa817bfd14bf1c97c9188afa36f4750130fcdf3f400eca9fa8" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -3515,48 +1844,6 @@ dependencies = [ "bitflags", ] -[[package]] -name = "writeable" -version = "0.6.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ea2f10b9bb0928dfb1b42b65e1f9e36f7f54dbdf08457afefb38afcdec4fa2bb" - -[[package]] -name = "x25519-dalek" -version = "2.0.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c7e468321c81fb07fa7f4c636c3972b9100f0346e5b6a9f2bd0603a52f7ed277" -dependencies = [ - "curve25519-dalek", - "rand_core 0.6.4", - "serde", - "zeroize", -] - -[[package]] -name = "yoke" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5f41bb01b8226ef4bfd589436a297c53d118f65921786300e427be8d487695cc" -dependencies = [ - "serde", - "stable_deref_trait", - "yoke-derive", - "zerofrom", -] - -[[package]] -name = "yoke-derive" -version = "0.8.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "38da3c9736e16c5d3c8c597a9aaa5d1fa565d0532ae05e27c24aa62fb32c0ab6" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", - "synstructure", -] - [[package]] name = "zerocopy" version = "0.8.26" @@ -3574,28 +1861,7 @@ checksum = "9ecf5b4cc5364572d7f4c329661bcc82724222973f2cab6f050a4e5c22f75181" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", -] - -[[package]] -name = "zerofrom" -version = "0.1.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "50cc42e0333e05660c3587f3bf9d0478688e15d870fab3346451ce7f8c9fbea5" -dependencies = [ - "zerofrom-derive", -] - -[[package]] -name = "zerofrom-derive" -version = "0.1.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d71e5d6e06ab090c67b5e44993ec16b72dcbaabc526db883a360057678b48502" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", - "synstructure", + "syn", ] [[package]] @@ -3615,38 +1881,5 @@ checksum = "ce36e65b0d2999d2aafac989fb249189a141aee1f53c612c1f37d72631959f69" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", -] - -[[package]] -name = "zerotrie" -version = "0.2.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "36f0bbd478583f79edad978b407914f61b2972f5af6fa089686016be8f9af595" -dependencies = [ - "displaydoc", - "yoke", - "zerofrom", -] - -[[package]] -name = "zerovec" -version = "0.11.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e7aa2bd55086f1ab526693ecbe444205da57e25f4489879da80635a46d90e73b" -dependencies = [ - "yoke", - "zerofrom", - "zerovec-derive", -] - -[[package]] -name = "zerovec-derive" -version = "0.11.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5b96237efa0c878c64bd89c436f661be4e46b2f3eff1ebb976f7ef2321d2f58f" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", + "syn", ] diff --git a/Cargo.toml b/Cargo.toml index b2492c48..83ac135d 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,44 +1,27 @@ [workspace] resolver = "2" members = [ - "api", "backup-shard", "btp", "ql-fsm", "ql-rpc", "ql-runtime", "ql-wire", - "quantum-link-macros", ] [workspace.package] homepage = "https://github.com/Foundation-Devices/foundation-api" [workspace.dependencies] -# blockchain commons -bc-components = { version = "0.28.0" } -bc-envelope = { version = "0.37.0" } -bc-xid = { version = "0.16.0" } dcbor = { version = "0.23.3" } -gstp = { version = "0.11.0" } - -chrono = "0.4" bytes = "1" getrandom = { version = "0.2" } insta = { version = "1.43.2" } -thiserror = { version = "2" } rkyv = { version = "0.8" } # workspace crates backup-shard = { path = "backup-shard" } btp = { path = "btp" } -foundation-api = { path = "api" } -quantum-link-macros = { path = "quantum-link-macros" } ql-fsm = { path = "ql-fsm" } ql-rpc = { path = "ql-rpc" } ql-wire = { path = "ql-wire" } - -[patch.crates-io] -pqcrypto-traits = { git = "https://github.com/Foundation-Devices/pqcrypto", rev = "ebadf71214f67cb970242fa1053b4acb65767737" } -pqcrypto-mldsa = { git = "https://github.com/Foundation-Devices/pqcrypto", rev = "ebadf71214f67cb970242fa1053b4acb65767737" } -pqcrypto-mlkem = { git = "https://github.com/Foundation-Devices/pqcrypto", rev = "ebadf71214f67cb970242fa1053b4acb65767737" } diff --git a/Justfile b/Justfile index 45492c2b..71889c0a 100644 --- a/Justfile +++ b/Justfile @@ -1,15 +1,3 @@ # Run clippy on all targets and features, treating warnings as errors clippy: cargo clippy --all-targets --all-features -- -D warnings - -# Run golden/snapshot tests -golden: - cargo test -p foundation-api --test golden_tests - -# Update golden/snapshot tests (accept all new snapshots) -golden-update: - INSTA_UPDATE=always cargo test -p foundation-api --test golden_tests - -# Review pending golden/snapshot changes interactively (requires cargo-insta) -golden-review: - cargo insta review \ No newline at end of file diff --git a/README.md b/README.md index 2e786904..0d280b22 100644 --- a/README.md +++ b/README.md @@ -1,14 +1,15 @@ # Foundation API -This monorepo contains the core crates for a device-to-device API using Blockchain Commons' GSTP +This monorepo contains the core crates for Foundation device-to-device protocols. ## Crates -- **abstracted**: Abstractions of the BLE and SE chips -- **api**: The API - contains predefined QL messages -- **api-demo**: Tokio-based demo of device-to-device communication - **btp**: Beefcake Transfer Protocol for splitting messages into MTU sized chunks -- **quantum-link-macros**: Macros to easily turn Rust Structs and Enums into valid QL messages +- **backup-shard**: Magic backup shard encoding +- **ql-wire**: QuantumLink wire-format definitions +- **ql-fsm**: QuantumLink Sans-IO protocol finite state machine +- **ql-runtime**: QuantumLink async runtime +- **ql-rpc**: RPC modality layer over QuantumLink streams ## Development diff --git a/api/.gitignore b/api/.gitignore deleted file mode 100644 index 96ef6c0b..00000000 --- a/api/.gitignore +++ /dev/null @@ -1,2 +0,0 @@ -/target -Cargo.lock diff --git a/api/Cargo.toml b/api/Cargo.toml deleted file mode 100644 index 13e7ec3b..00000000 --- a/api/Cargo.toml +++ /dev/null @@ -1,31 +0,0 @@ -[package] -name = "foundation-api" -version = "2.0.0" -edition = "2021" -description = "Foundation API using Gordian Sealed Transaction Protocol (GSTP)." -authors = ["Wolf McNally, Blockchain Commons, Foundation Devices"] -repository = "https://github.com/Foundation-Devices/foundation-api" -readme = "README.md" -license = "Proprietary" - -[dependencies] -bc-envelope = { workspace = true } -bc-xid = { workspace = true } -rkyv = { workspace = true, optional = true } -flutter_rust_bridge = { version = "=2.11.1", optional = true } -quantum-link-macros = { workspace = true } -gstp = { workspace = true } -bc-components = { workspace = true } -dcbor = { workspace = true } -chrono = { workspace = true } -thiserror = { workspace = true } - -[dev-dependencies] -insta = { workspace = true } - -[features] -keyos = ["rkyv"] -envoy = ["flutter_rust_bridge"] - -[lints.rust] -unexpected_cfgs = { level = "warn", check-cfg = ['cfg(frb_expand)'] } diff --git a/api/src/api/backup.rs b/api/src/api/backup.rs deleted file mode 100644 index ec27c35f..00000000 --- a/api/src/api/backup.rs +++ /dev/null @@ -1,264 +0,0 @@ -use quantum_link_macros::quantum_link; - -#[quantum_link] -#[repr(transparent)] -pub struct Shard(pub Vec); - -#[quantum_link] -#[repr(transparent)] -pub struct SeedFingerprint(pub [u8; 32]); - -#[quantum_link] -pub struct BackupShardRequest { - #[n(0)] - pub shard: Shard, -} - -#[quantum_link] -pub enum BackupShardResponse { - #[n(0)] - Success, - #[n(1)] - Error { - #[n(0)] - error: String, - }, -} - -#[quantum_link] -pub struct RestoreShardRequest { - #[n(0)] - pub seed_fingerprint: SeedFingerprint, - #[n(1)] - pub timestamp: Option, -} - -#[quantum_link] -pub enum RestoreShardResponse { - #[n(0)] - Success { - #[n(0)] - shard: Shard, - }, - #[n(1)] - Error { - #[n(0)] - error: String, - }, - #[n(2)] - NotFound, -} - -#[quantum_link] -pub struct EnvoyMagicBackupEnabledRequest {} - -#[quantum_link] -pub struct EnvoyMagicBackupEnabledResponse { - #[n(0)] - pub enabled: bool, -} - -#[quantum_link] -pub struct PrimeMagicBackupEnabled { - #[n(0)] - pub enabled: bool, - #[n(1)] - pub seed_fingerprint: SeedFingerprint, -} - -#[quantum_link] -pub struct PrimeMagicBackupStatusRequest { - #[n(0)] - pub seed_fingerprint: SeedFingerprint, - #[n(1)] - pub timestamp: Option, -} - -#[quantum_link] -pub struct PrimeMagicBackupStatusResponse { - #[n(0)] - pub shard_backup_found: bool, -} - -// -// MAGIC BACKUPS -// - -#[quantum_link] -#[derive(Eq)] -pub struct BackupChunk { - #[n(0)] - pub chunk_index: u32, - #[n(1)] - pub total_chunks: u32, - #[n(2)] - pub data: Vec, -} - -impl BackupChunk { - pub fn is_last(&self) -> bool { - self.chunk_index == self.total_chunks - 1 - } -} - -// -// CREATING BACKUP -// - -// from prime -> envoy -#[quantum_link] -pub enum CreateMagicBackupEvent { - #[n(0)] - Start(StartMagicBackup), - #[n(1)] - Chunk(BackupChunk), -} - -#[quantum_link] -pub struct StartMagicBackup { - #[n(0)] - pub seed_fingerprint: SeedFingerprint, - #[n(1)] - pub total_chunks: u32, - #[n(2)] - pub hash: [u8; 32], -} - -// envoy -> prime -// error can be sent at any time -// success is expected at the end of the flow -#[quantum_link] -pub enum CreateMagicBackupResult { - #[n(0)] - Success, - #[n(1)] - Error { - #[n(0)] - error: String, - }, -} - -// -// RESTORING BACKUP -// - -#[quantum_link] -pub struct RestoreMagicBackupRequest { - #[n(0)] - pub seed_fingerprint: SeedFingerprint, - /// if 0, then go from start - #[n(1)] - pub resume_from_chunk: u32, -} - -#[quantum_link] -pub enum RestoreMagicBackupEvent { - // there is no backup found from the provided fingerprint - #[n(0)] - NotFound, - // envoy found a backup and is beginning transmission - #[n(1)] - Starting(BackupMetadata), - // a backup chunk - #[n(2)] - Chunk(BackupChunk), - // envoy failed - #[n(3)] - Error { - #[n(0)] - error: String, - }, -} - -#[quantum_link] -#[derive(Eq)] -pub struct BackupMetadata { - #[n(0)] - pub total_chunks: u32, -} - -// sent from prime -> envoy -#[quantum_link] -pub enum RestoreMagicBackupResult { - #[n(0)] - Success, - #[n(1)] - Error { - #[n(0)] - error: String, - }, -} - -// -// MAGIC BACKUPS V2 -// - -#[quantum_link] -pub struct CreateMagicBackupV2 { - #[n(0)] - pub timestamp: u64, - /// Backup identifier (SHA-256 hash). - #[n(1)] - pub hash: Vec, - /// ML-DSA-44 public key. - #[n(2)] - pub pubkey: Vec, - /// Encrypted backup payload. - #[n(3)] - pub data: Vec, - /// ML-DSA-44 client signature. - #[n(4)] - pub client_signature: Vec, -} - -#[quantum_link] -pub struct GetMagicBackupV2 { - #[n(0)] - pub key: Vec, - #[n(1)] - pub timestamp: u64, - /// ML-DSA-44 signature. - #[n(2)] - pub signature: Vec, -} - -#[quantum_link] -pub struct DeleteMagicBackupV2 { - #[n(0)] - pub key: Vec, - #[n(1)] - pub timestamp: u64, - /// ML-DSA-44 signature. - #[n(2)] - pub signature: Vec, -} - -// prime -> envoy -#[quantum_link] -pub enum MagicBackupRequestV2 { - #[n(0)] - Create(CreateMagicBackupV2), - #[n(1)] - Get(GetMagicBackupV2), - #[n(2)] - Delete(DeleteMagicBackupV2), -} - -// envoy -> prime -#[quantum_link] -pub enum MagicBackupResponseV2 { - #[n(0)] - Created, - #[n(1)] - Backup { - #[n(0)] - data: Vec, - }, - #[n(2)] - Deleted, - #[n(3)] - Error { - #[n(0)] - error: String, - }, -} diff --git a/api/src/api/bitcoin.rs b/api/src/api/bitcoin.rs deleted file mode 100644 index 821f280e..00000000 --- a/api/src/api/bitcoin.rs +++ /dev/null @@ -1,32 +0,0 @@ -use quantum_link_macros::quantum_link; - -#[quantum_link] -pub struct SignPsbt { - #[n(0)] - pub account_id: String, - #[n(1)] - pub psbt: Vec, -} - -#[quantum_link] -pub struct AccountUpdate { - #[n(0)] - pub account_id: String, - #[n(1)] - pub update: Vec, -} - -#[quantum_link] -pub struct BroadcastTransaction { - #[n(0)] - pub account_id: String, - #[n(1)] - pub psbt: Vec, -} - -// If None, there's no passphrase, hide passphrased accounts -#[quantum_link] -pub struct ApplyPassphrase { - #[n(0)] - pub fingerprint: Option, -} diff --git a/api/src/api/firmware.rs b/api/src/api/firmware.rs deleted file mode 100644 index e3ad9280..00000000 --- a/api/src/api/firmware.rs +++ /dev/null @@ -1,116 +0,0 @@ -use quantum_link_macros::quantum_link; - -// From Prime to Envoy -#[quantum_link] -pub struct FirmwareUpdateCheckRequest { - #[n(0)] - pub current_version: String, -} - -// From Envoy to Prime -#[quantum_link] -pub enum FirmwareUpdateCheckResponse { - #[n(0)] - Available(FirmwareUpdateAvailable), - #[n(1)] - NotAvailable, -} - -#[quantum_link] -pub struct FirmwareUpdateAvailable { - #[n(0)] - pub version: String, - #[n(1)] - pub changelog: String, - #[n(2)] - pub timestamp: u32, - #[n(3)] - pub total_size: u32, - #[n(4)] - pub patch_count: u8, -} - -// From Prime to Envoy -#[quantum_link] -pub struct FirmwareFetchRequest { - #[n(0)] - pub current_version: String, - #[n(1)] - pub chunk_offset: Option, -} - -// From Envoy to Prime -#[quantum_link] -pub enum FirmwareFetchEvent { - // there is no update available from the provided prime version - #[n(0)] - UpdateNotAvailable, - // envoy has found an update, and will begin transmission - #[n(1)] - Starting(FirmwareUpdateAvailable), - // envoy is downloading the update - #[n(2)] - Downloading, - // envoy is sending a chunk for an update patch - #[n(3)] - Chunk(FirmwareChunk), - // envoy failed - #[n(5)] - Error { - #[n(0)] - error: String, - }, -} - -#[quantum_link] -#[derive(Eq)] -pub struct FirmwareChunk { - #[n(0)] - pub patch_index: u8, - #[n(1)] - pub total_patches: u8, - #[n(2)] - pub chunk_index: u16, - #[n(3)] - pub total_chunks: u16, - #[n(4)] - pub data: Vec, -} - -impl FirmwareChunk { - pub fn is_last(&self) -> bool { - self.patch_index == self.total_patches - 1 && self.chunk_index == self.total_chunks - 1 - } -} - -#[quantum_link] -pub enum FirmwareInstallEvent { - #[n(0)] - UpdateVerified, - #[n(1)] - Installing, - #[n(2)] - Rebooting, - #[n(3)] - Success { - #[n(0)] - installed_version: String, - }, - #[n(4)] - Error { - #[n(0)] - error: String, - #[n(1)] - stage: InstallErrorStage, - }, -} - -#[quantum_link] -pub enum InstallErrorStage { - #[n(0)] - Download, - #[n(1)] - Verify, - #[n(2)] - Install, -} diff --git a/api/src/api/fx.rs b/api/src/api/fx.rs deleted file mode 100644 index 3f2eb816..00000000 --- a/api/src/api/fx.rs +++ /dev/null @@ -1,34 +0,0 @@ -use quantum_link_macros::quantum_link; - -#[quantum_link] -pub struct ExchangeRate { - #[n(0)] - pub currency_code: String, - #[n(1)] - pub rate: f32, - #[n(2)] - pub timestamp: u64, -} - -#[quantum_link] -pub struct ExchangeRateHistory { - #[n(0)] - pub history: Vec, - #[n(1)] - pub currency_code: String, -} - -#[quantum_link] -pub struct PricePoint { - #[n(0)] - pub rate: f32, - #[n(1)] - pub timestamp: u64, -} - -/// Prime → Envoy. ISO-4217 code; sent on settings change and on every reconnect. -#[quantum_link] -pub struct PrimeFiatPreference { - #[n(0)] - pub currency_code: String, -} diff --git a/api/src/api/message.rs b/api/src/api/message.rs deleted file mode 100644 index cfe8fd7a..00000000 --- a/api/src/api/message.rs +++ /dev/null @@ -1,149 +0,0 @@ -use quantum_link_macros::quantum_link; - -use super::onboarding::OnboardingState; -use crate::{ - backup::{ - BackupShardRequest, BackupShardResponse, CreateMagicBackupEvent, CreateMagicBackupResult, - EnvoyMagicBackupEnabledRequest, EnvoyMagicBackupEnabledResponse, MagicBackupRequestV2, - MagicBackupResponseV2, PrimeMagicBackupEnabled, PrimeMagicBackupStatusRequest, - PrimeMagicBackupStatusResponse, RestoreMagicBackupEvent, RestoreMagicBackupRequest, - RestoreMagicBackupResult, RestoreShardRequest, RestoreShardResponse, - }, - bitcoin::*, - firmware::{ - FirmwareFetchEvent, FirmwareFetchRequest, FirmwareInstallEvent, FirmwareUpdateCheckRequest, - FirmwareUpdateCheckResponse, - }, - fx::{ExchangeRate, ExchangeRateHistory, PrimeFiatPreference}, - pairing::{PairingRequest, PairingResponse, UnpairingRequest, UnpairingResponse}, - scv::SecurityCheck, - status::{ - DeviceNameUpdate, DeviceStatus, EnvoyStatus, Heartbeat, TimezoneRequest, TimezoneResponse, - }, -}; - -// Bump this every time there is a significant change -pub const PROTOCOL_VERSION: u8 = 1; - -#[quantum_link] -pub struct EnvoyMessage { - #[n(0)] - pub message: QuantumLinkMessage, - #[n(1)] - pub timestamp: u32, - #[n(2)] - pub protocol_version: Option, // This being None is implicit v0 -} - -#[quantum_link] -pub struct PassportMessage { - #[n(0)] - pub message: QuantumLinkMessage, - #[n(1)] - pub status: DeviceStatus, - #[n(2)] - pub protocol_version: Option, -} - -#[quantum_link] -pub enum QuantumLinkMessage { - #[n(0)] - ExchangeRate(ExchangeRate), - #[n(1)] - ExchangeRateHistory(ExchangeRateHistory), - - #[n(2)] - FirmwareUpdateCheckRequest(FirmwareUpdateCheckRequest), - #[n(3)] - FirmwareUpdateCheckResponse(FirmwareUpdateCheckResponse), - #[n(4)] - FirmwareFetchRequest(FirmwareFetchRequest), - #[n(5)] - FirmwareFetchEvent(FirmwareFetchEvent), - #[n(6)] - FirmwareInstallEvent(FirmwareInstallEvent), - - #[n(7)] - DeviceStatus(DeviceStatus), - #[n(8)] - EnvoyStatus(EnvoyStatus), - - #[n(9)] - PairingRequest(PairingRequest), - #[n(10)] - PairingResponse(PairingResponse), - - #[n(11)] - SecurityCheck(SecurityCheck), - #[n(12)] - OnboardingState(OnboardingState), - - #[n(13)] - SignPsbt(SignPsbt), - #[n(14)] - BroadcastTransaction(BroadcastTransaction), - #[n(15)] - AccountUpdate(AccountUpdate), - #[n(16)] - ApplyPassphrase(ApplyPassphrase), - - #[n(17)] - EnvoyMagicBackupEnabledRequest(EnvoyMagicBackupEnabledRequest), - #[n(18)] - EnvoyMagicBackupEnabledResponse(EnvoyMagicBackupEnabledResponse), - - #[n(19)] - PrimeMagicBackupEnabled(PrimeMagicBackupEnabled), - - #[n(20)] - PrimeMagicBackupStatusRequest(PrimeMagicBackupStatusRequest), - #[n(21)] - PrimeMagicBackupStatusResponse(PrimeMagicBackupStatusResponse), - - #[n(22)] - BackupShardRequest(BackupShardRequest), - #[n(23)] - BackupShardResponse(BackupShardResponse), - - #[n(24)] - RestoreShardRequest(RestoreShardRequest), - #[n(25)] - RestoreShardResponse(RestoreShardResponse), - - #[n(26)] - CreateMagicBackupEvent(CreateMagicBackupEvent), - #[n(27)] - CreateMagicBackupResult(CreateMagicBackupResult), - - #[n(28)] - RestoreMagicBackupRequest(RestoreMagicBackupRequest), - #[n(29)] - RestoreMagicBackupEvent(RestoreMagicBackupEvent), - #[n(30)] - RestoreMagicBackupResult(RestoreMagicBackupResult), - - #[n(31)] - Heartbeat(Heartbeat), - - #[n(33)] - TimezoneRequest(TimezoneRequest), - #[n(34)] - TimezoneResponse(TimezoneResponse), - - #[n(35)] - UnpairingRequest(UnpairingRequest), - #[n(36)] - UnpairingResponse(UnpairingResponse), - - #[n(37)] - DeviceNameUpdate(DeviceNameUpdate), - - #[n(38)] - MagicBackupRequestV2(MagicBackupRequestV2), - #[n(39)] - MagicBackupResponseV2(MagicBackupResponseV2), - - // Skipped tags (e.g. #[n(32)]) are intentional and must not be reused. - #[n(40)] - PrimeFiatPreference(PrimeFiatPreference), -} diff --git a/api/src/api/mod.rs b/api/src/api/mod.rs deleted file mode 100644 index a362499b..00000000 --- a/api/src/api/mod.rs +++ /dev/null @@ -1,13 +0,0 @@ -pub mod backup; -pub mod bitcoin; -pub mod firmware; -pub mod fx; -pub mod message; -pub mod onboarding; -pub mod pairing; -pub mod passport; -pub mod quantum_link; -pub mod scv; -pub mod status; -#[cfg(test)] -pub mod tests; diff --git a/api/src/api/onboarding.rs b/api/src/api/onboarding.rs deleted file mode 100644 index 8e85fb06..00000000 --- a/api/src/api/onboarding.rs +++ /dev/null @@ -1,47 +0,0 @@ -use quantum_link_macros::quantum_link; - -#[quantum_link] -pub enum OnboardingState { - #[n(0)] - SecurityChecked, - #[n(1)] - SecurityCheckFailed, - - #[n(2)] - FirmwareUpdateScreen, - - /// pin - #[n(3)] - SecuringDevice, - /// pin - #[n(4)] - DeviceSecured, - - #[n(5)] - WalletCreationScreen, - #[n(6)] - CreatingWallet, - #[n(7)] - WalletCreated, - - #[n(8)] - MagicBackupScreen, - #[n(9)] - CreatingMagicBackup, - #[n(10)] - MagicBackupCreated, - - #[n(11)] - CreatingManualBackup, - #[n(12)] - CreatingKeycardBackup, - - #[n(13)] - WritingDownSeedWords, - #[n(14)] - ConnectingWallet, - #[n(15)] - WalletConected, - #[n(16)] - Completed, -} diff --git a/api/src/api/pairing.rs b/api/src/api/pairing.rs deleted file mode 100644 index ccc94ea4..00000000 --- a/api/src/api/pairing.rs +++ /dev/null @@ -1,39 +0,0 @@ -use quantum_link_macros::quantum_link; - -use crate::{ - api::passport::{PassportFirmwareVersion, PassportModel, PassportSerial}, - passport::PassportColor, -}; - -#[quantum_link] -pub struct PairingResponse { - #[n(0)] - pub passport_model: PassportModel, - #[n(1)] - pub passport_firmware_version: PassportFirmwareVersion, - #[n(2)] - pub passport_serial: PassportSerial, - #[n(3)] - pub passport_color: PassportColor, - #[n(4)] - pub onboarding_complete: bool, - #[n(5)] - pub device_name: Option, -} - -#[quantum_link] -pub struct PairingRequest { - #[n(0)] - pub xid_document: Vec, - #[n(1)] - pub device_name: String, -} - -#[quantum_link] -pub struct UnpairingRequest {} - -#[quantum_link] -pub struct UnpairingResponse { - #[n(0)] - pub success: bool, -} diff --git a/api/src/api/passport.rs b/api/src/api/passport.rs deleted file mode 100644 index 84c2a788..00000000 --- a/api/src/api/passport.rs +++ /dev/null @@ -1,25 +0,0 @@ -use quantum_link_macros::quantum_link; - -#[quantum_link] -pub enum PassportModel { - #[n(0)] - Gen1, - #[n(1)] - Gen2, - #[n(2)] - Prime, -} - -#[quantum_link] -pub struct PassportFirmwareVersion(pub String); - -#[quantum_link] -pub struct PassportSerial(pub String); - -#[quantum_link] -pub enum PassportColor { - #[n(0)] - Light, - #[n(1)] - Dark, -} diff --git a/api/src/api/quantum_link.rs b/api/src/api/quantum_link.rs deleted file mode 100644 index 21b187fe..00000000 --- a/api/src/api/quantum_link.rs +++ /dev/null @@ -1,431 +0,0 @@ -use std::time::Duration; - -use bc_components::{EncapsulationScheme, PrivateKeys, PublicKeys, SignatureScheme, ARID}; -use bc_envelope::{ - prelude::{CBORCase, CBOR}, - Envelope, EventBehavior, Expression, ExpressionBehavior, Function, -}; -use bc_xid::XIDDocument; -use chrono::{DateTime, Utc}; -use dcbor::Date; -use gstp::{SealedEvent, SealedEventBehavior}; - -use crate::message::{EnvoyMessage, PassportMessage}; - -pub const QUANTUM_LINK: Function = Function::new_static_named("quantumLink"); -pub const EXPIRATION_DURATION: Duration = Duration::from_secs(60); - -#[derive(Debug, Copy, Clone, PartialEq, Eq)] -pub enum ReplayCheck { - Fresh, - Replay, - Expired, -} - -/// Storage for tracking received ARIDs to prevent replay attacks -#[derive(Debug, Default, Clone)] -pub struct ARIDCache { - cache: Vec<(ARID, DateTime)>, -} - -impl ARIDCache { - pub fn new() -> Self { - Self { cache: Vec::new() } - } - - /// Check if ARID has been seen before and store it until the event expires. - pub fn check_and_store( - &mut self, - arid: &ARID, - expires_at: DateTime, - now: DateTime, - ) -> ReplayCheck { - // Clean up expired entries first - self.cache.retain(|(_, expires_at)| now < *expires_at); - - if now >= expires_at { - return ReplayCheck::Expired; - } - - // Check if ARID already exists (replay attack) - if self.cache.iter().any(|(id, _)| id == arid) { - return ReplayCheck::Replay; - } - - self.cache.push((*arid, expires_at)); - ReplayCheck::Fresh - } - - /// Get the number of stored ARIDs - pub fn len(&self) -> usize { - self.cache.len() - } - - pub fn is_empty(&self) -> bool { - self.cache.len() == 0 - } - - /// Clear all stored ARIDs - pub fn clear(&mut self) { - self.cache.clear(); - } -} - -#[derive(Debug, thiserror::Error)] -pub enum QlError { - #[error(transparent)] - Cbor(#[from] dcbor::Error), - #[error(transparent)] - Envelope(#[from] bc_envelope::Error), - #[error(transparent)] - Gstp(#[from] gstp::Error), - - #[error("envelope did not contain leaf")] - NotLeaf, - #[error("missing date")] - MissingDate, - #[error("invalid function")] - InvalidFunction, - #[error("replay attack")] - ReplayAttack, - #[error("expired")] - Expired, - #[error("date too far in the future")] - FutureDated, -} - -pub trait QuantumLink: Into + TryFrom { - fn encode(self) -> Expression { - let cbor: CBOR = self.into(); - let envelope = Envelope::new(cbor); - Expression::new(QUANTUM_LINK).with_parameter("ql", envelope) - } - - fn decode(expression: &Expression) -> Result { - if expression.function() != &QUANTUM_LINK { - return Err(QlError::InvalidFunction); - } - let envelope = expression.object_for_parameter("ql")?; - let cbor = envelope.as_leaf().ok_or(QlError::NotLeaf)?; - - let message = Self::try_from(cbor)?; - Ok(message) - } - - fn seal( - self, - (sender_pk, sender_xid): (&PrivateKeys, &XIDDocument), - recipient: &XIDDocument, - ) -> Envelope { - let valid_until = Date::with_duration_from_now(EXPIRATION_DURATION); - - let event: SealedEvent = - SealedEvent::new(QuantumLink::encode(self), ARID::new(), sender_xid) - .with_date(&valid_until); - event - .to_envelope(Some(&valid_until), Some(sender_pk), Some(recipient)) - .unwrap() - } - - fn unseal( - envelope: &Envelope, - private_keys: &PrivateKeys, - ) -> Result<(Expression, XIDDocument), QlError> { - let now = Utc::now(); - let event: SealedEvent = - SealedEvent::try_from_envelope(envelope, None, Some(&Date::from(now)), private_keys)?; - let expires_at = event.date().ok_or(QlError::MissingDate)?.datetime(); - validate_expires_at(expires_at, now)?; - - let expression = event.content().clone(); - Ok((expression, event.sender().clone())) - } - - fn unseal_with_replay_check( - envelope: &Envelope, - private_keys: &PrivateKeys, - arid_cache: &mut ARIDCache, - ) -> Result<(Expression, XIDDocument), QlError> { - let now = Utc::now(); - let event: SealedEvent = - SealedEvent::try_from_envelope(envelope, None, Some(&Date::from(now)), private_keys)?; - - let arid = event.id(); - let expires_at = event.date().ok_or(QlError::MissingDate)?.datetime(); - validate_expires_at(expires_at, now)?; - - match arid_cache.check_and_store(&arid, expires_at, now) { - ReplayCheck::Fresh => {} - ReplayCheck::Replay => return Err(QlError::ReplayAttack), - ReplayCheck::Expired => return Err(QlError::Expired), - } - - let expression = event.content().clone(); - Ok((expression, event.sender().clone())) - } - - fn unseal_passport_message_with_replay_check( - envelope: &Envelope, - private_keys: &PrivateKeys, - arid_cache: &mut ARIDCache, - ) -> Result<(PassportMessage, XIDDocument), QlError> { - let (expression, sender) = - PassportMessage::unseal_with_replay_check(envelope, private_keys, arid_cache)?; - Ok((PassportMessage::decode(&expression)?, sender)) - } - - fn unseal_envoy_message_with_replay_check( - envelope: &Envelope, - private_keys: &PrivateKeys, - arid_cache: &mut ARIDCache, - ) -> Result<(EnvoyMessage, XIDDocument), QlError> { - let (expression, sender) = - EnvoyMessage::unseal_with_replay_check(envelope, private_keys, arid_cache)?; - Ok((EnvoyMessage::decode(&expression)?, sender)) - } -} - -impl QuantumLink for T where T: Into + TryFrom {} - -#[derive(Debug, Clone)] -#[cfg_attr(feature = "envoy", flutter_rust_bridge::frb(opaque))] -pub struct QuantumLinkIdentity { - pub private_keys: Option, - pub xid_document: XIDDocument, -} - -impl QuantumLinkIdentity { - pub fn generate() -> Self { - let (signing_private_key, signing_public_key) = SignatureScheme::MLDSA44.keypair(); - let (encapsulation_private_key, encapsulation_public_key) = - EncapsulationScheme::MLKEM512.keypair(); - - let private_keys = PrivateKeys::with_keys(signing_private_key, encapsulation_private_key); - let public_keys = PublicKeys::new(signing_public_key, encapsulation_public_key); - - let xid_document = XIDDocument::from(public_keys); - - QuantumLinkIdentity { - private_keys: Some(private_keys), - xid_document, - } - } - - pub fn to_bytes(&self) -> Vec { - let mut map = bc_envelope::prelude::Map::new(); - map.insert(CBOR::from("xid_document"), self.clone().xid_document); - if self.private_keys.is_some() { - map.insert( - CBOR::from("private_keys"), - self.clone().private_keys.unwrap(), - ); - } - - CBOR::from(map).to_cbor_data() - } - - pub fn from_bytes(bytes: &[u8]) -> dcbor::Result { - let cbor = CBOR::try_from_data(bytes)?; - let case = cbor.into_case(); - - let CBORCase::Map(map) = case else { - return Err(dcbor::Error::WrongType); - }; - - Ok(QuantumLinkIdentity { - xid_document: map.get("xid_document").ok_or(dcbor::Error::MissingMapKey)?, - private_keys: map.get("private_keys"), - }) - } -} - -fn expiration_duration() -> chrono::Duration { - chrono::Duration::from_std(EXPIRATION_DURATION).expect("expiration duration must fit chrono") -} - -fn validate_expires_at(expires_at: DateTime, now: DateTime) -> Result<(), QlError> { - if now >= expires_at { - return Err(QlError::Expired); - } - - if expires_at > now + expiration_duration() { - return Err(QlError::FutureDated); - } - - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - use crate::{ - api::{ - message::{QuantumLinkMessage, PROTOCOL_VERSION}, - quantum_link::QuantumLink, - }, - fx::ExchangeRate, - message::EnvoyMessage, - quantum_link::{ARIDCache, QlError, QuantumLinkIdentity}, - }; - - #[test] - fn accepts_fresh_and_rejects_immediate_replay() { - let envoy = QuantumLinkIdentity::generate(); - let passport = QuantumLinkIdentity::generate(); - let mut arid_cache = ARIDCache::new(); - - let original_message = exchange_rate_envoy_message(); - let envelope = QuantumLink::seal( - original_message.clone(), - (envoy.private_keys.as_ref().unwrap(), &envoy.xid_document), - &passport.xid_document, - ); - - let (decoded, _sender) = EnvoyMessage::unseal_envoy_message_with_replay_check( - &envelope, - &passport.private_keys.clone().unwrap(), - &mut arid_cache, - ) - .unwrap(); - assert_exchange_rate_matches(&original_message, &decoded); - - let result2 = EnvoyMessage::unseal_envoy_message_with_replay_check( - &envelope, - &passport.private_keys.unwrap(), - &mut arid_cache, - ); - assert!(matches!(result2, Err(QlError::ReplayAttack))); - } - - #[test] - fn rejects_expired_envelope() { - let envoy = QuantumLinkIdentity::generate(); - let passport = QuantumLinkIdentity::generate(); - let mut arid_cache = ARIDCache::new(); - let expired_at = Utc::now() - chrono::Duration::seconds(1); - - let envelope = seal_envoy_message_with_expiration( - exchange_rate_envoy_message(), - &envoy, - &passport, - expired_at, - ); - - let replay_checked_result = EnvoyMessage::unseal_envoy_message_with_replay_check( - &envelope, - &passport.private_keys.unwrap(), - &mut arid_cache, - ); - assert!(matches!(replay_checked_result, Err(QlError::Expired))); - } - - #[test] - fn rejects_future_dated_envelope() { - let envoy = QuantumLinkIdentity::generate(); - let passport = QuantumLinkIdentity::generate(); - let mut arid_cache = ARIDCache::new(); - let expires_at = Utc::now() + expiration_duration() + chrono::Duration::seconds(1); - - let envelope = seal_envoy_message_with_expiration( - exchange_rate_envoy_message(), - &envoy, - &passport, - expires_at, - ); - - let result = EnvoyMessage::unseal_envoy_message_with_replay_check( - &envelope, - &passport.private_keys.unwrap(), - &mut arid_cache, - ); - assert!(matches!(result, Err(QlError::FutureDated))); - } - - #[test] - fn arid_cache_reports_replay_and_expiration() { - let mut cache = ARIDCache::new(); - let arid1 = ARID::new(); - let arid2 = ARID::new(); - - let start = chrono::Utc::now(); - let expires_at = start + expiration_duration(); - - assert_eq!( - cache.check_and_store(&arid1, expires_at, start), - ReplayCheck::Fresh - ); - assert_eq!( - cache.check_and_store(&arid1, expires_at, start), - ReplayCheck::Replay - ); - - let after_expiration = expires_at + chrono::Duration::seconds(1); - assert_eq!( - cache.check_and_store(&arid1, expires_at, after_expiration), - ReplayCheck::Expired - ); - - assert_eq!( - cache.check_and_store( - &arid2, - after_expiration + expiration_duration(), - after_expiration, - ), - ReplayCheck::Fresh - ); - assert_eq!(cache.len(), 1); - assert!(!cache.cache.iter().any(|(id, _)| id == &arid1)); - } - - fn exchange_rate_envoy_message() -> EnvoyMessage { - let fx_rate = ExchangeRate { - currency_code: String::from("USD"), - rate: 0.85, - timestamp: 0, - }; - - EnvoyMessage { - message: QuantumLinkMessage::ExchangeRate(fx_rate), - timestamp: 123456, - protocol_version: Some(PROTOCOL_VERSION), - } - } - - fn assert_exchange_rate_matches(expected: &EnvoyMessage, actual: &EnvoyMessage) { - let expected_rate = match &expected.message { - QuantumLinkMessage::ExchangeRate(rate) => rate, - _ => panic!("Expected ExchangeRate message"), - }; - let actual_rate = match &actual.message { - QuantumLinkMessage::ExchangeRate(rate) => rate, - _ => panic!("Expected ExchangeRate message"), - }; - - assert_eq!(actual.timestamp, expected.timestamp); - assert_eq!(actual.protocol_version, expected.protocol_version); - assert_eq!(actual_rate.rate, expected_rate.rate); - } - - fn seal_envoy_message_with_expiration( - message: EnvoyMessage, - sender: &QuantumLinkIdentity, - recipient: &QuantumLinkIdentity, - expires_at: DateTime, - ) -> Envelope { - let valid_until = Date::from(expires_at); - - let event: SealedEvent = SealedEvent::new( - QuantumLink::encode(message), - ARID::new(), - &sender.xid_document, - ) - .with_date(&valid_until); - event - .to_envelope( - Some(&valid_until), - Some(sender.private_keys.as_ref().unwrap()), - Some(&recipient.xid_document), - ) - .unwrap() - } -} diff --git a/api/src/api/scv.rs b/api/src/api/scv.rs deleted file mode 100644 index 0cec6624..00000000 --- a/api/src/api/scv.rs +++ /dev/null @@ -1,52 +0,0 @@ -use quantum_link_macros::quantum_link; - -#[quantum_link] -pub enum SecurityCheck { - // Envoy to Prime: Initial challenge - #[n(0)] - ChallengeRequest(ChallengeRequest), - - // Prime to Envoy: Response to the challenge - #[n(1)] - ChallengeResponse(ChallengeResponseResult), - - // Envoy to Prime: Verification result - // only send if ChallengeResponse was successful - #[n(2)] - VerificationResult(VerificationResult), -} - -#[quantum_link] -pub struct ChallengeRequest { - #[n(0)] - pub data: Vec, -} - -#[quantum_link] -pub enum ChallengeResponseResult { - #[n(0)] - Success { - #[n(0)] - data: Vec, - }, - #[n(1)] - Error { - #[n(0)] - error: String, - }, -} - -#[quantum_link] -pub enum VerificationResult { - #[n(0)] - Success, - // Error due to Envoy not being able to perform the verification - #[n(1)] - Error { - #[n(0)] - error: String, - }, - // Actual failure indicating device has been tampered with - #[n(2)] - Failure, -} diff --git a/api/src/api/status.rs b/api/src/api/status.rs deleted file mode 100644 index f564cec2..00000000 --- a/api/src/api/status.rs +++ /dev/null @@ -1,35 +0,0 @@ -use quantum_link_macros::quantum_link; - -#[quantum_link] -pub struct DeviceStatus { - #[n(0)] - pub version: String, - #[n(1)] - pub battery_level: u8, -} - -#[quantum_link] -pub struct EnvoyStatus { - #[n(0)] - pub version: String, -} - -#[quantum_link] -pub struct Heartbeat {} - -#[quantum_link] -pub struct TimezoneRequest {} - -#[quantum_link] -pub struct TimezoneResponse { - #[n(0)] - pub offset_minutes: i32, - #[n(1)] - pub zone: String, -} - -#[quantum_link] -pub struct DeviceNameUpdate { - #[n(0)] - pub device_name: String, -} diff --git a/api/src/api/tests.rs b/api/src/api/tests.rs deleted file mode 100644 index b8e635b8..00000000 --- a/api/src/api/tests.rs +++ /dev/null @@ -1,414 +0,0 @@ -use dcbor::{CBORCase, CBOR}; -use quantum_link_macros::Cbor; - -#[derive(Debug, Clone, PartialEq, Cbor)] -pub struct TestStruct { - #[n(0)] - pub name: String, - #[n(1)] - pub value: u64, - #[n(2)] - pub enabled: bool, -} - -#[derive(Debug, Clone, PartialEq, Cbor)] -pub struct TestWithVec { - #[n(0)] - pub items: Vec, - #[n(1)] - pub label: String, -} - -#[derive(Debug, Clone, PartialEq, Cbor)] -pub struct TestWithArray { - #[n(0)] - pub hash: [u8; 32], - #[n(1)] - pub id: u64, -} - -#[derive(Debug, Clone, PartialEq, Cbor)] -pub enum TestEnumTuple { - #[n(0)] - First(TestStruct), - #[n(1)] - Second(TestWithVec), -} - -#[derive(Debug, Clone, PartialEq, Cbor)] -pub enum TestEnumStruct { - #[n(0)] - VariantA { - #[n(0)] - count: u64, - #[n(1)] - active: bool, - }, - #[n(1)] - VariantB { - #[n(0)] - message: String, - }, -} - -#[derive(Debug, Clone, PartialEq, Cbor)] -pub enum TestEnumUnit { - #[n(0)] - Empty, - #[n(1)] - WithData(TestStruct), -} - -#[derive(Debug, Clone, PartialEq, Cbor)] -pub enum TestEnumMixed { - #[n(0)] - Unit, - #[n(1)] - Tuple(TestStruct), - #[n(2)] - Struct { - #[n(0)] - field1: String, - #[n(1)] - field2: u64, - }, -} - -#[test] -fn struct_roundtrip() { - let original = TestStruct { - name: "test".to_string(), - value: 42, - enabled: true, - }; - - let cbor: CBOR = original.clone().into(); - let recovered: TestStruct = cbor.try_into().unwrap(); - - assert_eq!(original, recovered); -} - -#[test] -fn struct_with_vec_roundtrip() { - let original = TestWithVec { - items: vec![1, 2, 3, 4, 5], - label: "data".to_string(), - }; - - let cbor: CBOR = original.clone().into(); - let recovered: TestWithVec = cbor.try_into().unwrap(); - - assert_eq!(original, recovered); -} - -#[test] -fn struct_with_array_roundtrip() { - let original = TestWithArray { - hash: [ - 0xde, 0xad, 0xbe, 0xef, 0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77, 0x88, 0x99, - 0xaa, 0xbb, 0xcc, 0xdd, 0xee, 0xff, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08, - 0x09, 0x0a, 0x0b, 0x0c, - ], - id: 12345, - }; - - let cbor: CBOR = original.clone().into(); - let recovered: TestWithArray = cbor.try_into().unwrap(); - - assert_eq!(original, recovered); -} - -#[test] -fn byte_string_encoding() { - let test = TestWithVec { - items: vec![1, 2, 3], - label: "test".to_string(), - }; - - let cbor: CBOR = test.into(); - let case = cbor.into_case(); - - match case { - CBORCase::Map(map) => { - let items_cbor: CBOR = map.get(0).unwrap(); - let items_case = items_cbor.into_case(); - assert!( - matches!(items_case, CBORCase::ByteString(_)), - "Vec should be encoded as byte string" - ); - } - _ => panic!("Expected CBOR map"), - } -} - -#[test] -fn array_byte_string_encoding() { - let test = TestWithArray { - hash: [0u8; 32], - id: 1, - }; - - let cbor: CBOR = test.into(); - let case = cbor.into_case(); - - match case { - CBORCase::Map(map) => { - let hash_cbor: CBOR = map.get(0).unwrap(); - let hash_case = hash_cbor.into_case(); - assert!( - matches!(hash_case, CBORCase::ByteString(_)), - "[u8; N] should be encoded as byte string" - ); - } - _ => panic!("Expected CBOR map"), - } -} - -#[test] -fn enum_tuple_roundtrip() { - let test_struct = TestStruct { - name: "inner".to_string(), - value: 100, - enabled: false, - }; - - let original = TestEnumTuple::First(test_struct); - let cbor: CBOR = original.clone().into(); - let recovered: TestEnumTuple = cbor.try_into().unwrap(); - - assert_eq!(original, recovered); - - let test_vec = TestWithVec { - items: vec![255, 128, 0], - label: "bytes".to_string(), - }; - - let original2 = TestEnumTuple::Second(test_vec); - let cbor2: CBOR = original2.clone().into(); - let recovered2: TestEnumTuple = cbor2.try_into().unwrap(); - - assert_eq!(original2, recovered2); -} - -#[test] -fn enum_struct_roundtrip() { - let original_a = TestEnumStruct::VariantA { - count: 999, - active: true, - }; - - let cbor_a: CBOR = original_a.clone().into(); - let recovered_a: TestEnumStruct = cbor_a.try_into().unwrap(); - - assert_eq!(original_a, recovered_a); - - let original_b = TestEnumStruct::VariantB { - message: "hello world".to_string(), - }; - - let cbor_b: CBOR = original_b.clone().into(); - let recovered_b: TestEnumStruct = cbor_b.try_into().unwrap(); - - assert_eq!(original_b, recovered_b); -} - -#[test] -fn enum_unit_roundtrip() { - let original_empty = TestEnumUnit::Empty; - let cbor: CBOR = original_empty.clone().into(); - let recovered: TestEnumUnit = cbor.try_into().unwrap(); - - assert_eq!(original_empty, recovered); - - let test_struct = TestStruct { - name: "with data".to_string(), - value: 123, - enabled: true, - }; - - let original_with_data = TestEnumUnit::WithData(test_struct); - let cbor2: CBOR = original_with_data.clone().into(); - let recovered2: TestEnumUnit = cbor2.try_into().unwrap(); - - assert_eq!(original_with_data, recovered2); -} - -#[test] -fn enum_mixed_roundtrip() { - let unit = TestEnumMixed::Unit; - let cbor: CBOR = unit.clone().into(); - let recovered: TestEnumMixed = cbor.try_into().unwrap(); - assert_eq!(unit, recovered); - - let tuple = TestEnumMixed::Tuple(TestStruct { - name: "tuple".to_string(), - value: 50, - enabled: false, - }); - let cbor: CBOR = tuple.clone().into(); - let recovered: TestEnumMixed = cbor.try_into().unwrap(); - assert_eq!(tuple, recovered); - - let struct_var = TestEnumMixed::Struct { - field1: "struct variant".to_string(), - field2: 9999, - }; - let cbor: CBOR = struct_var.clone().into(); - let recovered: TestEnumMixed = cbor.try_into().unwrap(); - assert_eq!(struct_var, recovered); -} - -#[test] -fn cbor_structure() { - let test = TestStruct { - name: "check".to_string(), - value: 7, - enabled: true, - }; - - let cbor: CBOR = test.into(); - let case = cbor.into_case(); - - match case { - CBORCase::Map(map) => { - assert_eq!(map.len(), 3); - - assert!(map.get::(0).is_some()); - assert!(map.get::(1).is_some()); - assert!(map.get::(2).is_some()); - } - _ => panic!("Expected CBOR map"), - } -} - -#[test] -fn enum_cbor_structure() { - let test_struct = TestStruct { - name: "test".to_string(), - value: 1, - enabled: true, - }; - - let variant = TestEnumTuple::First(test_struct); - let cbor: CBOR = variant.into(); - let case = cbor.into_case(); - - match case { - CBORCase::Array(arr) => { - assert_eq!(arr.len(), 2); - - let index: u64 = arr.first().unwrap().clone().try_into().unwrap(); - assert_eq!(index, 0); - } - _ => panic!("Expected CBOR array for enum"), - } -} - -#[test] -fn enum_tuple_vs_struct_encoding() { - #[derive(Debug, Clone, PartialEq, Cbor)] - pub struct InnerData { - #[n(0)] - pub count: u64, - #[n(1)] - pub active: bool, - } - - #[derive(Debug, Clone, PartialEq, Cbor)] - pub enum EnumWithTupleStruct { - #[n(0)] - Variant(InnerData), - } - - #[derive(Debug, Clone, PartialEq, Cbor)] - pub enum EnumWithStructFields { - #[n(0)] - Variant { - #[n(0)] - count: u64, - #[n(1)] - active: bool, - }, - } - - let tuple_enum = EnumWithTupleStruct::Variant(InnerData { - count: 42, - active: true, - }); - - let struct_enum = EnumWithStructFields::Variant { - count: 42, - active: true, - }; - - let tuple_cbor: CBOR = tuple_enum.into(); - let struct_cbor: CBOR = struct_enum.into(); - - let tuple_bytes = tuple_cbor.to_cbor_data(); - let struct_bytes = struct_cbor.to_cbor_data(); - - assert_eq!( - tuple_bytes, struct_bytes, - "Enum with tuple(struct) should serialize the same as enum with struct fields" - ); -} - -#[test] -fn newtype() { - #[derive(Debug, Clone, PartialEq, Cbor)] - struct NewType(String); - - let value = NewType(String::from("yes")); - let cbor: CBOR = value.clone().into(); - let case = cbor.clone().into_case(); - - match case { - CBORCase::Text(_) => {} - _ => panic!("invalid case"), - } - - assert_eq!(value, NewType::try_from(cbor).unwrap()) -} - -#[test] -fn option_array() { - #[derive(Debug, Clone, PartialEq, Cbor)] - struct OptionArray { - #[n(0)] - arr: Option<[u8; 10]>, - #[n(1)] - vec: Option>, - } - - let a = [10; 10]; - let b = vec![12; 4]; - let value = OptionArray { - arr: Some(a), - vec: Some(b.clone()), - }; - let cbor: CBOR = value.clone().into(); - let case = cbor.clone().into_case(); - - match case { - CBORCase::Map(map) => { - assert_eq!(map.len(), 2); - let arr: CBOR = map.get(0).unwrap(); - match arr.into_case() { - CBORCase::ByteString(bytes) => { - assert_eq!(bytes.data(), &a) - } - _ => panic!("expected bytestring"), - } - let vec: CBOR = map.get(1).unwrap(); - match vec.into_case() { - CBORCase::ByteString(bytes) => { - assert_eq!(bytes.data(), &b) - } - _ => panic!("expected bytestring"), - } - } - _ => panic!("Expected CBOR array for enum"), - } - - assert_eq!(value, OptionArray::try_from(cbor).unwrap()) -} diff --git a/api/src/lib.rs b/api/src/lib.rs deleted file mode 100644 index e6ecd816..00000000 --- a/api/src/lib.rs +++ /dev/null @@ -1,11 +0,0 @@ -pub mod api; -pub use api::*; - -/// Marker trait for types that have a Cbor derive (structs and enums, not primitives). -/// This is used to enforce that enum tuple variants wrap Cbor-derived types. -pub(crate) trait CborMarker {} - -pub use bc_components; -pub use bc_envelope; -pub use bc_xid; -pub use dcbor; diff --git a/api/tests/golden_tests.rs b/api/tests/golden_tests.rs deleted file mode 100644 index e240c9e3..00000000 --- a/api/tests/golden_tests.rs +++ /dev/null @@ -1,550 +0,0 @@ -//! golden/snapshot tests for QuantumLinkMessage codec -//! -//! to update snapshots when serialization intentionally changes: -//! ``` -//! INSTA_UPDATE=always cargo test -//! ``` - -use dcbor::CBOR; -use foundation_api::{ - backup::*, bitcoin::*, firmware::*, fx::*, message::*, onboarding::*, pairing::*, passport::*, - scv::*, status::*, -}; - -/// convert a message to hex-encoded CBOR bytes -fn to_hex(message: &QuantumLinkMessage) -> String { - let cbor: CBOR = message.clone().into(); - let bytes = cbor.to_cbor_data(); - bytes.iter().map(|b| format!("{b:02x}")).collect::() -} - -/// decode hex-encoded CBOR bytes back to a message -fn from_hex(hex: &str) -> QuantumLinkMessage { - let bytes: Vec = (0..hex.len()) - .step_by(2) - .map(|i| u8::from_str_radix(&hex[i..i + 2], 16).unwrap()) - .collect(); - let cbor = CBOR::try_from_data(&bytes).unwrap(); - QuantumLinkMessage::try_from(cbor).unwrap() -} - -macro_rules! assert_golden { - ($message:expr) => {{ - let message = $message; - let hex = to_hex(&message); - insta::assert_snapshot!(hex.clone()); - - let decoded = from_hex(&hex); - assert_eq!(message, decoded, "roundtrip decode failed"); - }}; -} - -#[test] -fn golden_exchange_rate() { - assert_golden!(QuantumLinkMessage::ExchangeRate(ExchangeRate { - currency_code: "USD".to_string(), - rate: 42_000.5, - timestamp: 1700000000, - })); -} - -#[test] -fn golden_exchange_rate_history() { - assert_golden!(QuantumLinkMessage::ExchangeRateHistory( - ExchangeRateHistory { - history: vec![ - PricePoint { - rate: 41000.0, - timestamp: 1699999900, - }, - PricePoint { - rate: 42000.0, - timestamp: 1700000000, - }, - ], - currency_code: "EUR".to_string(), - } - )); -} - -#[test] -fn golden_firmware_update_check_request() { - assert_golden!(QuantumLinkMessage::FirmwareUpdateCheckRequest( - FirmwareUpdateCheckRequest { - current_version: "2.4.0".to_string(), - }, - )); -} - -#[test] -fn golden_firmware_update_check_response_available() { - assert_golden!(QuantumLinkMessage::FirmwareUpdateCheckResponse( - FirmwareUpdateCheckResponse::Available(FirmwareUpdateAvailable { - version: "2.5.0".to_string(), - changelog: "Bug fixes".to_string(), - timestamp: 1700000000, - total_size: 1024000, - patch_count: 3, - }), - )); -} - -#[test] -fn golden_firmware_update_check_response_not_available() { - assert_golden!(QuantumLinkMessage::FirmwareUpdateCheckResponse( - FirmwareUpdateCheckResponse::NotAvailable, - )); -} - -#[test] -fn golden_firmware_fetch_request() { - assert_golden!(QuantumLinkMessage::FirmwareFetchRequest( - FirmwareFetchRequest { - current_version: "2.4.0".to_string(), - chunk_offset: None - }, - )); -} - -#[test] -fn golden_firmware_fetch_event_not_available() { - assert_golden!(QuantumLinkMessage::FirmwareFetchEvent( - FirmwareFetchEvent::UpdateNotAvailable, - )); -} - -#[test] -fn golden_firmware_fetch_event_starting() { - assert_golden!(QuantumLinkMessage::FirmwareFetchEvent( - FirmwareFetchEvent::Starting(FirmwareUpdateAvailable { - version: "2.5.0".to_string(), - changelog: "New features".to_string(), - timestamp: 1700000000, - total_size: 2048000, - patch_count: 5, - }), - )); -} - -#[test] -fn golden_firmware_fetch_event_downloading() { - assert_golden!(QuantumLinkMessage::FirmwareFetchEvent( - FirmwareFetchEvent::Downloading, - )); -} - -#[test] -fn golden_firmware_fetch_event_chunk() { - assert_golden!(QuantumLinkMessage::FirmwareFetchEvent( - FirmwareFetchEvent::Chunk(FirmwareChunk { - patch_index: 0, - total_patches: 3, - chunk_index: 5, - total_chunks: 100, - data: vec![0xde, 0xad, 0xbe, 0xef], - }), - )); -} - -#[test] -fn golden_firmware_fetch_event_error() { - assert_golden!(QuantumLinkMessage::FirmwareFetchEvent( - FirmwareFetchEvent::Error { - error: "Download failed".to_string(), - }, - )); -} - -#[test] -fn golden_firmware_update_result_update_verified() { - assert_golden!(QuantumLinkMessage::FirmwareInstallEvent( - FirmwareInstallEvent::UpdateVerified, - )); -} - -#[test] -fn golden_firmware_update_result_installing() { - assert_golden!(QuantumLinkMessage::FirmwareInstallEvent( - FirmwareInstallEvent::Installing, - )); -} - -#[test] -fn golden_firmware_update_result_rebooting() { - assert_golden!(QuantumLinkMessage::FirmwareInstallEvent( - FirmwareInstallEvent::Rebooting, - )); -} - -#[test] -fn golden_firmware_update_result_success() { - assert_golden!(QuantumLinkMessage::FirmwareInstallEvent( - FirmwareInstallEvent::Success { - installed_version: "2.5.0".to_string(), - }, - )); -} - -#[test] -fn golden_firmware_update_result_error_verify() { - assert_golden!(QuantumLinkMessage::FirmwareInstallEvent( - FirmwareInstallEvent::Error { - error: "Signature verification failed".to_string(), - stage: InstallErrorStage::Verify, - }, - )); -} - -#[test] -fn golden_firmware_update_result_error_install() { - assert_golden!(QuantumLinkMessage::FirmwareInstallEvent( - FirmwareInstallEvent::Error { - error: "Installation failed".to_string(), - stage: InstallErrorStage::Install, - }, - )); -} - -#[test] -fn golden_device_status() { - assert_golden!(QuantumLinkMessage::DeviceStatus(DeviceStatus { - battery_level: 85, - version: "2.4.0".to_string(), - })); -} - -#[test] -fn golden_device_status_updating() { - assert_golden!(QuantumLinkMessage::DeviceStatus(DeviceStatus { - battery_level: 90, - version: "2.4.0".to_string(), - })); -} - -#[test] -fn golden_envoy_status() { - assert_golden!(QuantumLinkMessage::EnvoyStatus(EnvoyStatus { - version: "1.0.0".to_string(), - })); -} - -#[test] -fn golden_pairing_request() { - assert_golden!(QuantumLinkMessage::PairingRequest(PairingRequest { - xid_document: vec![0x01, 0x02, 0x03, 0x04], - device_name: "My iPhone".to_string(), - })); -} - -#[test] -fn golden_pairing_response() { - assert_golden!(QuantumLinkMessage::PairingResponse(PairingResponse { - passport_model: PassportModel::Prime, - passport_firmware_version: PassportFirmwareVersion("2.4.0".to_string()), - passport_serial: PassportSerial("ABC123".to_string()), - passport_color: PassportColor::Dark, - onboarding_complete: true, - device_name: Some("Passport Prime".to_string()), - })); -} - -#[test] -fn golden_onboarding_state_firmware_update_screen() { - assert_golden!(QuantumLinkMessage::OnboardingState( - OnboardingState::FirmwareUpdateScreen, - )); -} - -#[test] -fn golden_onboarding_state_completed() { - assert_golden!(QuantumLinkMessage::OnboardingState( - OnboardingState::Completed, - )); -} - -#[test] -fn golden_sign_psbt() { - assert_golden!(QuantumLinkMessage::SignPsbt(SignPsbt { - account_id: "account-1".to_string(), - psbt: vec![0x70, 0x73, 0x62, 0x74, 0xff], - })); -} - -#[test] -fn golden_broadcast_transaction() { - assert_golden!(QuantumLinkMessage::BroadcastTransaction( - BroadcastTransaction { - account_id: "account-1".to_string(), - psbt: vec![0x70, 0x73, 0x62, 0x74, 0xff], - }, - )); -} - -#[test] -fn golden_account_update() { - assert_golden!(QuantumLinkMessage::AccountUpdate(AccountUpdate { - account_id: "account-1".to_string(), - update: vec![0x01, 0x02, 0x03], - })); -} - -#[test] -fn golden_apply_passphrase_some() { - assert_golden!(QuantumLinkMessage::ApplyPassphrase(ApplyPassphrase { - fingerprint: Some("abc123".to_string()), - })); -} - -#[test] -fn golden_apply_passphrase_none() { - assert_golden!(QuantumLinkMessage::ApplyPassphrase(ApplyPassphrase { - fingerprint: None, - })); -} - -#[test] -fn golden_security_check_challenge_request() { - assert_golden!(QuantumLinkMessage::SecurityCheck( - SecurityCheck::ChallengeRequest(ChallengeRequest { - data: vec![0xca, 0xfe, 0xba, 0xbe], - }), - )); -} - -#[test] -fn golden_security_check_challenge_response_success() { - assert_golden!(QuantumLinkMessage::SecurityCheck( - SecurityCheck::ChallengeResponse(ChallengeResponseResult::Success { - data: vec![0xde, 0xad, 0xbe, 0xef], - }), - )); -} - -#[test] -fn golden_security_check_challenge_response_error() { - assert_golden!(QuantumLinkMessage::SecurityCheck( - SecurityCheck::ChallengeResponse(ChallengeResponseResult::Error { - error: "Invalid signature".to_string(), - }), - )); -} - -#[test] -fn golden_security_check_verification_success() { - assert_golden!(QuantumLinkMessage::SecurityCheck( - SecurityCheck::VerificationResult(VerificationResult::Success), - )); -} - -#[test] -fn golden_security_check_verification_error() { - assert_golden!(QuantumLinkMessage::SecurityCheck( - SecurityCheck::VerificationResult(VerificationResult::Error { - error: "Verification failed".to_string(), - }), - )); -} - -#[test] -fn golden_envoy_magic_backup_enabled_request() { - assert_golden!(QuantumLinkMessage::EnvoyMagicBackupEnabledRequest( - EnvoyMagicBackupEnabledRequest {}, - )); -} - -#[test] -fn golden_envoy_magic_backup_enabled_response() { - assert_golden!(QuantumLinkMessage::EnvoyMagicBackupEnabledResponse( - EnvoyMagicBackupEnabledResponse { enabled: true }, - )); -} - -#[test] -fn golden_prime_magic_backup_enabled() { - assert_golden!(QuantumLinkMessage::PrimeMagicBackupEnabled( - PrimeMagicBackupEnabled { - enabled: true, - seed_fingerprint: SeedFingerprint([0x42; 32]), - }, - )); -} - -#[test] -fn golden_prime_magic_backup_status_request() { - assert_golden!(QuantumLinkMessage::PrimeMagicBackupStatusRequest( - PrimeMagicBackupStatusRequest { - seed_fingerprint: SeedFingerprint([0xab; 32]), - timestamp: None, - }, - )); -} - -#[test] -fn golden_prime_magic_backup_status_response() { - assert_golden!(QuantumLinkMessage::PrimeMagicBackupStatusResponse( - PrimeMagicBackupStatusResponse { - shard_backup_found: true, - }, - )); -} - -#[test] -fn golden_backup_shard_request() { - assert_golden!(QuantumLinkMessage::BackupShardRequest(BackupShardRequest { - shard: Shard(vec![0x01, 0x02, 0x03, 0x04, 0x05]), - })); -} - -#[test] -fn golden_backup_shard_response_success() { - assert_golden!(QuantumLinkMessage::BackupShardResponse( - BackupShardResponse::Success, - )); -} - -#[test] -fn golden_backup_shard_response_error() { - assert_golden!(QuantumLinkMessage::BackupShardResponse( - BackupShardResponse::Error { - error: "Storage full".to_string(), - }, - )); -} - -#[test] -fn golden_restore_shard_request() { - assert_golden!(QuantumLinkMessage::RestoreShardRequest( - RestoreShardRequest { - seed_fingerprint: SeedFingerprint([0xcd; 32]), - timestamp: None, - }, - )); -} - -#[test] -fn golden_restore_shard_response_success() { - assert_golden!(QuantumLinkMessage::RestoreShardResponse( - RestoreShardResponse::Success { - shard: Shard(vec![0x0a, 0x0b, 0x0c]), - }, - )); -} - -#[test] -fn golden_restore_shard_response_error() { - assert_golden!(QuantumLinkMessage::RestoreShardResponse( - RestoreShardResponse::Error { - error: "Not found".to_string(), - }, - )); -} - -#[test] -fn golden_restore_shard_response_not_found() { - assert_golden!(QuantumLinkMessage::RestoreShardResponse( - RestoreShardResponse::NotFound, - )); -} - -#[test] -fn golden_create_magic_backup_event_start() { - assert_golden!(QuantumLinkMessage::CreateMagicBackupEvent( - CreateMagicBackupEvent::Start(StartMagicBackup { - seed_fingerprint: SeedFingerprint([0xef; 32]), - total_chunks: 100, - hash: [0xaa; 32], - }), - )); -} - -#[test] -fn golden_create_magic_backup_event_chunk() { - assert_golden!(QuantumLinkMessage::CreateMagicBackupEvent( - CreateMagicBackupEvent::Chunk(BackupChunk { - chunk_index: 5, - total_chunks: 100, - data: vec![0x11, 0x22, 0x33], - }), - )); -} - -#[test] -fn golden_create_magic_backup_result_success() { - assert_golden!(QuantumLinkMessage::CreateMagicBackupResult( - CreateMagicBackupResult::Success, - )); -} - -#[test] -fn golden_create_magic_backup_result_error() { - assert_golden!(QuantumLinkMessage::CreateMagicBackupResult( - CreateMagicBackupResult::Error { - error: "Upload failed".to_string(), - }, - )); -} - -#[test] -fn golden_restore_magic_backup_request() { - assert_golden!(QuantumLinkMessage::RestoreMagicBackupRequest( - RestoreMagicBackupRequest { - seed_fingerprint: SeedFingerprint([0xbb; 32]), - resume_from_chunk: 50, - }, - )); -} - -#[test] -fn golden_restore_magic_backup_event_no_backup() { - assert_golden!(QuantumLinkMessage::RestoreMagicBackupEvent( - RestoreMagicBackupEvent::NotFound, - )); -} - -#[test] -fn golden_restore_magic_backup_event_starting() { - assert_golden!(QuantumLinkMessage::RestoreMagicBackupEvent( - RestoreMagicBackupEvent::Starting(BackupMetadata { total_chunks: 200 }), - )); -} - -#[test] -fn golden_restore_magic_backup_event_chunk() { - assert_golden!(QuantumLinkMessage::RestoreMagicBackupEvent( - RestoreMagicBackupEvent::Chunk(BackupChunk { - chunk_index: 10, - total_chunks: 50, - data: vec![0xaa, 0xbb, 0xcc, 0xdd], - }), - )); -} - -#[test] -fn golden_restore_magic_backup_event_error() { - assert_golden!(QuantumLinkMessage::RestoreMagicBackupEvent( - RestoreMagicBackupEvent::Error { - error: "Network error".to_string(), - }, - )); -} - -#[test] -fn golden_restore_magic_backup_result_success() { - assert_golden!(QuantumLinkMessage::RestoreMagicBackupResult( - RestoreMagicBackupResult::Success, - )); -} - -#[test] -fn golden_restore_magic_backup_result_error() { - assert_golden!(QuantumLinkMessage::RestoreMagicBackupResult( - RestoreMagicBackupResult::Error { - error: "Checksum mismatch".to_string(), - }, - )); -} - -#[test] -fn golden_heartbeat() { - assert_golden!(QuantumLinkMessage::Heartbeat(Heartbeat {})) -} diff --git a/api/tests/snapshots/golden_tests__golden_account_update.snap b/api/tests/snapshots/golden_tests__golden_account_update.snap deleted file mode 100644 index c4158ee2..00000000 --- a/api/tests/snapshots/golden_tests__golden_account_update.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820fa200696163636f756e742d310143010203 diff --git a/api/tests/snapshots/golden_tests__golden_apply_passphrase_none.snap b/api/tests/snapshots/golden_tests__golden_apply_passphrase_none.snap deleted file mode 100644 index d18af859..00000000 --- a/api/tests/snapshots/golden_tests__golden_apply_passphrase_none.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8210a0 diff --git a/api/tests/snapshots/golden_tests__golden_apply_passphrase_some.snap b/api/tests/snapshots/golden_tests__golden_apply_passphrase_some.snap deleted file mode 100644 index fe18e395..00000000 --- a/api/tests/snapshots/golden_tests__golden_apply_passphrase_some.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8210a10066616263313233 diff --git a/api/tests/snapshots/golden_tests__golden_backup_shard_request.snap b/api/tests/snapshots/golden_tests__golden_backup_shard_request.snap deleted file mode 100644 index f77f942f..00000000 --- a/api/tests/snapshots/golden_tests__golden_backup_shard_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8216a100450102030405 diff --git a/api/tests/snapshots/golden_tests__golden_backup_shard_response_error.snap b/api/tests/snapshots/golden_tests__golden_backup_shard_response_error.snap deleted file mode 100644 index 736377b6..00000000 --- a/api/tests/snapshots/golden_tests__golden_backup_shard_response_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82178201a1006c53746f726167652066756c6c diff --git a/api/tests/snapshots/golden_tests__golden_backup_shard_response_success.snap b/api/tests/snapshots/golden_tests__golden_backup_shard_response_success.snap deleted file mode 100644 index 18b44bec..00000000 --- a/api/tests/snapshots/golden_tests__golden_backup_shard_response_success.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82178100 diff --git a/api/tests/snapshots/golden_tests__golden_broadcast_transaction.snap b/api/tests/snapshots/golden_tests__golden_broadcast_transaction.snap deleted file mode 100644 index f240d2d4..00000000 --- a/api/tests/snapshots/golden_tests__golden_broadcast_transaction.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820ea200696163636f756e742d31014570736274ff diff --git a/api/tests/snapshots/golden_tests__golden_create_magic_backup_event_chunk.snap b/api/tests/snapshots/golden_tests__golden_create_magic_backup_event_chunk.snap deleted file mode 100644 index bfe41e5a..00000000 --- a/api/tests/snapshots/golden_tests__golden_create_magic_backup_event_chunk.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181a8201a300050118640243112233 diff --git a/api/tests/snapshots/golden_tests__golden_create_magic_backup_event_start.snap b/api/tests/snapshots/golden_tests__golden_create_magic_backup_event_start.snap deleted file mode 100644 index adc08fb4..00000000 --- a/api/tests/snapshots/golden_tests__golden_create_magic_backup_event_start.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181a8200a3005820efefefefefefefefefefefefefefefefefefefefefefefefefefefefefefefef011864025820aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa diff --git a/api/tests/snapshots/golden_tests__golden_create_magic_backup_result_error.snap b/api/tests/snapshots/golden_tests__golden_create_magic_backup_result_error.snap deleted file mode 100644 index 7aebdbcd..00000000 --- a/api/tests/snapshots/golden_tests__golden_create_magic_backup_result_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181b8201a1006d55706c6f6164206661696c6564 diff --git a/api/tests/snapshots/golden_tests__golden_create_magic_backup_result_success.snap b/api/tests/snapshots/golden_tests__golden_create_magic_backup_result_success.snap deleted file mode 100644 index 0c21e2e2..00000000 --- a/api/tests/snapshots/golden_tests__golden_create_magic_backup_result_success.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181b8100 diff --git a/api/tests/snapshots/golden_tests__golden_device_status.snap b/api/tests/snapshots/golden_tests__golden_device_status.snap deleted file mode 100644 index 733ccebc..00000000 --- a/api/tests/snapshots/golden_tests__golden_device_status.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8207a20065322e342e30011855 diff --git a/api/tests/snapshots/golden_tests__golden_device_status_updating.snap b/api/tests/snapshots/golden_tests__golden_device_status_updating.snap deleted file mode 100644 index 0e69f1d5..00000000 --- a/api/tests/snapshots/golden_tests__golden_device_status_updating.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8207a20065322e342e3001185a diff --git a/api/tests/snapshots/golden_tests__golden_envoy_magic_backup_enabled_request.snap b/api/tests/snapshots/golden_tests__golden_envoy_magic_backup_enabled_request.snap deleted file mode 100644 index 153c3310..00000000 --- a/api/tests/snapshots/golden_tests__golden_envoy_magic_backup_enabled_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8211a0 diff --git a/api/tests/snapshots/golden_tests__golden_envoy_magic_backup_enabled_response.snap b/api/tests/snapshots/golden_tests__golden_envoy_magic_backup_enabled_response.snap deleted file mode 100644 index 2d62aef6..00000000 --- a/api/tests/snapshots/golden_tests__golden_envoy_magic_backup_enabled_response.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8212a100f5 diff --git a/api/tests/snapshots/golden_tests__golden_envoy_status.snap b/api/tests/snapshots/golden_tests__golden_envoy_status.snap deleted file mode 100644 index df6e905d..00000000 --- a/api/tests/snapshots/golden_tests__golden_envoy_status.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8208a10065312e302e30 diff --git a/api/tests/snapshots/golden_tests__golden_exchange_rate.snap b/api/tests/snapshots/golden_tests__golden_exchange_rate.snap deleted file mode 100644 index 40d10651..00000000 --- a/api/tests/snapshots/golden_tests__golden_exchange_rate.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8200a3006355534401fa47241080021a6553f100 diff --git a/api/tests/snapshots/golden_tests__golden_exchange_rate_history.snap b/api/tests/snapshots/golden_tests__golden_exchange_rate_history.snap deleted file mode 100644 index c727066f..00000000 --- a/api/tests/snapshots/golden_tests__golden_exchange_rate_history.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8201a20082a20019a028011a6553f09ca20019a410011a6553f1000163455552 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_chunk.snap b/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_chunk.snap deleted file mode 100644 index 3c50383b..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_chunk.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82058203a50000010302050318640444deadbeef diff --git a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_downloading.snap b/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_downloading.snap deleted file mode 100644 index 87593212..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_downloading.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82058102 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_error.snap b/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_error.snap deleted file mode 100644 index 12c49123..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82058205a1006f446f776e6c6f6164206661696c6564 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_not_available.snap b/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_not_available.snap deleted file mode 100644 index 1035adb8..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_not_available.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82058100 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_starting.snap b/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_starting.snap deleted file mode 100644 index 7538c4fe..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_fetch_event_starting.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82058201a50065322e352e30016c4e6577206665617475726573021a6553f100031a001f40000405 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_fetch_request.snap b/api/tests/snapshots/golden_tests__golden_firmware_fetch_request.snap deleted file mode 100644 index 0cddae61..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_fetch_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8204a10065322e342e30 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_check_request.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_check_request.snap deleted file mode 100644 index 6d75041a..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_check_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8202a10065322e342e30 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_check_response_available.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_check_response_available.snap deleted file mode 100644 index b6c888fc..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_check_response_available.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82038200a50065322e352e300169427567206669786573021a6553f100031a000fa0000403 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_check_response_not_available.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_check_response_not_available.snap deleted file mode 100644 index 34404d73..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_check_response_not_available.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82038101 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_result_error.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_result_error.snap deleted file mode 100644 index c9c71376..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_result_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82068201a10073496e7374616c6c6174696f6e206661696c6564 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_result_error_install.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_result_error_install.snap deleted file mode 100644 index a2accd32..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_result_error_install.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82068204a20073496e7374616c6c6174696f6e206661696c6564018102 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_result_error_verify.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_result_error_verify.snap deleted file mode 100644 index 5a9de85f..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_result_error_verify.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82068204a200781d5369676e617475726520766572696669636174696f6e206661696c6564018101 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_result_installing.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_result_installing.snap deleted file mode 100644 index 2465e2d4..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_result_installing.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82068101 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_result_rebooting.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_result_rebooting.snap deleted file mode 100644 index efe79a65..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_result_rebooting.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82068102 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_result_success.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_result_success.snap deleted file mode 100644 index 83432358..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_result_success.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82068203a10065322e352e30 diff --git a/api/tests/snapshots/golden_tests__golden_firmware_update_result_update_verified.snap b/api/tests/snapshots/golden_tests__golden_firmware_update_result_update_verified.snap deleted file mode 100644 index 0d16e030..00000000 --- a/api/tests/snapshots/golden_tests__golden_firmware_update_result_update_verified.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82068100 diff --git a/api/tests/snapshots/golden_tests__golden_heartbeat.snap b/api/tests/snapshots/golden_tests__golden_heartbeat.snap deleted file mode 100644 index d4da48e6..00000000 --- a/api/tests/snapshots/golden_tests__golden_heartbeat.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181fa0 diff --git a/api/tests/snapshots/golden_tests__golden_onboarding_state_completed.snap b/api/tests/snapshots/golden_tests__golden_onboarding_state_completed.snap deleted file mode 100644 index 0ded2ecf..00000000 --- a/api/tests/snapshots/golden_tests__golden_onboarding_state_completed.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820c8110 diff --git a/api/tests/snapshots/golden_tests__golden_onboarding_state_firmware_update_screen.snap b/api/tests/snapshots/golden_tests__golden_onboarding_state_firmware_update_screen.snap deleted file mode 100644 index fce20b19..00000000 --- a/api/tests/snapshots/golden_tests__golden_onboarding_state_firmware_update_screen.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820c8102 diff --git a/api/tests/snapshots/golden_tests__golden_pairing_request.snap b/api/tests/snapshots/golden_tests__golden_pairing_request.snap deleted file mode 100644 index b1773d3a..00000000 --- a/api/tests/snapshots/golden_tests__golden_pairing_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8209a200440102030401694d79206950686f6e65 diff --git a/api/tests/snapshots/golden_tests__golden_pairing_response.snap b/api/tests/snapshots/golden_tests__golden_pairing_response.snap deleted file mode 100644 index 7bd30893..00000000 --- a/api/tests/snapshots/golden_tests__golden_pairing_response.snap +++ /dev/null @@ -1,6 +0,0 @@ ---- -source: api/tests/golden_tests.rs -assertion_line: 241 -expression: hex.clone() ---- -820aa60081020165322e342e30026641424331323303810104f5056e50617373706f7274205072696d65 diff --git a/api/tests/snapshots/golden_tests__golden_prime_magic_backup_enabled.snap b/api/tests/snapshots/golden_tests__golden_prime_magic_backup_enabled.snap deleted file mode 100644 index fdd857e8..00000000 --- a/api/tests/snapshots/golden_tests__golden_prime_magic_backup_enabled.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8213a200f50158204242424242424242424242424242424242424242424242424242424242424242 diff --git a/api/tests/snapshots/golden_tests__golden_prime_magic_backup_status_request.snap b/api/tests/snapshots/golden_tests__golden_prime_magic_backup_status_request.snap deleted file mode 100644 index 4085bf08..00000000 --- a/api/tests/snapshots/golden_tests__golden_prime_magic_backup_status_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8214a1005820abababababababababababababababababababababababababababababababab diff --git a/api/tests/snapshots/golden_tests__golden_prime_magic_backup_status_response.snap b/api/tests/snapshots/golden_tests__golden_prime_magic_backup_status_response.snap deleted file mode 100644 index 55f5557d..00000000 --- a/api/tests/snapshots/golden_tests__golden_prime_magic_backup_status_response.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8215a100f5 diff --git a/api/tests/snapshots/golden_tests__golden_raw_data.snap b/api/tests/snapshots/golden_tests__golden_raw_data.snap deleted file mode 100644 index 8c311b84..00000000 --- a/api/tests/snapshots/golden_tests__golden_raw_data.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -821864a10044feedface diff --git a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_chunk.snap b/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_chunk.snap deleted file mode 100644 index c71f6b94..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_chunk.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181d8202a3000a0118320244aabbccdd diff --git a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_error.snap b/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_error.snap deleted file mode 100644 index 2e846ecc..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181d8203a1006d4e6574776f726b206572726f72 diff --git a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_no_backup.snap b/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_no_backup.snap deleted file mode 100644 index 945baddc..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_no_backup.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181d8100 diff --git a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_starting.snap b/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_starting.snap deleted file mode 100644 index e645a697..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_event_starting.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181d8201a10018c8 diff --git a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_request.snap b/api/tests/snapshots/golden_tests__golden_restore_magic_backup_request.snap deleted file mode 100644 index abc49fbb..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181ca2005820bbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbbb011832 diff --git a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_result_error.snap b/api/tests/snapshots/golden_tests__golden_restore_magic_backup_result_error.snap deleted file mode 100644 index 1e47f70c..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_result_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181e8201a10071436865636b73756d206d69736d61746368 diff --git a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_result_success.snap b/api/tests/snapshots/golden_tests__golden_restore_magic_backup_result_success.snap deleted file mode 100644 index a2580d80..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_magic_backup_result_success.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -82181e8100 diff --git a/api/tests/snapshots/golden_tests__golden_restore_shard_request.snap b/api/tests/snapshots/golden_tests__golden_restore_shard_request.snap deleted file mode 100644 index 566178e0..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_shard_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -821818a1005820cdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcdcd diff --git a/api/tests/snapshots/golden_tests__golden_restore_shard_response_error.snap b/api/tests/snapshots/golden_tests__golden_restore_shard_response_error.snap deleted file mode 100644 index b9595b45..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_shard_response_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8218198201a100694e6f7420666f756e64 diff --git a/api/tests/snapshots/golden_tests__golden_restore_shard_response_not_found.snap b/api/tests/snapshots/golden_tests__golden_restore_shard_response_not_found.snap deleted file mode 100644 index 04619ea3..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_shard_response_not_found.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8218198102 diff --git a/api/tests/snapshots/golden_tests__golden_restore_shard_response_success.snap b/api/tests/snapshots/golden_tests__golden_restore_shard_response_success.snap deleted file mode 100644 index af0ffb04..00000000 --- a/api/tests/snapshots/golden_tests__golden_restore_shard_response_success.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -8218198200a100430a0b0c diff --git a/api/tests/snapshots/golden_tests__golden_security_check_challenge_request.snap b/api/tests/snapshots/golden_tests__golden_security_check_challenge_request.snap deleted file mode 100644 index d7b98fd7..00000000 --- a/api/tests/snapshots/golden_tests__golden_security_check_challenge_request.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820b8200a10044cafebabe diff --git a/api/tests/snapshots/golden_tests__golden_security_check_challenge_response_error.snap b/api/tests/snapshots/golden_tests__golden_security_check_challenge_response_error.snap deleted file mode 100644 index 5945e5ea..00000000 --- a/api/tests/snapshots/golden_tests__golden_security_check_challenge_response_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820b82018201a10071496e76616c6964207369676e6174757265 diff --git a/api/tests/snapshots/golden_tests__golden_security_check_challenge_response_success.snap b/api/tests/snapshots/golden_tests__golden_security_check_challenge_response_success.snap deleted file mode 100644 index 1df2e59e..00000000 --- a/api/tests/snapshots/golden_tests__golden_security_check_challenge_response_success.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820b82018200a10044deadbeef diff --git a/api/tests/snapshots/golden_tests__golden_security_check_verification_error.snap b/api/tests/snapshots/golden_tests__golden_security_check_verification_error.snap deleted file mode 100644 index 83b577a7..00000000 --- a/api/tests/snapshots/golden_tests__golden_security_check_verification_error.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820b82028201a10073566572696669636174696f6e206661696c6564 diff --git a/api/tests/snapshots/golden_tests__golden_security_check_verification_success.snap b/api/tests/snapshots/golden_tests__golden_security_check_verification_success.snap deleted file mode 100644 index 98d43c14..00000000 --- a/api/tests/snapshots/golden_tests__golden_security_check_verification_success.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820b82028100 diff --git a/api/tests/snapshots/golden_tests__golden_sign_psbt.snap b/api/tests/snapshots/golden_tests__golden_sign_psbt.snap deleted file mode 100644 index 5379861b..00000000 --- a/api/tests/snapshots/golden_tests__golden_sign_psbt.snap +++ /dev/null @@ -1,5 +0,0 @@ ---- -source: api/tests/golden_tests.rs -expression: hex.clone() ---- -820da200696163636f756e742d31014570736274ff diff --git a/quantum-link-macros/Cargo.toml b/quantum-link-macros/Cargo.toml deleted file mode 100644 index 6debf69c..00000000 --- a/quantum-link-macros/Cargo.toml +++ /dev/null @@ -1,14 +0,0 @@ -[package] -name = "quantum-link-macros" -version = "0.1.0" -edition = "2021" -homepage.workspace = true - -[dependencies] -#foundation-api = { workspace = true } -quote = "^1" -syn = { version = "^2.0.5", features = ["full", "extra-traits"] } -proc-macro2 = "1" - -[lib] -proc-macro = true \ No newline at end of file diff --git a/quantum-link-macros/src/lib.rs b/quantum-link-macros/src/lib.rs deleted file mode 100644 index 1552432f..00000000 --- a/quantum-link-macros/src/lib.rs +++ /dev/null @@ -1,632 +0,0 @@ -use proc_macro::TokenStream; -use proc_macro2::TokenStream as TokenStream2; -use quote::quote; -use syn::{ - parse_macro_input, spanned::Spanned, Attribute, Data, DataStruct, DeriveInput, Fields, Lit, - Meta, Type, Visibility, -}; - -#[proc_macro_attribute] -pub fn quantum_link(_metadata: TokenStream, input: TokenStream) -> TokenStream { - let input: DeriveInput = syn::parse(input).unwrap(); - - if let Err(e) = validate_visibility(&input) { - return e.to_compile_error().into(); - } - - let expanded = quote! { - #[derive(Clone, Debug, PartialEq, quantum_link_macros::Cbor)] - #[cfg_attr(feature = "keyos", derive(rkyv::Archive, rkyv::Serialize, rkyv::Deserialize))] - #[cfg_attr(feature = "envoy", flutter_rust_bridge::frb(non_opaque))] - #input - }; - - TokenStream::from(expanded) -} - -fn validate_visibility(input: &DeriveInput) -> syn::Result<()> { - if !matches!(input.vis, Visibility::Public(_)) { - return Err(syn::Error::new( - input.ident.span(), - "quantum link types must be public", - )); - } - if let Data::Struct(data_struct) = &input.data { - match &data_struct.fields { - Fields::Named(fields) => { - for field in &fields.named { - if !matches!(field.vis, Visibility::Public(_)) { - return Err(syn::Error::new( - field.ident.as_ref().unwrap().span(), - "fields in quantum link structs must be public", - )); - } - } - } - Fields::Unnamed(fields) => { - for field in fields.unnamed.iter() { - if !matches!(field.vis, Visibility::Public(_)) { - return Err(syn::Error::new( - field.span(), - "fields in quantum link structs must be public", - )); - } - } - } - Fields::Unit => {} - } - } - - Ok(()) -} - -/// derive macro generates -/// - From for CBOR -/// - TryFrom for T -#[proc_macro_derive(Cbor, attributes(n))] -pub fn derive_cbor(input: TokenStream) -> TokenStream { - let input = parse_macro_input!(input as DeriveInput); - - derive_cbor_impl(input) - .unwrap_or_else(|e| e.to_compile_error()) - .into() -} - -fn derive_cbor_impl(input: DeriveInput) -> syn::Result { - let name = &input.ident; - let generics = &input.generics; - let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); - - let (into_impl, try_from_impl) = match &input.data { - Data::Struct(data_struct) => generate_struct_impls(&data_struct.fields, name)?, - Data::Enum(data_enum) => { - let into_body = generate_enum_into_cbor(name, &data_enum.variants)?; - let try_from_body = generate_enum_try_from_cbor(&data_enum.variants)?; - (into_body, try_from_body) - } - Data::Union(data_union) => { - return Err(syn::Error::new( - data_union.union_token.span(), - "unions not supported", - )); - } - }; - - // only auto-impl for non-tuple structs - let cbor_marker_impl = match &input.data { - Data::Struct(DataStruct { - fields: Fields::Unnamed(_), - .. - }) => quote! {}, - _ => { - quote! { - impl #impl_generics crate::CborMarker for #name #ty_generics #where_clause {} - } - } - }; - - Ok(quote! { - impl #impl_generics From<#name #ty_generics> for dcbor::CBOR #where_clause { - fn from(value: #name #ty_generics) -> dcbor::CBOR { - #into_impl - } - } - - impl #impl_generics TryFrom for #name #ty_generics #where_clause { - type Error = dcbor::Error; - - fn try_from(cbor: dcbor::CBOR) -> dcbor::Result { - #try_from_impl - } - } - - #cbor_marker_impl - }) -} - -// -// struct -// - -fn generate_struct_impls( - fields: &Fields, - name: &syn::Ident, -) -> syn::Result<(TokenStream2, TokenStream2)> { - match fields { - Fields::Named(fields) => { - let into_body = generate_named_struct_into_cbor(&fields.named)?; - let try_from_body = generate_named_struct_try_from_cbor(&fields.named)?; - Ok((into_body, try_from_body)) - } - Fields::Unnamed(fields) => { - if fields.unnamed.len() != 1 { - return Err(syn::Error::new( - fields.span(), - "only single-field tuple structs (newtypes) are supported", - )); - } - let (into_body, try_from_body) = - generate_newtype_struct_impls(fields.unnamed.first().unwrap())?; - Ok((into_body, try_from_body)) - } - Fields::Unit => Err(syn::Error::new(name.span(), "unit structs not supported")), - } -} - -fn generate_named_struct_into_cbor( - fields: &syn::punctuated::Punctuated, -) -> syn::Result { - check_duplicate_indices(fields)?; - - let mut field_insertions = Vec::new(); - - for field in fields { - let field_name = field.ident.as_ref().unwrap(); - let field_type = &field.ty; - let index = get_field_index(&field.attrs) - .ok_or_else(|| syn::Error::new(field.span(), "missing #[n(x)] attribute"))?; - - if let Some(inner) = get_option_inner(field_type) { - let cbor_value = gen_to_cbor(&inner, quote! { val }); - field_insertions.push(quote! { - if let Some(val) = value.#field_name { - map.insert(dcbor::CBOR::from(#index), #cbor_value); - } - }); - } else { - let insertion = gen_map_insert(index, field_type, quote! { value.#field_name }); - field_insertions.push(insertion); - } - } - - Ok(quote! { - let mut map = dcbor::Map::new(); - #(#field_insertions)* - dcbor::CBOR::from(map) - }) -} - -fn generate_named_struct_try_from_cbor( - fields: &syn::punctuated::Punctuated, -) -> syn::Result { - let mut field_extractions = Vec::new(); - let mut field_names = Vec::new(); - - for field in fields { - let field_name = field.ident.as_ref().unwrap(); - let field_type = &field.ty; - let index = get_field_index(&field.attrs) - .ok_or_else(|| syn::Error::new(field.span(), "missing #[n(x)] attribute"))?; - - let extraction = if let Some(inner) = get_option_inner(field_type) { - let value = gen_map_get_optional(index, &inner, quote! { map }); - quote! { let #field_name: #field_type = #value; } - } else { - let value = gen_map_get_required(index, field_type, quote! { map }); - quote! { let #field_name: #field_type = #value; } - }; - - field_extractions.push(extraction); - field_names.push(field_name); - } - - Ok(quote! { - let case = cbor.into_case(); - let dcbor::CBORCase::Map(map) = case else { - return Err(dcbor::Error::WrongType); - }; - - #(#field_extractions)* - - Ok(Self { - #(#field_names),* - }) - }) -} - -fn generate_newtype_struct_impls(field: &syn::Field) -> syn::Result<(TokenStream2, TokenStream2)> { - let field_type = &field.ty; - - if get_field_index(&field.attrs).is_some() { - return Err(syn::Error::new( - field.span(), - "newtype structs cannot have #[n(x)] attribute; use a named struct instead", - )); - } - - let into_body = gen_to_cbor(field_type, quote! { value.0 }); - let from_value = gen_from_cbor(field_type, quote! { cbor }); - let try_from_body = quote! { Ok(Self(#from_value)) }; - - Ok((into_body, try_from_body)) -} - -// -// enum -// - -fn generate_enum_into_cbor( - enum_name: &syn::Ident, - variants: &syn::punctuated::Punctuated, -) -> syn::Result { - check_duplicate_indices(variants)?; - - let mut variant_arms = Vec::new(); - - for variant in variants { - let variant_name = &variant.ident; - let variant_index = get_field_index(&variant.attrs) - .ok_or_else(|| syn::Error::new(variant.span(), "missing #[n(x)] attribute"))?; - - let arm = match &variant.fields { - Fields::Unit => { - quote! { - #enum_name::#variant_name => { - dcbor::CBOR::from(vec![dcbor::CBOR::from(#variant_index)]) - } - } - } - Fields::Unnamed(fields) => { - generate_tuple_variant_into_cbor(enum_name, variant_name, variant_index, fields)? - } - Fields::Named(fields) => { - generate_struct_variant_into_cbor(enum_name, variant_name, variant_index, fields)? - } - }; - - variant_arms.push(arm); - } - - Ok(quote! { - match value { - #(#variant_arms)* - } - }) -} - -fn generate_tuple_variant_into_cbor( - enum_name: &syn::Ident, - variant_name: &syn::Ident, - variant_index: u64, - fields: &syn::FieldsUnnamed, -) -> syn::Result { - if fields.unnamed.len() != 1 { - return Err(syn::Error::new( - fields.span(), - "tuple variants must have exactly one field", - )); - } - - let field = fields.unnamed.first().unwrap(); - let field_type = &field.ty; - - if get_field_index(&field.attrs).is_some() { - return Err(syn::Error::new( - field.span(), - "tuple variant fields cannot have #[n(x)] attribute; use a struct variant instead", - )); - } - - Ok(quote! { - #enum_name::#variant_name(inner) => { - const _: fn() = || { - fn assert_cbor_marker() {} - assert_cbor_marker::<#field_type>(); - }; - dcbor::CBOR::from(vec![ - dcbor::CBOR::from(#variant_index), - dcbor::CBOR::from(inner), - ]) - } - }) -} - -fn generate_struct_variant_into_cbor( - enum_name: &syn::Ident, - variant_name: &syn::Ident, - variant_index: u64, - fields: &syn::FieldsNamed, -) -> syn::Result { - check_duplicate_indices(&fields.named)?; - - let mut field_names = Vec::new(); - let mut field_insertions = Vec::new(); - - for field in &fields.named { - let field_name = field.ident.as_ref().unwrap(); - let field_type = &field.ty; - let field_index = get_field_index(&field.attrs) - .ok_or_else(|| syn::Error::new(field.span(), "missing #[n(x)] attribute"))?; - - field_names.push(field_name); - - let cbor_value = gen_to_cbor(field_type, quote! { #field_name }); - field_insertions.push(quote! { - inner_map.insert(dcbor::CBOR::from(#field_index), #cbor_value); - }); - } - - Ok(quote! { - #enum_name::#variant_name { #(#field_names),* } => { - let mut inner_map = dcbor::Map::new(); - #(#field_insertions)* - - dcbor::CBOR::from(vec![ - dcbor::CBOR::from(#variant_index), - dcbor::CBOR::from(inner_map), - ]) - } - }) -} - -fn generate_enum_try_from_cbor( - variants: &syn::punctuated::Punctuated, -) -> syn::Result { - let mut variant_arms = Vec::new(); - - for variant in variants { - let variant_name = &variant.ident; - let variant_index = get_field_index(&variant.attrs) - .ok_or_else(|| syn::Error::new(variant.span(), "missing #[n(x)] attribute"))?; - - let arm = match &variant.fields { - Fields::Unit => { - quote! { #variant_index => Ok(Self::#variant_name), } - } - Fields::Unnamed(fields) => { - generate_tuple_variant_try_from_cbor(variant_name, variant_index, fields)? - } - Fields::Named(fields) => { - generate_struct_variant_try_from_cbor(variant_name, variant_index, fields)? - } - }; - - variant_arms.push(arm); - } - - Ok(quote! { - let case = cbor.into_case(); - let dcbor::CBORCase::Array(arr) = case else { - return Err(dcbor::Error::WrongType); - }; - - let variant_index: u64 = >::try_from( - arr.get(0).ok_or(dcbor::Error::WrongType)?.clone() - )?; - - match variant_index { - #(#variant_arms)* - _ => Err(dcbor::Error::WrongType), - } - }) -} - -fn generate_tuple_variant_try_from_cbor( - variant_name: &syn::Ident, - variant_index: u64, - fields: &syn::FieldsUnnamed, -) -> syn::Result { - if fields.unnamed.len() != 1 { - return Err(syn::Error::new( - fields.span(), - "tuple variants must have exactly one field", - )); - } - - let field = fields.unnamed.first().unwrap(); - let field_type = &field.ty; - - if get_field_index(&field.attrs).is_some() { - return Err(syn::Error::new( - field.span(), - "tuple variant fields cannot have #[n(x)] attribute; use a struct variant instead", - )); - } - - Ok(quote! { - #variant_index => { - let variant_data = arr.get(1).ok_or(dcbor::Error::WrongType)?; - let inner: #field_type = variant_data.clone().try_into()?; - Ok(Self::#variant_name(inner)) - } - }) -} - -fn generate_struct_variant_try_from_cbor( - variant_name: &syn::Ident, - variant_index: u64, - fields: &syn::FieldsNamed, -) -> syn::Result { - let mut field_extractions = Vec::new(); - let mut field_names = Vec::new(); - - for field in &fields.named { - let field_name = field.ident.as_ref().unwrap(); - let field_type = &field.ty; - let field_index = get_field_index(&field.attrs) - .ok_or_else(|| syn::Error::new(field.span(), "missing #[n(x)] attribute"))?; - - let extraction = if let Some(inner) = get_option_inner(field_type) { - let value = gen_map_get_optional(field_index, &inner, quote! { inner_map }); - quote! { let #field_name: #field_type = #value; } - } else { - let value = gen_map_get_required(field_index, field_type, quote! { inner_map }); - quote! { let #field_name: #field_type = #value; } - }; - - field_extractions.push(extraction); - field_names.push(field_name); - } - - Ok(quote! { - #variant_index => { - let variant_data = arr.get(1).ok_or(dcbor::Error::WrongType)?; - let inner_case = variant_data.clone().into_case(); - let dcbor::CBORCase::Map(inner_map) = inner_case else { - return Err(dcbor::Error::WrongType); - }; - - #(#field_extractions)* - - Ok(Self::#variant_name { - #(#field_names),* - }) - } - }) -} - -// -// helpers -// - -fn gen_to_cbor(field_type: &Type, value: TokenStream2) -> TokenStream2 { - if is_vec_u8(field_type) || is_u8_array(field_type) { - quote! { dcbor::CBOR::to_byte_string(#value) } - } else { - quote! { dcbor::CBOR::from(#value) } - } -} - -fn gen_from_cbor(field_type: &Type, cbor: TokenStream2) -> TokenStream2 { - if is_vec_u8(field_type) { - quote! { #cbor.try_into_byte_string()?.to_vec() } - } else if is_u8_array(field_type) { - gen_byte_array_from_cbor(field_type, cbor) - } else { - quote! { #cbor.try_into()? } - } -} - -fn gen_byte_array_from_cbor(field_type: &Type, cbor: TokenStream2) -> TokenStream2 { - quote! {{ - let bytes = #cbor.try_into_byte_string()?; - <#field_type>::try_from(bytes.as_ref()) - .map_err(|_| dcbor::Error::OutOfRange)? - }} -} - -fn gen_map_insert(index: u64, field_type: &Type, value: TokenStream2) -> TokenStream2 { - let cbor_value = gen_to_cbor(field_type, value); - quote! { - map.insert(dcbor::CBOR::from(#index), #cbor_value); - } -} - -fn gen_map_get_required(index: u64, field_type: &Type, map: TokenStream2) -> TokenStream2 { - let cbor_expr = quote! { - #map.get::(#index) - .ok_or(dcbor::Error::MissingMapKey)? - }; - gen_from_cbor(field_type, cbor_expr) -} - -fn gen_map_get_optional(index: u64, inner_type: &Type, map: TokenStream2) -> TokenStream2 { - let value_expr = if is_vec_u8(inner_type) { - quote! { field_cbor.try_into_byte_string()?.to_vec() } - } else if is_u8_array(inner_type) { - gen_byte_array_from_cbor(inner_type, quote! { field_cbor }) - } else { - quote! { field_cbor.try_into()? } - }; - - quote! { - match #map.get::(#index) { - Some(field_cbor) => Some(#value_expr), - None => None, - } - } -} - -fn get_field_index(attrs: &[Attribute]) -> Option { - for attr in attrs { - if attr.path().is_ident("n") { - if let Meta::List(meta_list) = &attr.meta { - let tokens = meta_list.tokens.clone(); - if let Ok(Lit::Int(lit_int)) = syn::parse2::(tokens) { - return lit_int.base10_parse().ok(); - } - } - } - } - None -} - -trait Indexed: Spanned { - fn index_attrs(&self) -> &[Attribute]; -} - -impl Indexed for syn::Field { - fn index_attrs(&self) -> &[Attribute] { - &self.attrs - } -} - -impl Indexed for syn::Variant { - fn index_attrs(&self) -> &[Attribute] { - &self.attrs - } -} - -fn check_duplicate_indices<'a>( - items: impl IntoIterator, -) -> syn::Result<()> { - let mut seen: std::collections::HashMap = - std::collections::HashMap::new(); - - for item in items { - if let Some(index) = get_field_index(item.index_attrs()) { - if let Some(&prev_span) = seen.get(&index) { - let mut err = - syn::Error::new(item.span(), format!("duplicate #[n({index})] attribute")); - err.combine(syn::Error::new(prev_span, "first use of this index")); - return Err(err); - } - seen.insert(index, item.span()); - } - } - - Ok(()) -} - -fn get_option_inner(ty: &Type) -> Option { - let Type::Path(p) = ty else { return None }; - let seg = p.path.segments.last().filter(|s| s.ident == "Option")?; - let syn::PathArguments::AngleBracketed(args) = &seg.arguments else { - return None; - }; - let syn::GenericArgument::Type(inner) = args.args.first()? else { - return None; - }; - Some(inner.clone()) -} - -fn is_vec_u8(ty: &Type) -> bool { - let Type::Path(p) = ty else { return false }; - let Some(seg) = p.path.segments.last().filter(|s| s.ident == "Vec") else { - return false; - }; - let syn::PathArguments::AngleBracketed(args) = &seg.arguments else { - return false; - }; - let Some(syn::GenericArgument::Type(Type::Path(inner))) = args.args.first() else { - return false; - }; - inner - .path - .segments - .last() - .map(|s| s.ident == "u8") - .unwrap_or(false) -} - -fn is_u8_array(ty: &Type) -> bool { - let Type::Array(array) = ty else { return false }; - let Type::Path(p) = &*array.elem else { - return false; - }; - p.path - .segments - .last() - .map(|s| s.ident == "u8") - .unwrap_or(false) -} From 6f95798b864fe255a00fc39b79fb11a5188d75de Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 22 Jun 2026 11:20:58 -0400 Subject: [PATCH 07/59] ql: open stream service id --- ql-fsm/src/fsm.rs | 21 ++++------ ql-fsm/src/lib.rs | 25 ++++++++---- ql-fsm/src/session/mod.rs | 54 ++++++++++++------------ ql-fsm/src/session/state.rs | 8 ++-- ql-fsm/src/session/stream_ops.rs | 15 +++++-- ql-fsm/src/session/tests.rs | 45 +++++++++++++------- ql-fsm/src/tests/proptest.rs | 15 ++++--- ql-fsm/src/tests/session.rs | 61 +++++++++++++++++----------- ql-runtime/src/command.rs | 4 +- ql-runtime/src/driver/mod.rs | 35 ++++++++++------ ql-runtime/src/handle/mod.rs | 18 +++++--- ql-runtime/src/tests/handshake.rs | 7 +++- ql-runtime/src/tests/mod.rs | 11 ++++- ql-runtime/src/tests/session.rs | 16 +++++--- ql-runtime/src/tests/stream.rs | 22 +++++----- ql-wire/src/encrypted/mod.rs | 2 + ql-wire/src/encrypted/service_id.rs | 34 ++++++++++++++++ ql-wire/src/encrypted/stream_data.rs | 9 ++-- 18 files changed, 261 insertions(+), 141 deletions(-) create mode 100644 ql-wire/src/encrypted/service_id.rs diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index 036a336e..75adbb6b 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -1,13 +1,14 @@ use std::{collections::VecDeque, time::Instant}; use bytes::Bytes; -use ql_wire::{self as wire, QlCrypto, RouteId, SessionCloseCode, StreamId, WireDecode}; +use ql_wire::{self as wire, QlCrypto, SessionCloseCode, StreamId, WireDecode}; use crate::{ handshake, session::{self, SessionEvent, TerminalFrame}, state::LinkState, - Event, NoPeerError, NoSessionError, OutboundWrite, QlFsm, ReceiveError, StreamError, WriteId, + Event, NoPeerError, NoSessionError, OpenStreamParams, OutboundWrite, QlFsm, ReceiveError, + StreamError, WriteId, }; pub struct EventSink<'a> { @@ -30,14 +31,8 @@ impl session::EventSink for EventSink<'_> { SessionEvent::Unpaired => { self.termination = Some(TerminalFrame::Unpair); } - SessionEvent::Opened { - stream_id, - route_id, - } => { - self.events.push_back(Event::Opened { - stream_id, - route_id, - }); + SessionEvent::Opened(stream_id) => { + self.events.push_back(Event::Opened(stream_id)); } SessionEvent::Readable(stream_id) => { self.events.push_back(Event::Readable(stream_id)); @@ -238,11 +233,13 @@ pub fn close_session(fsm: &mut QlFsm, code: SessionCloseCode) { pub fn open_stream( fsm: &mut QlFsm, - route_id: RouteId, + params: OpenStreamParams, ) -> Result, NoSessionError> { let QlFsm { state, events, .. } = fsm; let conn = state.link.connected_mut_or_err()?; - let inner = conn.session.open_stream(route_id, EventSink::new(events))?; + let inner = + conn.session + .open_stream(params.service_id, params.route_id, EventSink::new(events))?; Ok(crate::StreamOps { inner }) } diff --git a/ql-fsm/src/lib.rs b/ql-fsm/src/lib.rs index 3067efdb..0ff2c68d 100644 --- a/ql-fsm/src/lib.rs +++ b/ql-fsm/src/lib.rs @@ -36,8 +36,8 @@ pub use bytes::Bytes; pub use error::*; pub use pairing::PairingInvite; use ql_wire::{ - PairingToken, PeerBundle, QlCrypto, QlIdentity, RouteId, SessionClose, SessionCloseCode, - StreamClose, StreamId, + PairingToken, PeerBundle, QlCrypto, QlIdentity, RouteId, ServiceId, SessionClose, + SessionCloseCode, StreamClose, StreamHeader, StreamId, }; pub use session::{SessionEvent, StreamReadIter, StreamWriter}; @@ -67,10 +67,7 @@ pub enum Event { /// the peer changed lifecycle state PeerStatusChanged(PeerStatus), /// a stream was opened - Opened { - stream_id: StreamId, - route_id: RouteId, - }, + Opened(StreamId), /// a stream has bytes ready to read Readable(StreamId), /// a stream has room for more local writes @@ -114,6 +111,10 @@ impl StreamOps<'_> { self.inner.stream_id() } + pub fn header(&self) -> &StreamHeader { + self.inner.header() + } + /// returns the readable stream bytes as owned `Bytes` views without consuming them pub fn read(&self) -> StreamReadIter<'_> { self.inner.read() @@ -183,6 +184,11 @@ impl Default for QlFsmConfig { } } +pub struct OpenStreamParams { + pub service_id: ServiceId, + pub route_id: RouteId, +} + /// synchronous driver for peer binding, handshake, and encrypted streams pub struct QlFsm { config: QlFsmConfig, @@ -318,8 +324,11 @@ impl QlFsm { } /// opens a new outgoing stream - pub fn open_stream(&mut self, route_id: RouteId) -> Result, NoSessionError> { - fsm::open_stream(self, route_id) + pub fn open_stream( + &mut self, + params: OpenStreamParams, + ) -> Result, NoSessionError> { + fsm::open_stream(self, params) } /// returns a facade for an open stream diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index 55187757..b3821076 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -18,9 +18,9 @@ use std::time::{Duration, Instant}; use bytes::Bytes; use indexmap::IndexMap; use ql_wire::{ - CloseTarget, RecordAck, RecordSeq, RouteId, SessionClose, SessionCloseCode, SessionFrame, - SessionRecordBuilder, StreamClose, StreamData, StreamHeader, StreamId, StreamWindow, VarInt, - WireError, + CloseTarget, RecordAck, RecordSeq, RouteId, ServiceId, SessionClose, SessionCloseCode, + SessionFrame, SessionRecordBuilder, StreamClose, StreamData, StreamHeader, StreamId, + StreamWindow, VarInt, WireError, }; use self::{ @@ -67,10 +67,7 @@ impl Default for SessionConfig { #[derive(Debug, Clone, PartialEq, Eq)] pub enum SessionEvent { - Opened { - stream_id: StreamId, - route_id: RouteId, - }, + Opened(StreamId), Readable(StreamId), Writable(StreamId), Finished(StreamId), @@ -132,6 +129,7 @@ impl SessionFsm { pub fn open_stream( &mut self, + service_id: ServiceId, route_id: RouteId, sink: E, ) -> Result, NoSessionError> @@ -148,7 +146,10 @@ impl SessionFsm { stream_id, StreamState::new( StreamRole::Initiator, - Some(route_id), + Some(StreamHeader { + service_id, + route_id, + }), self.config.stream_receive_buffer_size, self.config.initial_peer_stream_receive_window, ), @@ -166,10 +167,15 @@ impl SessionFsm { E: EventSink, { self.ensure_session_open()?; - let Some(stream_index) = self.state.streams.get_index_of(&stream_id) else { + let Some(stream_index) = (|| { + let index = self.state.streams.get_index_of(&stream_id)?; + // Event::Opened only fires after we receive the first frame of a stream + // prevent early access to streams + let _ = self.state.streams[index].header.as_ref()?; + Some(index) + })() else { return Err(StreamError::MissingStream); }; - Ok(StreamOps::new(self, stream_id, stream_index, sink)) } @@ -538,7 +544,7 @@ impl SessionFsm { stream_id, offset, header: if matches!(stream.role, StreamRole::Initiator) && candidate.offset == 0 { - stream.route_id.map(|route_id| StreamHeader { route_id }) + stream.header } else { None }, @@ -654,15 +660,15 @@ impl SessionFsm { let readable_before = stream.readable_bytes(); let was_finished = matches!(stream.inbound_state, InboundState::Finished); - let opened_route = match (stream.role, stream.route_id, header, frame_offset) { + let opened = match (stream.role, stream.header, header, frame_offset) { (StreamRole::Responder, None, Some(header), 0) => { - stream.route_id = Some(header.route_id); - Some(header.route_id) + stream.header = Some(header); + true } (StreamRole::Initiator, _, Some(_), _) | (StreamRole::Responder, None, Some(_), _) | (StreamRole::Responder, None, None, 0) => return Err(()), - _ => None, + _ => false, }; match stream.inbound_state { @@ -678,11 +684,8 @@ impl SessionFsm { // retransmitted data for an already-finished stream is fine as long as it stays // within the finalized byte range and any repeated FIN lands on that same offset. if (!frame.fin || frame_end == final_offset) && frame_end <= final_offset { - if let Some(route_id) = opened_route { - sink.emit(SessionEvent::Opened { - stream_id, - route_id, - }); + if opened { + sink.emit(SessionEvent::Opened(stream_id)); if readable_before > 0 { sink.emit(SessionEvent::Readable(stream_id)); } else { @@ -702,18 +705,15 @@ impl SessionFsm { stream.inbound_state = InboundState::Finished; } - if let Some(route_id) = opened_route { - sink.emit(SessionEvent::Opened { - stream_id, - route_id, - }); + if opened { + sink.emit(SessionEvent::Opened(stream_id)); } - if stream.route_id.is_some() && readable_before == 0 && stream.readable_bytes() > 0 { + if stream.header.is_some() && readable_before == 0 && stream.readable_bytes() > 0 { sink.emit(SessionEvent::Readable(stream_id)); } - if stream.route_id.is_some() + if stream.header.is_some() && !was_finished && matches!(stream.inbound_state, InboundState::Finished) && stream.readable_bytes() == 0 diff --git a/ql-fsm/src/session/state.rs b/ql-fsm/src/session/state.rs index b63140a1..01fac93a 100644 --- a/ql-fsm/src/session/state.rs +++ b/ql-fsm/src/session/state.rs @@ -1,7 +1,7 @@ use std::time::Instant; use indexmap::IndexMap; -use ql_wire::{CloseTarget, RecordSeq, RouteId, SessionClose, StreamClose, StreamId}; +use ql_wire::{CloseTarget, RecordSeq, SessionClose, StreamClose, StreamHeader, StreamId}; use super::{ ack_tracker::AckTracker, remote_stream_history::RemoteStreamHistory, stream_rx::StreamRx, @@ -45,7 +45,7 @@ pub enum TerminalFrame { #[derive(Debug)] pub struct StreamState { pub role: StreamRole, - pub route_id: Option, + pub header: Option, pub rx: StreamRx, pub tx: StreamTx, pub pending_close: Option, @@ -59,14 +59,14 @@ pub struct StreamState { impl StreamState { pub fn new( role: StreamRole, - route_id: Option, + route_id: Option, receive_buffer_size: u32, initial_peer_stream_receive_window: u32, ) -> Self { let receive_buffer_size = receive_buffer_size as usize; Self { role, - route_id, + header: route_id, tx: StreamTx::new(), pending_close: None, peer_max_offset: u64::from(initial_peer_stream_receive_window), diff --git a/ql-fsm/src/session/stream_ops.rs b/ql-fsm/src/session/stream_ops.rs index 548189b7..9af8d439 100644 --- a/ql-fsm/src/session/stream_ops.rs +++ b/ql-fsm/src/session/stream_ops.rs @@ -1,4 +1,4 @@ -use ql_wire::{CloseTarget, StreamClose, StreamCloseCode, StreamId}; +use ql_wire::{CloseTarget, StreamClose, StreamCloseCode, StreamHeader, StreamId}; use super::{ state::{InboundState, StreamState}, @@ -31,11 +31,18 @@ impl<'a, E: EventSink> StreamOps<'a, E> { } } - /// returns this stream's identifier + /// returns this stream's identifier + #[inline] pub fn stream_id(&self) -> StreamId { self.stream_id } + /// returns the streams details + #[inline] + pub fn header(&self) -> &StreamHeader { + self.stream().header.as_ref().unwrap() + } + /// returns the readable stream bytes as owned `Bytes` views without consuming them pub fn read(&self) -> StreamReadIter<'_> { self.stream().rx.bytes() @@ -58,7 +65,7 @@ impl<'a, E: EventSink> StreamOps<'a, E> { if stream.recv_limit() > stream.advertised_max_offset { stream.pending_window = true; } - stream.route_id.is_some() + stream.header.is_some() && matches!(stream.inbound_state, InboundState::Finished) && stream.readable_bytes() == 0 }; @@ -92,10 +99,12 @@ impl<'a, E: EventSink> StreamOps<'a, E> { self.reap_on_drop = true; } + #[inline] fn stream(&self) -> &StreamState { &self.session.state.streams[self.stream_index] } + #[inline] fn stream_mut(&mut self) -> &mut StreamState { &mut self.session.state.streams[self.stream_index] } diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index f1f29879..bfcbfb18 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -3,8 +3,8 @@ use std::time::{Duration, Instant}; use bytes::Bytes; use ql_wire::{ decode_session_frames, parse_session_frames, CloseTarget, RecordAck, RecordSeq, RouteId, - SessionFrame, SessionRecordBuilder, StreamClose, StreamCloseCode, StreamData, StreamHeader, - StreamId, VarInt, QID, + ServiceId, SessionFrame, SessionRecordBuilder, StreamClose, StreamCloseCode, StreamData, + StreamHeader, StreamId, VarInt, QID, }; use super::{SessionConfig, SessionEvent, SessionFsm}; @@ -36,18 +36,19 @@ const TIMEOUT: StreamCloseCode = StreamCloseCode(2); fn header(value: u64) -> StreamHeader { StreamHeader { route_id: route_id(value), + service_id: ServiceId([0; 16]), } } +// todo: remove fn opened(stream_id: StreamId) -> SessionEvent { - SessionEvent::Opened { - stream_id, - route_id: route_id(1), - } + SessionEvent::Opened(stream_id) } fn open_stream_id(fsm: &mut SessionFsm) -> StreamId { - fsm.open_stream(route_id(1), |_| {}).unwrap().stream_id() + fsm.open_stream(ServiceId([0; 16]), route_id(1), |_| {}) + .unwrap() + .stream_id() } fn write_stream_bytes(fsm: &mut SessionFsm, stream_id: StreamId, bytes: &[u8]) -> usize { @@ -158,21 +159,32 @@ fn retransmit_uses_new_record_seq() { #[test] fn lost_record_on_one_stream_does_not_block_another_stream() { + const PAYLOAD_LEN: usize = 40; + let now = Instant::now(); let mut fsm = SessionFsm::new( SessionConfig { - record_max_size: 80 + SessionRecordBuilder::MIN_CAPACITY, + record_max_size: SessionRecordBuilder::MIN_CAPACITY + + 1 // discriminator byte + + StreamData::>::MIN_WIRE_SIZE + + PAYLOAD_LEN, ..SessionConfig::default() }, now, ); let stream_id_a = open_stream_id(&mut fsm); let stream_id_b = open_stream_id(&mut fsm); - let payload_a = vec![b'a'; 40]; - let payload_b = vec![b'b'; 40]; + let payload_a = vec![b'a'; PAYLOAD_LEN]; + let payload_b = vec![b'b'; PAYLOAD_LEN]; - assert_eq!(write_stream_bytes(&mut fsm, stream_id_a, &payload_a), 40); - assert_eq!(write_stream_bytes(&mut fsm, stream_id_b, &payload_b), 40); + assert_eq!( + write_stream_bytes(&mut fsm, stream_id_a, &payload_a), + PAYLOAD_LEN + ); + assert_eq!( + write_stream_bytes(&mut fsm, stream_id_b, &payload_b), + PAYLOAD_LEN + ); let (first_seq, first) = next_outbound(&mut fsm, now).unwrap(); let (second_seq, _second) = next_outbound(&mut fsm, now + Duration::from_millis(1)).unwrap(); @@ -439,7 +451,7 @@ fn stream_ids_follow_even_odd_xid_ordering() { }, now, ) - .open_stream(route_id(1), |_| {}) + .open_stream(ServiceId([0; 16]), route_id(1), |_| {}) .unwrap() .stream_id(); let odd_id = SessionFsm::new( @@ -449,7 +461,7 @@ fn stream_ids_follow_even_odd_xid_ordering() { }, now, ) - .open_stream(route_id(1), |_| {}) + .open_stream(ServiceId([0; 16]), route_id(1), |_| {}) .unwrap() .stream_id(); @@ -791,7 +803,10 @@ fn sparse_out_of_order_ack_ranges_page_and_quiesce() { let now = Instant::now(); let sender_config = SessionConfig { local_parity: StreamParity::Even, - record_max_size: SessionRecordBuilder::MIN_CAPACITY + 40, + record_max_size: SessionRecordBuilder::MIN_CAPACITY + + 1 // discriminator byte + + StreamData::>::MIN_WIRE_SIZE + + 10, // keeps stream-data records tiny enough to force ACK paging ack_delay: Duration::from_millis(5), retransmit_timeout: Duration::from_millis(25), stream_send_buffer_size: 8 * 1024, diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs index bc97ca77..50547325 100644 --- a/ql-fsm/src/tests/proptest.rs +++ b/ql-fsm/src/tests/proptest.rs @@ -7,14 +7,10 @@ extern crate proptest as proptest_crate; use bytes::Bytes; use proptest_crate::{collection::vec, prelude::*, test_runner::TestCaseResult}; -use ql_wire::{CloseTarget, StreamCloseCode, StreamId, WireError}; +use ql_wire::{CloseTarget, RouteId, ServiceId, StreamCloseCode, StreamId, WireError}; use super::*; - -fn test_route_id() -> ql_wire::RouteId { - ql_wire::RouteId::from_u32(1) -} -use crate::{state::LinkState, Event, PeerStatus, ReceiveError, WriteId}; +use crate::{state::LinkState, Event, OpenStreamParams, PeerStatus, ReceiveError, WriteId}; const SLOT_COUNT: usize = 4; @@ -285,7 +281,10 @@ impl Runner { .harness .node_mut(*side) .fsm - .open_stream(test_route_id()) + .open_stream(OpenStreamParams { + service_id: ServiceId([1; 16]), + route_id: RouteId::from(1u32), + }) .ok() .map(|stream| stream.stream_id()); if let Some(stream_id) = stream_id { @@ -424,7 +423,7 @@ impl Runner { } self.events[side.idx()].note_peer_status(status); } - Event::Opened { stream_id, .. } => { + Event::Opened(stream_id) => { prop_assert!( self.known_streams.contains(&stream_id), "side {side:?} emitted Opened for unknown stream {stream_id:?}" diff --git a/ql-fsm/src/tests/session.rs b/ql-fsm/src/tests/session.rs index c55e51c1..cdedd1f4 100644 --- a/ql-fsm/src/tests/session.rs +++ b/ql-fsm/src/tests/session.rs @@ -1,28 +1,25 @@ use std::time::Duration; use bytes::Bytes; -use ql_wire::{RouteId, SessionClose, StreamId, VarInt}; +use ql_wire::{RouteId, ServiceId, SessionClose, StreamId, VarInt}; use super::*; -use crate::{state::LinkState, CommitReadError, Event, NoSessionError, PeerStatus, StreamError}; +use crate::{ + state::LinkState, CommitReadError, Event, NoSessionError, OpenStreamParams, PeerStatus, + StreamError, +}; fn stream_id(value: u32) -> StreamId { StreamId(VarInt::from_u32(value)) } -fn route_id(value: u32) -> RouteId { - RouteId::from_u32(value) -} - -fn opened(stream_id: StreamId) -> Event { - Event::Opened { - stream_id, - route_id: route_id(1), - } -} - fn open_stream_id(fsm: &mut QlFsm) -> StreamId { - fsm.open_stream(route_id(1)).unwrap().stream_id() + fsm.open_stream(OpenStreamParams { + service_id: ServiceId([1; 16]), + route_id: RouteId::from(1u32), + }) + .unwrap() + .stream_id() } fn write_stream_bytes( @@ -77,7 +74,7 @@ fn connected_fsms_deliver_stream_data() { harness.pump(); - assert_eq!(harness.take_event(Side::B), Some(opened(stream_id))); + assert_eq!(harness.take_event(Side::B), Some(Event::Opened(stream_id))); assert_eq!( harness.take_event(Side::B), Some(Event::Readable(stream_id)) @@ -126,7 +123,7 @@ fn session_retransmit_uses_new_record_seq() { harness.on_timer(Side::B); harness.pump(); - assert_eq!(harness.take_event(Side::B), Some(opened(stream_id))); + assert_eq!(harness.take_event(Side::B), Some(Event::Opened(stream_id))); assert_eq!( harness.take_event(Side::B), Some(Event::Readable(stream_id)) @@ -169,7 +166,10 @@ fn simultaneous_opens_use_even_and_odd_stream_ids() { harness.pump(); - assert_eq!(harness.take_event(Side::A), Some(opened(stream_id_b))); + assert_eq!( + harness.take_event(Side::A), + Some(Event::Opened(stream_id_b)) + ); assert_eq!( harness.take_event(Side::A), Some(Event::Readable(stream_id_b)) @@ -178,7 +178,10 @@ fn simultaneous_opens_use_even_and_odd_stream_ids() { read_stream_all(&mut harness.a.fsm, stream_id_b), b"from-b".to_vec() ); - assert_eq!(harness.take_event(Side::B), Some(opened(stream_id_a))); + assert_eq!( + harness.take_event(Side::B), + Some(Event::Opened(stream_id_a)) + ); assert_eq!( harness.take_event(Side::B), Some(Event::Readable(stream_id_a)) @@ -195,7 +198,10 @@ fn disconnected_stream_operations_fail_with_no_session() { let missing = stream_id(0); assert!(matches!( - harness.a.fsm.open_stream(route_id(1)), + harness.a.fsm.open_stream(OpenStreamParams { + service_id: ServiceId([1; 16]), + route_id: RouteId::from(1u32), + }), Err(NoSessionError) )); assert_eq!( @@ -278,7 +284,7 @@ fn returned_session_write_is_reissued_with_new_record_seq() { harness.deliver(Side::B, reissued.record); harness.pump(); - assert_eq!(harness.take_event(Side::B), Some(opened(stream_id))); + assert_eq!(harness.take_event(Side::B), Some(Event::Opened(stream_id))); assert_eq!( harness.take_event(Side::B), Some(Event::Readable(stream_id)) @@ -365,7 +371,10 @@ fn close_session_disconnects_locally() { )); assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); assert!(matches!( - harness.a.fsm.open_stream(route_id(1)), + harness.a.fsm.open_stream(OpenStreamParams { + service_id: ServiceId([1; 16]), + route_id: RouteId::from(1u32), + }), Err(NoSessionError) )); assert_eq!(harness.a.fsm.queue_ping(), Err(NoSessionError)); @@ -395,7 +404,10 @@ fn unpair_clears_bound_peer_and_emits_unpair_frame() { ); assert!(harness.a.fsm.peer().is_none()); assert!(matches!( - harness.a.fsm.open_stream(route_id(1)), + harness.a.fsm.open_stream(OpenStreamParams { + service_id: ServiceId([1; 16]), + route_id: RouteId::from(1u32), + }), Err(NoSessionError) )); assert_eq!(harness.a.fsm.queue_ping(), Err(NoSessionError)); @@ -422,7 +434,10 @@ fn inbound_unpair_clears_remote_peer_binding() { ); assert!(harness.b.fsm.peer().is_none()); assert!(matches!( - harness.b.fsm.open_stream(route_id(1)), + harness.b.fsm.open_stream(OpenStreamParams { + service_id: ServiceId([1; 16]), + route_id: RouteId::from(1u32), + }), Err(NoSessionError) )); assert!(matches!(harness.connect_ik(Side::B), Err(NoPeerError))); diff --git a/ql-runtime/src/command.rs b/ql-runtime/src/command.rs index 4a47a45e..1d4857ca 100644 --- a/ql-runtime/src/command.rs +++ b/ql-runtime/src/command.rs @@ -1,6 +1,7 @@ use ql_fsm::{NoSessionError, PairingInvite}; use ql_wire::{ - CloseTarget, PairingToken, PeerBundle, RouteId, SessionCloseCode, StreamCloseCode, StreamId, + CloseTarget, PairingToken, PeerBundle, RouteId, ServiceId, SessionCloseCode, StreamCloseCode, + StreamId, }; use crate::{StreamReader, StreamWriter}; @@ -18,6 +19,7 @@ pub enum Command { invite: PairingInvite, }, OpenStream { + service_id: ServiceId, route_id: RouteId, start: oneshot::Sender>, }, diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 35de1bf0..2f7e4e11 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -16,7 +16,7 @@ use std::{ use async_channel::Recv; use futures_lite::future::{poll_fn, yield_now}; use ql_fsm::{Event, QlFsm, WriteId}; -use ql_wire::{CloseTarget, StreamCloseCode, StreamId}; +use ql_wire::{CloseTarget, StreamCloseCode, StreamHeader, StreamId}; use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; use crate::{ @@ -210,7 +210,11 @@ impl DriverState { log::info!("unpairing peer"); fsm.unpair(); } - Command::OpenStream { route_id, start } => { + Command::OpenStream { + service_id, + route_id, + start, + } => { log::info!("open stream requested: route_id={route_id}"); let Some(runtime_tx) = self.runtime_tx.upgrade() else { log::warn!("open stream aborted: runtime channel unavailable"); @@ -218,7 +222,10 @@ impl DriverState { return; }; - let mut stream_ops = match fsm.open_stream(route_id) { + let mut stream_ops = match fsm.open_stream(ql_fsm::OpenStreamParams { + service_id, + route_id, + }) { Ok(stream_ops) => stream_ops, Err(error) => { log::warn!("open stream failed: route_id={route_id}"); @@ -227,7 +234,7 @@ impl DriverState { } }; let stream_id = stream_ops.stream_id(); - log::info!("open stream allocated: route_id={route_id} stream_id={stream_id}"); + log::info!("open stream allocated: service_id={service_id} route_id={route_id} stream_id={stream_id}"); let (reader, writer, reader_io, writer_io) = io::new_stream( stream_id, CloseTarget::Return, @@ -314,12 +321,9 @@ impl DriverState { } platform.handle_peer_status(peer, status); } - Event::Opened { - stream_id, - route_id, - } => { - log::info!("inbound stream opened: stream_id={stream_id} route_id={route_id}"); - self.handle_opened_stream(fsm, platform, stream_id, route_id); + Event::Opened(stream_id) => { + log::info!("inbound stream opened: stream_id={stream_id}"); + self.handle_opened_stream(fsm, platform, stream_id); } Event::Readable(stream_id) => { log::trace!("stream readable: stream_id={stream_id}"); @@ -358,7 +362,6 @@ impl DriverState { fsm: &mut QlFsm, platform: &P, stream_id: StreamId, - route_id: ql_wire::RouteId, ) { let Some(runtime_tx) = self.runtime_tx.upgrade() else { log::warn!( @@ -386,12 +389,20 @@ impl DriverState { ), ); + let stream = fsm.stream(stream_id).unwrap(); + let StreamHeader { + service_id, + route_id, + } = *stream.header(); + log::info!( - "delivering inbound stream to platform: stream_id={stream_id} route_id={route_id}" + "delivering inbound stream to platform: service_id={service_id} route_id={route_id} stream_id={stream_id}", ); + platform.handle_inbound(QlStream { stream_id, route_id, + service_id, writer, reader, }); diff --git a/ql-runtime/src/handle/mod.rs b/ql-runtime/src/handle/mod.rs index 1782c17a..4275d277 100644 --- a/ql-runtime/src/handle/mod.rs +++ b/ql-runtime/src/handle/mod.rs @@ -1,13 +1,15 @@ -use ql_fsm::{NoSessionError, PairingInvite}; -use ql_wire::{PairingToken, PeerBundle, RouteId, SessionCloseCode, StreamId}; +use ql_fsm::{NoSessionError, OpenStreamParams, PairingInvite}; +use ql_wire::{PairingToken, PeerBundle, RouteId, ServiceId, SessionCloseCode, StreamId}; use crate::command::Command; pub use crate::io::{StreamReader, StreamWriter}; #[derive(Debug)] pub struct QlStream { - pub stream_id: StreamId, + pub service_id: ServiceId, pub route_id: RouteId, + pub stream_id: StreamId, + pub writer: StreamWriter, pub reader: StreamReader, } @@ -54,11 +56,16 @@ impl RuntimeHandle { } /// opens a new stream on the active encrypted session - pub async fn open_stream(&self, route_id: RouteId) -> Result { + pub async fn open_stream(&self, params: OpenStreamParams) -> Result { let (start_tx, start_rx) = oneshot::channel(); + let OpenStreamParams { + service_id, + route_id, + } = params; self.send(Command::OpenStream { route_id, + service_id, start: start_tx, }); @@ -66,8 +73,9 @@ impl RuntimeHandle { let (stream_id, reader, writer) = start_rx.await.unwrap()?; Ok(QlStream { - stream_id, route_id, + service_id, + stream_id, writer, reader, }) diff --git a/ql-runtime/src/tests/handshake.rs b/ql-runtime/src/tests/handshake.rs index 65731bbc..e2e8b030 100644 --- a/ql-runtime/src/tests/handshake.rs +++ b/ql-runtime/src/tests/handshake.rs @@ -18,7 +18,10 @@ async fn opening_stream_requires_connection() { run_local_test(async { let pair = TestPair::new(default_runtime_config()); assert!(matches!( - pair.side(Side::A).handle.open_stream(test_route_id()).await, + pair.side(Side::A) + .handle + .open_stream(test_open_stream_params()) + .await, Err(NoSessionError) )); }) @@ -85,7 +88,7 @@ async fn rejected_session_write_is_reissued() { request }); - let mut stream = handle_a.open_stream(test_route_id()).await.unwrap(); + let mut stream = handle_a.open_stream(test_open_stream_params()).await.unwrap(); stream .writer .write(Bytes::from_static(b"retry")) diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index af368738..c3e6996b 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -11,11 +11,11 @@ use std::{ use async_channel::{Receiver, Sender}; use futures_lite::Stream; -use ql_fsm::PeerStatus; +use ql_fsm::{OpenStreamParams, PeerStatus}; use ql_wire::{ generate_identity, test_identities, MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, Nonce, PairingToken, PeerBundle, QlAead, QlHash, QlIdentity, QlKem, QlRandom, - RecordHeader, RecordType, RouteId, SessionKey, SoftwareCrypto, WireDecode, QID, + RecordHeader, RecordType, RouteId, ServiceId, SessionKey, SoftwareCrypto, WireDecode, QID, }; use tokio::{task::LocalSet, time::Sleep}; @@ -66,6 +66,13 @@ fn test_route_id() -> RouteId { RouteId::from_u32(1) } +fn test_open_stream_params() -> OpenStreamParams { + OpenStreamParams { + service_id: ServiceId([1; 16]), + route_id: test_route_id(), + } +} + #[derive(Debug, Clone)] struct WriteStats { active: Arc, diff --git a/ql-runtime/src/tests/session.rs b/ql-runtime/src/tests/session.rs index ec351e35..67d44b96 100644 --- a/ql-runtime/src/tests/session.rs +++ b/ql-runtime/src/tests/session.rs @@ -37,7 +37,7 @@ async fn close_session_aborts_active_streams_and_allows_reconnect() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); stream @@ -102,7 +102,7 @@ async fn unpair_aborts_active_streams_and_prevents_reconnect() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); stream @@ -126,11 +126,17 @@ async fn unpair_aborts_active_streams_and_prevents_reconnect() { .unwrap(); assert!(matches!( - pair.side(Side::A).handle.open_stream(test_route_id()).await, + pair.side(Side::A) + .handle + .open_stream(test_open_stream_params()) + .await, Err(NoSessionError) )); assert!(matches!( - pair.side(Side::B).handle.open_stream(test_route_id()).await, + pair.side(Side::B) + .handle + .open_stream(test_open_stream_params()) + .await, Err(NoSessionError) )); @@ -195,7 +201,7 @@ async fn session_timeout_disconnects_and_fails_pending_open() { drop_flag.store(true, Ordering::Relaxed); - let mut pending = handle_a.open_stream(test_route_id()).await.unwrap(); + let mut pending = handle_a.open_stream(test_open_stream_params()).await.unwrap(); let err = pending.writer.finish().await.unwrap_err(); assert!(matches!(err, QlStreamError::NoSession)); diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index 176711c8..72d8269b 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -30,7 +30,7 @@ async fn open_stream_duplex_happy_path() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); stream @@ -90,7 +90,7 @@ async fn reader_respects_max_len() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); stream @@ -128,7 +128,7 @@ async fn large_stream_payload_round_trips() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); stream @@ -168,7 +168,7 @@ async fn dropping_responder_closes_initiator_response() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); let err = stream.writer.finish().await.unwrap_err(); @@ -220,7 +220,7 @@ async fn dropping_inbound_reader_cancels_remote_writer() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); stream.writer.finish().await.unwrap(); @@ -256,7 +256,7 @@ async fn closing_initiator_reader_preserves_initiator_writer() { let stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); let mut writer = stream.writer; @@ -322,7 +322,7 @@ async fn max_concurrent_message_writes_is_respected() { for i in 0..4u8 { let handle = handle_a.clone(); tasks.push(tokio::task::spawn_local(async move { - let mut stream = handle.open_stream(test_route_id()).await.unwrap(); + let mut stream = handle.open_stream(test_open_stream_params()).await.unwrap(); stream.writer.write(Bytes::from(vec![i; 8])).await.unwrap(); stream.writer.finish().await.unwrap(); assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), None); @@ -396,7 +396,7 @@ async fn stream_round_trip_survives_encrypted_packet_drops() { received_request }); - let mut stream = handle_a.open_stream(test_route_id()).await.unwrap(); + let mut stream = handle_a.open_stream(test_open_stream_params()).await.unwrap(); stream .writer .write(Bytes::from(request_payload.clone())) @@ -497,7 +497,7 @@ async fn multi_megabyte_stream_survives_asymmetric_loss_and_delay() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); for (index, chunk) in payload.chunks(chunk_len).enumerate() { @@ -603,7 +603,7 @@ async fn reproducer_writer_stalls_after_reverse_path_impairment() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); for chunk in payload.chunks(chunk_len) { @@ -651,7 +651,7 @@ async fn responder_drains_multiple_local_chunks_per_writable_wake() { let mut stream = pair .side(Side::A) .handle - .open_stream(test_route_id()) + .open_stream(test_open_stream_params()) .await .unwrap(); stream diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index 563f9ded..4dea9b14 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -7,6 +7,7 @@ mod ack; mod builder; mod close; mod route_id; +mod service_id; mod stream_close; mod stream_data; mod stream_id; @@ -16,6 +17,7 @@ pub use ack::*; pub use builder::*; pub use close::*; pub use route_id::*; +pub use service_id::*; pub use stream_close::*; pub use stream_data::*; pub use stream_id::*; diff --git a/ql-wire/src/encrypted/service_id.rs b/ql-wire/src/encrypted/service_id.rs new file mode 100644 index 00000000..a05a8438 --- /dev/null +++ b/ql-wire/src/encrypted/service_id.rs @@ -0,0 +1,34 @@ +use crate::{ByteSlice, Reader, WireDecode, WireEncode, WireError}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct ServiceId(pub [u8; 16]); + +impl ServiceId { + pub const ENCODED_LEN: usize = size_of::(); +} + +impl WireEncode for ServiceId { + fn encoded_len(&self) -> usize { + Self::ENCODED_LEN + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl WireDecode for ServiceId { + fn decode(reader: &mut Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +impl std::fmt::Display for ServiceId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + for byte in self.0 { + write!(f, "{byte:02x}")?; + } + Ok(()) + } +} diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index 9174fe5a..85dd1baf 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -1,6 +1,6 @@ use bytes::Buf; -use super::{RouteId, StreamId}; +use super::{RouteId, ServiceId, StreamId}; use crate::{codec, BufView, ByteSlice, VarInt, WireDecode, WireEncode, WireError}; /// carries bytes for a stream and may finish that sending direction. @@ -104,16 +104,18 @@ impl WireEncode for StreamData { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct StreamHeader { + pub service_id: ServiceId, pub route_id: RouteId, } impl StreamHeader { - pub const MAX_WIRE_SIZE: usize = RouteId::MAX_ENCODED_LEN; + pub const MAX_WIRE_SIZE: usize = ServiceId::ENCODED_LEN + RouteId::MAX_ENCODED_LEN; } impl WireDecode for StreamHeader { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { + service_id: reader.decode()?, route_id: reader.decode()?, }) } @@ -121,10 +123,11 @@ impl WireDecode for StreamHeader { impl WireEncode for StreamHeader { fn encoded_len(&self) -> usize { - self.route_id.encoded_len() + self.service_id.encoded_len() + self.route_id.encoded_len() } fn encode(&self, out: &mut W) { + self.service_id.encode(out); self.route_id.encode(out); } } From 00972ae0da4295c96ae3fc082cdd1bad9b13f9f2 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 22 Jun 2026 16:38:19 -0400 Subject: [PATCH 08/59] fsm: store handshake id --- ql-fsm/src/handshake/ik.rs | 11 ++++++++--- ql-fsm/src/handshake/kk.rs | 15 +++++++++------ ql-fsm/src/handshake/mod.rs | 26 ++++++++++++++++++++++++-- ql-fsm/src/handshake/xx.rs | 12 +++++++++--- ql-fsm/src/state.rs | 1 + ql-fsm/src/tests/mod.rs | 6 ++++-- 6 files changed, 55 insertions(+), 16 deletions(-) diff --git a/ql-fsm/src/handshake/ik.rs b/ql-fsm/src/handshake/ik.rs index 7e6ebd1e..41bc6192 100644 --- a/ql-fsm/src/handshake/ik.rs +++ b/ql-fsm/src/handshake/ik.rs @@ -64,7 +64,7 @@ pub fn handle_ik1( .finalize(crypto) .map_err(ReceiveError::InvalidIkHandshake)?, ); - finish_handshake(fsm, transport, remote_bundle)?; + finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle)?; fsm.state.handshake = None; enqueue_handshake(fsm, QlHandshakeRecord::Ik2(outbound)); Ok(()) @@ -99,16 +99,21 @@ pub fn handle_ik2( .finalize(crypto) .map_err(ReceiveError::InvalidIkHandshake)?, ); - finish_handshake(fsm, transport, remote_bundle) + finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle) } pub fn should_ignore_inbound(fsm: &QlFsm, message: &Ik1) -> bool { match &fsm.state.link { LinkState::Idle - | LinkState::Connected(_) | LinkState::KkInitiator(_) | LinkState::XxInitiator(_) | LinkState::XxResponder(_) => false, + LinkState::Connected(_) => super::is_connected_replay( + fsm, + message.meta.handshake_id, + message.header.sender, + message.header.recipient, + ), LinkState::IkInitiator(state) => { if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(message.header.sender) { return false; diff --git a/ql-fsm/src/handshake/kk.rs b/ql-fsm/src/handshake/kk.rs index e78c8a6d..683a77b7 100644 --- a/ql-fsm/src/handshake/kk.rs +++ b/ql-fsm/src/handshake/kk.rs @@ -63,7 +63,7 @@ pub fn handle_kk1( .finalize(crypto) .map_err(ReceiveError::InvalidKkHandshake)?, ); - finish_handshake(fsm, transport, remote_bundle)?; + finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle)?; fsm.state.handshake = None; enqueue_handshake(fsm, QlHandshakeRecord::Kk2(outbound)); Ok(()) @@ -98,15 +98,18 @@ pub fn handle_kk2( .finalize(crypto) .map_err(ReceiveError::InvalidKkHandshake)?, ); - finish_handshake(fsm, transport, remote_bundle) + finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle) } pub fn should_ignore_inbound(fsm: &QlFsm, message: &Kk1) -> bool { match &fsm.state.link { - LinkState::Idle - | LinkState::Connected(_) - | LinkState::XxInitiator(_) - | LinkState::XxResponder(_) => false, + LinkState::Idle | LinkState::XxInitiator(_) | LinkState::XxResponder(_) => false, + LinkState::Connected(_) => super::is_connected_replay( + fsm, + message.meta.handshake_id, + message.header.sender, + message.header.recipient, + ), LinkState::IkInitiator(_) => true, LinkState::KkInitiator(state) => { if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(message.header.sender) { diff --git a/ql-fsm/src/handshake/mod.rs b/ql-fsm/src/handshake/mod.rs index 1881f66e..8187eee5 100644 --- a/ql-fsm/src/handshake/mod.rs +++ b/ql-fsm/src/handshake/mod.rs @@ -2,7 +2,9 @@ mod ik; mod kk; mod xx; -use ql_wire::{self as wire, EphemeralPublicKey, HandshakeMeta, QlCrypto, QlHandshakeRecord}; +use ql_wire::{ + self as wire, EphemeralPublicKey, HandshakeId, HandshakeMeta, QlCrypto, QlHandshakeRecord, QID, +}; use crate::{ fsm::emit_peer_status, @@ -92,6 +94,7 @@ pub fn next_handshake_deadline(fsm: &QlFsm) -> Option { pub fn finish_handshake( fsm: &mut QlFsm, + handshake_id: HandshakeId, transport: SessionTransport, remote_bundle: wire::PeerBundle, ) -> Result<(), ReceiveError> { @@ -124,7 +127,11 @@ pub fn finish_handshake( }, fsm.state.now, ); - fsm.state.link = LinkState::Connected(ConnectedState { transport, session }); + fsm.state.link = LinkState::Connected(ConnectedState { + handshake_id, + transport, + session, + }); emit_peer_status(fsm, fsm.state.link.status()); Ok(()) } @@ -138,3 +145,18 @@ pub fn reset_connected_session_if_needed(fsm: &mut QlFsm) { fn local_start_wins(local: &EphemeralPublicKey, inbound: &EphemeralPublicKey) -> bool { local.mlkem_public_key.as_bytes() <= inbound.mlkem_public_key.as_bytes() } + +fn is_connected_replay( + fsm: &QlFsm, + handshake_id: HandshakeId, + sender: QID, + recipient: QID, +) -> bool { + let LinkState::Connected(connected) = &fsm.state.link else { + return false; + }; + + connected.handshake_id == handshake_id + && recipient == fsm.identity.qid + && fsm.state.peer.as_ref().map(|peer| peer.qid) == Some(sender) +} diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs index c9a289e0..34b511ac 100644 --- a/ql-fsm/src/handshake/xx.rs +++ b/ql-fsm/src/handshake/xx.rs @@ -146,7 +146,7 @@ pub fn handle_xx3( .finalize(crypto) .map_err(ReceiveError::InvalidXxHandshake)?, ); - finish_handshake(fsm, transport, remote_bundle) + finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle) } pub fn handle_xx4( @@ -178,7 +178,7 @@ pub fn handle_xx4( .finalize(crypto) .map_err(ReceiveError::InvalidXxHandshake)?, ); - finish_handshake(fsm, transport, remote_bundle) + finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle) } pub fn disarm_pairing(fsm: &mut QlFsm) { @@ -190,7 +190,13 @@ pub fn disarm_pairing(fsm: &mut QlFsm) { pub fn should_ignore_inbound(fsm: &QlFsm, crypto: &impl QlCrypto, message: &Xx1) -> bool { match &fsm.state.link { - LinkState::Idle | LinkState::Connected(_) => false, + LinkState::Idle => false, + LinkState::Connected(_) => super::is_connected_replay( + fsm, + message.meta.handshake_id, + message.header.sender, + message.header.recipient, + ), LinkState::IkInitiator(_) | LinkState::KkInitiator(_) | LinkState::XxResponder(_) => true, LinkState::XxInitiator(state) => { if state.handshake.pairing_id(crypto) != message.pairing_id { diff --git a/ql-fsm/src/state.rs b/ql-fsm/src/state.rs index 8268bc16..c7657719 100644 --- a/ql-fsm/src/state.rs +++ b/ql-fsm/src/state.rs @@ -51,6 +51,7 @@ pub enum LinkState { } pub struct ConnectedState { + pub handshake_id: HandshakeId, pub transport: SessionTransport, pub session: SessionFsm, } diff --git a/ql-fsm/src/tests/mod.rs b/ql-fsm/src/tests/mod.rs index c4d005e4..88ed8d8a 100644 --- a/ql-fsm/src/tests/mod.rs +++ b/ql-fsm/src/tests/mod.rs @@ -5,8 +5,8 @@ mod session; use std::time::{Duration, Instant}; use ql_wire::{ - self, generate_identity, test_identities, ConnectionId, PairingToken, QlCrypto, SessionKey, - SoftwareCrypto, TransportParams, QID, + self, generate_identity, test_identities, ConnectionId, HandshakeId, PairingToken, QlCrypto, + SessionKey, SoftwareCrypto, TransportParams, QID, }; use crate::{ @@ -102,6 +102,7 @@ impl Harness { let b_to_a_conn = ConnectionId::from_data([0xB2; ConnectionId::SIZE]); harness.a.fsm.state.link = LinkState::Connected(ConnectedState { + handshake_id: HandshakeId(0), transport: SessionTransport { tx_key: a_to_b_key.clone(), rx_key: b_to_a_key.clone(), @@ -118,6 +119,7 @@ impl Harness { session: SessionFsm::new(session_config(&harness, true), harness.now), }); harness.b.fsm.state.link = LinkState::Connected(ConnectedState { + handshake_id: HandshakeId(0), transport: SessionTransport { tx_key: b_to_a_key, rx_key: a_to_b_key, From 4b20c72ab6ea2a8d7f8d10727519709f09b3bbf1 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 22 Jun 2026 15:42:58 -0400 Subject: [PATCH 09/59] wire: helper macros + cleanup --- ql-fsm/src/tests/mod.rs | 8 +-- ql-runtime/src/tests/handshake.rs | 5 +- ql-runtime/src/tests/session.rs | 5 +- ql-runtime/src/tests/stream.rs | 5 +- ql-wire/src/bytes.rs | 54 +----------------- ql-wire/src/crypto.rs | 12 +++- ql-wire/src/encrypted/ack.rs | 20 ------- ql-wire/src/encrypted/mod.rs | 23 +++++--- ql-wire/src/encrypted/route_id.rs | 55 ------------------ ql-wire/src/encrypted/stream_id.rs | 35 ------------ ql-wire/src/encrypted_message.rs | 62 +------------------- ql-wire/src/handshake/mod.rs | 21 +++---- ql-wire/src/handshake/pairing.rs | 48 +--------------- ql-wire/src/header.rs | 72 +---------------------- ql-wire/src/identity.rs | 23 +------- ql-wire/src/lib.rs | 4 +- ql-wire/src/macros.rs | 92 ++++++++++++++++++++++++++++++ ql-wire/src/nonce.rs | 13 ----- ql-wire/src/pq.rs | 12 +--- ql-wire/src/qid.rs | 32 +---------- ql-wire/src/testing.rs | 12 ++-- ql-wire/src/tests.rs | 80 ++------------------------ 22 files changed, 170 insertions(+), 523 deletions(-) delete mode 100644 ql-wire/src/encrypted/route_id.rs delete mode 100644 ql-wire/src/encrypted/stream_id.rs create mode 100644 ql-wire/src/macros.rs delete mode 100644 ql-wire/src/nonce.rs diff --git a/ql-fsm/src/tests/mod.rs b/ql-fsm/src/tests/mod.rs index 88ed8d8a..2df7d14b 100644 --- a/ql-fsm/src/tests/mod.rs +++ b/ql-fsm/src/tests/mod.rs @@ -96,10 +96,10 @@ impl Harness { fn connected(config: QlFsmConfig) -> Self { let mut harness = Self::paired_known(config); - let a_to_b_key = SessionKey::from_data([7; SessionKey::SIZE]); - let b_to_a_key = SessionKey::from_data([9; SessionKey::SIZE]); - let a_to_b_conn = ConnectionId::from_data([0xA1; ConnectionId::SIZE]); - let b_to_a_conn = ConnectionId::from_data([0xB2; ConnectionId::SIZE]); + let a_to_b_key = SessionKey([7; SessionKey::SIZE]); + let b_to_a_key = SessionKey([9; SessionKey::SIZE]); + let a_to_b_conn = ConnectionId([0xA1; ConnectionId::SIZE]); + let b_to_a_conn = ConnectionId([0xB2; ConnectionId::SIZE]); harness.a.fsm.state.link = LinkState::Connected(ConnectedState { handshake_id: HandshakeId(0), diff --git a/ql-runtime/src/tests/handshake.rs b/ql-runtime/src/tests/handshake.rs index e2e8b030..b2666f02 100644 --- a/ql-runtime/src/tests/handshake.rs +++ b/ql-runtime/src/tests/handshake.rs @@ -88,7 +88,10 @@ async fn rejected_session_write_is_reissued() { request }); - let mut stream = handle_a.open_stream(test_open_stream_params()).await.unwrap(); + let mut stream = handle_a + .open_stream(test_open_stream_params()) + .await + .unwrap(); stream .writer .write(Bytes::from_static(b"retry")) diff --git a/ql-runtime/src/tests/session.rs b/ql-runtime/src/tests/session.rs index 67d44b96..ae58d933 100644 --- a/ql-runtime/src/tests/session.rs +++ b/ql-runtime/src/tests/session.rs @@ -201,7 +201,10 @@ async fn session_timeout_disconnects_and_fails_pending_open() { drop_flag.store(true, Ordering::Relaxed); - let mut pending = handle_a.open_stream(test_open_stream_params()).await.unwrap(); + let mut pending = handle_a + .open_stream(test_open_stream_params()) + .await + .unwrap(); let err = pending.writer.finish().await.unwrap_err(); assert!(matches!(err, QlStreamError::NoSession)); diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index 72d8269b..8dee4f18 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -396,7 +396,10 @@ async fn stream_round_trip_survives_encrypted_packet_drops() { received_request }); - let mut stream = handle_a.open_stream(test_open_stream_params()).await.unwrap(); + let mut stream = handle_a + .open_stream(test_open_stream_params()) + .await + .unwrap(); stream .writer .write(Bytes::from(request_payload.clone())) diff --git a/ql-wire/src/bytes.rs b/ql-wire/src/bytes.rs index 9fecf5ea..09cd1202 100644 --- a/ql-wire/src/bytes.rs +++ b/ql-wire/src/bytes.rs @@ -1,4 +1,4 @@ -use core::ops::{Deref, DerefMut}; +use core::ops::Deref; use bytes::{Buf, Bytes}; @@ -10,11 +10,6 @@ pub trait ByteSlice: Deref + Sized { fn split_at(self, mid: usize) -> Result<(Self, Self), Self>; } -/// A mutable reference to bytes. -pub trait ByteSliceMut: ByteSlice + DerefMut {} - -impl ByteSliceMut for B where B: ByteSlice + DerefMut {} - impl ByteSlice for &[u8] { #[inline] fn split_at(self, mid: usize) -> Result<(Self, Self), Self> { @@ -126,50 +121,3 @@ impl BufView for Bytes { self.as_ref() } } - -#[cfg(test)] -mod tests { - use bytes::Buf; - - use super::{BufView, ByteSlice, ByteSliceMut}; - - #[test] - fn shared_slice_split_at() { - let bytes: &[u8] = b"abcdef"; - let (left, right) = ByteSlice::split_at(bytes, 2).unwrap(); - assert_eq!(left, b"ab"); - assert_eq!(right, b"cdef"); - } - - #[test] - fn mutable_slice_split_at() { - let mut bytes = *b"abcdef"; - let (left, right) = ByteSlice::split_at(&mut bytes[..], 2).unwrap(); - assert_eq!(left, b"ab"); - assert_eq!(right, b"cdef"); - } - - #[test] - fn mutable_split_trait_is_implemented() { - fn assert_split_mut(_value: T) {} - - let mut bytes = [0u8; 4]; - assert_split_mut(&mut bytes[..]); - } - - #[test] - fn split_at_rejects_out_of_bounds_index() { - let bytes: &[u8] = b"abcdef"; - assert!(ByteSlice::split_at(bytes, 7).is_err()); - } - - #[test] - fn slice_buf_view_is_contiguous() { - let bytes: &[u8] = b"abcdef"; - let mut buf = bytes.buf(); - assert_eq!(buf.remaining(), 6); - assert_eq!(buf.chunk(), b"abcdef"); - buf.advance(6); - assert!(!buf.has_remaining()); - } -} diff --git a/ql-wire/src/crypto.rs b/ql-wire/src/crypto.rs index 96ace383..0f00bd80 100644 --- a/ql-wire/src/crypto.rs +++ b/ql-wire/src/crypto.rs @@ -1,8 +1,18 @@ use crate::{ - MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, Nonce, SessionKey, + MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, SessionKey, ENCRYPTED_MESSAGE_AUTH_SIZE, }; +crate::array_wrapper!(Nonce, 12); + +impl Nonce { + pub fn from_counter(counter: u64) -> Self { + let mut nonce = [0u8; Self::SIZE]; + nonce[4..].copy_from_slice(&counter.to_le_bytes()); + Self(nonce) + } +} + pub trait QlRandom { fn fill_random_bytes(&self, out: &mut [u8]); } diff --git a/ql-wire/src/encrypted/ack.rs b/ql-wire/src/encrypted/ack.rs index 2eb34b33..7459a608 100644 --- a/ql-wire/src/encrypted/ack.rs +++ b/ql-wire/src/encrypted/ack.rs @@ -319,26 +319,6 @@ mod tests { ); } - #[test] - fn builder_matches_from_ranges() { - let mut builder = RecordAckBuilder::new(); - assert!(builder - .try_push_range(ack_range(95, 100), usize::MAX) - .unwrap()); - assert!(builder - .try_push_range(ack_range(90, 92), usize::MAX) - .unwrap()); - assert!(builder - .try_push_range(ack_range(80, 80), usize::MAX) - .unwrap()); - - assert_eq!( - builder.build().unwrap(), - RecordAck::from_ranges([ack_range(95, 100), ack_range(90, 92), ack_range(80, 80)]) - .unwrap() - ); - } - #[test] fn builder_stops_when_budget_is_exhausted() { let first_only = RecordAck::from_ranges([ack_range(95, 100)]).unwrap(); diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index 4dea9b14..8a12f145 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -1,28 +1,27 @@ use crate::{ - codec, encrypted_message::EncryptedMessage, BufView, ByteSlice, Nonce, QlCrypto, Reader, - SessionHeader, SessionKey, WireDecode, WireEncode, WireError, + codec, encrypted_message::EncryptedMessage, varint_wrapper, BufView, ByteSlice, Nonce, + QlCrypto, Reader, SessionHeader, SessionKey, WireDecode, WireEncode, WireError, }; mod ack; mod builder; mod close; -mod route_id; mod service_id; mod stream_close; mod stream_data; -mod stream_id; mod stream_window; pub use ack::*; pub use builder::*; pub use close::*; -pub use route_id::*; pub use service_id::*; pub use stream_close::*; pub use stream_data::*; -pub use stream_id::*; pub use stream_window::*; +varint_wrapper!(RouteId); +varint_wrapper!(StreamId); + #[derive(Debug, Clone, PartialEq, Eq)] pub enum SessionFrame { // todo: do we need ping as explicit frame? @@ -176,5 +175,15 @@ pub fn decrypt_record>( ) -> Result { let aad = header.aad(); let nonce = Nonce::from_counter(header.seq.into_inner()); - encrypted.decrypt_in_place(crypto, session_key, &nonce, &aad) + let mut ciphertext = encrypted.ciphertext; + if !crypto.aes256_gcm_decrypt( + session_key, + &nonce, + &aad, + ciphertext.as_mut(), + &encrypted.auth, + ) { + return Err(WireError::DecryptFailed); + } + Ok(ciphertext) } diff --git a/ql-wire/src/encrypted/route_id.rs b/ql-wire/src/encrypted/route_id.rs deleted file mode 100644 index 6b91a521..00000000 --- a/ql-wire/src/encrypted/route_id.rs +++ /dev/null @@ -1,55 +0,0 @@ -use crate::{ByteSlice, Reader, VarInt, VarIntBoundsExceeded, WireDecode, WireEncode, WireError}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -#[repr(transparent)] -pub struct RouteId(pub VarInt); - -impl RouteId { - pub const MAX_ENCODED_LEN: usize = VarInt::MAX_SIZE; - - pub const fn from_u32(value: u32) -> Self { - Self(VarInt::from_u32(value)) - } - - pub fn from_u64(value: u64) -> Result { - Ok(Self(VarInt::from_u64(value)?)) - } - - pub const fn into_inner(self) -> u64 { - self.0.into_inner() - } -} - -impl WireEncode for RouteId { - fn encoded_len(&self) -> usize { - self.0.size() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl WireDecode for RouteId { - fn decode(reader: &mut Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -impl From for RouteId { - fn from(value: VarInt) -> Self { - Self(value) - } -} - -impl From for RouteId { - fn from(value: u32) -> Self { - Self::from_u32(value) - } -} - -impl std::fmt::Display for RouteId { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.0) - } -} diff --git a/ql-wire/src/encrypted/stream_id.rs b/ql-wire/src/encrypted/stream_id.rs deleted file mode 100644 index 07002259..00000000 --- a/ql-wire/src/encrypted/stream_id.rs +++ /dev/null @@ -1,35 +0,0 @@ -use crate::{ByteSlice, Reader, VarInt, WireDecode, WireEncode, WireError}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -#[repr(transparent)] -pub struct StreamId(pub VarInt); - -impl StreamId { - pub const MAX_ENCODED_LEN: usize = VarInt::MAX_SIZE; - - pub const fn into_inner(self) -> u64 { - self.0.into_inner() - } -} - -impl WireEncode for StreamId { - fn encoded_len(&self) -> usize { - self.0.size() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl WireDecode for StreamId { - fn decode(reader: &mut Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -impl std::fmt::Display for StreamId { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.0) - } -} diff --git a/ql-wire/src/encrypted_message.rs b/ql-wire/src/encrypted_message.rs index 9e11d3d0..58b72591 100644 --- a/ql-wire/src/encrypted_message.rs +++ b/ql-wire/src/encrypted_message.rs @@ -1,7 +1,4 @@ -use crate::{ - codec, ByteSlice, Nonce, QlCrypto, SessionKey, WireDecode, WireEncode, WireError, - ENCRYPTED_MESSAGE_AUTH_SIZE, -}; +use crate::{codec, ByteSlice, WireDecode, WireEncode, WireError, ENCRYPTED_MESSAGE_AUTH_SIZE}; #[derive(Debug, Clone, PartialEq, Eq)] pub struct EncryptedMessage { @@ -10,9 +7,6 @@ pub struct EncryptedMessage { } impl EncryptedMessage { - pub const AUTH_SIZE: usize = ENCRYPTED_MESSAGE_AUTH_SIZE; - pub const HEADER_LEN: usize = Self::AUTH_SIZE; - pub fn into_owned(self) -> EncryptedMessage> where B: ByteSlice, @@ -33,25 +27,9 @@ impl WireDecode for EncryptedMessage { } } -impl> EncryptedMessage { - pub fn decrypt( - &self, - crypto: &impl QlCrypto, - key: &SessionKey, - nonce: &Nonce, - aad: &[u8], - ) -> Result, WireError> { - let mut plaintext = self.ciphertext.as_ref().to_vec(); - if !crypto.aes256_gcm_decrypt(key, nonce, aad, &mut plaintext, &self.auth) { - return Err(WireError::DecryptFailed); - } - Ok(plaintext) - } -} - impl> WireEncode for EncryptedMessage { fn encoded_len(&self) -> usize { - Self::HEADER_LEN + self.ciphertext.as_ref().len() + ENCRYPTED_MESSAGE_AUTH_SIZE + self.ciphertext.as_ref().len() } fn encode(&self, out: &mut W) { @@ -59,39 +37,3 @@ impl> WireEncode for EncryptedMessage { self.ciphertext.as_ref().encode(out); } } - -impl> EncryptedMessage { - pub fn decrypt_in_place( - mut self, - crypto: &impl QlCrypto, - key: &SessionKey, - nonce: &Nonce, - aad: &[u8], - ) -> Result { - let ciphertext = self.ciphertext.as_mut(); - if !crypto.aes256_gcm_decrypt(key, nonce, aad, ciphertext, &self.auth) { - return Err(WireError::DecryptFailed); - } - Ok(self.ciphertext) - } -} - -impl EncryptedMessage> { - pub fn encrypt( - crypto: &impl QlCrypto, - key: &SessionKey, - mut plaintext: Vec, - nonce: &Nonce, - aad: &[u8], - ) -> Self { - let auth = crypto.aes256_gcm_encrypt(key, nonce, aad, &mut plaintext); - Self { - auth, - ciphertext: plaintext, - } - } - - pub fn decode(bytes: &[u8]) -> Result { - Ok(EncryptedMessage::decode_exact(bytes)?.into_owned()) - } -} diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index a9b7cf87..1ac9aabe 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -117,8 +117,6 @@ impl codec::WireDecode for EncryptedMlKemCiphertext { pub struct EncryptedPeerBundle(pub Box<[u8]>); impl EncryptedPeerBundle { - pub const MAX_WIRE_SIZE: usize = PeerBundle::MAX_WIRE_SIZE + ENCRYPTED_MESSAGE_AUTH_SIZE; - pub fn as_bytes(&self) -> &[u8] { self.0.as_ref() } @@ -137,9 +135,6 @@ impl WireEncode for EncryptedPeerBundle { impl codec::WireDecode for EncryptedPeerBundle { fn decode(reader: &mut codec::Reader) -> Result { let data = reader.take_rest(); - if data.len() > Self::MAX_WIRE_SIZE { - return Err(WireError::InvalidPayload); - } Ok(Self(data.to_vec().into_boxed_slice())) } } @@ -309,8 +304,8 @@ impl SymmetricState { fn split_for_role(&self, crypto: &impl QlCrypto, role: Role) -> (SessionKey, SessionKey) { let temp_key = hmac_sha256(crypto, &self.chaining_key, &[&[]]); - let k1 = SessionKey::from_data(hmac_sha256(crypto, &temp_key, &[&[1]])); - let k2 = SessionKey::from_data(hmac_sha256(crypto, &temp_key, &[k1.as_bytes(), &[2]])); + let k1 = SessionKey(hmac_sha256(crypto, &temp_key, &[&[1]])); + let k2 = SessionKey(hmac_sha256(crypto, &temp_key, &[k1.as_bytes(), &[2]])); match role { Role::Initiator => (k1, k2), Role::Responder => (k2, k1), @@ -472,7 +467,8 @@ fn decrypt_peer_bundle( ) -> Result { let plaintext = symmetric.decrypt_and_hash(crypto, bundle.as_bytes())?; let bundle = PeerBundle::decode_exact(plaintext.as_slice())?; - if !bundle.qid_matches_public_key(crypto) { + let peer_qid = QID::derive(crypto, &bundle.mlkem_public_key); + if peer_qid != bundle.qid { return Err(WireError::InvalidRemoteBundle); } Ok(bundle) @@ -536,10 +532,7 @@ fn derive_connection_ids( let mut responder_rx = [0u8; ConnectionId::SIZE]; initiator_rx.copy_from_slice(&initiator[..ConnectionId::SIZE]); responder_rx.copy_from_slice(&responder[..ConnectionId::SIZE]); - ( - ConnectionId::from_data(initiator_rx), - ConnectionId::from_data(responder_rx), - ) + (ConnectionId(initiator_rx), ConnectionId(responder_rx)) } fn hkdf2( @@ -550,7 +543,7 @@ fn hkdf2( let temp_key = hmac_sha256(crypto, chaining_key, &[input_key_material]); let out1 = hmac_sha256(crypto, &temp_key, &[&[1]]); let out2 = hmac_sha256(crypto, &temp_key, &[&out1, &[2]]); - (out1, SessionKey::from_data(out2)) + (out1, SessionKey(out2)) } fn hkdf3( @@ -562,7 +555,7 @@ fn hkdf3( let out1 = hmac_sha256(crypto, &temp_key, &[&[1]]); let out2 = hmac_sha256(crypto, &temp_key, &[&out1, &[2]]); let out3 = hmac_sha256(crypto, &temp_key, &[&out2, &[3]]); - (out1, out2, SessionKey::from_data(out3)) + (out1, out2, SessionKey(out3)) } fn hmac_sha256(crypto: &impl QlCrypto, key: &[u8], parts: &[&[u8]]) -> [u8; 32] { diff --git a/ql-wire/src/handshake/pairing.rs b/ql-wire/src/handshake/pairing.rs index 237f066b..25227ad9 100644 --- a/ql-wire/src/handshake/pairing.rs +++ b/ql-wire/src/handshake/pairing.rs @@ -1,17 +1,13 @@ use std::fmt::{self, Display, Formatter}; -use crate::{codec, ByteSlice, QlCrypto, WireEncode, WireError}; +use crate::QlCrypto; const PAIRING_ID_DOMAIN: &[u8] = b"ql-wire:pairing-id:v1"; const PAIRING_PSK_DOMAIN: &[u8] = b"ql-wire:pairing-psk:v1"; -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -#[repr(transparent)] -pub struct PairingToken(pub [u8; Self::SIZE]); +crate::array_wrapper!(PairingToken, 16); impl PairingToken { - pub const SIZE: usize = 16; - pub fn id(&self, crypto: &impl QlCrypto) -> PairingId { let hash = crypto.sha256(&[PAIRING_ID_DOMAIN, &self.0]); let mut id = [0u8; PairingId::SIZE]; @@ -33,29 +29,7 @@ impl Display for PairingToken { } } -impl WireEncode for PairingToken { - fn encoded_len(&self) -> usize { - Self::SIZE - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl codec::WireDecode for PairingToken { - fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -#[repr(transparent)] -pub struct PairingId(pub [u8; Self::SIZE]); - -impl PairingId { - pub const SIZE: usize = 16; -} +crate::array_wrapper!(PairingId, 16); impl Display for PairingId { fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { @@ -65,19 +39,3 @@ impl Display for PairingId { Ok(()) } } - -impl WireEncode for PairingId { - fn encoded_len(&self) -> usize { - Self::SIZE - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl codec::WireDecode for PairingId { - fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index 88764c0a..a186fdcf 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -1,8 +1,6 @@ use ::bytes::BufMut; -use crate::{ - codec, ByteSlice, VarInt, VarIntBoundsExceeded, WireEncode, WireError, QL_WIRE_VERSION, -}; +use crate::{codec, ByteSlice, WireEncode, WireError, QL_WIRE_VERSION}; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct SessionHeader { @@ -10,73 +8,9 @@ pub struct SessionHeader { pub seq: RecordSeq, } -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -#[repr(transparent)] -pub struct RecordSeq(pub VarInt); +crate::varint_wrapper!(RecordSeq); -impl RecordSeq { - pub const MAX_ENCODED_LEN: usize = VarInt::MAX_SIZE; - - pub const fn from_u32(value: u32) -> Self { - Self(VarInt::from_u32(value)) - } - - pub fn from_u64(value: u64) -> Result { - Ok(Self(VarInt::from_u64(value)?)) - } - - pub const fn into_inner(self) -> u64 { - self.0.into_inner() - } -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -#[repr(transparent)] -pub struct ConnectionId(pub [u8; Self::SIZE]); - -impl ConnectionId { - pub const SIZE: usize = 16; - - pub const fn from_data(data: [u8; Self::SIZE]) -> Self { - Self(data) - } - - pub const fn as_bytes(&self) -> &[u8; Self::SIZE] { - &self.0 - } -} - -impl codec::WireDecode for RecordSeq { - fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -impl WireEncode for RecordSeq { - fn encoded_len(&self) -> usize { - self.0.size() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl codec::WireDecode for ConnectionId { - fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self::from_data(reader.decode()?)) - } -} - -impl WireEncode for ConnectionId { - fn encoded_len(&self) -> usize { - Self::SIZE - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} +crate::array_wrapper!(ConnectionId, 16); impl SessionHeader { pub const MAX_ENCODED_LEN: usize = ConnectionId::SIZE + RecordSeq::MAX_ENCODED_LEN; diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index bdc54b2a..8423d8cb 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -18,11 +18,6 @@ impl PeerBundle { pub const VERSION: u16 = 1; pub const FIXED_WIRE_SIZE: usize = size_of::() + QID::SIZE + size_of::() + MlKemPublicKey::SIZE; - pub const MAX_WIRE_SIZE: usize = Self::FIXED_WIRE_SIZE + VarInt::MAX_SIZE + QlName::MAX_LEN; - - pub fn qid_matches_public_key(&self, crypto: &impl QlHash) -> bool { - self.qid.matches_public_key(crypto, &self.mlkem_public_key) - } } impl WireEncode for PeerBundle { @@ -82,17 +77,6 @@ impl QlIdentity { }) } - #[must_use] - pub fn with_capabilities(mut self, capabilities: u32) -> Self { - self.capabilities = capabilities; - self - } - - pub fn with_name(mut self, name: impl Into) -> Result { - self.name = QlName::new(name)?; - Ok(self) - } - pub fn bundle(&self) -> PeerBundle { PeerBundle { version: PeerBundle::VERSION, @@ -134,11 +118,8 @@ pub fn generate_identity( crypto: &impl QlCrypto, name: impl Into, ) -> Result { - let MlKemKeyPair { - private: mlkem_private_key, - public: mlkem_public_key, - } = crypto.mlkem_generate_keypair(); - QlIdentity::new(crypto, mlkem_private_key, mlkem_public_key, name) + let MlKemKeyPair { private, public } = crypto.mlkem_generate_keypair(); + QlIdentity::new(crypto, private, public, name) } #[derive(Debug, Clone, PartialEq, Eq)] diff --git a/ql-wire/src/lib.rs b/ql-wire/src/lib.rs index 1713745a..9381afd7 100644 --- a/ql-wire/src/lib.rs +++ b/ql-wire/src/lib.rs @@ -13,7 +13,7 @@ mod error; mod handshake; mod header; mod identity; -mod nonce; +mod macros; mod pq; mod qid; mod record; @@ -30,7 +30,7 @@ pub use error::*; pub use handshake::*; pub use header::*; pub use identity::*; -pub use nonce::*; +pub(crate) use macros::*; pub use pq::*; pub use qid::*; pub use record::*; diff --git a/ql-wire/src/macros.rs b/ql-wire/src/macros.rs new file mode 100644 index 00000000..a0f52b03 --- /dev/null +++ b/ql-wire/src/macros.rs @@ -0,0 +1,92 @@ +macro_rules! varint_wrapper { + ($name:ident) => { + #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] + #[repr(transparent)] + pub struct $name(pub $crate::VarInt); + + impl $name { + pub const MAX_ENCODED_LEN: usize = $crate::VarInt::MAX_SIZE; + + pub const fn from_u32(value: u32) -> Self { + Self($crate::VarInt::from_u32(value)) + } + + pub fn from_u64(value: u64) -> Result { + Ok(Self($crate::VarInt::from_u64(value)?)) + } + + pub const fn into_inner(self) -> u64 { + self.0.into_inner() + } + } + + impl $crate::WireEncode for $name { + fn encoded_len(&self) -> usize { + self.0.size() + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } + } + + impl $crate::WireDecode for $name { + fn decode(reader: &mut $crate::Reader) -> Result { + Ok(Self(reader.decode()?)) + } + } + + impl From<$crate::VarInt> for $name { + fn from(value: $crate::VarInt) -> Self { + Self(value) + } + } + + impl From for $name { + fn from(value: u32) -> Self { + Self::from_u32(value) + } + } + + impl std::fmt::Display for $name { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } + } + }; +} + +macro_rules! array_wrapper { + ($name:ident, $size:expr) => { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + #[repr(transparent)] + pub struct $name(pub [u8; Self::SIZE]); + + impl $name { + pub const SIZE: usize = $size; + + pub const fn as_bytes(&self) -> &[u8; Self::SIZE] { + &self.0 + } + } + + impl $crate::WireEncode for $name { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } + } + + impl $crate::codec::WireDecode for $name { + fn decode(reader: &mut $crate::codec::Reader) -> Result { + Ok(Self(reader.decode()?)) + } + } + }; +} + +pub(crate) use array_wrapper; +pub(crate) use varint_wrapper; diff --git a/ql-wire/src/nonce.rs b/ql-wire/src/nonce.rs deleted file mode 100644 index c7e6d793..00000000 --- a/ql-wire/src/nonce.rs +++ /dev/null @@ -1,13 +0,0 @@ -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -#[repr(transparent)] -pub struct Nonce(pub [u8; Self::SIZE]); - -impl Nonce { - pub const SIZE: usize = 12; - - pub fn from_counter(counter: u64) -> Self { - let mut nonce = [0u8; Self::SIZE]; - nonce[4..].copy_from_slice(&counter.to_le_bytes()); - Self(nonce) - } -} diff --git a/ql-wire/src/pq.rs b/ql-wire/src/pq.rs index 327ef7c4..8aa513bc 100644 --- a/ql-wire/src/pq.rs +++ b/ql-wire/src/pq.rs @@ -11,19 +11,11 @@ const ML_KEM_1024_PRIVATE_KEY_SIZE: usize = 3168; const ML_KEM_1024_CIPHERTEXT_SIZE: usize = 1568; #[derive(Debug, Clone, PartialEq, Eq, Hash)] -pub struct SessionKey([u8; Self::SIZE]); +pub struct SessionKey(pub [u8; Self::SIZE]); impl SessionKey { pub const SIZE: usize = ML_KEM_1024_SHARED_SECRET_SIZE; - pub const fn from_data(data: [u8; Self::SIZE]) -> Self { - Self(data) - } - - pub const fn data(&self) -> &[u8; Self::SIZE] { - &self.0 - } - pub const fn as_bytes(&self) -> &[u8; Self::SIZE] { &self.0 } @@ -53,7 +45,7 @@ impl WireEncode for SessionKey { impl codec::WireDecode for SessionKey { fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self::from_data(reader.decode()?)) + Ok(Self(reader.decode()?)) } } diff --git a/ql-wire/src/qid.rs b/ql-wire/src/qid.rs index 55c6684f..04f4e759 100644 --- a/ql-wire/src/qid.rs +++ b/ql-wire/src/qid.rs @@ -1,12 +1,8 @@ -use crate::{codec, ByteSlice, MlKemPublicKey, QlHash, WireEncode, WireError, ML_KEM_SUITE_TAG}; +use crate::{MlKemPublicKey, QlHash, ML_KEM_SUITE_TAG}; -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -#[repr(transparent)] -pub struct QID(pub [u8; Self::SIZE]); +crate::array_wrapper!(QID, 16); impl QID { - pub const SIZE: usize = 16; - pub fn derive(crypto: &impl QlHash, mlkem_public_key: &MlKemPublicKey) -> Self { let digest = crypto.sha256(&[ b"quantum-link qid v1", @@ -17,28 +13,4 @@ impl QID { qid.copy_from_slice(&digest[..Self::SIZE]); Self(qid) } - - pub fn matches_public_key( - &self, - crypto: &impl QlHash, - mlkem_public_key: &MlKemPublicKey, - ) -> bool { - *self == Self::derive(crypto, mlkem_public_key) - } -} - -impl WireEncode for QID { - fn encoded_len(&self) -> usize { - Self::SIZE - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl codec::WireDecode for QID { - fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } } diff --git a/ql-wire/src/testing.rs b/ql-wire/src/testing.rs index a1223c12..adc2e093 100644 --- a/ql-wire/src/testing.rs +++ b/ql-wire/src/testing.rs @@ -44,7 +44,7 @@ impl QlAead for SoftwareCrypto { aad: &[u8], buffer: &mut [u8], ) -> [u8; ENCRYPTED_MESSAGE_AUTH_SIZE] { - let key: AesGcm256Key = (*key.data()).into(); + let key: AesGcm256Key = (*key.as_bytes()).into(); let plaintext = buffer.to_vec(); let mut auth = [0u8; ENCRYPTED_MESSAGE_AUTH_SIZE]; key.encrypt( @@ -66,7 +66,7 @@ impl QlAead for SoftwareCrypto { buffer: &mut [u8], auth_tag: &[u8; ENCRYPTED_MESSAGE_AUTH_SIZE], ) -> bool { - let key: AesGcm256Key = (*key.data()).into(); + let key: AesGcm256Key = (*key.as_bytes()).into(); let ciphertext = buffer.to_vec(); key.decrypt(buffer, (&nonce.0).into(), aad, &ciphertext, auth_tag.into()) .is_ok() @@ -97,7 +97,7 @@ impl QlKem for SoftwareCrypto { shared.copy_from_slice(shared_value.as_slice()); ( MlKemCiphertext::new(Box::new(ciphertext)), - SessionKey::from_data(shared), + SessionKey(shared), ) } @@ -111,7 +111,7 @@ impl QlKem for SoftwareCrypto { let shared = mlkem1024::decapsulate(&private_key, &ciphertext); let mut out = [0u8; SessionKey::SIZE]; out.copy_from_slice(shared.as_slice()); - SessionKey::from_data(out) + SessionKey(out) } } @@ -161,7 +161,7 @@ impl QlKem for NoopCrypto { fn mlkem_encapsulate(&self, _public_key: &MlKemPublicKey) -> (MlKemCiphertext, SessionKey) { ( MlKemCiphertext::new(Box::new([0; MlKemCiphertext::SIZE])), - SessionKey::from_data([0; SessionKey::SIZE]), + SessionKey([0; SessionKey::SIZE]), ) } @@ -170,7 +170,7 @@ impl QlKem for NoopCrypto { _private_key: &MlKemPrivateKey, _ciphertext: &MlKemCiphertext, ) -> SessionKey { - SessionKey::from_data([0; SessionKey::SIZE]) + SessionKey([0; SessionKey::SIZE]) } } diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index a09a36b8..acd69f6b 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -86,9 +86,8 @@ fn encrypt_record( #[test] fn peer_bundle_round_trip() { let crypto = SoftwareCrypto; - let identity = generate_identity(&crypto, "alice") - .unwrap() - .with_capabilities(0x55aa_33cc); + let mut identity = generate_identity(&crypto, "alice").unwrap(); + identity.capabilities = 1231; let bundle = identity.bundle(); let encoded = bundle.encode_vec(); @@ -111,44 +110,6 @@ fn identity_name_validation() { )); } -#[test] -fn qid_derives_from_mlkem_public_key() { - let crypto = SoftwareCrypto; - let public_key = MlKemPublicKey::new(Box::new([42; MlKemPublicKey::SIZE])); - let qid = QID::derive(&crypto, &public_key); - - let digest = crypto.sha256(&[ - b"quantum-link qid v1", - ML_KEM_SUITE_TAG, - public_key.as_bytes(), - ]); - let mut expected = [0u8; QID::SIZE]; - expected.copy_from_slice(&digest[..QID::SIZE]); - - assert_eq!(qid, QID(expected)); - assert!(qid.matches_public_key(&crypto, &public_key)); -} - -#[test] -fn qid_changes_when_mlkem_public_key_changes() { - let crypto = SoftwareCrypto; - let first = MlKemPublicKey::new(Box::new([1; MlKemPublicKey::SIZE])); - let second = MlKemPublicKey::new(Box::new([2; MlKemPublicKey::SIZE])); - - assert_ne!(QID::derive(&crypto, &first), QID::derive(&crypto, &second)); -} - -#[test] -fn peer_bundle_detects_tampered_qid() { - let crypto = SoftwareCrypto; - let identity = generate_identity(&crypto, "alice").unwrap(); - let mut bundle = identity.bundle(); - - bundle.qid = qid(9); - - assert!(!bundle.qid_matches_public_key(&crypto)); -} - #[test] fn handshake_record_round_trip_supports_ik_kk_and_xx() { let ik = QlHandshakeRecord::Ik1(Ik1 { @@ -758,7 +719,7 @@ fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { fn encrypted_session_record_round_trip_uses_connection_id_header() { let crypto = SoftwareCrypto; let header = SessionHeader { - connection_id: ConnectionId::from_data([0x44; ConnectionId::SIZE]), + connection_id: ConnectionId([0x44; ConnectionId::SIZE]), seq: record_seq(11), }; let body = vec![ @@ -787,7 +748,7 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { code: SessionCloseCode::TIMEOUT, }), ]; - let session_key = SessionKey::from_data([7; SessionKey::SIZE]); + let session_key = SessionKey([7; SessionKey::SIZE]); let record = encrypt_record(&crypto, header, &session_key, &body); let bytes = encode_record_vec(RecordType::Session, &record); @@ -807,7 +768,7 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { assert_eq!(decode_session_frames(&decrypted).unwrap(), body); let wrong_header = SessionHeader { - connection_id: ConnectionId::from_data([0x99; ConnectionId::SIZE]), + connection_id: ConnectionId([0x99; ConnectionId::SIZE]), seq: header.seq, }; assert_eq!( @@ -825,37 +786,6 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { ); } -#[test] -fn session_varint_fields_expand_at_expected_boundaries() { - let short_header = SessionHeader { - connection_id: ConnectionId::from_data([0x11; ConnectionId::SIZE]), - seq: record_seq(63), - }; - let long_header = SessionHeader { - connection_id: ConnectionId::from_data([0x11; ConnectionId::SIZE]), - seq: record_seq(64), - }; - - assert_eq!(short_header.encode_vec().len(), ConnectionId::SIZE + 1); - assert_eq!(long_header.encode_vec().len(), ConnectionId::SIZE + 2); - - let frame = StreamData { - stream_id: stream_id(64), - offset: varint(16_384), - header: None, - fin: true, - bytes: b"abc".to_vec(), - }; - let encoded = frame.encode_vec(); - - assert_eq!( - StreamData::decode_exact(encoded.as_slice()) - .unwrap() - .into_owned(), - frame - ); -} - #[test] fn protocol_record_size_breakdown() { fn print_size(label: &str, size: usize) { From b4284cb7b89b31b9b3c9d5cd066f48b0b8b4dbe7 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 08:03:42 -0400 Subject: [PATCH 10/59] fsm: cleanup --- ql-fsm/src/handshake/ik.rs | 20 ++++++++++-------- ql-fsm/src/handshake/kk.rs | 20 ++++++++++-------- ql-fsm/src/handshake/mod.rs | 15 ++++++++++++++ ql-fsm/src/handshake/xx.rs | 20 ++++++++++-------- ql-fsm/src/state.rs | 41 +++++-------------------------------- 5 files changed, 53 insertions(+), 63 deletions(-) diff --git a/ql-fsm/src/handshake/ik.rs b/ql-fsm/src/handshake/ik.rs index 41bc6192..327891a5 100644 --- a/ql-fsm/src/handshake/ik.rs +++ b/ql-fsm/src/handshake/ik.rs @@ -1,10 +1,10 @@ use ql_wire::{self as wire, Ik1, Ik2, PeerBundle, QlCrypto, QlHandshakeRecord}; use super::{ - emit_peer_status, enqueue_handshake, finish_handshake, reset_connected_session_if_needed, + emit_peer_status, enqueue_handshake, establish_session, reset_connected_session_if_needed, }; use crate::{ - state::{IkInitiatorState, LinkState, SessionTransport}, + state::{InitiatorState, LinkState}, QlFsm, ReceiveError, }; @@ -18,7 +18,7 @@ pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle ); let message = handshake.write_1(crypto, meta).unwrap(); - fsm.state.link = LinkState::IkInitiator(IkInitiatorState { + fsm.state.link = LinkState::IkInitiator(InitiatorState { handshake_id: meta.handshake_id, initial_ephemeral: message.ephemeral.clone(), handshake, @@ -59,12 +59,13 @@ pub fn handle_ik1( let outbound = handshake .write_2(crypto, message.meta) .map_err(ReceiveError::InvalidIkHandshake)?; - let (transport, remote_bundle) = SessionTransport::from_finalized( + establish_session( + fsm, + message.meta.handshake_id, handshake .finalize(crypto) .map_err(ReceiveError::InvalidIkHandshake)?, - ); - finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle)?; + )?; fsm.state.handshake = None; enqueue_handshake(fsm, QlHandshakeRecord::Ik2(outbound)); Ok(()) @@ -93,13 +94,14 @@ pub fn handle_ik2( let LinkState::IkInitiator(state) = fsm.state.link.take() else { unreachable!("active IK initiator was checked above"); }; - let (transport, remote_bundle) = SessionTransport::from_finalized( + establish_session( + fsm, + message.meta.handshake_id, state .handshake .finalize(crypto) .map_err(ReceiveError::InvalidIkHandshake)?, - ); - finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle) + ) } pub fn should_ignore_inbound(fsm: &QlFsm, message: &Ik1) -> bool { diff --git a/ql-fsm/src/handshake/kk.rs b/ql-fsm/src/handshake/kk.rs index 683a77b7..21eaa787 100644 --- a/ql-fsm/src/handshake/kk.rs +++ b/ql-fsm/src/handshake/kk.rs @@ -1,10 +1,10 @@ use ql_wire::{self as wire, Kk1, Kk2, PeerBundle, QlCrypto, QlHandshakeRecord}; use super::{ - emit_peer_status, enqueue_handshake, finish_handshake, reset_connected_session_if_needed, + emit_peer_status, enqueue_handshake, establish_session, reset_connected_session_if_needed, }; use crate::{ - state::{KkInitiatorState, LinkState, SessionTransport}, + state::{InitiatorState, LinkState}, QlFsm, ReceiveError, }; @@ -18,7 +18,7 @@ pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle ); let message = handshake.write_1(crypto, meta).unwrap(); - fsm.state.link = LinkState::KkInitiator(KkInitiatorState { + fsm.state.link = LinkState::KkInitiator(InitiatorState { handshake_id: meta.handshake_id, initial_ephemeral: message.ephemeral.clone(), handshake, @@ -58,12 +58,13 @@ pub fn handle_kk1( let outbound = handshake .write_2(crypto, message.meta) .map_err(ReceiveError::InvalidKkHandshake)?; - let (transport, remote_bundle) = SessionTransport::from_finalized( + establish_session( + fsm, + message.meta.handshake_id, handshake .finalize(crypto) .map_err(ReceiveError::InvalidKkHandshake)?, - ); - finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle)?; + )?; fsm.state.handshake = None; enqueue_handshake(fsm, QlHandshakeRecord::Kk2(outbound)); Ok(()) @@ -92,13 +93,14 @@ pub fn handle_kk2( let LinkState::KkInitiator(state) = fsm.state.link.take() else { unreachable!("active KK initiator was checked above"); }; - let (transport, remote_bundle) = SessionTransport::from_finalized( + establish_session( + fsm, + message.meta.handshake_id, state .handshake .finalize(crypto) .map_err(ReceiveError::InvalidKkHandshake)?, - ); - finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle) + ) } pub fn should_ignore_inbound(fsm: &QlFsm, message: &Kk1) -> bool { diff --git a/ql-fsm/src/handshake/mod.rs b/ql-fsm/src/handshake/mod.rs index 8187eee5..2e36ee65 100644 --- a/ql-fsm/src/handshake/mod.rs +++ b/ql-fsm/src/handshake/mod.rs @@ -136,6 +136,21 @@ pub fn finish_handshake( Ok(()) } +pub fn establish_session( + fsm: &mut QlFsm, + handshake_id: HandshakeId, + finalized: wire::FinalizedHandshake, +) -> Result<(), ReceiveError> { + let transport = SessionTransport { + tx_key: finalized.tx_key, + rx_key: finalized.rx_key, + tx_connection_id: finalized.tx_connection_id, + rx_connection_id: finalized.rx_connection_id, + remote_transport_params: finalized.remote_transport_params, + }; + finish_handshake(fsm, handshake_id, transport, finalized.remote_bundle) +} + pub fn reset_connected_session_if_needed(fsm: &mut QlFsm) { if matches!(fsm.state.link, LinkState::Connected(_)) { fsm.state.link = LinkState::Idle; diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs index 34b511ac..4aa634ef 100644 --- a/ql-fsm/src/handshake/xx.rs +++ b/ql-fsm/src/handshake/xx.rs @@ -1,10 +1,10 @@ use ql_wire::{self as wire, PairingToken, QlCrypto, QlHandshakeRecord, Xx1, Xx2, Xx3, Xx4, QID}; use super::{ - emit_peer_status, enqueue_handshake, finish_handshake, reset_connected_session_if_needed, + emit_peer_status, enqueue_handshake, establish_session, reset_connected_session_if_needed, }; use crate::{ - state::{LinkState, SessionTransport, XxInitiatorState, XxResponderState}, + state::{InitiatorState, LinkState, XxResponderState}, QlFsm, ReceiveError, }; @@ -24,7 +24,7 @@ pub fn start_initiator( ); let message = handshake.write_1(crypto, meta).unwrap(); - fsm.state.link = LinkState::XxInitiator(XxInitiatorState { + fsm.state.link = LinkState::XxInitiator(InitiatorState { handshake_id: meta.handshake_id, initial_ephemeral: message.ephemeral.clone(), handshake, @@ -140,13 +140,14 @@ pub fn handle_xx3( .map_err(ReceiveError::InvalidXxHandshake)?; fsm.state.handshake = None; enqueue_handshake(fsm, QlHandshakeRecord::Xx4(outbound)); - let (transport, remote_bundle) = SessionTransport::from_finalized( + establish_session( + fsm, + message.meta.handshake_id, state .handshake .finalize(crypto) .map_err(ReceiveError::InvalidXxHandshake)?, - ); - finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle) + ) } pub fn handle_xx4( @@ -172,13 +173,14 @@ pub fn handle_xx4( let LinkState::XxInitiator(state) = fsm.state.link.take() else { unreachable!("active XX initiator was checked above"); }; - let (transport, remote_bundle) = SessionTransport::from_finalized( + establish_session( + fsm, + message.meta.handshake_id, state .handshake .finalize(crypto) .map_err(ReceiveError::InvalidXxHandshake)?, - ); - finish_handshake(fsm, message.meta.handshake_id, transport, remote_bundle) + ) } pub fn disarm_pairing(fsm: &mut QlFsm) { diff --git a/ql-fsm/src/state.rs b/ql-fsm/src/state.rs index c7657719..d57661bd 100644 --- a/ql-fsm/src/state.rs +++ b/ql-fsm/src/state.rs @@ -25,27 +25,12 @@ pub struct SessionTransport { pub remote_transport_params: TransportParams, } -impl SessionTransport { - pub fn from_finalized(finalized: ql_wire::FinalizedHandshake) -> (Self, PeerBundle) { - ( - Self { - tx_key: finalized.tx_key, - rx_key: finalized.rx_key, - tx_connection_id: finalized.tx_connection_id, - rx_connection_id: finalized.rx_connection_id, - remote_transport_params: finalized.remote_transport_params, - }, - finalized.remote_bundle, - ) - } -} - #[allow(clippy::large_enum_variant)] pub enum LinkState { Idle, - IkInitiator(IkInitiatorState), - KkInitiator(KkInitiatorState), - XxInitiator(XxInitiatorState), + IkInitiator(InitiatorState), + KkInitiator(InitiatorState), + XxInitiator(InitiatorState), XxResponder(XxResponderState), Connected(ConnectedState), } @@ -57,24 +42,8 @@ pub struct ConnectedState { } #[derive(Debug, Clone)] -pub struct IkInitiatorState { - pub handshake: IkHandshake, - pub handshake_id: HandshakeId, - pub deadline: Instant, - pub initial_ephemeral: EphemeralPublicKey, -} - -#[derive(Debug, Clone)] -pub struct KkInitiatorState { - pub handshake: KkHandshake, - pub handshake_id: HandshakeId, - pub deadline: Instant, - pub initial_ephemeral: EphemeralPublicKey, -} - -#[derive(Debug, Clone)] -pub struct XxInitiatorState { - pub handshake: XxHandshake, +pub struct InitiatorState { + pub handshake: H, pub handshake_id: HandshakeId, pub deadline: Instant, pub initial_ephemeral: EphemeralPublicKey, From c388687be3542713dd12be9791124806505006c5 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 08:11:40 -0400 Subject: [PATCH 11/59] wire: smaller ConnectionId --- ql-wire/src/header.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index a186fdcf..b5d9cf16 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -10,7 +10,7 @@ pub struct SessionHeader { crate::varint_wrapper!(RecordSeq); -crate::array_wrapper!(ConnectionId, 16); +crate::array_wrapper!(ConnectionId, 8); impl SessionHeader { pub const MAX_ENCODED_LEN: usize = ConnectionId::SIZE + RecordSeq::MAX_ENCODED_LEN; From d3f3a2c6b227504113676ad0d0a9dad7f96ec9b2 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 09:15:40 -0400 Subject: [PATCH 12/59] ql-rpc: scope routes by service id --- ql-rpc/src/lib.rs | 2 ++ ql-rpc/src/router/builder.rs | 43 ++++++++++++++++++-------------- ql-rpc/src/router/mod.rs | 37 ++++++++++++++++++++------- ql-rpc/src/rpc/mod.rs | 7 ++++-- ql-rpc/src/rpc/progress/codec.rs | 5 +++- ql-rpc/src/service_id.rs | 19 ++++++++++++++ ql-rpc/src/stream.rs | 3 ++- ql-runtime/src/rpc/adapter.rs | 14 +++++++++-- ql-runtime/src/rpc/mod.rs | 28 +++++++++++++++------ ql-runtime/src/tests/rpc.rs | 11 +++++++- 10 files changed, 126 insertions(+), 43 deletions(-) create mode 100644 ql-rpc/src/service_id.rs diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index efea0250..04507a16 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -9,6 +9,7 @@ mod framed_value; mod route_id; mod router; mod rpc; +mod service_id; mod stream; pub use chunk_queue::ChunkQueue; @@ -18,6 +19,7 @@ use framed_value::*; pub use route_id::RouteId; pub use router::*; pub use rpc::*; +pub use service_id::ServiceId; pub use stream::*; #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] diff --git a/ql-rpc/src/router/builder.rs b/ql-rpc/src/router/builder.rs index b59a84e6..d3dc3d55 100644 --- a/ql-rpc/src/router/builder.rs +++ b/ql-rpc/src/router/builder.rs @@ -11,6 +11,7 @@ use crate::{ request::{server::*, Request as RequestRpc}, subscription::{server::*, Subscription as SubscriptionRpc}, upload::{server::*, Upload as UploadRpc}, + RouteKey, }; pub struct LocalRoutes; @@ -50,7 +51,7 @@ where } pub fn build(mut self, state: S) -> Router { - self.routes.sort_by_key(|entry| entry.route_id); + self.routes.sort_by_key(|entry| entry.key); self.routes.shrink_to_fit(); Router { config: self.config, @@ -60,11 +61,15 @@ where } } - fn add_route(mut self, route_id: crate::RouteId, route: RouteFn) -> Self { - if self.routes.iter().any(|entry| entry.route_id == route_id) { - panic!("duplicate rpc route {}", route_id.into_inner()); + fn add_route(mut self, key: RouteKey, route: RouteFn) -> Self { + if self.routes.iter().any(|entry| entry.key == key) { + panic!( + "duplicate rpc route {} for service {:?}", + key.route_id.into_inner(), + key.service_id.into_inner() + ); } - self.routes.push(RouteEntry::new(route_id, route)); + self.routes.push(RouteEntry::new(key, route)); self } } @@ -79,7 +84,7 @@ where M: RequestRpc + 'static, S: RequestHandlerLocal + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_request_inner::( state, @@ -97,7 +102,7 @@ where M: NotificationRpc + 'static, S: NotificationHandlerLocal + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_notification_inner::( state, @@ -115,7 +120,7 @@ where M: DuplexRpc + 'static, S: DuplexHandlerLocal + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_duplex_inner::( state, @@ -132,7 +137,7 @@ where M: DownloadRpc + 'static, S: DownloadHandlerLocal + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_download_inner::( state, @@ -150,7 +155,7 @@ where M: SubscriptionRpc + 'static, S: SubscriptionHandlerLocal + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_subscription_inner::( state, @@ -168,7 +173,7 @@ where M: ProgressRpc + 'static, S: ProgressHandlerLocal + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_progress_inner::( state, @@ -186,7 +191,7 @@ where M: UploadRpc + 'static, S: UploadHandlerLocal + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_upload_inner::( state, @@ -213,7 +218,7 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_request_inner::( state, @@ -234,7 +239,7 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_notification_inner::( state, @@ -256,7 +261,7 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_duplex_inner::( state, @@ -276,7 +281,7 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_download_inner::( state, @@ -297,7 +302,7 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_subscription_inner::( state, @@ -318,7 +323,7 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_progress_inner::( state, @@ -339,7 +344,7 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(M::ROUTE, |spawner, state, config, stream| { + self.add_route(RouteKey::new::(), |spawner, state, config, stream| { let (reader, writer) = stream.split(); spawner.spawn(handle_upload_inner::( state, diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index 31e973ac..3e9cb67b 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -1,4 +1,4 @@ -use crate::{RouteId, StreamCloseCode}; +use crate::{RouteId, ServiceId, StreamCloseCode}; mod builder; mod config; @@ -30,11 +30,27 @@ where routes: Vec>, } +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +pub struct RouteKey { + pub service_id: ServiceId, + pub route_id: RouteId, +} + +impl RouteKey { + pub const fn new() -> Self { + Self { + service_id: R::SERVICE, + route_id: R::ROUTE, + } + } + +} + struct RouteEntry where Sp: Spawner, { - route_id: RouteId, + key: RouteKey, route: RouteFn, } @@ -42,8 +58,8 @@ impl RouteEntry where Sp: Spawner, { - fn new(route_id: RouteId, route: RouteFn) -> Self { - Self { route_id, route } + fn new(key: RouteKey, route: RouteFn) -> Self { + Self { key, route } } } @@ -68,11 +84,10 @@ where } pub fn handle(&self, stream: St) -> Option<(RouteId, Sp::Handle)> { + let service_id = stream.service_id()?; let route_id = stream.route_id()?; - let Ok(index) = self - .routes - .binary_search_by_key(&route_id, |entry| entry.route_id) - else { + let key = RouteKey { service_id, route_id }; + let Ok(index) = self.routes.binary_search_by_key(&key, |entry| entry.key) else { close_stream(stream, StreamCloseCode::UNKNOWN_ROUTE); return None; }; @@ -84,6 +99,10 @@ where } pub fn route_ids(&self) -> impl ExactSizeIterator + '_ { - self.routes.iter().map(|entry| entry.route_id) + self.routes.iter().map(|entry| entry.key.route_id) + } + + pub fn route_keys(&self) -> impl ExactSizeIterator + '_ { + self.routes.iter().map(|entry| entry.key) } } diff --git a/ql-rpc/src/rpc/mod.rs b/ql-rpc/src/rpc/mod.rs index 2d84f050..0c7d9788 100644 --- a/ql-rpc/src/rpc/mod.rs +++ b/ql-rpc/src/rpc/mod.rs @@ -5,7 +5,7 @@ //! route dispatch uses [`crate::RouteId`] and the submodules provide the matching //! client and server helpers for encoding, decoding, and handler glue -use crate::RouteId; +use crate::{RouteId, ServiceId}; pub mod download; pub mod duplex; @@ -18,7 +18,10 @@ pub mod upload; mod utils; pub trait Route { - /// route used to dispatch this rpc family + /// service used to scope this rpc route. + const SERVICE: ServiceId; + + /// route used to dispatch this rpc family within [`Self::SERVICE`]. const ROUTE: RouteId; } diff --git a/ql-rpc/src/rpc/progress/codec.rs b/ql-rpc/src/rpc/progress/codec.rs index a0dc1b8c..eb4a4b31 100644 --- a/ql-rpc/src/rpc/progress/codec.rs +++ b/ql-rpc/src/rpc/progress/codec.rs @@ -93,11 +93,14 @@ mod tests { use bytes::Bytes; use super::{encode_progress, encode_response, ReadStep, ResponseReader}; - use crate::{progress::Progress, Route, RouteId}; + use crate::{progress::Progress, Route, RouteId, ServiceId}; + + const TEST_SERVICE: ServiceId = ServiceId::from_bytes([7; 16]); struct Watch; impl Route for Watch { + const SERVICE: ServiceId = TEST_SERVICE; const ROUTE: RouteId = RouteId::from_u32(11); } diff --git a/ql-rpc/src/service_id.rs b/ql-rpc/src/service_id.rs new file mode 100644 index 00000000..62b25983 --- /dev/null +++ b/ql-rpc/src/service_id.rs @@ -0,0 +1,19 @@ +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct ServiceId(pub [u8; 16]); + +impl ServiceId { + pub const fn from_bytes(value: [u8; 16]) -> Self { + Self(value) + } + + pub const fn into_inner(self) -> [u8; 16] { + self.0 + } +} + +impl From<[u8; 16]> for ServiceId { + fn from(value: [u8; 16]) -> Self { + Self::from_bytes(value) + } +} diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index f6174efd..0344f3a1 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -5,13 +5,14 @@ use std::{ use bytes::Bytes; -use crate::{RouteId, StreamCloseCode}; +use crate::{RouteId, ServiceId, StreamCloseCode}; pub trait RpcStream { type Error: StreamError; type Reader: RpcRead; type Writer: RpcWrite; + fn service_id(&self) -> Option; fn route_id(&self) -> Option; fn split(self) -> (Self::Reader, Self::Writer); } diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index a7347602..b592a5c3 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -1,8 +1,10 @@ use std::task::{Context, Poll}; use bytes::Bytes; -use ql_rpc::{RouteId, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError}; -use ql_wire::{RouteId as WireRouteId, StreamCloseCode as WireStreamCloseCode}; +use ql_rpc::{RouteId, RpcRead, RpcStream, RpcWrite, ServiceId, StreamCloseCode, StreamError}; +use ql_wire::{ + RouteId as WireRouteId, ServiceId as WireServiceId, StreamCloseCode as WireStreamCloseCode, +}; use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; @@ -11,6 +13,10 @@ impl RpcStream for QlStream { type Reader = StreamReader; type Writer = StreamWriter; + fn service_id(&self) -> Option { + Some(ServiceId::from_bytes(self.service_id.0)) + } + fn route_id(&self) -> Option { let route_id = u32::try_from(self.route_id.into_inner()).ok()?; Some(RouteId::from_u32(route_id)) @@ -61,6 +67,10 @@ pub(super) fn to_wire_route_id(route_id: RouteId) -> WireRouteId { WireRouteId::from_u32(route_id.into_inner()) } +pub(super) fn to_wire_service_id(service_id: ServiceId) -> WireServiceId { + WireServiceId(service_id.into_inner()) +} + pub(super) fn to_wire_close_code(code: StreamCloseCode) -> WireStreamCloseCode { WireStreamCloseCode(code.into_inner()) } diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index d8be02c5..b05f2a47 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -9,6 +9,7 @@ mod subscription; mod upload; use bytes::Bytes; +use ql_fsm::OpenStreamParams; use ql_rpc::{ download::{self as rpc_download, Download as DownloadRpc}, duplex::{self as rpc_duplex, Duplex as DuplexRpc}, @@ -35,7 +36,7 @@ impl RpcHandle { notification::encode_notification::(event, &mut payload); let mut stream = self .inner - .open_stream(adapter::to_wire_route_id(M::ROUTE)) + .open_stream(open_stream_params(M::SERVICE, M::ROUTE)) .await?; stream.reader.close(ql_wire::StreamCloseCode::CANCELLED); stream.writer.write(Bytes::from(payload)).await?; @@ -49,7 +50,7 @@ impl RpcHandle { { let mut payload = Vec::new(); request::encode_request::(request, &mut payload); - let response = self.start_request(M::ROUTE, payload).await?; + let response = self.start_request(M::SERVICE, M::ROUTE, payload).await?; Ok(request::read_response::(response).await?) } @@ -62,7 +63,7 @@ impl RpcHandle { { let mut payload = Vec::new(); rpc_subscription::encode_request::(request, &mut payload); - let response = self.start_request(M::ROUTE, payload).await?; + let response = self.start_request(M::SERVICE, M::ROUTE, payload).await?; Ok(Subscription { inner: rpc_subscription::SubscriptionCall::new(response), }) @@ -77,7 +78,7 @@ impl RpcHandle { { let mut payload = Vec::new(); rpc_download::encode_request::(request, &mut payload); - let response = self.start_request(M::ROUTE, payload).await?; + let response = self.start_request(M::SERVICE, M::ROUTE, payload).await?; Ok(DownloadCall { inner: rpc_download::DownloadCall::new(response), }) @@ -92,7 +93,7 @@ impl RpcHandle { { let mut payload = Vec::new(); rpc_progress::encode_request::(request, &mut payload); - let response = self.start_request(M::ROUTE, payload).await?; + let response = self.start_request(M::SERVICE, M::ROUTE, payload).await?; Ok(ProgressCall { inner: rpc_progress::ProgressCall::new(response), }) @@ -106,7 +107,7 @@ impl RpcHandle { rpc_upload::encode_request::(request, &mut payload); let mut stream = self .inner - .open_stream(adapter::to_wire_route_id(M::ROUTE)) + .open_stream(open_stream_params(M::SERVICE, M::ROUTE)) .await?; stream.writer.write(Bytes::from(payload)).await?; Ok(UploadCall { @@ -120,7 +121,7 @@ impl RpcHandle { { let stream = self .inner - .open_stream(adapter::to_wire_route_id(M::ROUTE)) + .open_stream(open_stream_params(M::SERVICE, M::ROUTE)) .await?; Ok(DuplexCall { sender: DuplexSender { @@ -140,15 +141,26 @@ impl RpcHandle { async fn start_request( &self, + service_id: ql_rpc::ServiceId, route_id: ql_rpc::RouteId, payload: Vec, ) -> Result> { let mut stream = self .inner - .open_stream(adapter::to_wire_route_id(route_id)) + .open_stream(open_stream_params(service_id, route_id)) .await?; stream.writer.write(Bytes::from(payload)).await?; stream.writer.finish().await?; Ok(stream.reader) } } + +fn open_stream_params( + service_id: ql_rpc::ServiceId, + route_id: ql_rpc::RouteId, +) -> OpenStreamParams { + OpenStreamParams { + service_id: adapter::to_wire_service_id(service_id), + route_id: adapter::to_wire_route_id(route_id), + } +} diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index 3244c587..3d51d447 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -12,7 +12,7 @@ use futures_lite::StreamExt; use ql_rpc::{ DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, NotificationHandlerLocal, ProgressHandlerLocal, ProgressResponder, RequestHandler, - RequestHandlerLocal, Response, RouteId, SendSpawner, Spawner, StreamCloseCode, + RequestHandlerLocal, Response, RouteId, SendSpawner, ServiceId, Spawner, StreamCloseCode, SubscriptionHandlerLocal, SubscriptionResponder, UploadHandlerLocal, UploadReader, UploadResponder, }; @@ -20,6 +20,8 @@ use ql_rpc::{ use super::*; use crate::{rpc::RpcError, QlStream, StreamWriter}; +const TEST_SERVICE: ServiceId = ServiceId::from_bytes([7; 16]); + #[derive(Debug, Clone, Copy)] struct TokioLocalSpawner; @@ -55,6 +57,7 @@ impl SendSpawner for TokioSendSpawner { struct Echo; impl ql_rpc::Route for Echo { + const SERVICE: ServiceId = TEST_SERVICE; const ROUTE: RouteId = RouteId::from_u32(51); } @@ -68,6 +71,7 @@ impl ql_rpc::request::Request for Echo { struct Feed; impl ql_rpc::Route for Feed { + const SERVICE: ServiceId = TEST_SERVICE; const ROUTE: RouteId = RouteId::from_u32(52); } @@ -80,6 +84,7 @@ impl ql_rpc::subscription::Subscription for Feed { struct Notice; impl ql_rpc::Route for Notice { + const SERVICE: ServiceId = TEST_SERVICE; const ROUTE: RouteId = RouteId::from_u32(521); } @@ -91,6 +96,7 @@ impl ql_rpc::notification::Notification for Notice { struct Download; impl ql_rpc::Route for Download { + const SERVICE: ServiceId = TEST_SERVICE; const ROUTE: RouteId = RouteId::from_u32(53); } @@ -104,6 +110,7 @@ impl ql_rpc::progress::Progress for Download { struct BlobDownload; impl ql_rpc::Route for BlobDownload { + const SERVICE: ServiceId = TEST_SERVICE; const ROUTE: RouteId = RouteId::from_u32(54); } @@ -117,6 +124,7 @@ impl ql_rpc::download::Download for BlobDownload { struct BlobUpload; impl ql_rpc::Route for BlobUpload { + const SERVICE: ServiceId = TEST_SERVICE; const ROUTE: RouteId = RouteId::from_u32(55); } @@ -130,6 +138,7 @@ impl ql_rpc::upload::Upload for BlobUpload { struct Chat; impl ql_rpc::Route for Chat { + const SERVICE: ServiceId = TEST_SERVICE; const ROUTE: RouteId = RouteId::from_u32(56); } From 2b065956cae17f46fb90eb710473762700e3b1c0 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 09:23:17 -0400 Subject: [PATCH 13/59] ql-common: extract shared primitives --- Cargo.lock | 6 + Cargo.toml | 2 + ql-common/Cargo.toml | 8 + ql-common/src/lib.rs | 57 +++++++ ql-common/src/varint.rs | 177 ++++++++++++++++++++ ql-fsm/src/session/ack_tracker.rs | 8 +- ql-fsm/src/session/mod.rs | 6 +- ql-fsm/src/session/remote_stream_history.rs | 1 + ql-fsm/src/session/stream_parity.rs | 4 +- ql-fsm/src/session/tests.rs | 6 +- ql-rpc/Cargo.toml | 1 + ql-rpc/src/lib.rs | 30 +--- ql-rpc/src/route_id.rs | 19 --- ql-rpc/src/router/builder.rs | 4 +- ql-rpc/src/router/mod.rs | 6 +- ql-rpc/src/rpc/progress/codec.rs | 2 +- ql-rpc/src/service_id.rs | 19 --- ql-runtime/src/rpc/adapter.rs | 30 +--- ql-runtime/src/rpc/mod.rs | 6 +- ql-runtime/src/tests/rpc.rs | 2 +- ql-wire/Cargo.toml | 1 + ql-wire/src/crypto.rs | 2 +- ql-wire/src/encrypted/ack.rs | 7 +- ql-wire/src/encrypted/builder.rs | 2 +- ql-wire/src/encrypted/mod.rs | 13 +- ql-wire/src/encrypted/service_id.rs | 34 ---- ql-wire/src/encrypted/stream_close.rs | 17 +- ql-wire/src/encrypted/stream_data.rs | 8 +- ql-wire/src/handshake/mod.rs | 6 +- ql-wire/src/handshake/pairing.rs | 4 +- ql-wire/src/header.rs | 6 +- ql-wire/src/identity.rs | 6 +- ql-wire/src/lib.rs | 8 +- ql-wire/src/macros.rs | 51 ++---- ql-wire/src/qid.rs | 24 ++- ql-wire/src/tests.rs | 2 +- ql-wire/src/varint.rs | 131 +-------------- 37 files changed, 349 insertions(+), 367 deletions(-) create mode 100644 ql-common/Cargo.toml create mode 100644 ql-common/src/lib.rs create mode 100644 ql-common/src/varint.rs delete mode 100644 ql-rpc/src/route_id.rs delete mode 100644 ql-rpc/src/service_id.rs delete mode 100644 ql-wire/src/encrypted/service_id.rs diff --git a/Cargo.lock b/Cargo.lock index c1e30d37..008d5fed 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1050,6 +1050,10 @@ dependencies = [ "syn", ] +[[package]] +name = "ql-common" +version = "0.1.0" + [[package]] name = "ql-fsm" version = "0.1.0" @@ -1065,6 +1069,7 @@ name = "ql-rpc" version = "0.1.0" dependencies = [ "bytes", + "ql-common", "trait-variant", ] @@ -1095,6 +1100,7 @@ dependencies = [ "getrandom 0.2.16", "libcrux-aesgcm", "libcrux-ml-kem", + "ql-common", "sha2", ] diff --git a/Cargo.toml b/Cargo.toml index 83ac135d..7522bcd0 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -3,6 +3,7 @@ resolver = "2" members = [ "backup-shard", "btp", + "ql-common", "ql-fsm", "ql-rpc", "ql-runtime", @@ -22,6 +23,7 @@ rkyv = { version = "0.8" } # workspace crates backup-shard = { path = "backup-shard" } btp = { path = "btp" } +ql-common = { path = "ql-common" } ql-fsm = { path = "ql-fsm" } ql-rpc = { path = "ql-rpc" } ql-wire = { path = "ql-wire" } diff --git a/ql-common/Cargo.toml b/ql-common/Cargo.toml new file mode 100644 index 00000000..4a44b3b4 --- /dev/null +++ b/ql-common/Cargo.toml @@ -0,0 +1,8 @@ +[package] +name = "ql-common" +version = "0.1.0" +edition = "2021" +description = "QuantumLink shared primitive types" +license = "Proprietary" + +[dependencies] diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs new file mode 100644 index 00000000..c95f82a0 --- /dev/null +++ b/ql-common/src/lib.rs @@ -0,0 +1,57 @@ +//! Shared QuantumLink primitive types. + +mod varint; +pub use varint::*; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct StreamCloseCode(pub u16); + +impl StreamCloseCode { + /// the stream was aborted intentionally before graceful completion + pub const CANCELLED: Self = Self(0); + /// local internal error + pub const INTERNAL: Self = Self(1); + /// request was refused + pub const REFUSED: Self = Self(2); + /// operation timed out + pub const TIMEOUT: Self = Self(3); + /// configured limit was exceeded + pub const LIMIT: Self = Self(4); + /// route identifier was unknown + pub const UNKNOWN_ROUTE: Self = Self(5); +} + +impl std::fmt::Display for StreamCloseCode { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct QID(pub [u8; Self::SIZE]); + +impl QID { + pub const SIZE: usize = 16; +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct ServiceId(pub [u8; 16]); + +impl ServiceId { + pub const SIZE: usize = size_of::(); +} + +impl std::fmt::Display for ServiceId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + for byte in self.0 { + write!(f, "{byte:02x}")?; + } + Ok(()) + } +} + +varint_wrapper!(RouteId); +varint_wrapper!(StreamId); diff --git a/ql-common/src/varint.rs b/ql-common/src/varint.rs new file mode 100644 index 00000000..8ac64028 --- /dev/null +++ b/ql-common/src/varint.rs @@ -0,0 +1,177 @@ +use core::fmt; + +/// An integer less than 2^62 encoded with QUIC variable-length integer rules. +#[derive(Default, Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)] +pub struct VarInt(u64); + +impl VarInt { + /// The largest representable value. + pub const MAX: Self = Self((1u64 << 62) - 1); + /// The largest encoded value length. + pub const MAX_SIZE: usize = 8; + pub const MIN_SIZE: usize = 1; + + /// Construct a `VarInt` infallibly from a `u32`. + pub const fn from_u32(x: u32) -> Self { + Self(x as u64) + } + + /// Construct a `VarInt` from a `u64`. + pub const fn from_u64(x: u64) -> Result { + if x < (1u64 << 62) { + Ok(Self(x)) + } else { + Err(VarIntBoundsExceeded) + } + } + + /// Create a `VarInt` without checking the bounds. + /// + /// # Safety + /// + /// `x` must be less than 2^62. + pub const unsafe fn from_u64_unchecked(x: u64) -> Self { + Self(x) + } + + /// Extract the inner integer value. + pub const fn into_inner(self) -> u64 { + self.0 + } + + /// Return the number of bytes required to encode this value. + pub const fn size(self) -> usize { + let x = self.0; + if x < (1u64 << 6) { + 1 + } else if x < (1u64 << 14) { + 2 + } else if x < (1u64 << 30) { + 4 + } else { + 8 + } + } +} + +impl From for u64 { + fn from(value: VarInt) -> Self { + value.0 + } +} + +impl From for VarInt { + fn from(value: u8) -> Self { + Self(value.into()) + } +} + +impl From for VarInt { + fn from(value: u16) -> Self { + Self(value.into()) + } +} + +impl From for VarInt { + fn from(value: u32) -> Self { + Self(value.into()) + } +} + +impl TryFrom for VarInt { + type Error = VarIntBoundsExceeded; + + fn try_from(value: u64) -> Result { + Self::from_u64(value) + } +} + +impl TryFrom for VarInt { + type Error = VarIntBoundsExceeded; + + fn try_from(value: u128) -> Result { + Self::from_u64(value.try_into().map_err(|_| VarIntBoundsExceeded)?) + } +} + +impl TryFrom for VarInt { + type Error = VarIntBoundsExceeded; + + fn try_from(value: usize) -> Result { + Self::from_u64(value as u64) + } +} + +impl fmt::Debug for VarInt { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(f) + } +} + +impl fmt::Display for VarInt { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(f) + } +} + +#[derive(Debug, Copy, Clone, Eq, PartialEq)] +pub struct VarIntBoundsExceeded; + +impl fmt::Display for VarIntBoundsExceeded { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("value too large for varint encoding") + } +} + +impl std::error::Error for VarIntBoundsExceeded {} + +#[macro_export] +macro_rules! varint_wrapper { + ($name:ident) => { + #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] + #[repr(transparent)] + pub struct $name(pub $crate::VarInt); + + impl $name { + pub const MAX_ENCODED_LEN: usize = $crate::VarInt::MAX_SIZE; + + pub const fn from_u32(value: u32) -> Self { + Self($crate::VarInt::from_u32(value)) + } + + pub const fn from_u64(value: u64) -> Result { + match $crate::VarInt::from_u64(value) { + Ok(v) => Ok(Self(v)), + Err(e) => Err(e), + } + } + + /// Create this wrapper without checking the bounds. + /// + /// # Safety + /// + /// `value` must be less than 2^62. + pub const unsafe fn from_u64_unchecked(value: u64) -> Self { + Self(unsafe { $crate::VarInt::from_u64_unchecked(value) }) + } + } + + impl From<$crate::VarInt> for $name { + fn from(value: $crate::VarInt) -> Self { + Self(value) + } + } + + impl From for $name { + fn from(value: u32) -> Self { + Self::from_u32(value) + } + } + + impl std::fmt::Display for $name { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } + } + }; +} diff --git a/ql-fsm/src/session/ack_tracker.rs b/ql-fsm/src/session/ack_tracker.rs index a75b5c63..b22dac36 100644 --- a/ql-fsm/src/session/ack_tracker.rs +++ b/ql-fsm/src/session/ack_tracker.rs @@ -45,7 +45,7 @@ impl AckTracker { } pub fn insert(&mut self, seq: RecordSeq) -> ReceiveOutcome { - let seq = seq.into_inner(); + let seq = seq.0.into_inner(); let largest_accepted = self.accepted_records.max(); if largest_accepted.is_some_and(|largest| seq < self.accepted_cutoff(largest)) { return ReceiveOutcome::TooOld; @@ -167,8 +167,8 @@ fn to_ack_range(range: std::ops::Range) -> RangeInclusive { } fn from_ack_range(range: RangeInclusive) -> std::ops::Range { - let start = range.start().into_inner(); - let end = range.end().into_inner().checked_add(1).unwrap(); + let start = range.start().0.into_inner(); + let end = range.end().0.into_inner().checked_add(1).unwrap(); start..end } @@ -188,7 +188,7 @@ mod tests { pending_ack .ack .ranges() - .map(|range| (range.start().into_inner(), range.end().into_inner())) + .map(|range| (range.start().0.into_inner(), range.end().0.into_inner())) .collect() } diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index b3821076..df87352d 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -585,7 +585,7 @@ impl SessionFsm { .state .tracked_records .extract_if(.., |_, record| { - record.sent_at.is_some() && ack.contains(record.seq.into_inner()) + record.sent_at.is_some() && ack.contains(record.seq.0.into_inner()) }) .map(|(_, record)| record) .collect::>(); @@ -941,9 +941,10 @@ fn local_stream_was_opened( stream_id: StreamId, ) -> bool { local_parity.matches(stream_id) - && stream_id.into_inner() + && stream_id.0.into_inner() < local_parity .make_stream_id(next_stream_ordinal) + .0 .into_inner() } @@ -1034,6 +1035,7 @@ fn acknowledge_tracked_frame( #[track_caller] fn next_seq(seq: &mut RecordSeq) { *seq = seq + .0 .into_inner() .checked_add(1) .and_then(|next| RecordSeq::from_u64(next).ok()) diff --git a/ql-fsm/src/session/remote_stream_history.rs b/ql-fsm/src/session/remote_stream_history.rs index 76c1e8bb..0d126b18 100644 --- a/ql-fsm/src/session/remote_stream_history.rs +++ b/ql-fsm/src/session/remote_stream_history.rs @@ -28,6 +28,7 @@ impl RemoteStreamHistory { fn stream_ordinal(&self, stream_id: StreamId) -> Option { let delta = stream_id + .0 .into_inner() .checked_sub(u64::from(self.parity.first_stream_id()))?; if delta % 2 != 0 { diff --git a/ql-fsm/src/session/stream_parity.rs b/ql-fsm/src/session/stream_parity.rs index 70f60776..26846d79 100644 --- a/ql-fsm/src/session/stream_parity.rs +++ b/ql-fsm/src/session/stream_parity.rs @@ -23,8 +23,8 @@ impl StreamParity { pub const fn matches(self, stream_id: StreamId) -> bool { match self { - Self::Even => stream_id.into_inner() % 2 == 0, - Self::Odd => stream_id.into_inner() % 2 == 1, + Self::Even => stream_id.0.into_inner() % 2 == 0, + Self::Odd => stream_id.0.into_inner() % 2 == 1, } } diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index bfcbfb18..628efba3 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -465,8 +465,8 @@ fn stream_ids_follow_even_odd_xid_ordering() { .unwrap() .stream_id(); - assert_eq!(even_id.into_inner() % 2, 0); - assert_eq!(odd_id.into_inner() % 2, 1); + assert_eq!(even_id.0.into_inner() % 2, 0); + assert_eq!(odd_id.0.into_inner() % 2, 1); } #[test] @@ -837,7 +837,7 @@ fn sparse_out_of_order_ack_ranges_page_and_quiesce() { for (seq, record) in originals .iter() - .filter(|(seq, _)| seq.into_inner() % 2 == 1) + .filter(|(seq, _)| seq.0.into_inner() % 2 == 1) { let _ = receive_events(&mut receiver, now, *seq, record); } diff --git a/ql-rpc/Cargo.toml b/ql-rpc/Cargo.toml index 836e9d2b..40af473c 100644 --- a/ql-rpc/Cargo.toml +++ b/ql-rpc/Cargo.toml @@ -7,4 +7,5 @@ license = "Proprietary" [dependencies] bytes = { workspace = true } +ql-common = { workspace = true } trait-variant = { version = "0.1" } diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index 04507a16..c3e20364 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -3,44 +3,18 @@ //! QuantumLink RPC protocol mod chunk_queue; -pub(crate) mod codec; +mod codec; mod error; mod framed_value; -mod route_id; mod router; mod rpc; -mod service_id; mod stream; pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; use framed_value::*; -pub use route_id::RouteId; +pub use ql_common::{RouteId, ServiceId, StreamCloseCode}; pub use router::*; pub use rpc::*; -pub use service_id::ServiceId; pub use stream::*; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -#[repr(transparent)] -pub struct StreamCloseCode(pub u16); - -impl StreamCloseCode { - /// operation was cancelled - pub const CANCELLED: Self = Self(0); - /// local internal error - pub const INTERNAL: Self = Self(1); - /// request was refused - pub const REFUSED: Self = Self(2); - /// operation timed out - pub const TIMEOUT: Self = Self(3); - /// configured limit was exceeded - pub const LIMIT: Self = Self(4); - /// route identifier was unknown - pub const UNKNOWN_ROUTE: Self = Self(5); - - pub const fn into_inner(self) -> u16 { - self.0 - } -} diff --git a/ql-rpc/src/route_id.rs b/ql-rpc/src/route_id.rs deleted file mode 100644 index 1b054e74..00000000 --- a/ql-rpc/src/route_id.rs +++ /dev/null @@ -1,19 +0,0 @@ -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -#[repr(transparent)] -pub struct RouteId(pub u32); - -impl RouteId { - pub const fn from_u32(value: u32) -> Self { - Self(value) - } - - pub const fn into_inner(self) -> u32 { - self.0 - } -} - -impl From for RouteId { - fn from(value: u32) -> Self { - Self::from_u32(value) - } -} diff --git a/ql-rpc/src/router/builder.rs b/ql-rpc/src/router/builder.rs index d3dc3d55..764c16b5 100644 --- a/ql-rpc/src/router/builder.rs +++ b/ql-rpc/src/router/builder.rs @@ -65,8 +65,8 @@ where if self.routes.iter().any(|entry| entry.key == key) { panic!( "duplicate rpc route {} for service {:?}", - key.route_id.into_inner(), - key.service_id.into_inner() + key.route_id.0.into_inner(), + key.service_id.0 ); } self.routes.push(RouteEntry::new(key, route)); diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index 3e9cb67b..1ab7de3c 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -43,7 +43,6 @@ impl RouteKey { route_id: R::ROUTE, } } - } struct RouteEntry @@ -86,7 +85,10 @@ where pub fn handle(&self, stream: St) -> Option<(RouteId, Sp::Handle)> { let service_id = stream.service_id()?; let route_id = stream.route_id()?; - let key = RouteKey { service_id, route_id }; + let key = RouteKey { + service_id, + route_id, + }; let Ok(index) = self.routes.binary_search_by_key(&key, |entry| entry.key) else { close_stream(stream, StreamCloseCode::UNKNOWN_ROUTE); return None; diff --git a/ql-rpc/src/rpc/progress/codec.rs b/ql-rpc/src/rpc/progress/codec.rs index eb4a4b31..0b01302c 100644 --- a/ql-rpc/src/rpc/progress/codec.rs +++ b/ql-rpc/src/rpc/progress/codec.rs @@ -95,7 +95,7 @@ mod tests { use super::{encode_progress, encode_response, ReadStep, ResponseReader}; use crate::{progress::Progress, Route, RouteId, ServiceId}; - const TEST_SERVICE: ServiceId = ServiceId::from_bytes([7; 16]); + const TEST_SERVICE: ServiceId = ServiceId([7; 16]); struct Watch; diff --git a/ql-rpc/src/service_id.rs b/ql-rpc/src/service_id.rs deleted file mode 100644 index 62b25983..00000000 --- a/ql-rpc/src/service_id.rs +++ /dev/null @@ -1,19 +0,0 @@ -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -#[repr(transparent)] -pub struct ServiceId(pub [u8; 16]); - -impl ServiceId { - pub const fn from_bytes(value: [u8; 16]) -> Self { - Self(value) - } - - pub const fn into_inner(self) -> [u8; 16] { - self.0 - } -} - -impl From<[u8; 16]> for ServiceId { - fn from(value: [u8; 16]) -> Self { - Self::from_bytes(value) - } -} diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index b592a5c3..c33ac37c 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -2,9 +2,6 @@ use std::task::{Context, Poll}; use bytes::Bytes; use ql_rpc::{RouteId, RpcRead, RpcStream, RpcWrite, ServiceId, StreamCloseCode, StreamError}; -use ql_wire::{ - RouteId as WireRouteId, ServiceId as WireServiceId, StreamCloseCode as WireStreamCloseCode, -}; use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; @@ -14,12 +11,11 @@ impl RpcStream for QlStream { type Writer = StreamWriter; fn service_id(&self) -> Option { - Some(ServiceId::from_bytes(self.service_id.0)) + Some(self.service_id) } fn route_id(&self) -> Option { - let route_id = u32::try_from(self.route_id.into_inner()).ok()?; - Some(RouteId::from_u32(route_id)) + Some(self.route_id) } fn split(self) -> (Self::Reader, Self::Writer) { @@ -39,7 +35,7 @@ impl RpcRead for StreamReader { } fn close(self, code: StreamCloseCode) { - StreamReader::close(self, to_wire_close_code(code)); + StreamReader::close(self, code); } } @@ -59,34 +55,20 @@ impl RpcWrite for StreamWriter { } fn close(self, code: StreamCloseCode) { - StreamWriter::close(self, to_wire_close_code(code)); + StreamWriter::close(self, code); } } -pub(super) fn to_wire_route_id(route_id: RouteId) -> WireRouteId { - WireRouteId::from_u32(route_id.into_inner()) -} - -pub(super) fn to_wire_service_id(service_id: ServiceId) -> WireServiceId { - WireServiceId(service_id.into_inner()) -} - -pub(super) fn to_wire_close_code(code: StreamCloseCode) -> WireStreamCloseCode { - WireStreamCloseCode(code.into_inner()) -} - impl From for QlStreamError { fn from(code: StreamCloseCode) -> Self { - Self::StreamClosed { - code: WireStreamCloseCode(code.into_inner()), - } + Self::StreamClosed { code } } } impl StreamError for QlStreamError { fn close_code(&self) -> Option { match self { - QlStreamError::StreamClosed { code } => Some(StreamCloseCode(code.0)), + QlStreamError::StreamClosed { code } => Some(*code), QlStreamError::NoSession => None, } } diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index b05f2a47..828bf0ee 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -38,7 +38,7 @@ impl RpcHandle { .inner .open_stream(open_stream_params(M::SERVICE, M::ROUTE)) .await?; - stream.reader.close(ql_wire::StreamCloseCode::CANCELLED); + stream.reader.close(ql_rpc::StreamCloseCode::CANCELLED); stream.writer.write(Bytes::from(payload)).await?; stream.writer.finish().await?; Ok(()) @@ -160,7 +160,7 @@ fn open_stream_params( route_id: ql_rpc::RouteId, ) -> OpenStreamParams { OpenStreamParams { - service_id: adapter::to_wire_service_id(service_id), - route_id: adapter::to_wire_route_id(route_id), + service_id, + route_id, } } diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index 3d51d447..10c0b6b8 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -20,7 +20,7 @@ use ql_rpc::{ use super::*; use crate::{rpc::RpcError, QlStream, StreamWriter}; -const TEST_SERVICE: ServiceId = ServiceId::from_bytes([7; 16]); +const TEST_SERVICE: ServiceId = ServiceId([7; 16]); #[derive(Debug, Clone, Copy)] struct TokioLocalSpawner; diff --git a/ql-wire/Cargo.toml b/ql-wire/Cargo.toml index 399846cc..9ba6e3d1 100644 --- a/ql-wire/Cargo.toml +++ b/ql-wire/Cargo.toml @@ -16,6 +16,7 @@ test-utils = [ [dependencies] bytes = { workspace = true } getrandom = { workspace = true, optional = true } +ql-common = { workspace = true } libcrux-aesgcm = { version = "0.0.7", optional = true } libcrux-ml-kem = { version = "0.0.7", optional = true } sha2 = { version = "0.10", optional = true } diff --git a/ql-wire/src/crypto.rs b/ql-wire/src/crypto.rs index 0f00bd80..888a756a 100644 --- a/ql-wire/src/crypto.rs +++ b/ql-wire/src/crypto.rs @@ -3,7 +3,7 @@ use crate::{ ENCRYPTED_MESSAGE_AUTH_SIZE, }; -crate::array_wrapper!(Nonce, 12); +array_wrapper!(Nonce, 12); impl Nonce { pub fn from_counter(counter: u64) -> Self { diff --git a/ql-wire/src/encrypted/ack.rs b/ql-wire/src/encrypted/ack.rs index 7459a608..eaaa8bb3 100644 --- a/ql-wire/src/encrypted/ack.rs +++ b/ql-wire/src/encrypted/ack.rs @@ -46,7 +46,7 @@ impl RecordAck { pub fn ranges(&self) -> RecordAckRangeIter<'_> { RecordAckRangeIter { - largest_acked: self.largest_acked.into_inner(), + largest_acked: self.largest_acked.0.into_inner(), first_range_len: Some(self.first_range_len), previous_start: None, blocks: self.blocks.iter(), @@ -164,6 +164,7 @@ impl codec::WireDecode for RecordAck { { let mut previous_start = ack .largest_acked + .0 .into_inner() .checked_sub(ack.first_range_len.into_inner()) .ok_or(WireError::InvalidPayload)?; @@ -206,8 +207,8 @@ impl RecordAckBuilder { range: RangeInclusive, max_wire_size: usize, ) -> Result { - let start = range.start().into_inner(); - let end = range.end().into_inner(); + let start = range.start().0.into_inner(); + let end = range.end().0.into_inner(); if start > end { return Err(RecordAckRangeError::InvertedRange); } diff --git a/ql-wire/src/encrypted/builder.rs b/ql-wire/src/encrypted/builder.rs index 42933235..a60f8a57 100644 --- a/ql-wire/src/encrypted/builder.rs +++ b/ql-wire/src/encrypted/builder.rs @@ -114,7 +114,7 @@ impl SessionRecordBuilder { seq: self.seq, }; let aad = header.aad(); - let nonce = Nonce::from_counter(self.seq.into_inner()); + let nonce = Nonce::from_counter(self.seq.0.into_inner()); let auth = crypto.aes256_gcm_encrypt( session_key, &nonce, diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index 8a12f145..1f45ac89 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -1,12 +1,11 @@ use crate::{ - codec, encrypted_message::EncryptedMessage, varint_wrapper, BufView, ByteSlice, Nonce, - QlCrypto, Reader, SessionHeader, SessionKey, WireDecode, WireEncode, WireError, + codec, encrypted_message::EncryptedMessage, BufView, ByteSlice, Nonce, QlCrypto, Reader, + RouteId, ServiceId, SessionHeader, SessionKey, StreamId, WireDecode, WireEncode, WireError, }; mod ack; mod builder; mod close; -mod service_id; mod stream_close; mod stream_data; mod stream_window; @@ -14,13 +13,13 @@ mod stream_window; pub use ack::*; pub use builder::*; pub use close::*; -pub use service_id::*; pub use stream_close::*; pub use stream_data::*; pub use stream_window::*; -varint_wrapper!(RouteId); -varint_wrapper!(StreamId); +varint_wrapper_codec!(RouteId); +varint_wrapper_codec!(StreamId); +array_wrapper_codec!(ServiceId); #[derive(Debug, Clone, PartialEq, Eq)] pub enum SessionFrame { @@ -174,7 +173,7 @@ pub fn decrypt_record>( session_key: &SessionKey, ) -> Result { let aad = header.aad(); - let nonce = Nonce::from_counter(header.seq.into_inner()); + let nonce = Nonce::from_counter(header.seq.0.into_inner()); let mut ciphertext = encrypted.ciphertext; if !crypto.aes256_gcm_decrypt( session_key, diff --git a/ql-wire/src/encrypted/service_id.rs b/ql-wire/src/encrypted/service_id.rs deleted file mode 100644 index a05a8438..00000000 --- a/ql-wire/src/encrypted/service_id.rs +++ /dev/null @@ -1,34 +0,0 @@ -use crate::{ByteSlice, Reader, WireDecode, WireEncode, WireError}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -#[repr(transparent)] -pub struct ServiceId(pub [u8; 16]); - -impl ServiceId { - pub const ENCODED_LEN: usize = size_of::(); -} - -impl WireEncode for ServiceId { - fn encoded_len(&self) -> usize { - Self::ENCODED_LEN - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl WireDecode for ServiceId { - fn decode(reader: &mut Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -impl std::fmt::Display for ServiceId { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - for byte in self.0 { - write!(f, "{byte:02x}")?; - } - Ok(()) - } -} diff --git a/ql-wire/src/encrypted/stream_close.rs b/ql-wire/src/encrypted/stream_close.rs index 20ddb879..dffbada9 100644 --- a/ql-wire/src/encrypted/stream_close.rs +++ b/ql-wire/src/encrypted/stream_close.rs @@ -1,5 +1,5 @@ use super::StreamId; -use crate::{codec, ByteSlice, WireEncode, WireError}; +use crate::{codec, ByteSlice, StreamCloseCode, WireEncode, WireError}; /// aborts one or both lanes of a stream with a close code /// @@ -84,15 +84,6 @@ impl codec::WireDecode for CloseTarget { } } -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -#[repr(transparent)] -pub struct StreamCloseCode(pub u16); - -impl StreamCloseCode { - /// the stream was aborted intentionally before graceful completion - pub const CANCELLED: Self = Self(0); -} - impl codec::WireDecode for StreamCloseCode { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self(reader.decode()?)) @@ -108,9 +99,3 @@ impl WireEncode for StreamCloseCode { self.0.encode(out); } } - -impl std::fmt::Display for StreamCloseCode { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.0) - } -} diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index 85dd1baf..f789ad09 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -1,7 +1,9 @@ use bytes::Buf; -use super::{RouteId, ServiceId, StreamId}; -use crate::{codec, BufView, ByteSlice, VarInt, WireDecode, WireEncode, WireError}; +use crate::{ + codec, BufView, ByteSlice, RouteId, ServiceId, StreamId, VarInt, WireDecode, WireEncode, + WireError, +}; /// carries bytes for a stream and may finish that sending direction. #[derive(Debug, Clone, PartialEq, Eq)] @@ -109,7 +111,7 @@ pub struct StreamHeader { } impl StreamHeader { - pub const MAX_WIRE_SIZE: usize = ServiceId::ENCODED_LEN + RouteId::MAX_ENCODED_LEN; + pub const MAX_WIRE_SIZE: usize = ServiceId::SIZE + RouteId::MAX_ENCODED_LEN; } impl WireDecode for StreamHeader { diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index 1ac9aabe..b8f06248 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -1,6 +1,6 @@ use crate::{ - codec, ByteSlice, ConnectionId, HandshakeKind, MlKemCiphertext, MlKemKeyPair, MlKemPublicKey, - Nonce, PeerBundle, QlCrypto, SessionKey, WireDecode, WireEncode, WireError, + codec, derive_qid, ByteSlice, ConnectionId, HandshakeKind, MlKemCiphertext, MlKemKeyPair, + MlKemPublicKey, Nonce, PeerBundle, QlCrypto, SessionKey, WireDecode, WireEncode, WireError, ENCRYPTED_MESSAGE_AUTH_SIZE, QID, }; @@ -467,7 +467,7 @@ fn decrypt_peer_bundle( ) -> Result { let plaintext = symmetric.decrypt_and_hash(crypto, bundle.as_bytes())?; let bundle = PeerBundle::decode_exact(plaintext.as_slice())?; - let peer_qid = QID::derive(crypto, &bundle.mlkem_public_key); + let peer_qid = derive_qid(crypto, &bundle.mlkem_public_key); if peer_qid != bundle.qid { return Err(WireError::InvalidRemoteBundle); } diff --git a/ql-wire/src/handshake/pairing.rs b/ql-wire/src/handshake/pairing.rs index 25227ad9..ea305a0b 100644 --- a/ql-wire/src/handshake/pairing.rs +++ b/ql-wire/src/handshake/pairing.rs @@ -5,7 +5,7 @@ use crate::QlCrypto; const PAIRING_ID_DOMAIN: &[u8] = b"ql-wire:pairing-id:v1"; const PAIRING_PSK_DOMAIN: &[u8] = b"ql-wire:pairing-psk:v1"; -crate::array_wrapper!(PairingToken, 16); +array_wrapper!(PairingToken, 16); impl PairingToken { pub fn id(&self, crypto: &impl QlCrypto) -> PairingId { @@ -29,7 +29,7 @@ impl Display for PairingToken { } } -crate::array_wrapper!(PairingId, 16); +array_wrapper!(PairingId, 16); impl Display for PairingId { fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index b5d9cf16..bfe496f3 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -8,9 +8,11 @@ pub struct SessionHeader { pub seq: RecordSeq, } -crate::varint_wrapper!(RecordSeq); +ql_common::varint_wrapper!(RecordSeq); -crate::array_wrapper!(ConnectionId, 8); +varint_wrapper_codec!(RecordSeq); + +array_wrapper!(ConnectionId, 8); impl SessionHeader { pub const MAX_ENCODED_LEN: usize = ConnectionId::SIZE + RecordSeq::MAX_ENCODED_LEN; diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index 8423d8cb..e8c30c90 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -1,8 +1,8 @@ use std::ops::Deref; use crate::{ - codec, ByteSlice, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash, VarInt, - WireEncode, WireError, QID, + codec, derive_qid, ByteSlice, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash, + VarInt, WireEncode, WireError, QID, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -67,7 +67,7 @@ impl QlIdentity { name: impl Into, ) -> Result { let name = QlName::new(name)?; - let qid = QID::derive(crypto, &mlkem_public_key); + let qid = derive_qid(crypto, &mlkem_public_key); Ok(Self { qid, mlkem_private_key, diff --git a/ql-wire/src/lib.rs b/ql-wire/src/lib.rs index 9381afd7..56288eae 100644 --- a/ql-wire/src/lib.rs +++ b/ql-wire/src/lib.rs @@ -4,6 +4,8 @@ #![allow(clippy::too_many_arguments)] +#[macro_use] +mod macros; mod bytes; mod codec; mod crypto; @@ -13,7 +15,6 @@ mod error; mod handshake; mod header; mod identity; -mod macros; mod pq; mod qid; mod record; @@ -30,13 +31,14 @@ pub use error::*; pub use handshake::*; pub use header::*; pub use identity::*; -pub(crate) use macros::*; pub use pq::*; pub use qid::*; +pub use ql_common::{ + RouteId, ServiceId, StreamCloseCode, StreamId, VarInt, VarIntBoundsExceeded, QID, +}; pub use record::*; #[cfg(any(feature = "test-utils", test))] pub use testing::*; -pub use varint::*; pub const QL_WIRE_VERSION: u8 = 1; pub const ENCRYPTED_MESSAGE_AUTH_SIZE: usize = 16; diff --git a/ql-wire/src/macros.rs b/ql-wire/src/macros.rs index a0f52b03..8cd15a5a 100644 --- a/ql-wire/src/macros.rs +++ b/ql-wire/src/macros.rs @@ -1,25 +1,5 @@ -macro_rules! varint_wrapper { - ($name:ident) => { - #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] - #[repr(transparent)] - pub struct $name(pub $crate::VarInt); - - impl $name { - pub const MAX_ENCODED_LEN: usize = $crate::VarInt::MAX_SIZE; - - pub const fn from_u32(value: u32) -> Self { - Self($crate::VarInt::from_u32(value)) - } - - pub fn from_u64(value: u64) -> Result { - Ok(Self($crate::VarInt::from_u64(value)?)) - } - - pub const fn into_inner(self) -> u64 { - self.0.into_inner() - } - } - +macro_rules! varint_wrapper_codec { + ($name:ty) => { impl $crate::WireEncode for $name { fn encoded_len(&self) -> usize { self.0.size() @@ -32,25 +12,27 @@ macro_rules! varint_wrapper { impl $crate::WireDecode for $name { fn decode(reader: &mut $crate::Reader) -> Result { - Ok(Self(reader.decode()?)) + Ok(<$name>::from(reader.decode::<$crate::VarInt>()?)) } } + }; +} - impl From<$crate::VarInt> for $name { - fn from(value: $crate::VarInt) -> Self { - Self(value) +macro_rules! array_wrapper_codec { + ($name:ty) => { + impl $crate::WireEncode for $name { + fn encoded_len(&self) -> usize { + <$name>::SIZE } - } - impl From for $name { - fn from(value: u32) -> Self { - Self::from_u32(value) + fn encode(&self, out: &mut W) { + self.0.encode(out); } } - impl std::fmt::Display for $name { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.0) + impl $crate::codec::WireDecode for $name { + fn decode(reader: &mut $crate::codec::Reader) -> Result { + Ok(Self(reader.decode()?)) } } }; @@ -87,6 +69,3 @@ macro_rules! array_wrapper { } }; } - -pub(crate) use array_wrapper; -pub(crate) use varint_wrapper; diff --git a/ql-wire/src/qid.rs b/ql-wire/src/qid.rs index 04f4e759..b5149bda 100644 --- a/ql-wire/src/qid.rs +++ b/ql-wire/src/qid.rs @@ -1,16 +1,14 @@ -use crate::{MlKemPublicKey, QlHash, ML_KEM_SUITE_TAG}; +use crate::{MlKemPublicKey, QlHash, ML_KEM_SUITE_TAG, QID}; -crate::array_wrapper!(QID, 16); +array_wrapper_codec!(QID); -impl QID { - pub fn derive(crypto: &impl QlHash, mlkem_public_key: &MlKemPublicKey) -> Self { - let digest = crypto.sha256(&[ - b"quantum-link qid v1", - ML_KEM_SUITE_TAG, - mlkem_public_key.as_bytes(), - ]); - let mut qid = [0u8; Self::SIZE]; - qid.copy_from_slice(&digest[..Self::SIZE]); - Self(qid) - } +pub fn derive_qid(crypto: &impl QlHash, mlkem_public_key: &MlKemPublicKey) -> QID { + let digest = crypto.sha256(&[ + b"quantum-link qid v1", + ML_KEM_SUITE_TAG, + mlkem_public_key.as_bytes(), + ]); + let mut qid = [0u8; QID::SIZE]; + qid.copy_from_slice(&digest[..QID::SIZE]); + QID(qid) } diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index acd69f6b..6657ea31 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -778,7 +778,7 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { let wrong_seq_header = SessionHeader { connection_id: header.connection_id, - seq: record_seq(header.seq.into_inner() + 1), + seq: record_seq(header.seq.0.into_inner() + 1), }; assert_eq!( encrypted::decrypt_record(&crypto, &wrong_seq_header, encrypted, &session_key), diff --git a/ql-wire/src/varint.rs b/ql-wire/src/varint.rs index 7a39bd16..8c94d34c 100644 --- a/ql-wire/src/varint.rs +++ b/ql-wire/src/varint.rs @@ -1,135 +1,8 @@ -use core::fmt; - use bytes::BufMut; use crate::{ByteSlice, Reader, WireDecode, WireEncode, WireError}; -/// An integer less than 2^62 encoded with QUIC variable-length integer rules. -#[derive(Default, Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)] -pub struct VarInt(pub(crate) u64); - -impl VarInt { - /// The largest representable value. - pub const MAX: Self = Self((1u64 << 62) - 1); - /// The largest encoded value length. - pub const MAX_SIZE: usize = 8; - pub const MIN_SIZE: usize = 1; - - /// Construct a `VarInt` infallibly from a `u32`. - pub const fn from_u32(x: u32) -> Self { - Self(x as u64) - } - - /// Construct a `VarInt` from a `u64`. - pub fn from_u64(x: u64) -> Result { - if x < (1u64 << 62) { - Ok(Self(x)) - } else { - Err(VarIntBoundsExceeded) - } - } - - /// Create a `VarInt` without checking the bounds. - /// - /// # Safety - /// - /// `x` must be less than 2^62. - pub const unsafe fn from_u64_unchecked(x: u64) -> Self { - Self(x) - } - - /// Extract the inner integer value. - pub const fn into_inner(self) -> u64 { - self.0 - } - - /// Return the number of bytes required to encode this value. - pub const fn size(self) -> usize { - let x = self.0; - if x < (1u64 << 6) { - 1 - } else if x < (1u64 << 14) { - 2 - } else if x < (1u64 << 30) { - 4 - } else { - 8 - } - } -} - -impl From for u64 { - fn from(value: VarInt) -> Self { - value.0 - } -} - -impl From for VarInt { - fn from(value: u8) -> Self { - Self(value.into()) - } -} - -impl From for VarInt { - fn from(value: u16) -> Self { - Self(value.into()) - } -} - -impl From for VarInt { - fn from(value: u32) -> Self { - Self(value.into()) - } -} - -impl TryFrom for VarInt { - type Error = VarIntBoundsExceeded; - - fn try_from(value: u64) -> Result { - Self::from_u64(value) - } -} - -impl TryFrom for VarInt { - type Error = VarIntBoundsExceeded; - - fn try_from(value: u128) -> Result { - Self::from_u64(value.try_into().map_err(|_| VarIntBoundsExceeded)?) - } -} - -impl TryFrom for VarInt { - type Error = VarIntBoundsExceeded; - - fn try_from(value: usize) -> Result { - Self::from_u64(value as u64) - } -} - -impl fmt::Debug for VarInt { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.fmt(f) - } -} - -impl fmt::Display for VarInt { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.fmt(f) - } -} - -#[derive(Debug, Copy, Clone, Eq, PartialEq)] -pub struct VarIntBoundsExceeded; - -impl fmt::Display for VarIntBoundsExceeded { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.write_str("value too large for varint encoding") - } -} - -impl std::error::Error for VarIntBoundsExceeded {} - -impl WireDecode for VarInt { +impl WireDecode for ql_common::VarInt { fn decode(reader: &mut Reader) -> Result { let first = reader.decode::()?; let tag = first >> 6; @@ -162,7 +35,7 @@ impl WireDecode for VarInt { } } -impl WireEncode for VarInt { +impl WireEncode for ql_common::VarInt { fn encoded_len(&self) -> usize { self.size() } From b046cec3f9e949f82f3ee3eeda6b2ed9cf7b5655 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 11:09:53 -0400 Subject: [PATCH 14/59] ql: stream close code origin --- ql-common/src/lib.rs | 9 +++++++++ ql-rpc/src/lib.rs | 2 +- ql-runtime/src/driver/mod.rs | 17 +++++++++++++---- ql-runtime/src/error.rs | 9 ++++++--- ql-runtime/src/io/inner.rs | 16 ++++++++++------ ql-runtime/src/rpc/adapter.rs | 7 +++++-- ql-runtime/src/rpc/error.rs | 11 +++++++---- ql-runtime/src/tests/rpc.rs | 7 ++++--- ql-runtime/src/tests/stream.rs | 11 +++++++---- ql-wire/src/lib.rs | 3 ++- 10 files changed, 64 insertions(+), 28 deletions(-) diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index c95f82a0..ee46c003 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -28,6 +28,15 @@ impl std::fmt::Display for StreamCloseCode { } } +/// origin of a stream close: either we triggered it locally or the peer sent it. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum StreamCloseOrigin { + /// the close code originated from the peer + Peer, + /// the close code originated from local logic + Local, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] pub struct QID(pub [u8; Self::SIZE]); diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index c3e20364..94c24e01 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -14,7 +14,7 @@ pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; use framed_value::*; -pub use ql_common::{RouteId, ServiceId, StreamCloseCode}; +pub use ql_common::{RouteId, ServiceId, StreamCloseCode, StreamCloseOrigin}; pub use router::*; pub use rpc::*; pub use stream::*; diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 2f7e4e11..2b45f872 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -16,7 +16,7 @@ use std::{ use async_channel::Recv; use futures_lite::future::{poll_fn, yield_now}; use ql_fsm::{Event, QlFsm, WriteId}; -use ql_wire::{CloseTarget, StreamCloseCode, StreamHeader, StreamId}; +use ql_wire::{CloseTarget, StreamCloseCode, StreamCloseOrigin, StreamHeader, StreamId}; use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; use crate::{ @@ -488,10 +488,16 @@ impl DriverState { let stream = entry.get_mut(); if frame.target == CloseTarget::Both || frame.target == stream.inbound_target() { - stream.inbound_fail(QlStreamError::StreamClosed { code: frame.code }); + stream.inbound_fail(QlStreamError::StreamClosed { + code: frame.code, + origin: StreamCloseOrigin::Peer, + }); } if frame.target == CloseTarget::Both || frame.target == stream.outbound_target() { - stream.outbound_fail(QlStreamError::StreamClosed { code: frame.code }); + stream.outbound_fail(QlStreamError::StreamClosed { + code: frame.code, + origin: StreamCloseOrigin::Peer, + }); } Self::try_reap_stream(entry); } @@ -507,7 +513,10 @@ impl DriverState { return; }; let stream = entry.get_mut(); - stream.outbound_fail(QlStreamError::StreamClosed { code: frame.code }); + stream.outbound_fail(QlStreamError::StreamClosed { + code: frame.code, + origin: StreamCloseOrigin::Peer, + }); Self::try_reap_stream(entry); } diff --git a/ql-runtime/src/error.rs b/ql-runtime/src/error.rs index 5b74bcf8..55833e88 100644 --- a/ql-runtime/src/error.rs +++ b/ql-runtime/src/error.rs @@ -1,15 +1,18 @@ -use ql_wire::StreamCloseCode; +use ql_wire::{StreamCloseCode, StreamCloseOrigin}; #[derive(Debug, Clone, PartialEq, Eq)] pub enum QlStreamError { - StreamClosed { code: StreamCloseCode }, + StreamClosed { + code: StreamCloseCode, + origin: StreamCloseOrigin, + }, NoSession, } impl std::fmt::Display for QlStreamError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - Self::StreamClosed { code } => write!(f, "stream closed {code:?}"), + Self::StreamClosed { code, origin } => write!(f, "stream closed {code:?} ({origin:?})"), Self::NoSession => f.write_str("no session"), } } diff --git a/ql-runtime/src/io/inner.rs b/ql-runtime/src/io/inner.rs index 64df6ced..5d322aad 100644 --- a/ql-runtime/src/io/inner.rs +++ b/ql-runtime/src/io/inner.rs @@ -292,7 +292,7 @@ mod loom_tests { use bytes::Bytes; use loom::thread; - use ql_wire::StreamCloseCode; + use ql_wire::{StreamCloseCode, StreamCloseOrigin}; use super::*; use crate::{ @@ -409,6 +409,7 @@ mod loom_tests { thread::spawn(move || { shared.rx.fail(QlStreamError::StreamClosed { code: StreamCloseCode::CANCELLED, + origin: StreamCloseOrigin::Local, }) }) }; @@ -420,13 +421,13 @@ mod loom_tests { (Ok(Item::Chunk(bytes)), None) => { assert_eq!(bytes, Bytes::from_static(b"abc")); match shared.rx.pop() { - Ok(Item::Error(QlStreamError::StreamClosed { code })) => { + Ok(Item::Error(QlStreamError::StreamClosed { code, .. })) => { assert_eq!(code, StreamCloseCode::CANCELLED); } _ => panic!("expected terminal reader error"), } } - (Ok(Item::Error(QlStreamError::StreamClosed { code })), Some(bytes)) => { + (Ok(Item::Error(QlStreamError::StreamClosed { code, .. })), Some(bytes)) => { assert_eq!(code, StreamCloseCode::CANCELLED); assert_eq!(bytes, Bytes::from_static(b"abc")); assert!(matches!(shared.rx.pop(), Err(PopError))); @@ -500,6 +501,7 @@ mod loom_tests { thread::spawn(move || { let displaced = shared.tx.fail(QlStreamError::StreamClosed { code: StreamCloseCode::CANCELLED, + origin: StreamCloseOrigin::Local, }); assert_eq!(displaced.unwrap(), Some(Bytes::from_static(b"abc"))); }) @@ -510,7 +512,7 @@ mod loom_tests { assert!(TxInner::terminal_ready(shared.tx.load_state())); shared.tx.unregister_waiter(); match shared.tx.pop() { - Ok(Item::Error(QlStreamError::StreamClosed { code })) => { + Ok(Item::Error(QlStreamError::StreamClosed { code, .. })) => { assert_eq!(code, StreamCloseCode::CANCELLED); } _ => panic!("expected terminal writer error"), @@ -570,6 +572,7 @@ mod loom_tests { thread::spawn(move || { shared.tx.fail(QlStreamError::StreamClosed { code: StreamCloseCode::CANCELLED, + origin: StreamCloseOrigin::Local, }) }) }; @@ -594,7 +597,7 @@ mod loom_tests { } match shared.tx.pop() { - Ok(Item::Error(QlStreamError::StreamClosed { code })) => { + Ok(Item::Error(QlStreamError::StreamClosed { code, .. })) => { assert_eq!(code, StreamCloseCode::CANCELLED); } _ => panic!("expected terminal writer error"), @@ -616,6 +619,7 @@ mod loom_tests { thread::spawn(move || { shared.tx.fail(QlStreamError::StreamClosed { code: StreamCloseCode::CANCELLED, + origin: StreamCloseOrigin::Local, }) }) }; @@ -631,7 +635,7 @@ mod loom_tests { Ok(_) => { assert!(!TxInner::terminal_ok(shared.tx.load_state())); match shared.tx.pop() { - Ok(Item::Error(QlStreamError::StreamClosed { code })) => { + Ok(Item::Error(QlStreamError::StreamClosed { code, .. })) => { assert_eq!(code, StreamCloseCode::CANCELLED); } _ => panic!("expected terminal writer error"), diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index c33ac37c..ea7d32fd 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -61,14 +61,17 @@ impl RpcWrite for StreamWriter { impl From for QlStreamError { fn from(code: StreamCloseCode) -> Self { - Self::StreamClosed { code } + Self::StreamClosed { + code, + origin: ql_wire::StreamCloseOrigin::Local, + } } } impl StreamError for QlStreamError { fn close_code(&self) -> Option { match self { - QlStreamError::StreamClosed { code } => Some(*code), + QlStreamError::StreamClosed { code, .. } => Some(*code), QlStreamError::NoSession => None, } } diff --git a/ql-runtime/src/rpc/error.rs b/ql-runtime/src/rpc/error.rs index 4cc9e176..a7f322f5 100644 --- a/ql-runtime/src/rpc/error.rs +++ b/ql-runtime/src/rpc/error.rs @@ -5,7 +5,10 @@ use crate::QlStreamError; #[derive(Debug)] pub enum RpcError { NoSession, - Closed(ql_rpc::StreamCloseCode), + Closed { + code: ql_rpc::StreamCloseCode, + origin: ql_rpc::StreamCloseOrigin, + }, Protocol(ql_rpc::Error), Codec(E), } @@ -19,7 +22,7 @@ impl From for RpcError { impl From for RpcError { fn from(error: QlStreamError) -> Self { match error { - QlStreamError::StreamClosed { code } => Self::Closed(ql_rpc::StreamCloseCode(code.0)), + QlStreamError::StreamClosed { code, origin } => Self::Closed { code, origin }, QlStreamError::NoSession => Self::NoSession, } } @@ -57,7 +60,7 @@ where fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { Self::NoSession => write!(f, "no session"), - Self::Closed(code) => write!(f, "stream closed {code:?}"), + Self::Closed { code, origin } => write!(f, "stream closed {code:?} ({origin:?})"), Self::Protocol(error) => write!(f, "{error}"), Self::Codec(error) => write!(f, "{error}"), } @@ -73,7 +76,7 @@ where Self::Protocol(error) => Some(error), Self::Codec(error) => Some(error), RpcError::NoSession => None, - RpcError::Closed(_) => None, + RpcError::Closed { .. } => None, } } } diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index 10c0b6b8..c2c43e1f 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -13,8 +13,8 @@ use ql_rpc::{ DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, NotificationHandlerLocal, ProgressHandlerLocal, ProgressResponder, RequestHandler, RequestHandlerLocal, Response, RouteId, SendSpawner, ServiceId, Spawner, StreamCloseCode, - SubscriptionHandlerLocal, SubscriptionResponder, UploadHandlerLocal, UploadReader, - UploadResponder, + StreamCloseOrigin, SubscriptionHandlerLocal, SubscriptionResponder, UploadHandlerLocal, + UploadReader, UploadResponder, }; use super::*; @@ -330,7 +330,8 @@ async fn rpc_router_enforces_max_request_bytes() { let response = rpc.request::(&"hello".to_string()).await; assert!(matches!( response, - Err(RpcError::Closed(code)) if code == StreamCloseCode::LIMIT + Err(RpcError::Closed { code, origin }) + if code == StreamCloseCode::LIMIT && origin == StreamCloseOrigin::Peer )); tokio::time::timeout(Duration::from_secs(2), responder) diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index 8dee4f18..757f94c5 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -1,7 +1,7 @@ use std::time::Duration; use bytes::Bytes; -use ql_wire::StreamCloseCode; +use ql_wire::{StreamCloseCode, StreamCloseOrigin}; use super::*; use crate::QlStreamError; @@ -174,13 +174,15 @@ async fn dropping_responder_closes_initiator_response() { let err = stream.writer.finish().await.unwrap_err(); assert!(matches!( err, - QlStreamError::StreamClosed { code } if code == StreamCloseCode::CANCELLED + QlStreamError::StreamClosed { code, origin } + if code == StreamCloseCode::CANCELLED && origin == StreamCloseOrigin::Peer )); let err = next_chunk(&mut stream.reader).await.unwrap_err(); assert!(matches!( err, - QlStreamError::StreamClosed { code } if code == StreamCloseCode::CANCELLED + QlStreamError::StreamClosed { code, origin } + if code == StreamCloseCode::CANCELLED && origin == StreamCloseOrigin::Peer )); tokio::time::timeout(Duration::from_secs(2), responder) @@ -213,7 +215,8 @@ async fn dropping_inbound_reader_cancels_remote_writer() { let err = writer.finish().await.unwrap_err(); assert!(matches!( err, - QlStreamError::StreamClosed { code } if code == StreamCloseCode::CANCELLED + QlStreamError::StreamClosed { code, origin } + if code == StreamCloseCode::CANCELLED && origin == StreamCloseOrigin::Peer )); }); diff --git a/ql-wire/src/lib.rs b/ql-wire/src/lib.rs index 56288eae..2bb1e3e5 100644 --- a/ql-wire/src/lib.rs +++ b/ql-wire/src/lib.rs @@ -34,7 +34,8 @@ pub use identity::*; pub use pq::*; pub use qid::*; pub use ql_common::{ - RouteId, ServiceId, StreamCloseCode, StreamId, VarInt, VarIntBoundsExceeded, QID, + RouteId, ServiceId, StreamCloseCode, StreamCloseOrigin, StreamId, VarInt, VarIntBoundsExceeded, + QID, }; pub use record::*; #[cfg(any(feature = "test-utils", test))] From f599eae8c60aba62544d95c8cfc43182bd61937c Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 13:39:42 -0400 Subject: [PATCH 15/59] ql: simpler error types --- ql-rpc/src/chunk_queue.rs | 10 ++-- ql-rpc/src/error.rs | 77 +++++++++------------------ ql-rpc/src/framed_value.rs | 26 +++++---- ql-rpc/src/router/builder.rs | 24 ++++----- ql-rpc/src/rpc/download/client.rs | 21 ++++---- ql-rpc/src/rpc/download/server.rs | 10 ++-- ql-rpc/src/rpc/duplex/client.rs | 12 ++--- ql-rpc/src/rpc/duplex/codec.rs | 12 ++--- ql-rpc/src/rpc/notification/server.rs | 12 ++--- ql-rpc/src/rpc/parts.rs | 40 +++++++------- ql-rpc/src/rpc/progress/client.rs | 11 ++-- ql-rpc/src/rpc/progress/codec.rs | 23 ++++---- ql-rpc/src/rpc/progress/server.rs | 10 ++-- ql-rpc/src/rpc/request/client.rs | 10 ++-- ql-rpc/src/rpc/request/server.rs | 10 ++-- ql-rpc/src/rpc/subscription/client.rs | 12 ++--- ql-rpc/src/rpc/subscription/codec.rs | 8 +-- ql-rpc/src/rpc/subscription/server.rs | 11 ++-- ql-rpc/src/rpc/upload/client.rs | 12 ++--- ql-rpc/src/rpc/upload/server.rs | 20 +++---- ql-rpc/src/rpc/utils.rs | 46 ++++++++-------- ql-rpc/src/stream.rs | 16 ++---- ql-runtime/src/rpc/adapter.rs | 20 +------ ql-runtime/src/rpc/error.rs | 19 ++----- 24 files changed, 206 insertions(+), 266 deletions(-) diff --git a/ql-rpc/src/chunk_queue.rs b/ql-rpc/src/chunk_queue.rs index 33f62998..770002f2 100644 --- a/ql-rpc/src/chunk_queue.rs +++ b/ql-rpc/src/chunk_queue.rs @@ -2,7 +2,7 @@ use std::collections::VecDeque; use bytes::{Buf, Bytes}; -use crate::{CodecError, Error}; +use crate::Error; const LENGTH_SIZE: usize = 8; @@ -25,9 +25,9 @@ impl ChunkQueue { self.remaining } - pub fn expect_empty(&self) -> Result<(), CodecError> { + pub fn expect_empty(&self) -> Result<(), Error> { if self.remaining > 0 { - Err(CodecError::Rpc(Error::TrailingBytes)) + Err(Error::TrailingBytes) } else { Ok(()) } @@ -199,9 +199,9 @@ impl<'a> DrainBuf<'a> { } } - pub fn expect_empty(&self) -> Result<(), CodecError> { + pub fn expect_empty(&self) -> Result<(), Error> { if self.remaining > 0 { - Err(CodecError::Rpc(Error::TrailingBytes)) + Err(Error::TrailingBytes) } else { Ok(()) } diff --git a/ql-rpc/src/error.rs b/ql-rpc/src/error.rs index 7404a22e..5892e271 100644 --- a/ql-rpc/src/error.rs +++ b/ql-rpc/src/error.rs @@ -1,3 +1,5 @@ +use crate::StreamCloseCode; + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum Error { Truncated, @@ -21,54 +23,36 @@ impl std::fmt::Display for Error { impl std::error::Error for Error {} -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum CodecError { - Rpc(Error), - Codec(E), -} - -impl std::error::Error for CodecError -where - E: std::error::Error + 'static, -{ - fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { - match self { - CodecError::Rpc(e) => Some(e), - CodecError::Codec(e) => Some(e), - } - } - - fn cause(&self) -> Option<&dyn std::error::Error> { - self.source() - } -} - -impl std::fmt::Display for CodecError -where - E: std::fmt::Display, -{ - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { +impl Error { + pub const fn close_code(self) -> StreamCloseCode { match self { - CodecError::Rpc(e) => write!(f, "{e}"), - CodecError::Codec(e) => write!(f, "{e}"), + Self::LengthOverflow => StreamCloseCode::LIMIT, + Self::Truncated + | Self::UnexpectedFrameKind(_) + | Self::MissingResponse + | Self::TrailingBytes => StreamCloseCode::REFUSED, } } } -impl From for CodecError { - fn from(error: Error) -> Self { - Self::Rpc(error) - } -} - #[derive(Debug, Clone, PartialEq, Eq)] -pub enum CallError { +pub enum RpcError { Protocol(Error), Codec(C), Transport(T), } -impl std::fmt::Display for CallError +impl RpcError { + pub const fn close_code(&self) -> Option { + match self { + Self::Protocol(error) => Some(error.close_code()), + Self::Codec(_) => Some(StreamCloseCode::REFUSED), + Self::Transport(_) => None, + } + } +} + +impl std::fmt::Display for RpcError where C: std::fmt::Display, T: std::fmt::Display, @@ -82,31 +66,22 @@ where } } -impl std::error::Error for CallError +impl std::error::Error for RpcError where C: std::error::Error + 'static, T: std::error::Error + 'static, { fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { match self { - CallError::Protocol(error) => Some(error), - CallError::Codec(error) => Some(error), - CallError::Transport(error) => Some(error), + Self::Protocol(error) => Some(error), + Self::Codec(error) => Some(error), + Self::Transport(error) => Some(error), } } } -impl From for CallError { +impl From for RpcError { fn from(error: Error) -> Self { Self::Protocol(error) } } - -impl From> for CallError { - fn from(error: CodecError) -> Self { - match error { - CodecError::Rpc(error) => Self::Protocol(error), - CodecError::Codec(error) => Self::Codec(error), - } - } -} diff --git a/ql-rpc/src/framed_value.rs b/ql-rpc/src/framed_value.rs index 600357da..df9664c2 100644 --- a/ql-rpc/src/framed_value.rs +++ b/ql-rpc/src/framed_value.rs @@ -2,7 +2,7 @@ use std::marker::PhantomData; use bytes::Bytes; -use crate::{chunk_queue::ChunkQueue, CodecError, RpcCodec}; +use crate::{chunk_queue::ChunkQueue, RpcCodec, RpcError}; /// reads one length-delimited rpc value from buffered byte chunks pub struct FramedReader { @@ -35,23 +35,23 @@ impl FramedReader { self } - pub fn advance(self) -> Result, CodecError> { - match self.advance_prefix()? { + pub fn advance(self) -> Result, RpcError> { + match self.advance_prefix::()? { FramedPrefixStep::NeedMore(next) => Ok(FramedReadStep::NeedMore(next)), FramedPrefixStep::Value { value, bytes } => { - bytes.expect_empty()?; + bytes.expect_empty().map_err(RpcError::Protocol)?; Ok(FramedReadStep::Value(value)) } } } - pub fn advance_prefix(self) -> Result, CodecError> { + pub fn advance_prefix(self) -> Result, RpcError> { let mut this = self; - let Some(mut body) = this.bytes.try_take_part()? else { + let Some(mut body) = this.bytes.try_take_part().map_err(RpcError::Protocol)? else { return Ok(FramedPrefixStep::NeedMore(this)); }; - let value = T::decode_value(&mut body).map_err(CodecError::Codec)?; + let value = T::decode_value(&mut body).map_err(RpcError::Codec)?; drop(body); Ok(FramedPrefixStep::Value { value, @@ -74,7 +74,7 @@ mod tests { match FramedReader::>::default() .push(Bytes::from(encoded)) - .advance() + .advance::() .unwrap() { FramedReadStep::Value(value) => assert_eq!(value, b"hello".to_vec()), @@ -90,14 +90,18 @@ mod tests { let reader = match FramedReader::>::default() .push(encoded.slice(..4)) - .advance() + .advance::() .unwrap() { FramedReadStep::NeedMore(next) => next, _ => unreachable!(), }; - match reader.push(encoded.slice(4..)).advance().unwrap() { + match reader + .push(encoded.slice(4..)) + .advance::() + .unwrap() + { FramedReadStep::Value(value) => assert_eq!(value, b"hello".to_vec()), _ => unreachable!(), } @@ -111,7 +115,7 @@ mod tests { match FramedReader::>::default() .push(Bytes::from(encoded)) - .advance_prefix() + .advance_prefix::() .unwrap() { FramedPrefixStep::Value { value, mut bytes } => { diff --git a/ql-rpc/src/router/builder.rs b/ql-rpc/src/router/builder.rs index 764c16b5..9a78f812 100644 --- a/ql-rpc/src/router/builder.rs +++ b/ql-rpc/src/router/builder.rs @@ -92,7 +92,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -110,7 +110,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -145,7 +145,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -163,7 +163,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -181,7 +181,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -199,7 +199,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -226,7 +226,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -247,7 +247,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -289,7 +289,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -310,7 +310,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -331,7 +331,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } @@ -352,7 +352,7 @@ where reader, writer, S::handle, - S::handle_transport_error, + S::handle_error, )) }) } diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index 9a648181..e8d4c962 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -5,7 +5,7 @@ use bytes::{BufMut, Bytes}; use crate::{ download::{Download, PartReadStep}, rpc::parts::FrameKind, - CallError, FramedPrefixStep, FramedReader, RpcCodec, RpcRead, StreamCloseCode, + FramedPrefixStep, FramedReader, RpcCodec, RpcError, RpcRead, StreamCloseCode, }; pub struct DownloadCall @@ -49,7 +49,7 @@ where pub async fn start( mut self, - ) -> Result<(M::ResponseHeader, DownloadReader), CallError> { + ) -> Result<(M::ResponseHeader, DownloadReader), RpcError> { loop { let reader = self.reader.take().unwrap(); let reader = match reader.advance_prefix() { @@ -64,7 +64,7 @@ where )); } Ok(FramedPrefixStep::NeedMore(next)) => next, - Err(error) => return Err(error.into()), + Err(error) => return Err(error), }; let stream = self.stream.as_mut().unwrap(); @@ -73,7 +73,7 @@ where self.reader = Some(reader.push(chunk)); } Ok(None) => return Err(crate::Error::Truncated.into()), - Err(error) => return Err(CallError::Transport(error)), + Err(error) => return Err(RpcError::Transport(error)), } } } @@ -106,8 +106,7 @@ where { pub async fn next_part( &mut self, - ) -> Result)>, CallError> - { + ) -> Result)>, RpcError> { if self.stream.is_none() { return Ok(None); } @@ -134,7 +133,7 @@ where } } - pub async fn complete(mut self) -> Result<(), CallError> { + pub async fn complete(mut self) -> Result<(), RpcError> { match self.read_frame().await? { PartReadStep::Finish => { self.stream.take(); @@ -159,12 +158,12 @@ where async fn read_frame( &mut self, - ) -> Result, CallError> { + ) -> Result, RpcError> { loop { match self.reader.advance() { Ok(PartReadStep::NeedMore) => {} Ok(step) => return Ok(step), - Err(error) => return Err(error.into()), + Err(error) => return Err(error), } let stream = self.stream.as_mut().unwrap(); @@ -173,7 +172,7 @@ where self.reader.push(chunk); } Ok(None) => return Err(crate::Error::Truncated.into()), - Err(error) => return Err(CallError::Transport(error)), + Err(error) => return Err(RpcError::Transport(error)), } } } @@ -202,7 +201,7 @@ where M: Download, R: RpcRead, { - pub async fn read_chunk(&mut self) -> Result, CallError> { + pub async fn read_chunk(&mut self) -> Result, RpcError> { if self.finished { return Ok(None); } diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs index fcdcb047..cd6119b3 100644 --- a/ql-rpc/src/rpc/download/server.rs +++ b/ql-rpc/src/rpc/download/server.rs @@ -10,7 +10,7 @@ use crate::{ parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, read_eof_request, }, - write_bytes, RouterConfig, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError, + write_bytes, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, }; #[trait_variant::make(DownloadHandler: Send)] @@ -21,7 +21,7 @@ where { async fn handle(self, message: M::Request, download: DownloadStart); - fn handle_transport_error(&self, _error: &St::Error) {} + fn handle_error(&self, _error: &RpcError) {} } pub struct DownloadStart @@ -196,19 +196,19 @@ pub(crate) async fn handle_download_inner( mut reader: St::Reader, writer: St::Writer, handle: H, - handle_transport_error: E, + handle_error: E, ) where M: DownloadRpc + 'static, St: RpcStream + 'static, H: FnOnce(S, M::Request, DownloadStart) -> HF, HF: Future, - E: FnOnce(&S, &St::Error), + E: FnOnce(&S, &RpcError), { let request = match read_eof_request::(&mut reader, config).await { Ok(request) => request, Err(error) => { let code = error.close_code(); - handle_transport_error(&state, &error); + handle_error(&state, &error); if let Some(code) = code { reader.close(code); writer.close(code); diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index e76050a6..c152b08d 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -8,7 +8,7 @@ use bytes::Bytes; use crate::{ duplex::{codec, Duplex, EventReader, ReadStep}, - finish_bytes, write_bytes, CallError, RpcCodec, RpcRead, RpcWrite, StreamCloseCode, + finish_bytes, write_bytes, RpcCodec, RpcError, RpcRead, RpcWrite, StreamCloseCode, }; pub struct DuplexCall @@ -94,14 +94,14 @@ where } } - pub async fn next_event(&mut self) -> Option>> { + pub async fn next_event(&mut self) -> Option>> { poll_fn(|cx| self.poll_next_event(cx)).await } pub fn poll_next_event( &mut self, cx: &mut Context<'_>, - ) -> Poll>>> { + ) -> Poll>>> { if self.stream.is_none() { return Poll::Ready(None); } @@ -111,8 +111,8 @@ where Ok(ReadStep::Event(value)) => return Poll::Ready(Some(Ok(value))), Ok(ReadStep::NeedMore) => {} Err(error) => { - self.stream.take(); - return Poll::Ready(Some(Err(error.into()))); + self.stream.disarm(); + return Poll::Ready(Some(Err(error))); } } @@ -131,7 +131,7 @@ where } Poll::Ready(Err(error)) => { self.stream.take(); - return Poll::Ready(Some(Err(CallError::Transport(error)))); + return Poll::Ready(Some(Err(RpcError::Transport(error)))); } Poll::Pending => { return Poll::Pending; diff --git a/ql-rpc/src/rpc/duplex/codec.rs b/ql-rpc/src/rpc/duplex/codec.rs index 68bc87c7..0392dfd9 100644 --- a/ql-rpc/src/rpc/duplex/codec.rs +++ b/ql-rpc/src/rpc/duplex/codec.rs @@ -2,7 +2,7 @@ use std::marker::PhantomData; use bytes::{BufMut, Bytes}; -use crate::{codec, CodecError, RpcCodec}; +use crate::{codec, RpcCodec, RpcError}; pub fn encode_event(event: &T, out: &mut (impl BufMut + AsMut<[u8]>)) where @@ -39,13 +39,13 @@ impl EventReader { self.bytes.remaining() == 0 } - pub fn advance(&mut self) -> Result, CodecError> { - let Some(mut body) = self.bytes.try_take_part()? else { + pub fn advance(&mut self) -> Result, RpcError> { + let Some(mut body) = self.bytes.try_take_part().map_err(RpcError::Protocol)? else { return Ok(ReadStep::NeedMore); }; let value = { - let value = T::decode_value(&mut body).map_err(CodecError::Codec)?; + let value = T::decode_value(&mut body).map_err(RpcError::Codec)?; drop(body); value }; @@ -68,14 +68,14 @@ mod tests { let mut reader = EventReader::>::default(); reader.push(Bytes::from(encoded)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { ReadStep::Event(value) => { assert_eq!(value, b"one".to_vec()); } _ => unreachable!(), }; - match reader.advance().unwrap() { + match reader.advance::().unwrap() { ReadStep::Event(value) => { assert_eq!(value, b"two".to_vec()); assert!(reader.is_empty()); diff --git a/ql-rpc/src/rpc/notification/server.rs b/ql-rpc/src/rpc/notification/server.rs index c9a4fdba..8e0a33de 100644 --- a/ql-rpc/src/rpc/notification/server.rs +++ b/ql-rpc/src/rpc/notification/server.rs @@ -1,8 +1,8 @@ use std::future::Future; use crate::{ - notification::Notification as NotificationRpc, rpc::read_eof_request, RouterConfig, RpcRead, - RpcStream, RpcWrite, StreamCloseCode, StreamError, + notification::Notification as NotificationRpc, rpc::read_eof_request, RouterConfig, RpcError, + RpcRead, RpcStream, RpcWrite, StreamCloseCode, }; #[trait_variant::make(NotificationHandler: Send)] @@ -13,7 +13,7 @@ where { async fn handle(self, message: M::Payload); - fn handle_transport_error(&self, _error: &St::Error) {} + fn handle_error(&self, _error: &RpcError) {} } pub(crate) async fn handle_notification_inner( @@ -22,19 +22,19 @@ pub(crate) async fn handle_notification_inner( mut reader: St::Reader, writer: St::Writer, handle: H, - handle_transport_error: E, + handle_error: E, ) where M: NotificationRpc + 'static, St: RpcStream + 'static, H: FnOnce(S, M::Payload) -> HF, HF: Future, - E: FnOnce(&S, &St::Error), + E: FnOnce(&S, &RpcError), { let notification = match read_eof_request::(&mut reader, config).await { Ok(notification) => notification, Err(error) => { let code = error.close_code(); - handle_transport_error(&state, &error); + handle_error(&state, &error); if let Some(code) = code { reader.close(code); writer.close(code); diff --git a/ql-rpc/src/rpc/parts.rs b/ql-rpc/src/rpc/parts.rs index 47ff1e87..f5821680 100644 --- a/ql-rpc/src/rpc/parts.rs +++ b/ql-rpc/src/rpc/parts.rs @@ -2,7 +2,7 @@ use std::marker::PhantomData; use bytes::{BufMut, Bytes}; -use crate::{codec, ChunkQueue, CodecError, RpcCodec}; +use crate::{codec, ChunkQueue, RpcCodec, RpcError}; pub enum PartReadStep { NeedMore, @@ -44,7 +44,7 @@ impl PartFrameReader { self.bytes.push(chunk); } - pub fn advance(&mut self) -> Result, CodecError> { + pub fn advance(&mut self) -> Result, RpcError> { loop { match self.pending_frame.take() { PendingFrame::Body { remaining } => { @@ -73,18 +73,18 @@ impl PartFrameReader { match kind { FrameKind::PartHeader => { - let value = H::decode_value(&mut body).map_err(CodecError::Codec)?; + let value = H::decode_value(&mut body).map_err(RpcError::Codec)?; return Ok(PartReadStep::PartHeader(value)); } FrameKind::BodyChunk => unreachable!("body chunk is not a control frame"), FrameKind::EndPart => { - body.expect_empty()?; + body.expect_empty().map_err(RpcError::Protocol)?; return Ok(PartReadStep::EndPart); } FrameKind::Finish => { - body.expect_empty()?; + body.expect_empty().map_err(RpcError::Protocol)?; drop(body); - self.bytes.expect_empty()?; + self.bytes.expect_empty().map_err(RpcError::Protocol)?; return Ok(PartReadStep::Finish); } } @@ -93,12 +93,12 @@ impl PartFrameReader { let Some((kind, len)) = self .bytes .try_take_tagged_part_header() - .map_err(CodecError::Rpc)? + .map_err(RpcError::Protocol)? else { return Ok(PartReadStep::NeedMore); }; - let kind = FrameKind::try_from(kind).map_err(CodecError::Rpc)?; + let kind = FrameKind::try_from(kind).map_err(RpcError::Protocol)?; self.pending_frame = if kind == FrameKind::BodyChunk { PendingFrame::Body { remaining: len } } else { @@ -195,41 +195,41 @@ mod tests { let mut reader = PartFrameReader::>::new(Default::default()); reader.push(Bytes::from(encoded)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::PartHeader(value) => { assert_eq!(value, b"a.txt".to_vec()); } _ => unreachable!(), }; - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"hel")), _ => unreachable!(), }; - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"lo")), _ => unreachable!(), }; - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::EndPart => {} _ => unreachable!(), }; - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::PartHeader(value) => { assert_eq!(value, b"b.txt".to_vec()); } _ => unreachable!(), }; - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::EndPart => {} _ => unreachable!(), }; - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::Finish => {} _ => unreachable!(), } @@ -243,13 +243,13 @@ mod tests { let mut reader = PartFrameReader::>::new(Default::default()); reader.push(encoded.slice(..4)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::NeedMore => {} _ => unreachable!(), }; reader.push(encoded.slice(4..)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::PartHeader(value) => assert_eq!(value, b"a.txt".to_vec()), _ => unreachable!(), } @@ -263,19 +263,19 @@ mod tests { let mut reader = PartFrameReader::>::new(Default::default()); reader.push(encoded.slice(..9)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::NeedMore => {} _ => unreachable!(), }; reader.push(encoded.slice(9..11)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"he")), _ => unreachable!(), }; reader.push(encoded.slice(11..)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"llo")), _ => unreachable!(), }; diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs index c2218c97..b34bfaea 100644 --- a/ql-rpc/src/rpc/progress/client.rs +++ b/ql-rpc/src/rpc/progress/client.rs @@ -6,7 +6,7 @@ use std::{ use crate::{ progress::{Progress, ReadStep, ResponseReader}, - CallError, Error, RpcRead, StreamCloseCode, + Error, RpcError, RpcRead, StreamCloseCode, }; pub struct ProgressCall @@ -24,7 +24,7 @@ where { Invalid, Reading(ResponseReader), - Terminal(Result>), + Terminal(Result>), Done, } @@ -67,7 +67,8 @@ where } Ok(ReadStep::NeedMore) => {} Err(error) => { - self.state = State::Terminal(Err(error.into())); + self.stream.disarm(); + self.state = State::Terminal(Err(error)); return Poll::Ready(None); } } @@ -85,7 +86,7 @@ where return Poll::Ready(None); } Poll::Ready(Err(error)) => { - self.state = State::Terminal(Err(CallError::Transport(error))); + self.state = State::Terminal(Err(RpcError::Transport(error))); return Poll::Ready(None); } Poll::Pending => return Poll::Pending, @@ -126,7 +127,7 @@ where M: Progress, R: RpcRead, { - type Output = Result>; + type Output = Result>; fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { let this = self.get_mut(); diff --git a/ql-rpc/src/rpc/progress/codec.rs b/ql-rpc/src/rpc/progress/codec.rs index 0b01302c..b354e472 100644 --- a/ql-rpc/src/rpc/progress/codec.rs +++ b/ql-rpc/src/rpc/progress/codec.rs @@ -2,7 +2,7 @@ use std::marker::PhantomData; use bytes::{BufMut, Bytes}; -use crate::{codec, progress::Progress, CodecError, Error, RpcCodec}; +use crate::{codec, progress::Progress, Error, RpcCodec, RpcError}; pub enum ReadStep { NeedMore, @@ -29,8 +29,11 @@ impl ResponseReader { self.bytes.push(chunk); } - pub fn advance(&mut self) -> Result, CodecError> { - let Some((kind, mut body)) = self.bytes.try_take_tagged_part().map_err(CodecError::Rpc)? + pub fn advance(&mut self) -> Result, RpcError> { + let Some((kind, mut body)) = self + .bytes + .try_take_tagged_part() + .map_err(RpcError::Protocol)? else { return Ok(ReadStep::NeedMore); }; @@ -38,22 +41,22 @@ impl ResponseReader { match kind { x if x == FrameKind::Progress as u8 => { let value = { - let value = M::Progress::decode_value(&mut body).map_err(CodecError::Codec)?; + let value = M::Progress::decode_value(&mut body).map_err(RpcError::Codec)?; drop(body); value }; Ok(ReadStep::Progress(value)) } x if x == FrameKind::Response as u8 => { - let response = M::Response::decode_value(&mut body).map_err(CodecError::Codec)?; + let response = M::Response::decode_value(&mut body).map_err(RpcError::Codec)?; drop(body); if self.bytes.remaining() > 0 { - Err(CodecError::Rpc(Error::TrailingBytes)) + Err(RpcError::Protocol(Error::TrailingBytes)) } else { Ok(ReadStep::Response(response)) } } - other => Err(CodecError::Rpc(Error::UnexpectedFrameKind(other))), + other => Err(RpcError::Protocol(Error::UnexpectedFrameKind(other))), } } } @@ -120,13 +123,13 @@ mod tests { let mut reader = ResponseReader::::default(); reader.push(Bytes::from(encoded)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { ReadStep::Progress(value) => { assert_eq!(value, b"10%".to_vec()); } _ => unreachable!(), }; - match reader.advance().unwrap() { + match reader.advance::().unwrap() { ReadStep::Response(value) => assert_eq!(value, b"done".to_vec()), _ => unreachable!(), } @@ -140,7 +143,7 @@ mod tests { let mut reader = ResponseReader::::default(); reader.push(Bytes::from(encoded)); - match reader.advance().unwrap() { + match reader.advance::().unwrap() { ReadStep::Response(value) => assert_eq!(value, b"done".to_vec()), _ => unreachable!(), } diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index b94421cf..29991a33 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -6,7 +6,7 @@ use crate::{ finish_bytes, progress::{encode_progress, encode_response, Progress}, rpc::read_framed_request, - write_bytes, RouterConfig, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError, + write_bytes, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, }; #[trait_variant::make(ProgressHandler: Send)] @@ -17,7 +17,7 @@ where { async fn handle(self, request: M::Request, responder: ProgressResponder); - fn handle_transport_error(&self, _error: &St::Error) {} + fn handle_error(&self, _error: &RpcError) {} } pub struct ProgressResponder @@ -81,19 +81,19 @@ pub(crate) async fn handle_progress_inner( mut reader: St::Reader, writer: St::Writer, handle: H, - handle_transport_error: E, + handle_error: E, ) where M: Progress + 'static, St: RpcStream + 'static, H: FnOnce(S, M::Request, ProgressResponder) -> HF, HF: Future, - E: FnOnce(&S, &St::Error), + E: FnOnce(&S, &RpcError), { let request = match read_framed_request::(&mut reader, config).await { Ok(request) => request, Err(error) => { let code = error.close_code(); - handle_transport_error(&state, &error); + handle_error(&state, &error); if let Some(code) = code { reader.close(code); writer.close(code); diff --git a/ql-rpc/src/rpc/request/client.rs b/ql-rpc/src/rpc/request/client.rs index e7ffb845..6f605be9 100644 --- a/ql-rpc/src/rpc/request/client.rs +++ b/ql-rpc/src/rpc/request/client.rs @@ -1,6 +1,6 @@ use bytes::BufMut; -use crate::{read_bytes, request::Request, CallError, ChunkQueue, RpcCodec, RpcRead}; +use crate::{read_bytes, request::Request, ChunkQueue, RpcCodec, RpcError, RpcRead}; pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { request.encode_value(out) @@ -10,9 +10,7 @@ pub fn encode_response(response: &M::Response, out: &mut (impl BufMu response.encode_value(out) } -pub async fn read_response( - mut reader: R, -) -> Result> +pub async fn read_response(mut reader: R) -> Result> where M: Request, R: RpcRead, @@ -21,12 +19,12 @@ where while let Some(chunk) = read_bytes(&mut reader, usize::MAX) .await - .map_err(CallError::Transport)? + .map_err(RpcError::Transport)? { bytes.push(chunk); } - let value = M::Response::decode_value(&mut bytes).map_err(CallError::Codec)?; + let value = M::Response::decode_value(&mut bytes).map_err(RpcError::Codec)?; if bytes.remaining() > 0 { return Err(crate::Error::TrailingBytes.into()); } diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs index 5211cce2..69dada62 100644 --- a/ql-rpc/src/rpc/request/server.rs +++ b/ql-rpc/src/rpc/request/server.rs @@ -4,7 +4,7 @@ use bytes::Bytes; use crate::{ finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, RouterConfig, - RpcCodec, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError, + RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, }; #[trait_variant::make(RequestHandler: Send)] @@ -15,7 +15,7 @@ where { async fn handle(self, message: M::Request, responder: Response); - fn handle_transport_error(&self, _error: &St::Error) {} + fn handle_error(&self, _error: &RpcError) {} } pub struct Response @@ -71,19 +71,19 @@ pub(crate) async fn handle_request_inner( mut reader: St::Reader, writer: St::Writer, handle: H, - handle_transport_error: E, + handle_error: E, ) where M: RequestRpc + 'static, St: RpcStream + 'static, H: FnOnce(S, M::Request, Response) -> HF, HF: Future, - E: FnOnce(&S, &St::Error), + E: FnOnce(&S, &RpcError), { let request = match read_eof_request::(&mut reader, config).await { Ok(request) => request, Err(error) => { let code = error.close_code(); - handle_transport_error(&state, &error); + handle_error(&state, &error); if let Some(code) = code { reader.close(code); writer.close(code); diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index fe6aa5b1..6c10a17e 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -5,7 +5,7 @@ use std::{ use crate::{ subscription::{ReadStep, ResponseReader, Subscription}, - CallError, RpcRead, StreamCloseCode, + RpcError, RpcRead, StreamCloseCode, }; pub struct SubscriptionCall @@ -29,14 +29,14 @@ where } } - pub async fn next_event(&mut self) -> Option>> { + pub async fn next_event(&mut self) -> Option>> { poll_fn(|cx| self.poll_next_event(cx)).await } pub fn poll_next_event( &mut self, cx: &mut Context<'_>, - ) -> Poll>>> { + ) -> Poll>>> { if self.stream.is_none() { return Poll::Ready(None); } @@ -46,8 +46,8 @@ where Ok(ReadStep::Item(value)) => return Poll::Ready(Some(Ok(value))), Ok(ReadStep::NeedMore) => {} Err(error) => { - self.stream.take(); - return Poll::Ready(Some(Err(error.into()))); + self.stream.disarm(); + return Poll::Ready(Some(Err(error))); } } @@ -66,7 +66,7 @@ where } Poll::Ready(Err(error)) => { self.stream.take(); - return Poll::Ready(Some(Err(CallError::Transport(error)))); + return Poll::Ready(Some(Err(RpcError::Transport(error)))); } Poll::Pending => { return Poll::Pending; diff --git a/ql-rpc/src/rpc/subscription/codec.rs b/ql-rpc/src/rpc/subscription/codec.rs index bdd16209..2025d4d2 100644 --- a/ql-rpc/src/rpc/subscription/codec.rs +++ b/ql-rpc/src/rpc/subscription/codec.rs @@ -2,7 +2,7 @@ use std::marker::PhantomData; use bytes::{BufMut, Bytes}; -use crate::{codec, subscription::Subscription, CodecError, RpcCodec}; +use crate::{codec, subscription::Subscription, RpcCodec, RpcError}; pub fn encode_request( request: &M::Request, @@ -43,13 +43,13 @@ impl ResponseReader { self.bytes.remaining() == 0 } - pub fn advance(&mut self) -> Result, CodecError> { - let Some(mut body) = self.bytes.try_take_part()? else { + pub fn advance(&mut self) -> Result, RpcError> { + let Some(mut body) = self.bytes.try_take_part().map_err(RpcError::Protocol)? else { return Ok(ReadStep::NeedMore); }; let item = { - let item = M::Event::decode_value(&mut body).map_err(CodecError::Codec)?; + let item = M::Event::decode_value(&mut body).map_err(RpcError::Codec)?; drop(body); item }; diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs index 6dfdd4b0..6cfe9228 100644 --- a/ql-rpc/src/rpc/subscription/server.rs +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -4,8 +4,7 @@ use bytes::Bytes; use crate::{ codec, finish_bytes, rpc::read_eof_request, subscription::Subscription as SubscriptionRpc, - write_bytes, RouterConfig, RpcCodec, RpcRead, RpcStream, RpcWrite, StreamCloseCode, - StreamError, + write_bytes, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, }; #[trait_variant::make(SubscriptionHandler: Send)] @@ -20,7 +19,7 @@ where responder: SubscriptionResponder, ); - fn handle_transport_error(&self, _error: &St::Error) {} + fn handle_error(&self, _error: &RpcError) {} } pub struct SubscriptionResponder @@ -80,19 +79,19 @@ pub(crate) async fn handle_subscription_inner( mut reader: St::Reader, writer: St::Writer, handle: H, - handle_transport_error: E, + handle_error: E, ) where M: SubscriptionRpc + 'static, St: RpcStream + 'static, H: FnOnce(S, M::Request, SubscriptionResponder) -> HF, HF: Future, - E: FnOnce(&S, &St::Error), + E: FnOnce(&S, &RpcError), { let request = match read_eof_request::(&mut reader, config).await { Ok(request) => request, Err(error) => { let code = error.close_code(); - handle_transport_error(&state, &error); + handle_error(&state, &error); if let Some(code) = code { reader.close(code); writer.close(code); diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index b41dedcd..996089f1 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -4,7 +4,7 @@ use crate::{ finish_bytes, read_bytes, rpc::parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, upload::Upload, - write_bytes, CallError, ChunkQueue, RpcCodec, RpcRead, RpcWrite, StreamCloseCode, + write_bytes, ChunkQueue, RpcCodec, RpcError, RpcRead, RpcWrite, StreamCloseCode, }; pub struct UploadCall @@ -56,28 +56,28 @@ where }) } - pub async fn finish(mut self) -> Result> { + pub async fn finish(mut self) -> Result> { let mut writer = self.writer.take().unwrap(); let mut encoded = Vec::new(); encode_finish(&mut encoded); write_bytes(&mut writer, Bytes::from(encoded)) .await - .map_err(CallError::Transport)?; + .map_err(RpcError::Transport)?; finish_bytes(&mut writer) .await - .map_err(CallError::Transport)?; + .map_err(RpcError::Transport)?; let mut reader = self.reader.take().unwrap(); let mut bytes = ChunkQueue::default(); while let Some(chunk) = read_bytes(&mut reader, usize::MAX) .await - .map_err(CallError::Transport)? + .map_err(RpcError::Transport)? { bytes.push(chunk); } - let value = M::Response::decode_value(&mut bytes).map_err(CallError::Codec)?; + let value = M::Response::decode_value(&mut bytes).map_err(RpcError::Codec)?; if bytes.remaining() > 0 { return Err(crate::Error::TrailingBytes.into()); } diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index d2e6765b..36f67b19 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -8,7 +8,7 @@ use crate::{ parts::{FrameKind, PartFrameReader, PartReadStep}, read_framed_request_prefix, }, - RouterConfig, RpcRead, RpcStream, RpcWrite, StreamCloseCode, StreamError, Upload, + RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, Upload, }; #[trait_variant::make(UploadHandler: Send)] @@ -24,7 +24,7 @@ where responder: UploadResponder, ); - fn handle_transport_error(&self, _error: &St::Error) {} + fn handle_error(&self, _error: &RpcError) {} } pub struct UploadReader @@ -59,7 +59,7 @@ where { pub async fn next_part( &mut self, - ) -> Result)>, crate::CallError> + ) -> Result)>, crate::RpcError> { if self.stream.is_none() { return Ok(None); @@ -89,12 +89,12 @@ where async fn read_frame( &mut self, - ) -> Result, crate::CallError> { + ) -> Result, crate::RpcError> { loop { match self.reader.advance() { Ok(PartReadStep::NeedMore) => {} Ok(step) => return Ok(step), - Err(error) => return Err(error.into()), + Err(error) => return Err(error), } let stream = self.stream.as_mut().unwrap(); @@ -103,7 +103,7 @@ where self.reader.push(chunk); } Ok(None) => return Err(crate::Error::Truncated.into()), - Err(error) => return Err(crate::CallError::Transport(error)), + Err(error) => return Err(crate::RpcError::Transport(error)), } } } @@ -138,7 +138,7 @@ where { pub async fn read_chunk( &mut self, - ) -> Result, crate::CallError> { + ) -> Result, crate::RpcError> { if self.finished { return Ok(None); } @@ -203,7 +203,7 @@ pub(crate) async fn handle_upload_inner( mut reader: St::Reader, writer: St::Writer, handle: H, - handle_transport_error: E, + handle_error: E, ) where M: Upload + 'static, St: RpcStream + 'static, @@ -214,14 +214,14 @@ pub(crate) async fn handle_upload_inner( UploadResponder, ) -> HF, HF: Future, - E: FnOnce(&S, &St::Error), + E: FnOnce(&S, &RpcError), { let (request, buffered) = match read_framed_request_prefix::(&mut reader, config).await { Ok(value) => value, Err(error) => { let code = error.close_code(); - handle_transport_error(&state, &error); + handle_error(&state, &error); if let Some(code) = code { reader.close(code); writer.close(code); diff --git a/ql-rpc/src/rpc/utils.rs b/ql-rpc/src/rpc/utils.rs index bf5f49ea..fb33a6f5 100644 --- a/ql-rpc/src/rpc/utils.rs +++ b/ql-rpc/src/rpc/utils.rs @@ -1,13 +1,13 @@ use crate::{ - read_bytes, ChunkQueue, CodecError, FramedPrefixStep, FramedReadStep, FramedReader, - RouterConfig, RpcCodec, RpcRead, StreamCloseCode, + read_bytes, ChunkQueue, Error, FramedPrefixStep, FramedReadStep, FramedReader, RouterConfig, + RpcCodec, RpcError, RpcRead, }; /// reads one length-delimited value and rejects trailing bytes pub(crate) async fn read_framed_request( reader: &mut R, config: RouterConfig, -) -> Result +) -> Result> where T: RpcCodec, R: RpcRead, @@ -16,16 +16,15 @@ where let mut total_read = 0usize; let value = loop { - match value_reader.advance() { + match value_reader.advance::() { Ok(FramedReadStep::Value(value)) => break value, Ok(FramedReadStep::NeedMore(next)) => value_reader = next, - Err(CodecError::Rpc(_error)) => return Err(StreamCloseCode::REFUSED.into()), - Err(CodecError::Codec(_error)) => return Err(StreamCloseCode::REFUSED.into()), + Err(error) => return Err(error), } let remaining = config.max_request_bytes.saturating_sub(total_read); if remaining == 0 { - return Err(StreamCloseCode::LIMIT.into()); + return Err(RpcError::Protocol(Error::LengthOverflow)); } match read_bytes(reader, remaining).await { @@ -33,8 +32,8 @@ where total_read += chunk.len(); value_reader = value_reader.push(chunk); } - Ok(None) => return Err(StreamCloseCode::REFUSED.into()), - Err(error) => return Err(error), + Ok(None) => return Err(RpcError::Protocol(Error::Truncated)), + Err(error) => return Err(RpcError::Transport(error)), } }; @@ -42,9 +41,9 @@ where let probe = remaining.max(1); match read_bytes(reader, probe).await { Ok(None) => Ok(value), - Ok(Some(_)) if remaining == 0 => Err(StreamCloseCode::LIMIT.into()), - Ok(Some(_)) => Err(StreamCloseCode::REFUSED.into()), - Err(error) => Err(error), + Ok(Some(_)) if remaining == 0 => Err(RpcError::Protocol(Error::LengthOverflow)), + Ok(Some(_)) => Err(RpcError::Protocol(Error::TrailingBytes)), + Err(error) => Err(RpcError::Transport(error)), } } @@ -52,7 +51,7 @@ where pub(crate) async fn read_framed_request_prefix( reader: &mut R, config: RouterConfig, -) -> Result<(T, ChunkQueue), R::Error> +) -> Result<(T, ChunkQueue), RpcError> where T: RpcCodec, R: RpcRead, @@ -61,16 +60,15 @@ where let mut total_read = 0usize; loop { - match value_reader.advance_prefix() { + match value_reader.advance_prefix::() { Ok(FramedPrefixStep::Value { value, bytes }) => return Ok((value, bytes)), Ok(FramedPrefixStep::NeedMore(next)) => value_reader = next, - Err(CodecError::Rpc(_error)) => return Err(StreamCloseCode::REFUSED.into()), - Err(CodecError::Codec(_error)) => return Err(StreamCloseCode::REFUSED.into()), + Err(error) => return Err(error), } let remaining = config.max_request_bytes.saturating_sub(total_read); if remaining == 0 { - return Err(StreamCloseCode::LIMIT.into()); + return Err(RpcError::Protocol(Error::LengthOverflow)); } match read_bytes(reader, remaining).await { @@ -78,8 +76,8 @@ where total_read += chunk.len(); value_reader = value_reader.push(chunk); } - Ok(None) => return Err(StreamCloseCode::REFUSED.into()), - Err(error) => return Err(error), + Ok(None) => return Err(RpcError::Protocol(Error::Truncated)), + Err(error) => return Err(RpcError::Transport(error)), } } } @@ -88,7 +86,7 @@ where pub(crate) async fn read_eof_request( reader: &mut R, config: RouterConfig, -) -> Result +) -> Result> where T: RpcCodec, R: RpcRead, @@ -102,19 +100,19 @@ where match read_bytes(reader, probe).await { Ok(Some(chunk)) => { if chunk.len() > remaining { - return Err(StreamCloseCode::LIMIT.into()); + return Err(RpcError::Protocol(Error::LengthOverflow)); } total_read += chunk.len(); bytes.push(chunk); } Ok(None) => break, - Err(error) => return Err(error), + Err(error) => return Err(RpcError::Transport(error)), } } - let value = T::decode_value(&mut bytes).map_err(|_error| StreamCloseCode::REFUSED)?; + let value = T::decode_value(&mut bytes).map_err(RpcError::Codec)?; if bytes.remaining() > 0 { - return Err(StreamCloseCode::REFUSED.into()); + return Err(RpcError::Protocol(Error::TrailingBytes)); } Ok(value) } diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index 0344f3a1..0cceff88 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -8,7 +8,7 @@ use bytes::Bytes; use crate::{RouteId, ServiceId, StreamCloseCode}; pub trait RpcStream { - type Error: StreamError; + type Error; type Reader: RpcRead; type Writer: RpcWrite; @@ -18,7 +18,7 @@ pub trait RpcStream { } pub trait RpcRead { - type Error: StreamError; + type Error; /// reads inbound bytes until eof or error fn poll_read( @@ -32,7 +32,7 @@ pub trait RpcRead { } pub trait RpcWrite { - type Error: StreamError; + type Error; /// writes outbound bytes before finish or close fn poll_write( @@ -48,16 +48,6 @@ pub trait RpcWrite { fn close(self, code: StreamCloseCode); } -pub trait StreamError: From { - fn close_code(&self) -> Option; -} - -impl StreamError for StreamCloseCode { - fn close_code(&self) -> Option { - Some(*self) - } -} - pub async fn read_bytes(reader: &mut R, max_len: usize) -> Result, R::Error> where R: RpcRead, diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index ea7d32fd..53604dfb 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -1,7 +1,7 @@ use std::task::{Context, Poll}; use bytes::Bytes; -use ql_rpc::{RouteId, RpcRead, RpcStream, RpcWrite, ServiceId, StreamCloseCode, StreamError}; +use ql_rpc::{RouteId, RpcRead, RpcStream, RpcWrite, ServiceId, StreamCloseCode}; use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; @@ -58,21 +58,3 @@ impl RpcWrite for StreamWriter { StreamWriter::close(self, code); } } - -impl From for QlStreamError { - fn from(code: StreamCloseCode) -> Self { - Self::StreamClosed { - code, - origin: ql_wire::StreamCloseOrigin::Local, - } - } -} - -impl StreamError for QlStreamError { - fn close_code(&self) -> Option { - match self { - QlStreamError::StreamClosed { code, .. } => Some(*code), - QlStreamError::NoSession => None, - } - } -} diff --git a/ql-runtime/src/rpc/error.rs b/ql-runtime/src/rpc/error.rs index a7f322f5..e1d1f384 100644 --- a/ql-runtime/src/rpc/error.rs +++ b/ql-runtime/src/rpc/error.rs @@ -34,21 +34,12 @@ impl From for RpcError { } } -impl From> for RpcError { - fn from(error: ql_rpc::CodecError) -> Self { +impl From> for RpcError { + fn from(error: ql_rpc::RpcError) -> Self { match error { - ql_rpc::CodecError::Rpc(error) => Self::Protocol(error), - ql_rpc::CodecError::Codec(error) => Self::Codec(error), - } - } -} - -impl From> for RpcError { - fn from(error: ql_rpc::CallError) -> Self { - match error { - ql_rpc::CallError::Protocol(error) => Self::Protocol(error), - ql_rpc::CallError::Codec(error) => Self::Codec(error), - ql_rpc::CallError::Transport(error) => error.into(), + ql_rpc::RpcError::Protocol(error) => Self::Protocol(error), + ql_rpc::RpcError::Codec(error) => Self::Codec(error), + ql_rpc::RpcError::Transport(error) => error.into(), } } } From 6ea9651e3d94d9bf15319a5b76fd264826c2a0eb Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 13:58:02 -0400 Subject: [PATCH 16/59] ql: better stream close codes --- ql-common/src/lib.rs | 38 ++++++++++++++++++++------- ql-rpc/src/error.rs | 4 +-- ql-rpc/src/rpc/download/client.rs | 6 ++--- ql-rpc/src/rpc/download/server.rs | 6 ++--- ql-rpc/src/rpc/duplex/client.rs | 4 +-- ql-rpc/src/rpc/progress/client.rs | 2 +- ql-rpc/src/rpc/progress/server.rs | 2 +- ql-rpc/src/rpc/request/server.rs | 2 +- ql-rpc/src/rpc/subscription/client.rs | 2 +- ql-rpc/src/rpc/subscription/server.rs | 2 +- ql-rpc/src/rpc/upload/client.rs | 4 +-- ql-rpc/src/rpc/upload/server.rs | 4 +-- ql-runtime/src/driver/mod.rs | 6 ++--- ql-runtime/src/io/reader.rs | 4 +-- ql-runtime/src/io/writer.rs | 2 +- ql-runtime/src/tests/stream.rs | 6 ++--- 16 files changed, 57 insertions(+), 37 deletions(-) diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index ee46c003..61b6085a 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -8,23 +8,43 @@ pub use varint::*; pub struct StreamCloseCode(pub u16); impl StreamCloseCode { - /// the stream was aborted intentionally before graceful completion + /// operation was explicitly cancelled pub const CANCELLED: Self = Self(0); + /// local reader/writer/call handle was dropped before completion + pub const DROPPED: Self = Self(1); + /// session/connection became unavailable while the stream was active + pub const DISCONNECTED: Self = Self(2); /// local internal error - pub const INTERNAL: Self = Self(1); - /// request was refused - pub const REFUSED: Self = Self(2); + pub const INTERNAL: Self = Self(3); + /// malformed stream data, invalid framing, or invalid RPC sequence + pub const PROTOCOL: Self = Self(4); + /// application codec failed to encode/decode payload + pub const CODEC: Self = Self(5); + /// stream/request was intentionally refused before processing + pub const REFUSED: Self = Self(6); /// operation timed out - pub const TIMEOUT: Self = Self(3); - /// configured limit was exceeded - pub const LIMIT: Self = Self(4); + pub const TIMEOUT: Self = Self(7); + /// configured or encoded size limit was exceeded + pub const LIMIT: Self = Self(8); /// route identifier was unknown - pub const UNKNOWN_ROUTE: Self = Self(5); + pub const UNKNOWN_ROUTE: Self = Self(9); } impl std::fmt::Display for StreamCloseCode { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.0) + match *self { + Self::CANCELLED => f.write_str("cancelled"), + Self::DROPPED => f.write_str("dropped"), + Self::DISCONNECTED => f.write_str("disconnected"), + Self::INTERNAL => f.write_str("internal"), + Self::PROTOCOL => f.write_str("protocol"), + Self::CODEC => f.write_str("codec"), + Self::REFUSED => f.write_str("refused"), + Self::TIMEOUT => f.write_str("timeout"), + Self::LIMIT => f.write_str("limit"), + Self::UNKNOWN_ROUTE => f.write_str("unknown route"), + Self(code) => write!(f, "{code}"), + } } } diff --git a/ql-rpc/src/error.rs b/ql-rpc/src/error.rs index 5892e271..d6c9cd35 100644 --- a/ql-rpc/src/error.rs +++ b/ql-rpc/src/error.rs @@ -30,7 +30,7 @@ impl Error { Self::Truncated | Self::UnexpectedFrameKind(_) | Self::MissingResponse - | Self::TrailingBytes => StreamCloseCode::REFUSED, + | Self::TrailingBytes => StreamCloseCode::PROTOCOL, } } } @@ -46,7 +46,7 @@ impl RpcError { pub const fn close_code(&self) -> Option { match self { Self::Protocol(error) => Some(error.close_code()), - Self::Codec(_) => Some(StreamCloseCode::REFUSED), + Self::Codec(_) => Some(StreamCloseCode::CODEC), Self::Transport(_) => None, } } diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index e8d4c962..db93a24c 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -95,7 +95,7 @@ where R: RpcRead, { fn drop(&mut self) { - self.close_inner(StreamCloseCode::CANCELLED); + self.close_inner(StreamCloseCode::DROPPED); } } @@ -191,7 +191,7 @@ where { fn drop(&mut self) { if self.stream.is_some() { - self.close_inner(StreamCloseCode::CANCELLED); + self.close_inner(StreamCloseCode::DROPPED); } } } @@ -235,7 +235,7 @@ where { fn drop(&mut self) { if !self.finished { - self.parent.close_inner(StreamCloseCode::CANCELLED); + self.parent.close_inner(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs index cd6119b3..819301db 100644 --- a/ql-rpc/src/rpc/download/server.rs +++ b/ql-rpc/src/rpc/download/server.rs @@ -103,7 +103,7 @@ where { fn drop(&mut self) { if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::CANCELLED); + writer.close(StreamCloseCode::DROPPED); } } } @@ -149,7 +149,7 @@ where { fn drop(&mut self) { if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::CANCELLED); + writer.close(StreamCloseCode::DROPPED); } } } @@ -184,7 +184,7 @@ where fn drop(&mut self) { if !self.finished { if let Some(writer) = self.parent.writer.take() { - writer.close(StreamCloseCode::CANCELLED); + writer.close(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index c152b08d..4045e294 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -77,7 +77,7 @@ where { fn drop(&mut self) { if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::CANCELLED); + writer.close(StreamCloseCode::DROPPED); } } } @@ -158,7 +158,7 @@ where { fn drop(&mut self) { if self.stream.is_some() { - self.close_inner(StreamCloseCode::CANCELLED); + self.close_inner(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs index b34bfaea..651ac44f 100644 --- a/ql-rpc/src/rpc/progress/client.rs +++ b/ql-rpc/src/rpc/progress/client.rs @@ -117,7 +117,7 @@ where { fn drop(&mut self) { if matches!(self.state, State::Reading(_)) { - self.close_inner(StreamCloseCode::CANCELLED); + self.close_inner(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index 29991a33..61fc32be 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -70,7 +70,7 @@ where { fn drop(&mut self) { if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::CANCELLED); + writer.close(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs index 69dada62..ca5c1107 100644 --- a/ql-rpc/src/rpc/request/server.rs +++ b/ql-rpc/src/rpc/request/server.rs @@ -60,7 +60,7 @@ where { fn drop(&mut self) { if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::CANCELLED); + writer.close(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index 6c10a17e..4be7762a 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -93,7 +93,7 @@ where { fn drop(&mut self) { if self.stream.is_some() { - self.close_inner(StreamCloseCode::CANCELLED); + self.close_inner(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs index 6cfe9228..9fa3d3b9 100644 --- a/ql-rpc/src/rpc/subscription/server.rs +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -68,7 +68,7 @@ where { fn drop(&mut self) { if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::CANCELLED); + writer.close(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index 996089f1..73d64f25 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -101,7 +101,7 @@ where R: RpcRead, { fn drop(&mut self) { - self.close(StreamCloseCode::CANCELLED); + self.close(StreamCloseCode::DROPPED); } } @@ -136,7 +136,7 @@ where { fn drop(&mut self) { if !self.finished { - self.parent.close(StreamCloseCode::CANCELLED); + self.parent.close(StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index 36f67b19..a5eafc32 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -126,7 +126,7 @@ where { fn drop(&mut self) { if self.stream.is_some() { - self.close_inner(StreamCloseCode::CANCELLED); + self.close_inner(StreamCloseCode::DROPPED); } } } @@ -172,7 +172,7 @@ where { fn drop(&mut self) { if !self.finished { - self.parent.close_inner(StreamCloseCode::CANCELLED); + self.parent.close_inner(StreamCloseCode::DROPPED); } } } diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 2b45f872..8b366ad1 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -255,7 +255,7 @@ impl DriverState { stream.inbound_close(); stream.outbound_close(); } - stream_ops.close(CloseTarget::Both, StreamCloseCode::CANCELLED); + stream_ops.close(CloseTarget::Both, StreamCloseCode::DROPPED); drop(stream_ops); return; } @@ -368,7 +368,7 @@ impl DriverState { "dropping inbound stream because handle channel is unavailable: stream_id={stream_id}" ); if let Ok(mut stream) = fsm.stream(stream_id) { - stream.close(CloseTarget::Both, StreamCloseCode::CANCELLED); + stream.close(CloseTarget::Both, StreamCloseCode::DISCONNECTED); } return; }; @@ -456,7 +456,7 @@ impl DriverState { stream_ops.commit_read(accepted).unwrap(); } if peer_closed { - stream_ops.close(target, StreamCloseCode::CANCELLED); + stream_ops.close(target, StreamCloseCode::DROPPED); if let Entry::Occupied(entry) = self.streams.entry(stream_id) { Self::try_reap_stream(entry); } diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs index 8c40ccd3..17a0bca9 100644 --- a/ql-runtime/src/io/reader.rs +++ b/ql-runtime/src/io/reader.rs @@ -181,12 +181,12 @@ impl Drop for StreamReader { "byte reader drop close: stream_id={:?} target={:?} code={:?}", self.rx.stream_id(), self.target, - StreamCloseCode::CANCELLED + StreamCloseCode::DROPPED ); self.handle.try_send(Command::CloseStream { stream_id: self.rx.stream_id(), target: self.target, - code: StreamCloseCode::CANCELLED, + code: StreamCloseCode::DROPPED, }); } } diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs index cfad3196..f20d0814 100644 --- a/ql-runtime/src/io/writer.rs +++ b/ql-runtime/src/io/writer.rs @@ -215,7 +215,7 @@ impl StreamWriter { impl Drop for StreamWriter { fn drop(&mut self) { - self.close_inner(StreamCloseCode::CANCELLED); + self.close_inner(StreamCloseCode::DROPPED); } } diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index 757f94c5..6c925d2d 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -175,14 +175,14 @@ async fn dropping_responder_closes_initiator_response() { assert!(matches!( err, QlStreamError::StreamClosed { code, origin } - if code == StreamCloseCode::CANCELLED && origin == StreamCloseOrigin::Peer + if code == StreamCloseCode::DROPPED && origin == StreamCloseOrigin::Peer )); let err = next_chunk(&mut stream.reader).await.unwrap_err(); assert!(matches!( err, QlStreamError::StreamClosed { code, origin } - if code == StreamCloseCode::CANCELLED && origin == StreamCloseOrigin::Peer + if code == StreamCloseCode::DROPPED && origin == StreamCloseOrigin::Peer )); tokio::time::timeout(Duration::from_secs(2), responder) @@ -216,7 +216,7 @@ async fn dropping_inbound_reader_cancels_remote_writer() { assert!(matches!( err, QlStreamError::StreamClosed { code, origin } - if code == StreamCloseCode::CANCELLED && origin == StreamCloseOrigin::Peer + if code == StreamCloseCode::DROPPED && origin == StreamCloseOrigin::Peer )); }); From de6e69e84b5d9cce6b01c91d6f13d1561fab23f5 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 14:14:57 -0400 Subject: [PATCH 17/59] ql-rpc: drop wrappers --- ql-rpc/src/rpc/download/client.rs | 53 +++---------- ql-rpc/src/rpc/download/server.rs | 65 +++++----------- ql-rpc/src/rpc/duplex/client.rs | 59 ++++----------- ql-rpc/src/rpc/progress/client.rs | 28 ++----- ql-rpc/src/rpc/progress/server.rs | 30 ++------ ql-rpc/src/rpc/request/server.rs | 30 +++----- ql-rpc/src/rpc/subscription/client.rs | 33 +++----- ql-rpc/src/rpc/subscription/server.rs | 27 ++----- ql-rpc/src/rpc/upload/client.rs | 48 ++++-------- ql-rpc/src/rpc/upload/server.rs | 29 ++----- ql-rpc/src/stream.rs | 104 ++++++++++++++++++++++++++ 11 files changed, 213 insertions(+), 293 deletions(-) diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index db93a24c..a1d0304d 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -5,7 +5,7 @@ use bytes::{BufMut, Bytes}; use crate::{ download::{Download, PartReadStep}, rpc::parts::FrameKind, - FramedPrefixStep, FramedReader, RpcCodec, RpcError, RpcRead, StreamCloseCode, + DropCloseRead, FramedPrefixStep, FramedReader, RpcCodec, RpcError, RpcRead, StreamCloseCode, }; pub struct DownloadCall @@ -13,7 +13,7 @@ where M: Download, R: RpcRead, { - stream: Option, + stream: DropCloseRead, reader: Option>, } @@ -31,7 +31,7 @@ where M: Download, R: RpcRead, { - stream: Option, + stream: DropCloseRead, reader: crate::download::PartFrameReader, } @@ -42,7 +42,7 @@ where { pub fn new(stream: R) -> Self { Self { - stream: Some(stream), + stream: DropCloseRead::new(stream), reader: Some(FramedReader::default()), } } @@ -54,11 +54,10 @@ where let reader = self.reader.take().unwrap(); let reader = match reader.advance_prefix() { Ok(FramedPrefixStep::Value { value, bytes }) => { - let stream = self.stream.take().unwrap(); return Ok(( value, DownloadReader { - stream: Some(stream), + stream: self.stream, reader: crate::download::PartFrameReader::::new(bytes), }, )); @@ -67,8 +66,7 @@ where Err(error) => return Err(error), }; - let stream = self.stream.as_mut().unwrap(); - match poll_fn(|cx| stream.poll_read(usize::MAX, cx)).await { + match poll_fn(|cx| self.stream.poll_read(usize::MAX, cx)).await { Ok(Some(chunk)) => { self.reader = Some(reader.push(chunk)); } @@ -83,19 +81,7 @@ where } fn close_inner(&mut self, code: StreamCloseCode) { - if let Some(stream) = self.stream.take() { - stream.close(code); - } - } -} - -impl Drop for DownloadCall -where - M: Download, - R: RpcRead, -{ - fn drop(&mut self) { - self.close_inner(StreamCloseCode::DROPPED); + DropCloseRead::close(&mut self.stream, code); } } @@ -107,7 +93,7 @@ where pub async fn next_part( &mut self, ) -> Result)>, RpcError> { - if self.stream.is_none() { + if !self.stream.is_some() { return Ok(None); } @@ -120,7 +106,7 @@ where }, ))), PartReadStep::Finish => { - self.stream.take(); + self.stream.disarm(); Ok(None) } PartReadStep::BodyBytes(_) => { @@ -136,7 +122,7 @@ where pub async fn complete(mut self) -> Result<(), RpcError> { match self.read_frame().await? { PartReadStep::Finish => { - self.stream.take(); + self.stream.disarm(); Ok(()) } PartReadStep::PartHeader(_) => { @@ -166,8 +152,7 @@ where Err(error) => return Err(error), } - let stream = self.stream.as_mut().unwrap(); - match poll_fn(|cx| stream.poll_read(usize::MAX, cx)).await { + match poll_fn(|cx| self.stream.poll_read(usize::MAX, cx)).await { Ok(Some(chunk)) => { self.reader.push(chunk); } @@ -178,21 +163,7 @@ where } fn close_inner(&mut self, code: StreamCloseCode) { - if let Some(stream) = self.stream.take() { - stream.close(code); - } - } -} - -impl Drop for DownloadReader -where - M: Download, - R: RpcRead, -{ - fn drop(&mut self) { - if self.stream.is_some() { - self.close_inner(StreamCloseCode::DROPPED); - } + DropCloseRead::close(&mut self.stream, code); } } diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs index 819301db..142d7e51 100644 --- a/ql-rpc/src/rpc/download/server.rs +++ b/ql-rpc/src/rpc/download/server.rs @@ -10,7 +10,8 @@ use crate::{ parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, read_eof_request, }, - write_bytes, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, + write_bytes, DropCloseWrite, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, + StreamCloseCode, }; #[trait_variant::make(DownloadHandler: Send)] @@ -29,7 +30,7 @@ where M: DownloadRpc, W: RpcWrite, { - writer: Option, + writer: DropCloseWrite, marker: PhantomData M>, } @@ -38,7 +39,7 @@ where M: DownloadRpc, W: RpcWrite, { - writer: Option, + writer: DropCloseWrite, marker: PhantomData M>, } @@ -58,29 +59,29 @@ where { pub(crate) fn new(writer: W) -> Self { Self { - writer: Some(writer), + writer: DropCloseWrite::new(writer), marker: PhantomData, } } /// send the response header and begin streaming parts pub async fn start( - mut self, + self, response_header: M::ResponseHeader, ) -> Result, W::Error> { - let mut writer = self.writer.take().unwrap(); + let mut writer = self.writer; let mut encoded = Vec::new(); codec::encode_value_part(&response_header, &mut encoded); write_bytes(&mut writer, Bytes::from(encoded)).await?; Ok(DownloadWriter { - writer: Some(writer), + writer, marker: PhantomData, }) } /// send a header-only response and finish the stream - pub async fn complete(mut self, response_header: M::ResponseHeader) -> Result<(), W::Error> { - let mut writer = self.writer.take().unwrap(); + pub async fn complete(self, response_header: M::ResponseHeader) -> Result<(), W::Error> { + let mut writer = self.writer; let mut encoded = Vec::new(); codec::encode_value_part(&response_header, &mut encoded); encode_finish(&mut encoded); @@ -90,21 +91,7 @@ where /// close the stream with a transport code pub fn close(mut self, code: StreamCloseCode) { - if let Some(writer) = self.writer.take() { - writer.close(code); - } - } -} - -impl Drop for DownloadStart -where - M: DownloadRpc, - W: RpcWrite, -{ - fn drop(&mut self) { - if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::DROPPED); - } + DropCloseWrite::close(&mut self.writer, code); } } @@ -117,7 +104,7 @@ where &mut self, part_header: M::PartHeader, ) -> Result, W::Error> { - let writer = self.writer.as_mut().unwrap(); + let writer = &mut self.writer; let mut encoded = Vec::new(); encode_part_header(&part_header, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await?; @@ -127,8 +114,8 @@ where }) } - pub async fn finish(mut self) -> Result<(), W::Error> { - let mut writer = self.writer.take().unwrap(); + pub async fn finish(self) -> Result<(), W::Error> { + let mut writer = self.writer; let mut encoded = Vec::new(); encode_finish(&mut encoded); write_bytes(&mut writer, Bytes::from(encoded)).await?; @@ -136,21 +123,7 @@ where } pub fn close(mut self, code: StreamCloseCode) { - if let Some(writer) = self.writer.take() { - writer.close(code); - } - } -} - -impl Drop for DownloadWriter -where - M: DownloadRpc, - W: RpcWrite, -{ - fn drop(&mut self) { - if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::DROPPED); - } + DropCloseWrite::close(&mut self.writer, code); } } @@ -160,14 +133,14 @@ where W: RpcWrite, { pub async fn send(&mut self, bytes: Bytes) -> Result<(), W::Error> { - let writer = self.parent.writer.as_mut().unwrap(); + let writer = &mut self.parent.writer; let mut encoded = Vec::new(); encode_body_chunk(&bytes, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await } pub async fn finish(mut self) -> Result<(), W::Error> { - let writer = self.parent.writer.as_mut().unwrap(); + let writer = &mut self.parent.writer; let mut encoded = Vec::new(); encode_end_part(&mut encoded); write_bytes(writer, Bytes::from(encoded)).await?; @@ -183,9 +156,7 @@ where { fn drop(&mut self) { if !self.finished { - if let Some(writer) = self.parent.writer.take() { - writer.close(StreamCloseCode::DROPPED); - } + DropCloseWrite::close(&mut self.parent.writer, StreamCloseCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index 4045e294..2f785e40 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -8,7 +8,8 @@ use bytes::Bytes; use crate::{ duplex::{codec, Duplex, EventReader, ReadStep}, - finish_bytes, write_bytes, RpcCodec, RpcError, RpcRead, RpcWrite, StreamCloseCode, + finish_bytes, write_bytes, DropCloseRead, DropCloseWrite, RpcCodec, RpcError, RpcRead, + RpcWrite, StreamCloseCode, }; pub struct DuplexCall @@ -26,7 +27,7 @@ where T: RpcCodec, W: RpcWrite, { - writer: Option, + writer: DropCloseWrite, marker: PhantomData T>, } @@ -35,7 +36,7 @@ where T: RpcCodec, R: RpcRead, { - stream: Option, + stream: DropCloseRead, reader: EventReader, } @@ -46,39 +47,24 @@ where { pub fn new(writer: W) -> Self { Self { - writer: Some(writer), + writer: DropCloseWrite::new(writer), marker: PhantomData, } } pub async fn send(&mut self, event: &T) -> Result<(), W::Error> { - let writer = self.writer.as_mut().unwrap(); + let writer = &mut self.writer; let mut encoded = Vec::new(); codec::encode_event(event, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await } pub async fn finish(mut self) -> Result<(), W::Error> { - let mut writer = self.writer.take().unwrap(); - finish_bytes(&mut writer).await + finish_bytes(&mut self.writer).await } pub fn close(mut self, code: StreamCloseCode) { - if let Some(writer) = self.writer.take() { - writer.close(code); - } - } -} - -impl Drop for DuplexSender -where - T: RpcCodec, - W: RpcWrite, -{ - fn drop(&mut self) { - if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::DROPPED); - } + DropCloseWrite::close(&mut self.writer, code); } } @@ -89,7 +75,7 @@ where { pub fn new(stream: R) -> Self { Self { - stream: Some(stream), + stream: DropCloseRead::new(stream), reader: EventReader::default(), } } @@ -102,7 +88,7 @@ where &mut self, cx: &mut Context<'_>, ) -> Poll>>> { - if self.stream.is_none() { + if !self.stream.is_some() { return Poll::Ready(None); } @@ -116,21 +102,20 @@ where } } - let stream = self.stream.as_mut().unwrap(); - match stream.poll_read(usize::MAX, cx) { + match self.stream.poll_read(usize::MAX, cx) { Poll::Ready(Ok(Some(chunk))) => { self.reader.push(chunk); } Poll::Ready(Ok(None)) => { if self.reader.is_empty() { - self.stream.take(); + self.stream.disarm(); return Poll::Ready(None); } - self.stream.take(); + self.stream.disarm(); return Poll::Ready(Some(Err(crate::Error::Truncated.into()))); } Poll::Ready(Err(error)) => { - self.stream.take(); + self.stream.disarm(); return Poll::Ready(Some(Err(RpcError::Transport(error)))); } Poll::Pending => { @@ -145,20 +130,6 @@ where } fn close_inner(&mut self, code: StreamCloseCode) { - if let Some(stream) = self.stream.take() { - stream.close(code); - } - } -} - -impl Drop for DuplexReceiver -where - T: RpcCodec, - R: RpcRead, -{ - fn drop(&mut self) { - if self.stream.is_some() { - self.close_inner(StreamCloseCode::DROPPED); - } + DropCloseRead::close(&mut self.stream, code); } } diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs index 651ac44f..a1dbd75c 100644 --- a/ql-rpc/src/rpc/progress/client.rs +++ b/ql-rpc/src/rpc/progress/client.rs @@ -6,7 +6,7 @@ use std::{ use crate::{ progress::{Progress, ReadStep, ResponseReader}, - Error, RpcError, RpcRead, StreamCloseCode, + DropCloseRead, Error, RpcError, RpcRead, StreamCloseCode, }; pub struct ProgressCall @@ -14,7 +14,7 @@ where M: Progress, R: RpcRead, { - stream: Option, + stream: DropCloseRead, state: State, } @@ -42,7 +42,7 @@ where { pub fn new(stream: R) -> Self { Self { - stream: Some(stream), + stream: DropCloseRead::new(stream), state: State::Reading(ResponseReader::default()), } } @@ -62,6 +62,7 @@ where match reader.advance() { Ok(ReadStep::Progress(value)) => return Poll::Ready(Some(value)), Ok(ReadStep::Response(response)) => { + self.stream.disarm(); self.state = State::Terminal(Ok(response)); return Poll::Ready(None); } @@ -73,8 +74,7 @@ where } } - let stream = self.stream.as_mut().unwrap(); - match stream.poll_read(usize::MAX, cx) { + match self.stream.poll_read(usize::MAX, cx) { Poll::Ready(Ok(Some(chunk))) => { let State::Reading(reader) = &mut self.state else { panic!("invalid state"); @@ -82,10 +82,12 @@ where reader.push(chunk); } Poll::Ready(Ok(None)) => { + self.stream.disarm(); self.state = State::Terminal(Err(Error::MissingResponse.into())); return Poll::Ready(None); } Poll::Ready(Err(error)) => { + self.stream.disarm(); self.state = State::Terminal(Err(RpcError::Transport(error))); return Poll::Ready(None); } @@ -104,21 +106,7 @@ where fn close_inner(&mut self, code: StreamCloseCode) { self.state = State::Done; - if let Some(stream) = self.stream.take() { - stream.close(code); - } - } -} - -impl Drop for ProgressCall -where - M: Progress, - R: RpcRead, -{ - fn drop(&mut self) { - if matches!(self.state, State::Reading(_)) { - self.close_inner(StreamCloseCode::DROPPED); - } + DropCloseRead::close(&mut self.stream, code); } } diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index 61fc32be..40bbf20f 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -6,7 +6,8 @@ use crate::{ finish_bytes, progress::{encode_progress, encode_response, Progress}, rpc::read_framed_request, - write_bytes, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, + write_bytes, DropCloseWrite, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, + StreamCloseCode, }; #[trait_variant::make(ProgressHandler: Send)] @@ -25,7 +26,7 @@ where M: Progress, W: RpcWrite, { - writer: Option, + writer: DropCloseWrite, marker: PhantomData M>, } @@ -36,42 +37,27 @@ where { pub(crate) fn new(writer: W) -> Self { Self { - writer: Some(writer), + writer: DropCloseWrite::new(writer), marker: PhantomData, } } pub async fn send(&mut self, progress: M::Progress) -> Result<(), W::Error> { - let writer = self.writer.as_mut().unwrap(); + let writer = &mut self.writer; let mut encoded = Vec::new(); encode_progress::(&progress, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await } pub async fn finish(mut self, response: M::Response) -> Result<(), W::Error> { - let mut writer = self.writer.take().unwrap(); let mut encoded = Vec::new(); encode_response::(&response, &mut encoded); - write_bytes(&mut writer, Bytes::from(encoded)).await?; - finish_bytes(&mut writer).await + write_bytes(&mut self.writer, Bytes::from(encoded)).await?; + finish_bytes(&mut self.writer).await } pub fn close(mut self, code: StreamCloseCode) { - if let Some(writer) = self.writer.take() { - writer.close(code); - } - } -} - -impl Drop for ProgressResponder -where - M: Progress, - W: RpcWrite, -{ - fn drop(&mut self) { - if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::DROPPED); - } + DropCloseWrite::close(&mut self.writer, code); } } diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs index ca5c1107..55ab74bb 100644 --- a/ql-rpc/src/rpc/request/server.rs +++ b/ql-rpc/src/rpc/request/server.rs @@ -3,8 +3,9 @@ use std::{future::Future, marker::PhantomData}; use bytes::Bytes; use crate::{ - finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, RouterConfig, - RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, + finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, + DropCloseWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, + StreamCloseCode, }; #[trait_variant::make(RequestHandler: Send)] @@ -22,7 +23,7 @@ pub struct Response where W: RpcWrite, { - writer: Option, + writer: DropCloseWrite, marker: PhantomData T>, } @@ -33,35 +34,22 @@ where { pub(crate) fn new(writer: W) -> Self { Self { - writer: Some(writer), + writer: DropCloseWrite::new(writer), marker: PhantomData, } } pub async fn respond(mut self, response: T) -> Result<(), W::Error> { - let mut writer = self.writer.take().unwrap(); + let writer = &mut self.writer; let mut encoded = Vec::new(); response.encode_value(&mut encoded); - write_bytes(&mut writer, Bytes::from(encoded)).await?; - finish_bytes(&mut writer).await?; + write_bytes(writer, Bytes::from(encoded)).await?; + finish_bytes(writer).await?; Ok(()) } pub fn close(mut self, code: StreamCloseCode) { - if let Some(writer) = self.writer.take() { - writer.close(code); - } - } -} - -impl Drop for Response -where - W: RpcWrite, -{ - fn drop(&mut self) { - if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::DROPPED); - } + DropCloseWrite::close(&mut self.writer, code); } } diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index 4be7762a..c3636019 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -5,7 +5,7 @@ use std::{ use crate::{ subscription::{ReadStep, ResponseReader, Subscription}, - RpcError, RpcRead, StreamCloseCode, + DropCloseRead, RpcError, RpcRead, StreamCloseCode, }; pub struct SubscriptionCall @@ -13,7 +13,7 @@ where M: Subscription, R: RpcRead, { - stream: Option, + stream: DropCloseRead, reader: ResponseReader, } @@ -24,7 +24,7 @@ where { pub fn new(stream: R) -> Self { Self { - stream: Some(stream), + stream: DropCloseRead::new(stream), reader: ResponseReader::default(), } } @@ -37,7 +37,7 @@ where &mut self, cx: &mut Context<'_>, ) -> Poll>>> { - if self.stream.is_none() { + if !self.stream.is_some() { return Poll::Ready(None); } @@ -51,21 +51,20 @@ where } } - let stream = self.stream.as_mut().unwrap(); - match stream.poll_read(usize::MAX, cx) { + match self.stream.poll_read(usize::MAX, cx) { Poll::Ready(Ok(Some(chunk))) => { self.reader.push(chunk); } Poll::Ready(Ok(None)) => { if self.reader.is_empty() { - self.stream.take(); + self.stream.disarm(); return Poll::Ready(None); } - self.stream.take(); + self.stream.disarm(); return Poll::Ready(Some(Err(crate::Error::Truncated.into()))); } Poll::Ready(Err(error)) => { - self.stream.take(); + self.stream.disarm(); return Poll::Ready(Some(Err(RpcError::Transport(error)))); } Poll::Pending => { @@ -80,20 +79,6 @@ where } fn close_inner(&mut self, code: StreamCloseCode) { - if let Some(stream) = self.stream.take() { - stream.close(code); - } - } -} - -impl Drop for SubscriptionCall -where - M: Subscription, - R: RpcRead, -{ - fn drop(&mut self) { - if self.stream.is_some() { - self.close_inner(StreamCloseCode::DROPPED); - } + DropCloseRead::close(&mut self.stream, code); } } diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs index 9fa3d3b9..0f687a87 100644 --- a/ql-rpc/src/rpc/subscription/server.rs +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -4,7 +4,8 @@ use bytes::Bytes; use crate::{ codec, finish_bytes, rpc::read_eof_request, subscription::Subscription as SubscriptionRpc, - write_bytes, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, + write_bytes, DropCloseWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, + StreamCloseCode, }; #[trait_variant::make(SubscriptionHandler: Send)] @@ -26,7 +27,7 @@ pub struct SubscriptionResponder where W: RpcWrite, { - writer: Option, + writer: DropCloseWrite, marker: PhantomData T>, } @@ -37,13 +38,13 @@ where { pub(crate) fn new(writer: W) -> Self { Self { - writer: Some(writer), + writer: DropCloseWrite::new(writer), marker: PhantomData, } } pub async fn send(&mut self, event: T) -> Result<(), W::Error> { - let writer = self.writer.as_mut().unwrap(); + let writer = &mut self.writer; let mut encoded = Vec::new(); codec::encode_value_part(&event, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await?; @@ -51,25 +52,11 @@ where } pub async fn finish(mut self) -> Result<(), W::Error> { - let mut writer = self.writer.take().unwrap(); - finish_bytes(&mut writer).await + finish_bytes(&mut self.writer).await } pub fn close(mut self, code: StreamCloseCode) { - if let Some(writer) = self.writer.take() { - writer.close(code); - } - } -} - -impl Drop for SubscriptionResponder -where - W: RpcWrite, -{ - fn drop(&mut self) { - if let Some(writer) = self.writer.take() { - writer.close(StreamCloseCode::DROPPED); - } + DropCloseWrite::close(&mut self.writer, code); } } diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index 73d64f25..af2af6ab 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -4,7 +4,8 @@ use crate::{ finish_bytes, read_bytes, rpc::parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, upload::Upload, - write_bytes, ChunkQueue, RpcCodec, RpcError, RpcRead, RpcWrite, StreamCloseCode, + write_bytes, ChunkQueue, DropCloseRead, DropCloseWrite, RpcCodec, RpcError, RpcRead, RpcWrite, + StreamCloseCode, }; pub struct UploadCall @@ -13,8 +14,8 @@ where W: RpcWrite, R: RpcRead, { - writer: Option, - reader: Option, + writer: DropCloseWrite, + reader: DropCloseRead, marker: std::marker::PhantomData M>, } @@ -36,8 +37,8 @@ where { pub fn new(writer: W, reader: R) -> Self { Self { - writer: Some(writer), - reader: Some(reader), + writer: DropCloseWrite::new(writer), + reader: DropCloseRead::new(reader), marker: std::marker::PhantomData, } } @@ -46,7 +47,7 @@ where &mut self, part_header: M::PartHeader, ) -> Result, W::Error> { - let writer = self.writer.as_mut().unwrap(); + let writer = &mut self.writer; let mut encoded = Vec::new(); encode_part_header(&part_header, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await?; @@ -57,20 +58,18 @@ where } pub async fn finish(mut self) -> Result> { - let mut writer = self.writer.take().unwrap(); + let writer = &mut self.writer; let mut encoded = Vec::new(); encode_finish(&mut encoded); - write_bytes(&mut writer, Bytes::from(encoded)) - .await - .map_err(RpcError::Transport)?; - finish_bytes(&mut writer) + write_bytes(writer, Bytes::from(encoded)) .await .map_err(RpcError::Transport)?; + finish_bytes(writer).await.map_err(RpcError::Transport)?; - let mut reader = self.reader.take().unwrap(); + let reader = &mut self.reader; let mut bytes = ChunkQueue::default(); - while let Some(chunk) = read_bytes(&mut reader, usize::MAX) + while let Some(chunk) = read_bytes(reader, usize::MAX) .await .map_err(RpcError::Transport)? { @@ -85,23 +84,8 @@ where } fn close(&mut self, code: StreamCloseCode) { - if let Some(reader) = self.reader.take() { - reader.close(code); - } - if let Some(writer) = self.writer.take() { - writer.close(code); - } - } -} - -impl Drop for UploadCall -where - M: Upload, - W: RpcWrite, - R: RpcRead, -{ - fn drop(&mut self) { - self.close(StreamCloseCode::DROPPED); + DropCloseRead::close(&mut self.reader, code); + DropCloseWrite::close(&mut self.writer, code); } } @@ -112,14 +96,14 @@ where R: RpcRead, { pub async fn send(&mut self, bytes: Bytes) -> Result<(), W::Error> { - let writer = self.parent.writer.as_mut().unwrap(); + let writer = &mut self.parent.writer; let mut encoded = Vec::new(); encode_body_chunk(&bytes, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await } pub async fn finish(mut self) -> Result<(), W::Error> { - let writer = self.parent.writer.as_mut().unwrap(); + let writer = &mut self.parent.writer; let mut encoded = Vec::new(); encode_end_part(&mut encoded); write_bytes(writer, Bytes::from(encoded)).await?; diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index a5eafc32..37536a40 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -8,7 +8,7 @@ use crate::{ parts::{FrameKind, PartFrameReader, PartReadStep}, read_framed_request_prefix, }, - RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, Upload, + DropCloseRead, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, Upload, }; #[trait_variant::make(UploadHandler: Send)] @@ -32,7 +32,7 @@ where M: Upload, R: RpcRead, { - stream: Option, + stream: DropCloseRead, reader: PartFrameReader, } @@ -61,7 +61,7 @@ where &mut self, ) -> Result)>, crate::RpcError> { - if self.stream.is_none() { + if !self.stream.is_some() { return Ok(None); } @@ -74,7 +74,7 @@ where }, ))), PartReadStep::Finish => { - self.stream.take(); + self.stream.disarm(); Ok(None) } PartReadStep::BodyBytes(_) => { @@ -97,8 +97,7 @@ where Err(error) => return Err(error), } - let stream = self.stream.as_mut().unwrap(); - match poll_fn(|cx| stream.poll_read(usize::MAX, cx)).await { + match poll_fn(|cx| self.stream.poll_read(usize::MAX, cx)).await { Ok(Some(chunk)) => { self.reader.push(chunk); } @@ -113,21 +112,7 @@ where } fn close_inner(&mut self, code: StreamCloseCode) { - if let Some(stream) = self.stream.take() { - stream.close(code); - } - } -} - -impl Drop for UploadReader -where - M: Upload, - R: RpcRead, -{ - fn drop(&mut self) { - if self.stream.is_some() { - self.close_inner(StreamCloseCode::DROPPED); - } + DropCloseRead::close(&mut self.stream, code); } } @@ -234,7 +219,7 @@ pub(crate) async fn handle_upload_inner( state, request, UploadReader { - stream: Some(reader), + stream: DropCloseRead::new(reader), reader: PartFrameReader::new(buffered), }, UploadResponder::new(writer), diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index 0cceff88..172e50cb 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -78,3 +78,107 @@ where reader.close(code); writer.close(code); } + +pub(crate) use drop::*; +mod drop { + use super::*; + + pub struct DropCloseRead { + inner: Option, + } + + impl DropCloseRead { + pub fn new(reader: R) -> Self { + Self { + inner: Some(reader), + } + } + + #[inline] + pub fn is_some(&self) -> bool { + self.inner.is_some() + } + + #[inline] + pub fn disarm(&mut self) { + self.inner.take(); + } + + #[inline] + pub fn close(&mut self, code: StreamCloseCode) { + if let Some(reader) = self.inner.take() { + reader.close(code); + } + } + } + + impl RpcRead for DropCloseRead { + type Error = R::Error; + + #[track_caller] + fn poll_read( + &mut self, + max_len: usize, + cx: &mut Context<'_>, + ) -> Poll, Self::Error>> { + self.inner.as_mut().unwrap().poll_read(max_len, cx) + } + + fn close(mut self, code: StreamCloseCode) { + Self::close(&mut self, code); + } + } + + impl Drop for DropCloseRead { + fn drop(&mut self) { + self.close(StreamCloseCode::DROPPED); + } + } + + pub struct DropCloseWrite { + inner: Option, + } + + impl DropCloseWrite { + pub fn new(writer: W) -> Self { + Self { + inner: Some(writer), + } + } + + #[inline] + pub fn close(&mut self, code: StreamCloseCode) { + if let Some(writer) = self.inner.take() { + writer.close(code); + } + } + } + + impl RpcWrite for DropCloseWrite { + type Error = W::Error; + + #[track_caller] + fn poll_write( + &mut self, + bytes: &mut Bytes, + cx: &mut Context<'_>, + ) -> Poll> { + self.inner.as_mut().unwrap().poll_write(bytes, cx) + } + + #[track_caller] + fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll> { + self.inner.as_mut().unwrap().poll_finish(cx) + } + + fn close(mut self, code: StreamCloseCode) { + Self::close(&mut self, code); + } + } + + impl Drop for DropCloseWrite { + fn drop(&mut self) { + self.close(StreamCloseCode::DROPPED) + } + } +} From 7e6bff77260fa8c77ae531d15c26cb3762b704be Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 15:27:12 -0400 Subject: [PATCH 18/59] ql: remove read max len --- ql-rpc/src/chunk_queue.rs | 7 ++++ ql-rpc/src/framed_value.rs | 7 ++++ ql-rpc/src/rpc/download/client.rs | 4 +-- ql-rpc/src/rpc/duplex/client.rs | 2 +- ql-rpc/src/rpc/progress/client.rs | 2 +- ql-rpc/src/rpc/request/client.rs | 5 +-- ql-rpc/src/rpc/subscription/client.rs | 2 +- ql-rpc/src/rpc/upload/client.rs | 5 +-- ql-rpc/src/rpc/upload/server.rs | 2 +- ql-rpc/src/rpc/utils.rs | 46 ++++++++++++------------ ql-rpc/src/stream.rs | 18 +++------- ql-runtime/src/io/reader.rs | 50 +++++---------------------- ql-runtime/src/rpc/adapter.rs | 8 ++--- ql-runtime/src/tests/mod.rs | 11 ++---- ql-runtime/src/tests/stream.rs | 50 --------------------------- 15 files changed, 62 insertions(+), 157 deletions(-) diff --git a/ql-rpc/src/chunk_queue.rs b/ql-rpc/src/chunk_queue.rs index 770002f2..3a6f1559 100644 --- a/ql-rpc/src/chunk_queue.rs +++ b/ql-rpc/src/chunk_queue.rs @@ -33,6 +33,13 @@ impl ChunkQueue { } } + pub(crate) fn next_part_total_len(&self) -> Result, Error> { + let Some(len) = self.peek_next_part_len()? else { + return Ok(None); + }; + Ok(Some(len.saturating_add(LENGTH_SIZE))) + } + pub fn pop_front(&mut self, max_len: usize) -> Option { let front = self.chunks.front_mut()?; let chunk = if max_len >= front.len() { diff --git a/ql-rpc/src/framed_value.rs b/ql-rpc/src/framed_value.rs index df9664c2..bf96dc6f 100644 --- a/ql-rpc/src/framed_value.rs +++ b/ql-rpc/src/framed_value.rs @@ -35,6 +35,13 @@ impl FramedReader { self } + pub(crate) fn exceeds_total_len(&self, max_len: usize) -> Result { + let Some(len) = self.bytes.next_part_total_len()? else { + return Ok(self.bytes.remaining() > max_len); + }; + Ok(len > max_len) + } + pub fn advance(self) -> Result, RpcError> { match self.advance_prefix::()? { FramedPrefixStep::NeedMore(next) => Ok(FramedReadStep::NeedMore(next)), diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index a1d0304d..5cdc41ed 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -66,7 +66,7 @@ where Err(error) => return Err(error), }; - match poll_fn(|cx| self.stream.poll_read(usize::MAX, cx)).await { + match poll_fn(|cx| self.stream.poll_read(cx)).await { Ok(Some(chunk)) => { self.reader = Some(reader.push(chunk)); } @@ -152,7 +152,7 @@ where Err(error) => return Err(error), } - match poll_fn(|cx| self.stream.poll_read(usize::MAX, cx)).await { + match poll_fn(|cx| self.stream.poll_read(cx)).await { Ok(Some(chunk)) => { self.reader.push(chunk); } diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index 2f785e40..c5326431 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -102,7 +102,7 @@ where } } - match self.stream.poll_read(usize::MAX, cx) { + match self.stream.poll_read(cx) { Poll::Ready(Ok(Some(chunk))) => { self.reader.push(chunk); } diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs index a1dbd75c..1f134122 100644 --- a/ql-rpc/src/rpc/progress/client.rs +++ b/ql-rpc/src/rpc/progress/client.rs @@ -74,7 +74,7 @@ where } } - match self.stream.poll_read(usize::MAX, cx) { + match self.stream.poll_read(cx) { Poll::Ready(Ok(Some(chunk))) => { let State::Reading(reader) = &mut self.state else { panic!("invalid state"); diff --git a/ql-rpc/src/rpc/request/client.rs b/ql-rpc/src/rpc/request/client.rs index 6f605be9..0d91263b 100644 --- a/ql-rpc/src/rpc/request/client.rs +++ b/ql-rpc/src/rpc/request/client.rs @@ -17,10 +17,7 @@ where { let mut bytes = ChunkQueue::default(); - while let Some(chunk) = read_bytes(&mut reader, usize::MAX) - .await - .map_err(RpcError::Transport)? - { + while let Some(chunk) = read_bytes(&mut reader).await.map_err(RpcError::Transport)? { bytes.push(chunk); } diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index c3636019..670c1d19 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -51,7 +51,7 @@ where } } - match self.stream.poll_read(usize::MAX, cx) { + match self.stream.poll_read(cx) { Poll::Ready(Ok(Some(chunk))) => { self.reader.push(chunk); } diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index af2af6ab..1fca05f4 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -69,10 +69,7 @@ where let reader = &mut self.reader; let mut bytes = ChunkQueue::default(); - while let Some(chunk) = read_bytes(reader, usize::MAX) - .await - .map_err(RpcError::Transport)? - { + while let Some(chunk) = read_bytes(reader).await.map_err(RpcError::Transport)? { bytes.push(chunk); } diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index 37536a40..2f80987e 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -97,7 +97,7 @@ where Err(error) => return Err(error), } - match poll_fn(|cx| self.stream.poll_read(usize::MAX, cx)).await { + match poll_fn(|cx| self.stream.poll_read(cx)).await { Ok(Some(chunk)) => { self.reader.push(chunk); } diff --git a/ql-rpc/src/rpc/utils.rs b/ql-rpc/src/rpc/utils.rs index fb33a6f5..457fe982 100644 --- a/ql-rpc/src/rpc/utils.rs +++ b/ql-rpc/src/rpc/utils.rs @@ -13,8 +13,6 @@ where R: RpcRead, { let mut value_reader = FramedReader::::default(); - let mut total_read = 0usize; - let value = loop { match value_reader.advance::() { Ok(FramedReadStep::Value(value)) => break value, @@ -22,26 +20,18 @@ where Err(error) => return Err(error), } - let remaining = config.max_request_bytes.saturating_sub(total_read); - if remaining == 0 { - return Err(RpcError::Protocol(Error::LengthOverflow)); - } - - match read_bytes(reader, remaining).await { + match read_bytes(reader).await { Ok(Some(chunk)) => { - total_read += chunk.len(); value_reader = value_reader.push(chunk); + reject_oversized_frame(&value_reader, config)?; } Ok(None) => return Err(RpcError::Protocol(Error::Truncated)), Err(error) => return Err(RpcError::Transport(error)), } }; - let remaining = config.max_request_bytes.saturating_sub(total_read); - let probe = remaining.max(1); - match read_bytes(reader, probe).await { + match read_bytes(reader).await { Ok(None) => Ok(value), - Ok(Some(_)) if remaining == 0 => Err(RpcError::Protocol(Error::LengthOverflow)), Ok(Some(_)) => Err(RpcError::Protocol(Error::TrailingBytes)), Err(error) => Err(RpcError::Transport(error)), } @@ -57,8 +47,6 @@ where R: RpcRead, { let mut value_reader = FramedReader::::default(); - let mut total_read = 0usize; - loop { match value_reader.advance_prefix::() { Ok(FramedPrefixStep::Value { value, bytes }) => return Ok((value, bytes)), @@ -66,15 +54,10 @@ where Err(error) => return Err(error), } - let remaining = config.max_request_bytes.saturating_sub(total_read); - if remaining == 0 { - return Err(RpcError::Protocol(Error::LengthOverflow)); - } - - match read_bytes(reader, remaining).await { + match read_bytes(reader).await { Ok(Some(chunk)) => { - total_read += chunk.len(); value_reader = value_reader.push(chunk); + reject_oversized_frame(&value_reader, config)?; } Ok(None) => return Err(RpcError::Protocol(Error::Truncated)), Err(error) => return Err(RpcError::Transport(error)), @@ -96,8 +79,7 @@ where loop { let remaining = config.max_request_bytes.saturating_sub(total_read); - let probe = remaining.max(1); - match read_bytes(reader, probe).await { + match read_bytes(reader).await { Ok(Some(chunk)) => { if chunk.len() > remaining { return Err(RpcError::Protocol(Error::LengthOverflow)); @@ -116,3 +98,19 @@ where } Ok(value) } + +fn reject_oversized_frame( + value_reader: &FramedReader, + config: RouterConfig, +) -> Result<(), RpcError> +where + T: RpcCodec, +{ + if value_reader + .exceeds_total_len(config.max_request_bytes) + .map_err(RpcError::Protocol)? + { + return Err(RpcError::Protocol(Error::LengthOverflow)); + } + Ok(()) +} diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index 172e50cb..07334c22 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -21,11 +21,7 @@ pub trait RpcRead { type Error; /// reads inbound bytes until eof or error - fn poll_read( - &mut self, - max_len: usize, - cx: &mut Context<'_>, - ) -> Poll, Self::Error>>; + fn poll_read(&mut self, cx: &mut Context<'_>) -> Poll, Self::Error>>; /// aborts the read side fn close(self, code: StreamCloseCode); @@ -48,11 +44,11 @@ pub trait RpcWrite { fn close(self, code: StreamCloseCode); } -pub async fn read_bytes(reader: &mut R, max_len: usize) -> Result, R::Error> +pub async fn read_bytes(reader: &mut R) -> Result, R::Error> where R: RpcRead, { - poll_fn(|cx| reader.poll_read(max_len, cx)).await + poll_fn(|cx| reader.poll_read(cx)).await } pub async fn write_bytes(writer: &mut W, bytes: Bytes) -> Result<(), W::Error> @@ -116,12 +112,8 @@ mod drop { type Error = R::Error; #[track_caller] - fn poll_read( - &mut self, - max_len: usize, - cx: &mut Context<'_>, - ) -> Poll, Self::Error>> { - self.inner.as_mut().unwrap().poll_read(max_len, cx) + fn poll_read(&mut self, cx: &mut Context<'_>) -> Poll, Self::Error>> { + self.inner.as_mut().unwrap().poll_read(cx) } fn close(mut self, code: StreamCloseCode) { diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs index 17a0bca9..3f96f2d3 100644 --- a/ql-runtime/src/io/reader.rs +++ b/ql-runtime/src/io/reader.rs @@ -16,7 +16,6 @@ use crate::{command::Command, log, QlStreamError, RuntimeHandle}; pub struct StreamReader { rx: Rx, target: CloseTarget, - pending: Bytes, terminal: ReaderTerminalState, handle: RuntimeHandle, } @@ -46,7 +45,6 @@ impl StreamReader { Self { rx: shared, target, - pending: Bytes::new(), terminal: ReaderTerminalState::Open, handle, } @@ -54,21 +52,20 @@ impl StreamReader { pub fn poll_read( &mut self, - max_len: usize, cx: &mut Context<'_>, ) -> Poll, QlStreamError>> { if matches!(self.terminal, ReaderTerminalState::Delivered) { return Poll::Ready(Ok(None)); } - match self.try_read_ready(max_len) { + match self.try_read_ready() { Poll::Ready(result) => return Poll::Ready(result), Poll::Pending => {} } self.rx.register_waiter(cx.waker()); - match self.try_read_ready(max_len) { + match self.try_read_ready() { Poll::Ready(result) => { self.rx.unregister_waiter(); Poll::Ready(result) @@ -77,22 +74,9 @@ impl StreamReader { } } - fn try_read_ready(&mut self, max_len: usize) -> Poll, QlStreamError>> { - if !self.pending.is_empty() { - let pending = &mut self.pending; - let bytes = if pending.len() <= max_len { - std::mem::take(pending) - } else { - pending.split_to(max_len) - }; - self.handle.try_send(Command::PollInbound { - stream_id: self.rx.stream_id(), - }); - return Poll::Ready(Ok(Some(bytes))); - } - + fn try_read_ready(&mut self) -> Poll, QlStreamError>> { match self.rx.pop() { - Ok(Item::Chunk(mut bytes)) => { + Ok(Item::Chunk(bytes)) => { log::trace!( "byte reader received chunk: stream_id={} target={:?} len={}", self.rx.stream_id(), @@ -102,12 +86,7 @@ impl StreamReader { self.handle.try_send(Command::PollInbound { stream_id: self.rx.stream_id(), }); - if bytes.len() <= max_len { - return Poll::Ready(Ok(Some(bytes))); - } - let head = bytes.split_to(max_len); - self.pending = bytes; - Poll::Ready(Ok(Some(head))) + Poll::Ready(Ok(Some(bytes))) } Ok(Item::Error(error)) => { log::debug!( @@ -134,19 +113,8 @@ impl StreamReader { } } - pub fn poll_read_chunk( - &mut self, - cx: &mut Context<'_>, - ) -> Poll, QlStreamError>> { - self.poll_read(usize::MAX, cx) - } - - pub async fn read(&mut self, max_len: usize) -> Result, QlStreamError> { - poll_fn(|cx| self.poll_read(max_len, cx)).await - } - - pub async fn read_chunk(&mut self) -> Result, QlStreamError> { - self.read(usize::MAX).await + pub async fn read(&mut self) -> Result, QlStreamError> { + poll_fn(|cx| self.poll_read(cx)).await } pub fn close(mut self, code: StreamCloseCode) { @@ -216,7 +184,7 @@ mod loom_tests { }) }; - let first = reader.poll_read(usize::MAX, &mut cx); + let first = reader.poll_read(&mut cx); producer.join().unwrap(); match first { @@ -225,7 +193,7 @@ mod loom_tests { } Poll::Pending => { assert_eq!( - reader.poll_read(usize::MAX, &mut cx), + reader.poll_read(&mut cx), Poll::Ready(Ok(Some(Bytes::from_static(b"abc")))) ); } diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index 53604dfb..dd6e4813 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -26,12 +26,8 @@ impl RpcStream for QlStream { impl RpcRead for StreamReader { type Error = QlStreamError; - fn poll_read( - &mut self, - max_len: usize, - cx: &mut Context<'_>, - ) -> Poll, QlStreamError>> { - StreamReader::poll_read(self, max_len, cx) + fn poll_read(&mut self, cx: &mut Context<'_>) -> Poll, QlStreamError>> { + StreamReader::poll_read(self, cx) } fn close(self, code: StreamCloseCode) { diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index c3e6996b..814511d2 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -655,20 +655,13 @@ async fn read_all(mut stream: crate::StreamReader) -> Result, QlStreamEr Ok(data) } -async fn next_chunk_max( - stream: &mut crate::StreamReader, - max_len: usize, -) -> Result>, crate::QlStreamError> { +async fn next_chunk(stream: &mut crate::StreamReader) -> Result>, QlStreamError> { stream - .read(max_len) + .read() .await .map(|chunk| chunk.map(|bytes| bytes.to_vec())) } -async fn next_chunk(stream: &mut crate::StreamReader) -> Result>, QlStreamError> { - next_chunk_max(stream, usize::MAX).await -} - fn default_runtime_config() -> RuntimeConfig { RuntimeConfig { fsm: QlFsmConfig { diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index 6c925d2d..fb23665b 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -59,56 +59,6 @@ async fn open_stream_duplex_happy_path() { .await; } -#[tokio::test(flavor = "current_thread")] -async fn reader_respects_max_len() { - run_local_test(async { - let mut pair = TestPair::new(default_runtime_config()); - pair.connect_and_wait(Side::A).await; - let inbound_b = pair.take_inbound(Side::B); - - let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - let mut reader = inbound.reader; - - assert_eq!( - next_chunk_max(&mut reader, 2).await.unwrap(), - Some(vec![1, 2]) - ); - assert_eq!( - next_chunk_max(&mut reader, 2).await.unwrap(), - Some(vec![3, 4]) - ); - assert_eq!( - next_chunk_max(&mut reader, 2).await.unwrap(), - Some(vec![5, 6]) - ); - assert_eq!(next_chunk(&mut reader).await.unwrap(), None); - - inbound.writer.finish().await.unwrap(); - }); - - let mut stream = pair - .side(Side::A) - .handle - .open_stream(test_open_stream_params()) - .await - .unwrap(); - stream - .writer - .write(Bytes::from_static(&[1, 2, 3, 4, 5, 6])) - .await - .unwrap(); - stream.writer.finish().await.unwrap(); - assert_eq!(next_chunk(&mut stream.reader).await.unwrap(), None); - - tokio::time::timeout(Duration::from_secs(2), responder) - .await - .unwrap() - .unwrap(); - }) - .await; -} - #[tokio::test(flavor = "current_thread")] async fn large_stream_payload_round_trips() { run_local_test(async { From c35bafbb27e8d30a1b155ecf2cbfd9c8ab3d0f77 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 16:20:42 -0400 Subject: [PATCH 19/59] ql: StreamClose -> StreamReset --- QL_V2.md | 6 +- ql-common/src/lib.rs | 14 +-- ql-fsm/src/fsm.rs | 8 +- ql-fsm/src/lib.rs | 16 ++-- ql-fsm/src/session/mod.rs | 66 ++++++------- ql-fsm/src/session/state.rs | 20 ++-- ql-fsm/src/session/stream_ops.rs | 10 +- ql-fsm/src/session/tests.rs | 92 +++++++++---------- ql-fsm/src/session/tracked.rs | 4 +- ql-fsm/src/tests/proptest.rs | 70 +++++++------- ql-fsm/src/tests/session.rs | 5 +- ql-rpc/src/error.rs | 14 +-- ql-rpc/src/lib.rs | 2 +- ql-rpc/src/router/mod.rs | 6 +- ql-rpc/src/rpc/download/client.rs | 30 +++--- ql-rpc/src/rpc/download/server.rs | 27 +++--- ql-rpc/src/rpc/duplex/client.rs | 24 ++--- ql-rpc/src/rpc/notification/server.rs | 12 +-- ql-rpc/src/rpc/progress/client.rs | 14 +-- ql-rpc/src/rpc/progress/server.rs | 17 ++-- ql-rpc/src/rpc/request/server.rs | 17 ++-- ql-rpc/src/rpc/subscription/client.rs | 14 +-- ql-rpc/src/rpc/subscription/server.rs | 18 ++-- ql-rpc/src/rpc/upload/client.rs | 20 ++-- ql-rpc/src/rpc/upload/server.rs | 30 +++--- ql-rpc/src/stream.rs | 52 +++++------ ql-runtime/src/command.rs | 10 +- ql-runtime/src/driver/mod.rs | 58 ++++++------ ql-runtime/src/driver/state.rs | 14 +-- ql-runtime/src/driver/test.rs | 40 ++++---- ql-runtime/src/error.rs | 10 +- ql-runtime/src/io/inner.rs | 46 +++++----- ql-runtime/src/io/mod.rs | 6 +- ql-runtime/src/io/reader.rs | 28 +++--- ql-runtime/src/io/writer.rs | 24 ++--- ql-runtime/src/rpc/adapter.rs | 10 +- ql-runtime/src/rpc/download.rs | 12 +-- ql-runtime/src/rpc/duplex.rs | 8 +- ql-runtime/src/rpc/error.rs | 12 +-- ql-runtime/src/rpc/mod.rs | 2 +- ql-runtime/src/rpc/progress.rs | 4 +- ql-runtime/src/rpc/subscription.rs | 4 +- ql-runtime/src/tests/rpc.rs | 10 +- ql-runtime/src/tests/stream.rs | 16 ++-- ql-wire/src/encrypted/builder.rs | 8 +- ql-wire/src/encrypted/close.rs | 2 +- ql-wire/src/encrypted/mod.rs | 20 ++-- .../{stream_close.rs => stream_reset.rs} | 38 ++++---- ql-wire/src/lib.rs | 3 +- ql-wire/src/tests.rs | 6 +- 50 files changed, 496 insertions(+), 503 deletions(-) rename ql-wire/src/encrypted/{stream_close.rs => stream_reset.rs} (69%) diff --git a/QL_V2.md b/QL_V2.md index 0062c7e6..c6887364 100644 --- a/QL_V2.md +++ b/QL_V2.md @@ -65,7 +65,7 @@ Today, varints are used for: - `StreamData.bytes_len` - `StreamWindow.stream_id` - `StreamWindow.maximum_offset` -- `StreamClose.stream_id` +- `StreamReset.stream_id` ### Handshake records @@ -128,7 +128,7 @@ The visible session header is authenticated as AEAD AAD but is not encrypted. | `Unpair` | 1 byte | forget the currently bound peer and abort the session | | `Ack` | `4+` bytes | acknowledge received session records with ACK ranges | | `StreamWindow` | `3..17` bytes | extend per-stream send credit | -| `StreamClose` | `5..12` bytes | abort one stream lane or both lanes | +| `StreamReset` | `5..12` bytes | abort one stream lane or both lanes | | `Close` | 3 bytes | close the whole session | | `StreamData` | `5..34 + payload_len` bytes | carry stream bytes, optional opener route, and optional `fin` | @@ -340,7 +340,7 @@ Receive credit advances when the local application commits read bytes, not merel ## Close And Liveness -`StreamClose` aborts a stream early. Semantically it can target: +`StreamReset` aborts a stream early. Semantically it can target: - the origin lane - the return lane diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index 61b6085a..5dfd6851 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -5,9 +5,9 @@ pub use varint::*; #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] -pub struct StreamCloseCode(pub u16); +pub struct ResetCode(pub u16); -impl StreamCloseCode { +impl ResetCode { /// operation was explicitly cancelled pub const CANCELLED: Self = Self(0); /// local reader/writer/call handle was dropped before completion @@ -30,7 +30,7 @@ impl StreamCloseCode { pub const UNKNOWN_ROUTE: Self = Self(9); } -impl std::fmt::Display for StreamCloseCode { +impl std::fmt::Display for ResetCode { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match *self { Self::CANCELLED => f.write_str("cancelled"), @@ -48,12 +48,12 @@ impl std::fmt::Display for StreamCloseCode { } } -/// origin of a stream close: either we triggered it locally or the peer sent it. +/// origin of a stream reset: either we triggered it locally or the peer sent it. #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -pub enum StreamCloseOrigin { - /// the close code originated from the peer +pub enum ResetOrigin { + /// the reset code originated from the peer Peer, - /// the close code originated from local logic + /// the reset code originated from local logic Local, } diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index 75adbb6b..6a804af8 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -46,11 +46,11 @@ impl session::EventSink for EventSink<'_> { SessionEvent::OutboundFinished(stream_id) => { self.events.push_back(Event::OutboundFinished(stream_id)); } - SessionEvent::Closed(frame) => { - self.events.push_back(Event::Closed(frame)); + SessionEvent::Reset(frame) => { + self.events.push_back(Event::Reset(frame)); } - SessionEvent::WritableClosed(frame) => { - self.events.push_back(Event::WritableClosed(frame)); + SessionEvent::WritableReset(frame) => { + self.events.push_back(Event::WritableReset(frame)); } SessionEvent::SessionClosed(close) => { self.termination = Some(TerminalFrame::Close(close.clone())); diff --git a/ql-fsm/src/lib.rs b/ql-fsm/src/lib.rs index 0ff2c68d..d70fd9f1 100644 --- a/ql-fsm/src/lib.rs +++ b/ql-fsm/src/lib.rs @@ -37,7 +37,7 @@ pub use error::*; pub use pairing::PairingInvite; use ql_wire::{ PairingToken, PeerBundle, QlCrypto, QlIdentity, RouteId, ServiceId, SessionClose, - SessionCloseCode, StreamClose, StreamHeader, StreamId, + SessionCloseCode, StreamHeader, StreamId, StreamReset, }; pub use session::{SessionEvent, StreamReadIter, StreamWriter}; @@ -76,10 +76,10 @@ pub enum Event { Finished(StreamId), /// our local FIN was acknowledged by the peer at the session layer OutboundFinished(StreamId), - /// a stream was closed - Closed(StreamClose), - /// local writes on this stream are closed - WritableClosed(StreamClose), + /// a stream was reset + Reset(StreamReset), + /// local writes on this stream are reset + WritableReset(StreamReset), /// the encrypted session was closed /// /// session close is abortive and best-effort. the session ends immediately @@ -135,9 +135,9 @@ impl StreamOps<'_> { self.inner.writer() } - /// closes the origin lane, return lane, or both lanes of the stream - pub fn close(&mut self, target: ql_wire::CloseTarget, code: ql_wire::StreamCloseCode) { - self.inner.close(target, code); + /// resets the origin lane, return lane, or both lanes of the stream + pub fn reset(&mut self, target: ql_wire::ResetTarget, code: ql_wire::ResetCode) { + self.inner.reset(target, code); } } diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index df87352d..05724e89 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -18,8 +18,8 @@ use std::time::{Duration, Instant}; use bytes::Bytes; use indexmap::IndexMap; use ql_wire::{ - CloseTarget, RecordAck, RecordSeq, RouteId, ServiceId, SessionClose, SessionCloseCode, - SessionFrame, SessionRecordBuilder, StreamClose, StreamData, StreamHeader, StreamId, + RecordAck, RecordSeq, ResetTarget, RouteId, ServiceId, SessionClose, SessionCloseCode, + SessionFrame, SessionRecordBuilder, StreamData, StreamHeader, StreamId, StreamReset, StreamWindow, VarInt, WireError, }; @@ -72,8 +72,8 @@ pub enum SessionEvent { Writable(StreamId), Finished(StreamId), OutboundFinished(StreamId), - Closed(StreamClose), - WritableClosed(StreamClose), + Reset(StreamReset), + WritableReset(StreamReset), SessionClosed(SessionClose), Unpaired, } @@ -248,8 +248,8 @@ impl SessionFsm { } } SessionFrame::StreamWindow(frame) => self.handle_stream_window(&frame, sink), - SessionFrame::StreamClose(frame) => { - if self.handle_stream_close(&frame, sink).is_err() { + SessionFrame::StreamReset(frame) => { + if self.handle_stream_reset(&frame, sink).is_err() { self.close(SessionCloseCode::PROTOCOL, sink); return; } @@ -409,7 +409,7 @@ impl SessionFsm { sent_at: None, }; - self.push_next_pending_stream_close(&mut builder, &mut outbound); + self.push_next_pending_stream_reset(&mut builder, &mut outbound); if self.state.pending_ping && builder.push_ping() { self.state.pending_ping = false; @@ -449,7 +449,7 @@ impl SessionFsm { self.clear_streams(); } - fn push_next_pending_stream_close( + fn push_next_pending_stream_reset( &mut self, builder: &mut SessionRecordBuilder, outbound: &mut TrackedRecord, @@ -463,15 +463,15 @@ impl SessionFsm { for offset in 0..len { let index = (start + offset) % len; let stream = self.state.streams.get_index_mut(index).unwrap().1; - let Some(close) = stream.pending_close.as_ref() else { + let Some(reset) = stream.pending_reset.as_ref() else { continue; }; - if !builder.push_stream_close(close) { + if !builder.push_stream_reset(reset) { break; } - outbound.frames.push(TrackedFrame::StreamClose( - stream.pending_close.take().unwrap(), + outbound.frames.push(TrackedFrame::StreamReset( + stream.pending_reset.take().unwrap(), )); } } @@ -673,7 +673,7 @@ impl SessionFsm { match stream.inbound_state { InboundState::Open => {} - InboundState::Discarding | InboundState::Closed(_) => return Ok(()), + InboundState::Discarding | InboundState::Reset(_) => return Ok(()), InboundState::Finished => { // finished stream should always have a final offset let Some(final_offset) = stream.rx.final_offset() else { @@ -740,9 +740,9 @@ impl SessionFsm { } } - fn handle_stream_close( + fn handle_stream_reset( &mut self, - frame: &StreamClose, + frame: &StreamReset, sink: &mut impl EventSink, ) -> Result<(), ()> { let stream_id = frame.stream_id; @@ -757,26 +757,26 @@ impl SessionFsm { if Self::target_affects_inbound(stream.role, frame.target) && !matches!( stream.inbound_state, - InboundState::Closed(_) | InboundState::Discarding + InboundState::Reset(_) | InboundState::Discarding ) { - stream.inbound_state = InboundState::Closed(frame.clone()); + stream.inbound_state = InboundState::Reset(frame.clone()); stream.reset_recv(); - sink.emit(SessionEvent::Closed(frame.clone())); + sink.emit(SessionEvent::Reset(frame.clone())); } if Self::target_affects_outbound(stream.role, frame.target) && !matches!(stream.outbound_state, OutboundState::Closed) { stream.outbound_state = OutboundState::Closed; stream.tx.clear(); - stream.pending_close = None; - sink.emit(SessionEvent::WritableClosed(frame.clone())); + stream.pending_reset = None; + sink.emit(SessionEvent::WritableReset(frame.clone())); } self.try_reap_stream(frame.stream_id); Ok(()) } - fn apply_local_close_to_stream(stream: &mut StreamState, target: CloseTarget) { + fn apply_local_reset_to_stream(stream: &mut StreamState, target: ResetTarget) { if Self::target_affects_inbound(stream.role, target) { stream.inbound_state = InboundState::Discarding; stream.reset_recv(); @@ -787,12 +787,12 @@ impl SessionFsm { } } - fn target_affects_inbound(role: StreamRole, target: CloseTarget) -> bool { - matches!(target, CloseTarget::Both) || role.inbound_target() == target + fn target_affects_inbound(role: StreamRole, target: ResetTarget) -> bool { + matches!(target, ResetTarget::Both) || role.inbound_target() == target } - fn target_affects_outbound(role: StreamRole, target: CloseTarget) -> bool { - matches!(target, CloseTarget::Both) || role.outbound_target() == target + fn target_affects_outbound(role: StreamRole, target: ResetTarget) -> bool { + matches!(target, ResetTarget::Both) || role.outbound_target() == target } fn stream_is_reapable(&self, stream_id: StreamId, stream: &StreamState) -> bool { @@ -800,7 +800,7 @@ impl SessionFsm { record.window_updates.iter().any(|(id, _)| *id == stream_id) || record.frames.iter().any(|frame| match frame { TrackedFrame::StreamData(frame) => frame.stream_id == stream_id, - TrackedFrame::StreamClose(frame) => frame.stream_id == stream_id, + TrackedFrame::StreamReset(frame) => frame.stream_id == stream_id, }) }); if tracked_refs_stream { @@ -808,7 +808,7 @@ impl SessionFsm { } if !stream.tx.is_empty() - || stream.pending_close.is_some() + || stream.pending_reset.is_some() || stream.pending_window || stream.readable_bytes() > 0 || stream.rx.buffered_end_offset() > stream.rx.start_offset() @@ -818,7 +818,7 @@ impl SessionFsm { matches!( stream.inbound_state, - InboundState::Finished | InboundState::Closed(_) | InboundState::Discarding + InboundState::Finished | InboundState::Reset(_) | InboundState::Discarding ) && matches!( stream.outbound_state, OutboundState::Finished | OutboundState::Closed @@ -975,14 +975,14 @@ fn restore_tracked_record( fn requeue_tracked_frame(streams: &mut IndexMap, frame: TrackedFrame) { match frame { - TrackedFrame::StreamClose(close) => restore_stream_close(streams, close), + TrackedFrame::StreamReset(reset) => restore_stream_reset(streams, reset), TrackedFrame::StreamData(frame) => restore_stream_data(streams, frame), } } -fn restore_stream_close(streams: &mut IndexMap, close: StreamClose) { - if let Some(stream) = streams.get_mut(&close.stream_id) { - stream.pending_close = Some(close); +fn restore_stream_reset(streams: &mut IndexMap, reset: StreamReset) { + if let Some(stream) = streams.get_mut(&reset.stream_id) { + stream.pending_reset = Some(reset); } } @@ -1009,7 +1009,7 @@ fn acknowledge_tracked_frame( sink: &mut impl EventSink, ) { match frame { - TrackedFrame::StreamClose(_) => {} + TrackedFrame::StreamReset(_) => {} TrackedFrame::StreamData(frame) => { let stream_id = frame.stream_id; if let Some(stream) = streams.get_mut(&stream_id) { diff --git a/ql-fsm/src/session/state.rs b/ql-fsm/src/session/state.rs index 01fac93a..310d7f0f 100644 --- a/ql-fsm/src/session/state.rs +++ b/ql-fsm/src/session/state.rs @@ -1,7 +1,7 @@ use std::time::Instant; use indexmap::IndexMap; -use ql_wire::{CloseTarget, RecordSeq, SessionClose, StreamClose, StreamHeader, StreamId}; +use ql_wire::{RecordSeq, ResetTarget, SessionClose, StreamHeader, StreamId, StreamReset}; use super::{ ack_tracker::AckTracker, remote_stream_history::RemoteStreamHistory, stream_rx::StreamRx, @@ -48,7 +48,7 @@ pub struct StreamState { pub header: Option, pub rx: StreamRx, pub tx: StreamTx, - pub pending_close: Option, + pub pending_reset: Option, pub peer_max_offset: u64, pub outbound_state: OutboundState, pub inbound_state: InboundState, @@ -68,7 +68,7 @@ impl StreamState { role, header: route_id, tx: StreamTx::new(), - pending_close: None, + pending_reset: None, peer_max_offset: u64::from(initial_peer_stream_receive_window), outbound_state: OutboundState::Open, inbound_state: InboundState::Open, @@ -108,17 +108,17 @@ pub enum StreamRole { } impl StreamRole { - pub fn outbound_target(self) -> CloseTarget { + pub fn outbound_target(self) -> ResetTarget { match self { - Self::Initiator => CloseTarget::Origin, - Self::Responder => CloseTarget::Return, + Self::Initiator => ResetTarget::Origin, + Self::Responder => ResetTarget::Return, } } - pub fn inbound_target(self) -> CloseTarget { + pub fn inbound_target(self) -> ResetTarget { match self { - Self::Initiator => CloseTarget::Return, - Self::Responder => CloseTarget::Origin, + Self::Initiator => ResetTarget::Return, + Self::Responder => ResetTarget::Origin, } } } @@ -135,6 +135,6 @@ pub enum OutboundState { pub enum InboundState { Open, Finished, - Closed(StreamClose), + Reset(StreamReset), Discarding, } diff --git a/ql-fsm/src/session/stream_ops.rs b/ql-fsm/src/session/stream_ops.rs index 9af8d439..4b3c36ad 100644 --- a/ql-fsm/src/session/stream_ops.rs +++ b/ql-fsm/src/session/stream_ops.rs @@ -1,4 +1,4 @@ -use ql_wire::{CloseTarget, StreamClose, StreamCloseCode, StreamHeader, StreamId}; +use ql_wire::{ResetCode, ResetTarget, StreamHeader, StreamId, StreamReset}; use super::{ state::{InboundState, StreamState}, @@ -86,12 +86,12 @@ impl<'a, E: EventSink> StreamOps<'a, E> { Some(StreamWriter::new(stream, send_buffer_size)) } - /// closes the origin lane, return lane, or both lanes of the stream - pub fn close(&mut self, target: CloseTarget, code: StreamCloseCode) { + /// resets the origin lane, return lane, or both lanes of the stream + pub fn reset(&mut self, target: ResetTarget, code: ResetCode) { let stream_id = self.stream_id; let stream = self.stream_mut(); - SessionFsm::apply_local_close_to_stream(stream, target); - stream.pending_close = Some(StreamClose { + SessionFsm::apply_local_reset_to_stream(stream, target); + stream.pending_reset = Some(StreamReset { stream_id, target, code, diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index 628efba3..7a9dc06c 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -2,9 +2,9 @@ use std::time::{Duration, Instant}; use bytes::Bytes; use ql_wire::{ - decode_session_frames, parse_session_frames, CloseTarget, RecordAck, RecordSeq, RouteId, - ServiceId, SessionFrame, SessionRecordBuilder, StreamClose, StreamCloseCode, StreamData, - StreamHeader, StreamId, VarInt, QID, + decode_session_frames, parse_session_frames, RecordAck, RecordSeq, ResetCode, ResetTarget, + RouteId, ServiceId, SessionFrame, SessionRecordBuilder, StreamData, StreamHeader, StreamId, + StreamReset, VarInt, QID, }; use super::{SessionConfig, SessionEvent, SessionFsm}; @@ -30,8 +30,8 @@ fn record_ack(seq: RecordSeq) -> RecordAck { RecordAck::from_ranges([seq..=seq]).unwrap() } -const REFUSED: StreamCloseCode = StreamCloseCode(1); -const TIMEOUT: StreamCloseCode = StreamCloseCode(2); +const REFUSED: ResetCode = ResetCode(1); +const TIMEOUT: ResetCode = ResetCode(2); fn header(value: u64) -> StreamHeader { StreamHeader { @@ -414,21 +414,21 @@ fn inbound_empty_fin_emits_finished_immediately() { } #[test] -fn remote_stream_close_is_reliable_and_retried() { +fn remote_stream_reset_is_reliable_and_retried() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); let stream_id = open_stream_id(&mut fsm); fsm.stream(stream_id, |_| {}) .unwrap() - .close(CloseTarget::Both, StreamCloseCode::CANCELLED); + .reset(ResetTarget::Both, ResetCode::CANCELLED); let (write_id, builder) = fsm.take_next_write(now).unwrap(); - fsm.complete_write(now, write_id.expect("stream close should be tracked"), true); + fsm.complete_write(now, write_id.expect("stream reset should be tracked"), true); let first = decode_session_frames(builder.bytes()).unwrap(); assert!(matches!( first.as_slice(), - [SessionFrame::StreamClose(StreamClose { stream_id: id, .. })] if *id == stream_id + [SessionFrame::StreamReset(StreamReset { stream_id: id, .. })] if *id == stream_id )); let mut emit = |_| {}; @@ -488,22 +488,22 @@ fn duplicate_stream_data_is_not_redelivered() { } #[test] -fn duplicate_remote_close_after_reap_is_ignored() { +fn duplicate_remote_reset_after_reap_is_ignored() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let close = StreamClose { + let reset = StreamReset { stream_id: stream_id(1), - target: CloseTarget::Both, - code: StreamCloseCode(9), + target: ResetTarget::Both, + code: ResetCode(9), }; - let record = vec![SessionFrame::StreamClose(close.clone())]; + let record = vec![SessionFrame::StreamReset(reset.clone())]; let first = receive_events(&mut fsm, now, seq(1), &record); assert_eq!( first, vec![ - SessionEvent::Closed(close.clone()), - SessionEvent::WritableClosed(close), + SessionEvent::Reset(reset.clone()), + SessionEvent::WritableReset(reset), ] ); @@ -512,14 +512,14 @@ fn duplicate_remote_close_after_reap_is_ignored() { } #[test] -fn late_remote_stream_data_after_close_is_ignored() { +fn late_remote_stream_data_after_reset_is_ignored() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); let stream_id = stream_id(1); - let close = vec![SessionFrame::StreamClose(StreamClose { + let reset = vec![SessionFrame::StreamReset(StreamReset { stream_id, - target: CloseTarget::Both, - code: StreamCloseCode(9), + target: ResetTarget::Both, + code: ResetCode(9), })]; let data = vec![SessionFrame::StreamData(StreamData { stream_id, @@ -529,19 +529,19 @@ fn late_remote_stream_data_after_close_is_ignored() { bytes: b"hello".to_vec(), })]; - let first = receive_events(&mut fsm, now, seq(1), &close); + let first = receive_events(&mut fsm, now, seq(1), &reset); assert_eq!( first, vec![ - SessionEvent::Closed(StreamClose { + SessionEvent::Reset(StreamReset { stream_id, - target: CloseTarget::Both, - code: StreamCloseCode(9), + target: ResetTarget::Both, + code: ResetCode(9), }), - SessionEvent::WritableClosed(StreamClose { + SessionEvent::WritableReset(StreamReset { stream_id, - target: CloseTarget::Both, - code: StreamCloseCode(9), + target: ResetTarget::Both, + code: ResetCode(9), }), ] ); @@ -612,64 +612,64 @@ fn duplicate_finished_remote_data_before_read_is_ignored() { fn out_of_order_remote_stream_first_observations_still_open_once_each() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let close3 = vec![SessionFrame::StreamClose(StreamClose { + let reset3 = vec![SessionFrame::StreamReset(StreamReset { stream_id: stream_id(3), - target: CloseTarget::Both, + target: ResetTarget::Both, code: REFUSED, })]; - let close1 = vec![SessionFrame::StreamClose(StreamClose { + let reset1 = vec![SessionFrame::StreamReset(StreamReset { stream_id: stream_id(1), - target: CloseTarget::Both, + target: ResetTarget::Both, code: TIMEOUT, })]; - let first = receive_events(&mut fsm, now, seq(1), &close3); + let first = receive_events(&mut fsm, now, seq(1), &reset3); assert_eq!( first, vec![ - SessionEvent::Closed(StreamClose { + SessionEvent::Reset(StreamReset { stream_id: stream_id(3), - target: CloseTarget::Both, + target: ResetTarget::Both, code: REFUSED, }), - SessionEvent::WritableClosed(StreamClose { + SessionEvent::WritableReset(StreamReset { stream_id: stream_id(3), - target: CloseTarget::Both, + target: ResetTarget::Both, code: REFUSED, }), ] ); - let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &close1); + let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &reset1); assert_eq!( second, vec![ - SessionEvent::Closed(StreamClose { + SessionEvent::Reset(StreamReset { stream_id: stream_id(1), - target: CloseTarget::Both, + target: ResetTarget::Both, code: TIMEOUT, }), - SessionEvent::WritableClosed(StreamClose { + SessionEvent::WritableReset(StreamReset { stream_id: stream_id(1), - target: CloseTarget::Both, + target: ResetTarget::Both, code: TIMEOUT, }), ] ); - let third = receive_events(&mut fsm, now + Duration::from_millis(2), seq(3), &close3); + let third = receive_events(&mut fsm, now + Duration::from_millis(2), seq(3), &reset3); assert!(third.is_empty()); } #[test] -fn invalid_remote_stream_close_closes_session() { +fn invalid_remote_stream_reset_closes_session() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let invalid = vec![SessionFrame::StreamClose(StreamClose { + let invalid = vec![SessionFrame::StreamReset(StreamReset { stream_id: stream_id(0), - target: CloseTarget::Both, - code: StreamCloseCode(9), + target: ResetTarget::Both, + code: ResetCode(9), })]; let events = receive_events(&mut fsm, now, seq(1), &invalid); diff --git a/ql-fsm/src/session/tracked.rs b/ql-fsm/src/session/tracked.rs index 84317951..72439875 100644 --- a/ql-fsm/src/session/tracked.rs +++ b/ql-fsm/src/session/tracked.rs @@ -2,7 +2,7 @@ use std::time::Instant; -use ql_wire::{RecordAck, RecordSeq, StreamClose, StreamId}; +use ql_wire::{RecordAck, RecordSeq, StreamId, StreamReset}; #[derive(Debug, Clone)] pub struct TrackedRecord { @@ -17,7 +17,7 @@ pub struct TrackedRecord { #[derive(Debug, Clone)] pub enum TrackedFrame { StreamData(TrackedStreamData), - StreamClose(StreamClose), + StreamReset(StreamReset), } #[derive(Debug, Clone, Copy, PartialEq, Eq)] diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs index 50547325..ab68fae8 100644 --- a/ql-fsm/src/tests/proptest.rs +++ b/ql-fsm/src/tests/proptest.rs @@ -7,7 +7,7 @@ extern crate proptest as proptest_crate; use bytes::Bytes; use proptest_crate::{collection::vec, prelude::*, test_runner::TestCaseResult}; -use ql_wire::{CloseTarget, RouteId, ServiceId, StreamCloseCode, StreamId, WireError}; +use ql_wire::{ResetCode, ResetTarget, RouteId, ServiceId, StreamId, WireError}; use super::*; use crate::{state::LinkState, Event, OpenStreamParams, PeerStatus, ReceiveError, WriteId}; @@ -59,7 +59,7 @@ enum Action { side: Side, slot: usize, }, - Close { + Reset { side: Side, slot: usize, }, @@ -98,8 +98,8 @@ impl Action { Self::Finish { side, slot } } - fn close(side: Side, slot: usize) -> Self { - Self::Close { side, slot } + fn reset(side: Side, slot: usize) -> Self { + Self::Reset { side, slot } } } @@ -114,8 +114,8 @@ struct SideEventState { opened: BTreeSet, finished: BTreeSet, outbound_finished: BTreeSet, - writable_closed: BTreeSet, - closed: BTreeSet, + writable_reset: BTreeSet, + reset: BTreeSet, peer_statuses: Vec, last_peer_status: Option, session_epoch: usize, @@ -143,7 +143,7 @@ struct Runner { expected: [BTreeMap>; 2], received: [BTreeMap>; 2], finished_by: [BTreeSet; 2], - closed_by: [BTreeSet; 2], + reset_by: [BTreeSet; 2], } impl Runner { @@ -167,7 +167,7 @@ impl Runner { expected: [BTreeMap::new(), BTreeMap::new()], received: [BTreeMap::new(), BTreeMap::new()], finished_by: [BTreeSet::new(), BTreeSet::new()], - closed_by: [BTreeSet::new(), BTreeSet::new()], + reset_by: [BTreeSet::new(), BTreeSet::new()], } } @@ -199,7 +199,7 @@ impl Runner { expected: [BTreeMap::new(), BTreeMap::new()], received: [BTreeMap::new(), BTreeMap::new()], finished_by: [BTreeSet::new(), BTreeSet::new()], - closed_by: [BTreeSet::new(), BTreeSet::new()], + reset_by: [BTreeSet::new(), BTreeSet::new()], } } @@ -329,19 +329,19 @@ impl Runner { } } } - Action::Close { side, slot } => { + Action::Reset { side, slot } => { if let Some(stream_id) = self.slots[side.idx()][*slot] { - let closed = self + let reset = self .harness .node_mut(*side) .fsm .stream(stream_id) .is_ok_and(|mut stream| { - stream.close(CloseTarget::Both, StreamCloseCode::CANCELLED); + stream.reset(ResetTarget::Both, ResetCode::CANCELLED); true }); - if closed { - self.closed_by[side.idx()].insert(stream_id); + if reset { + self.reset_by[side.idx()].insert(stream_id); self.slots[side.idx()][*slot] = None; } } @@ -449,8 +449,8 @@ impl Runner { "side {side:?} emitted duplicate Finished for {stream_id:?}" ); prop_assert!( - !self.events[side.idx()].closed.contains(&stream_id), - "side {side:?} emitted Finished after Closed for {stream_id:?}" + !self.events[side.idx()].reset.contains(&stream_id), + "side {side:?} emitted Finished after Reset for {stream_id:?}" ); } Event::OutboundFinished(stream_id) => { @@ -463,27 +463,27 @@ impl Runner { "side {side:?} emitted duplicate OutboundFinished for {stream_id:?}" ); } - Event::Closed(frame) => { + Event::Reset(frame) => { prop_assert!( self.known_streams.contains(&frame.stream_id), - "side {side:?} emitted Closed for unknown stream {:?}", + "side {side:?} emitted Reset for unknown stream {:?}", frame.stream_id ); prop_assert!( - self.events[side.idx()].closed.insert(frame.stream_id), - "side {side:?} emitted duplicate Closed for {:?}", + self.events[side.idx()].reset.insert(frame.stream_id), + "side {side:?} emitted duplicate Reset for {:?}", frame.stream_id ); } - Event::WritableClosed(frame) => { + Event::WritableReset(frame) => { let stream_id = frame.stream_id; prop_assert!( self.known_streams.contains(&stream_id), - "side {side:?} emitted WritableClosed for unknown stream {stream_id:?}" + "side {side:?} emitted WritableReset for unknown stream {stream_id:?}" ); prop_assert!( - self.events[side.idx()].writable_closed.insert(stream_id), - "side {side:?} emitted duplicate WritableClosed for {stream_id:?}" + self.events[side.idx()].writable_reset.insert(stream_id), + "side {side:?} emitted duplicate WritableReset for {stream_id:?}" ); } Event::SessionClosed(_) => { @@ -590,18 +590,18 @@ impl Runner { for stream_id in &self.finished_by[side.idx()] { prop_assert!( self.events[opposite(side).idx()].finished.contains(stream_id) - || self.events[opposite(side).idx()].closed.contains(stream_id) + || self.events[opposite(side).idx()].reset.contains(stream_id) || !connected[opposite(side).idx()], - "side {side:?} finished {stream_id:?} but side {:?} saw neither Finished nor Closed", + "side {side:?} finished {stream_id:?} but side {:?} saw neither Finished nor Reset", opposite(side) ); } - for stream_id in &self.closed_by[side.idx()] { + for stream_id in &self.reset_by[side.idx()] { prop_assert!( - self.events[opposite(side).idx()].closed.contains(stream_id) + self.events[opposite(side).idx()].reset.contains(stream_id) || !connected[opposite(side).idx()], - "side {side:?} closed {stream_id:?} but side {:?} saw no Closed event", + "side {side:?} reset {stream_id:?} but side {:?} saw no Reset event", opposite(side) ); } @@ -634,8 +634,8 @@ impl Runner { events.opened.is_empty() && events.finished.is_empty() && events.outbound_finished.is_empty() - && events.closed.is_empty() - && events.writable_closed.is_empty() + && events.reset.is_empty() + && events.writable_reset.is_empty() }), "handshake-only property observed stream activity" ); @@ -706,8 +706,8 @@ impl Runner { } fn inbound_aborted(&self, side: Side, stream_id: StreamId) -> bool { - self.events[side.idx()].closed.contains(&stream_id) - || self.closed_by[side.idx()].contains(&stream_id) + self.events[side.idx()].reset.contains(&stream_id) + || self.reset_by[side.idx()].contains(&stream_id) } } @@ -870,7 +870,7 @@ fn connected_action_strategy() -> impl Strategy { side_usize_action(slot.clone(), Action::open_stream), side_usize_vec_action(slot.clone(), bytes, Action::write), side_usize_action(slot.clone(), Action::finish), - side_usize_action(slot, Action::close), + side_usize_action(slot, Action::reset), ] } @@ -915,7 +915,7 @@ fn terminal_action_strategy() -> impl Strategy { side_usize_action(slot.clone(), Action::open_stream), side_usize_vec_action(slot.clone(), bytes, Action::write), side_usize_action(slot.clone(), Action::finish), - side_usize_action(slot, Action::close), + side_usize_action(slot, Action::reset), side_action(Action::TakeNext), side_usize_action(queue_index.clone(), Action::confirm_taken), side_usize_action(queue_index.clone(), Action::reject_taken), diff --git a/ql-fsm/src/tests/session.rs b/ql-fsm/src/tests/session.rs index cdedd1f4..73aabe57 100644 --- a/ql-fsm/src/tests/session.rs +++ b/ql-fsm/src/tests/session.rs @@ -218,10 +218,7 @@ fn disconnected_stream_operations_fail_with_no_session() { ); assert_eq!( harness.a.fsm.stream(missing).map(|mut stream| { - stream.close( - ql_wire::CloseTarget::Both, - ql_wire::StreamCloseCode::CANCELLED, - ); + stream.reset(ql_wire::ResetTarget::Both, ql_wire::ResetCode::CANCELLED); }), Err(StreamError::NoSession) ); diff --git a/ql-rpc/src/error.rs b/ql-rpc/src/error.rs index d6c9cd35..8c47362d 100644 --- a/ql-rpc/src/error.rs +++ b/ql-rpc/src/error.rs @@ -1,4 +1,4 @@ -use crate::StreamCloseCode; +use crate::ResetCode; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum Error { @@ -24,13 +24,13 @@ impl std::fmt::Display for Error { impl std::error::Error for Error {} impl Error { - pub const fn close_code(self) -> StreamCloseCode { + pub const fn reset_code(self) -> ResetCode { match self { - Self::LengthOverflow => StreamCloseCode::LIMIT, + Self::LengthOverflow => ResetCode::LIMIT, Self::Truncated | Self::UnexpectedFrameKind(_) | Self::MissingResponse - | Self::TrailingBytes => StreamCloseCode::PROTOCOL, + | Self::TrailingBytes => ResetCode::PROTOCOL, } } } @@ -43,10 +43,10 @@ pub enum RpcError { } impl RpcError { - pub const fn close_code(&self) -> Option { + pub const fn reset_code(&self) -> Option { match self { - Self::Protocol(error) => Some(error.close_code()), - Self::Codec(_) => Some(StreamCloseCode::CODEC), + Self::Protocol(error) => Some(error.reset_code()), + Self::Codec(_) => Some(ResetCode::CODEC), Self::Transport(_) => None, } } diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index 94c24e01..90941bc7 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -14,7 +14,7 @@ pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; use framed_value::*; -pub use ql_common::{RouteId, ServiceId, StreamCloseCode, StreamCloseOrigin}; +pub use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId}; pub use router::*; pub use rpc::*; pub use stream::*; diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index 1ab7de3c..8bf5d541 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -1,4 +1,4 @@ -use crate::{RouteId, ServiceId, StreamCloseCode}; +use crate::{ResetCode, RouteId, ServiceId}; mod builder; mod config; @@ -9,7 +9,6 @@ pub use self::{ config::RouterConfig, mode::*, }; -use crate::{close_stream, RpcStream}; pub use crate::{ download::{DownloadHandler, DownloadHandlerLocal, DownloadStart, DownloadWriter}, duplex::{DuplexHandler, DuplexHandlerLocal, DuplexPeer}, @@ -19,6 +18,7 @@ pub use crate::{ subscription::{SubscriptionHandler, SubscriptionHandlerLocal, SubscriptionResponder}, upload::{UploadHandler, UploadHandlerLocal, UploadReader, UploadResponder}, }; +use crate::{reset_stream, RpcStream}; pub struct Router where @@ -90,7 +90,7 @@ where route_id, }; let Ok(index) = self.routes.binary_search_by_key(&key, |entry| entry.key) else { - close_stream(stream, StreamCloseCode::UNKNOWN_ROUTE); + reset_stream(stream, ResetCode::UNKNOWN_ROUTE); return None; }; let route = self.routes[index].route; diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index 5cdc41ed..b586ca15 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -5,7 +5,7 @@ use bytes::{BufMut, Bytes}; use crate::{ download::{Download, PartReadStep}, rpc::parts::FrameKind, - DropCloseRead, FramedPrefixStep, FramedReader, RpcCodec, RpcError, RpcRead, StreamCloseCode, + DropResetRead, FramedPrefixStep, FramedReader, ResetCode, RpcCodec, RpcError, RpcRead, }; pub struct DownloadCall @@ -13,7 +13,7 @@ where M: Download, R: RpcRead, { - stream: DropCloseRead, + stream: DropResetRead, reader: Option>, } @@ -31,7 +31,7 @@ where M: Download, R: RpcRead, { - stream: DropCloseRead, + stream: DropResetRead, reader: crate::download::PartFrameReader, } @@ -42,7 +42,7 @@ where { pub fn new(stream: R) -> Self { Self { - stream: DropCloseRead::new(stream), + stream: DropResetRead::new(stream), reader: Some(FramedReader::default()), } } @@ -76,12 +76,12 @@ where } } - pub fn close(mut self, code: StreamCloseCode) { - self.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); } - fn close_inner(&mut self, code: StreamCloseCode) { - DropCloseRead::close(&mut self.stream, code); + fn reset_inner(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.stream, code); } } @@ -138,8 +138,8 @@ where } } - pub fn close(mut self, code: StreamCloseCode) { - self.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); } async fn read_frame( @@ -162,8 +162,8 @@ where } } - fn close_inner(&mut self, code: StreamCloseCode) { - DropCloseRead::close(&mut self.stream, code); + fn reset_inner(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.stream, code); } } @@ -193,8 +193,8 @@ where } } - pub fn close(mut self, code: StreamCloseCode) { - self.parent.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.parent.reset_inner(code); self.finished = true; } } @@ -206,7 +206,7 @@ where { fn drop(&mut self) { if !self.finished { - self.parent.close_inner(StreamCloseCode::DROPPED); + self.parent.reset_inner(ResetCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs index 142d7e51..95e22822 100644 --- a/ql-rpc/src/rpc/download/server.rs +++ b/ql-rpc/src/rpc/download/server.rs @@ -10,8 +10,7 @@ use crate::{ parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, read_eof_request, }, - write_bytes, DropCloseWrite, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, - StreamCloseCode, + write_bytes, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(DownloadHandler: Send)] @@ -30,7 +29,7 @@ where M: DownloadRpc, W: RpcWrite, { - writer: DropCloseWrite, + writer: DropResetWrite, marker: PhantomData M>, } @@ -39,7 +38,7 @@ where M: DownloadRpc, W: RpcWrite, { - writer: DropCloseWrite, + writer: DropResetWrite, marker: PhantomData M>, } @@ -59,7 +58,7 @@ where { pub(crate) fn new(writer: W) -> Self { Self { - writer: DropCloseWrite::new(writer), + writer: DropResetWrite::new(writer), marker: PhantomData, } } @@ -89,9 +88,9 @@ where finish_bytes(&mut writer).await } - /// close the stream with a transport code - pub fn close(mut self, code: StreamCloseCode) { - DropCloseWrite::close(&mut self.writer, code); + /// reset the stream with a transport code + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); } } @@ -122,8 +121,8 @@ where finish_bytes(&mut writer).await } - pub fn close(mut self, code: StreamCloseCode) { - DropCloseWrite::close(&mut self.writer, code); + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); } } @@ -156,7 +155,7 @@ where { fn drop(&mut self) { if !self.finished { - DropCloseWrite::close(&mut self.parent.writer, StreamCloseCode::DROPPED); + DropResetWrite::reset(&mut self.parent.writer, ResetCode::DROPPED); } } } @@ -178,11 +177,11 @@ pub(crate) async fn handle_download_inner( let request = match read_eof_request::(&mut reader, config).await { Ok(request) => request, Err(error) => { - let code = error.close_code(); + let code = error.reset_code(); handle_error(&state, &error); if let Some(code) = code { - reader.close(code); - writer.close(code); + reader.reset(code); + writer.reset(code); } return; } diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index c5326431..7b6bb5f0 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -8,8 +8,8 @@ use bytes::Bytes; use crate::{ duplex::{codec, Duplex, EventReader, ReadStep}, - finish_bytes, write_bytes, DropCloseRead, DropCloseWrite, RpcCodec, RpcError, RpcRead, - RpcWrite, StreamCloseCode, + finish_bytes, write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcCodec, RpcError, + RpcRead, RpcWrite, }; pub struct DuplexCall @@ -27,7 +27,7 @@ where T: RpcCodec, W: RpcWrite, { - writer: DropCloseWrite, + writer: DropResetWrite, marker: PhantomData T>, } @@ -36,7 +36,7 @@ where T: RpcCodec, R: RpcRead, { - stream: DropCloseRead, + stream: DropResetRead, reader: EventReader, } @@ -47,7 +47,7 @@ where { pub fn new(writer: W) -> Self { Self { - writer: DropCloseWrite::new(writer), + writer: DropResetWrite::new(writer), marker: PhantomData, } } @@ -63,8 +63,8 @@ where finish_bytes(&mut self.writer).await } - pub fn close(mut self, code: StreamCloseCode) { - DropCloseWrite::close(&mut self.writer, code); + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); } } @@ -75,7 +75,7 @@ where { pub fn new(stream: R) -> Self { Self { - stream: DropCloseRead::new(stream), + stream: DropResetRead::new(stream), reader: EventReader::default(), } } @@ -125,11 +125,11 @@ where } } - pub fn close(mut self, code: StreamCloseCode) { - self.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); } - fn close_inner(&mut self, code: StreamCloseCode) { - DropCloseRead::close(&mut self.stream, code); + fn reset_inner(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.stream, code); } } diff --git a/ql-rpc/src/rpc/notification/server.rs b/ql-rpc/src/rpc/notification/server.rs index 8e0a33de..ae766532 100644 --- a/ql-rpc/src/rpc/notification/server.rs +++ b/ql-rpc/src/rpc/notification/server.rs @@ -1,8 +1,8 @@ use std::future::Future; use crate::{ - notification::Notification as NotificationRpc, rpc::read_eof_request, RouterConfig, RpcError, - RpcRead, RpcStream, RpcWrite, StreamCloseCode, + notification::Notification as NotificationRpc, rpc::read_eof_request, ResetCode, RouterConfig, + RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(NotificationHandler: Send)] @@ -33,16 +33,16 @@ pub(crate) async fn handle_notification_inner( let notification = match read_eof_request::(&mut reader, config).await { Ok(notification) => notification, Err(error) => { - let code = error.close_code(); + let code = error.reset_code(); handle_error(&state, &error); if let Some(code) = code { - reader.close(code); - writer.close(code); + reader.reset(code); + writer.reset(code); } return; } }; - writer.close(StreamCloseCode::CANCELLED); + writer.reset(ResetCode::CANCELLED); handle(state, notification).await; } diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs index 1f134122..755653c8 100644 --- a/ql-rpc/src/rpc/progress/client.rs +++ b/ql-rpc/src/rpc/progress/client.rs @@ -6,7 +6,7 @@ use std::{ use crate::{ progress::{Progress, ReadStep, ResponseReader}, - DropCloseRead, Error, RpcError, RpcRead, StreamCloseCode, + DropResetRead, Error, ResetCode, RpcError, RpcRead, }; pub struct ProgressCall @@ -14,7 +14,7 @@ where M: Progress, R: RpcRead, { - stream: DropCloseRead, + stream: DropResetRead, state: State, } @@ -42,7 +42,7 @@ where { pub fn new(stream: R) -> Self { Self { - stream: DropCloseRead::new(stream), + stream: DropResetRead::new(stream), state: State::Reading(ResponseReader::default()), } } @@ -100,13 +100,13 @@ where self.poll_step(cx) } - pub fn close(mut self, code: StreamCloseCode) { - self.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); } - fn close_inner(&mut self, code: StreamCloseCode) { + fn reset_inner(&mut self, code: ResetCode) { self.state = State::Done; - DropCloseRead::close(&mut self.stream, code); + DropResetRead::reset(&mut self.stream, code); } } diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index 40bbf20f..e396d679 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -6,8 +6,7 @@ use crate::{ finish_bytes, progress::{encode_progress, encode_response, Progress}, rpc::read_framed_request, - write_bytes, DropCloseWrite, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, - StreamCloseCode, + write_bytes, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(ProgressHandler: Send)] @@ -26,7 +25,7 @@ where M: Progress, W: RpcWrite, { - writer: DropCloseWrite, + writer: DropResetWrite, marker: PhantomData M>, } @@ -37,7 +36,7 @@ where { pub(crate) fn new(writer: W) -> Self { Self { - writer: DropCloseWrite::new(writer), + writer: DropResetWrite::new(writer), marker: PhantomData, } } @@ -56,8 +55,8 @@ where finish_bytes(&mut self.writer).await } - pub fn close(mut self, code: StreamCloseCode) { - DropCloseWrite::close(&mut self.writer, code); + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); } } @@ -78,11 +77,11 @@ pub(crate) async fn handle_progress_inner( let request = match read_framed_request::(&mut reader, config).await { Ok(request) => request, Err(error) => { - let code = error.close_code(); + let code = error.reset_code(); handle_error(&state, &error); if let Some(code) = code { - reader.close(code); - writer.close(code); + reader.reset(code); + writer.reset(code); } return; } diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs index 55ab74bb..ba810e7b 100644 --- a/ql-rpc/src/rpc/request/server.rs +++ b/ql-rpc/src/rpc/request/server.rs @@ -4,8 +4,7 @@ use bytes::Bytes; use crate::{ finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, - DropCloseWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, - StreamCloseCode, + DropResetWrite, ResetCode, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(RequestHandler: Send)] @@ -23,7 +22,7 @@ pub struct Response where W: RpcWrite, { - writer: DropCloseWrite, + writer: DropResetWrite, marker: PhantomData T>, } @@ -34,7 +33,7 @@ where { pub(crate) fn new(writer: W) -> Self { Self { - writer: DropCloseWrite::new(writer), + writer: DropResetWrite::new(writer), marker: PhantomData, } } @@ -48,8 +47,8 @@ where Ok(()) } - pub fn close(mut self, code: StreamCloseCode) { - DropCloseWrite::close(&mut self.writer, code); + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); } } @@ -70,11 +69,11 @@ pub(crate) async fn handle_request_inner( let request = match read_eof_request::(&mut reader, config).await { Ok(request) => request, Err(error) => { - let code = error.close_code(); + let code = error.reset_code(); handle_error(&state, &error); if let Some(code) = code { - reader.close(code); - writer.close(code); + reader.reset(code); + writer.reset(code); } return; } diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index 670c1d19..95259406 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -5,7 +5,7 @@ use std::{ use crate::{ subscription::{ReadStep, ResponseReader, Subscription}, - DropCloseRead, RpcError, RpcRead, StreamCloseCode, + DropResetRead, ResetCode, RpcError, RpcRead, }; pub struct SubscriptionCall @@ -13,7 +13,7 @@ where M: Subscription, R: RpcRead, { - stream: DropCloseRead, + stream: DropResetRead, reader: ResponseReader, } @@ -24,7 +24,7 @@ where { pub fn new(stream: R) -> Self { Self { - stream: DropCloseRead::new(stream), + stream: DropResetRead::new(stream), reader: ResponseReader::default(), } } @@ -74,11 +74,11 @@ where } } - pub fn close(mut self, code: StreamCloseCode) { - self.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); } - fn close_inner(&mut self, code: StreamCloseCode) { - DropCloseRead::close(&mut self.stream, code); + fn reset_inner(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.stream, code); } } diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs index 0f687a87..f6dcbd54 100644 --- a/ql-rpc/src/rpc/subscription/server.rs +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -4,8 +4,8 @@ use bytes::Bytes; use crate::{ codec, finish_bytes, rpc::read_eof_request, subscription::Subscription as SubscriptionRpc, - write_bytes, DropCloseWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, - StreamCloseCode, + write_bytes, DropResetWrite, ResetCode, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, + RpcWrite, }; #[trait_variant::make(SubscriptionHandler: Send)] @@ -27,7 +27,7 @@ pub struct SubscriptionResponder where W: RpcWrite, { - writer: DropCloseWrite, + writer: DropResetWrite, marker: PhantomData T>, } @@ -38,7 +38,7 @@ where { pub(crate) fn new(writer: W) -> Self { Self { - writer: DropCloseWrite::new(writer), + writer: DropResetWrite::new(writer), marker: PhantomData, } } @@ -55,8 +55,8 @@ where finish_bytes(&mut self.writer).await } - pub fn close(mut self, code: StreamCloseCode) { - DropCloseWrite::close(&mut self.writer, code); + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); } } @@ -77,11 +77,11 @@ pub(crate) async fn handle_subscription_inner( let request = match read_eof_request::(&mut reader, config).await { Ok(request) => request, Err(error) => { - let code = error.close_code(); + let code = error.reset_code(); handle_error(&state, &error); if let Some(code) = code { - reader.close(code); - writer.close(code); + reader.reset(code); + writer.reset(code); } return; } diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index 1fca05f4..e16833b3 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -4,8 +4,8 @@ use crate::{ finish_bytes, read_bytes, rpc::parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, upload::Upload, - write_bytes, ChunkQueue, DropCloseRead, DropCloseWrite, RpcCodec, RpcError, RpcRead, RpcWrite, - StreamCloseCode, + write_bytes, ChunkQueue, DropResetRead, DropResetWrite, ResetCode, RpcCodec, RpcError, RpcRead, + RpcWrite, }; pub struct UploadCall @@ -14,8 +14,8 @@ where W: RpcWrite, R: RpcRead, { - writer: DropCloseWrite, - reader: DropCloseRead, + writer: DropResetWrite, + reader: DropResetRead, marker: std::marker::PhantomData M>, } @@ -37,8 +37,8 @@ where { pub fn new(writer: W, reader: R) -> Self { Self { - writer: DropCloseWrite::new(writer), - reader: DropCloseRead::new(reader), + writer: DropResetWrite::new(writer), + reader: DropResetRead::new(reader), marker: std::marker::PhantomData, } } @@ -80,9 +80,9 @@ where Ok(value) } - fn close(&mut self, code: StreamCloseCode) { - DropCloseRead::close(&mut self.reader, code); - DropCloseWrite::close(&mut self.writer, code); + fn reset(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.reader, code); + DropResetWrite::reset(&mut self.writer, code); } } @@ -117,7 +117,7 @@ where { fn drop(&mut self) { if !self.finished { - self.parent.close(StreamCloseCode::DROPPED); + self.parent.reset(ResetCode::DROPPED); } } } diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index 2f80987e..82245494 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -8,7 +8,7 @@ use crate::{ parts::{FrameKind, PartFrameReader, PartReadStep}, read_framed_request_prefix, }, - DropCloseRead, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, StreamCloseCode, Upload, + DropResetRead, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, Upload, }; #[trait_variant::make(UploadHandler: Send)] @@ -32,7 +32,7 @@ where M: Upload, R: RpcRead, { - stream: DropCloseRead, + stream: DropResetRead, reader: PartFrameReader, } @@ -107,12 +107,12 @@ where } } - pub fn close(mut self, code: StreamCloseCode) { - self.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); } - fn close_inner(&mut self, code: StreamCloseCode) { - DropCloseRead::close(&mut self.stream, code); + fn reset_inner(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.stream, code); } } @@ -144,8 +144,8 @@ where } } - pub fn close(mut self, code: StreamCloseCode) { - self.parent.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.parent.reset_inner(code); self.finished = true; } } @@ -157,7 +157,7 @@ where { fn drop(&mut self) { if !self.finished { - self.parent.close_inner(StreamCloseCode::DROPPED); + self.parent.reset_inner(ResetCode::DROPPED); } } } @@ -177,8 +177,8 @@ where self.inner.respond(response).await } - pub fn close(self, code: StreamCloseCode) { - self.inner.close(code); + pub fn reset(self, code: ResetCode) { + self.inner.reset(code); } } @@ -205,11 +205,11 @@ pub(crate) async fn handle_upload_inner( match read_framed_request_prefix::(&mut reader, config).await { Ok(value) => value, Err(error) => { - let code = error.close_code(); + let code = error.reset_code(); handle_error(&state, &error); if let Some(code) = code { - reader.close(code); - writer.close(code); + reader.reset(code); + writer.reset(code); } return; } @@ -219,7 +219,7 @@ pub(crate) async fn handle_upload_inner( state, request, UploadReader { - stream: DropCloseRead::new(reader), + stream: DropResetRead::new(reader), reader: PartFrameReader::new(buffered), }, UploadResponder::new(writer), diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index 07334c22..d6cffecc 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -5,7 +5,7 @@ use std::{ use bytes::Bytes; -use crate::{RouteId, ServiceId, StreamCloseCode}; +use crate::{ResetCode, RouteId, ServiceId}; pub trait RpcStream { type Error; @@ -24,24 +24,24 @@ pub trait RpcRead { fn poll_read(&mut self, cx: &mut Context<'_>) -> Poll, Self::Error>>; /// aborts the read side - fn close(self, code: StreamCloseCode); + fn reset(self, code: ResetCode); } pub trait RpcWrite { type Error; - /// writes outbound bytes before finish or close + /// writes outbound bytes before finish or reset fn poll_write( &mut self, bytes: &mut Bytes, cx: &mut Context<'_>, ) -> Poll>; - /// completes the write side and must be polled until ready without further write or close calls + /// completes the write side and must be polled until ready without further write or reset calls fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll>; /// aborts the write side before finish - fn close(self, code: StreamCloseCode); + fn reset(self, code: ResetCode); } pub async fn read_bytes(reader: &mut R) -> Result, R::Error> @@ -66,24 +66,24 @@ where poll_fn(|cx| writer.poll_finish(cx)).await } -pub fn close_stream(stream: St, code: StreamCloseCode) +pub fn reset_stream(stream: St, code: ResetCode) where St: RpcStream, { let (reader, writer) = stream.split(); - reader.close(code); - writer.close(code); + reader.reset(code); + writer.reset(code); } pub(crate) use drop::*; mod drop { use super::*; - pub struct DropCloseRead { + pub struct DropResetRead { inner: Option, } - impl DropCloseRead { + impl DropResetRead { pub fn new(reader: R) -> Self { Self { inner: Some(reader), @@ -101,14 +101,14 @@ mod drop { } #[inline] - pub fn close(&mut self, code: StreamCloseCode) { + pub fn reset(&mut self, code: ResetCode) { if let Some(reader) = self.inner.take() { - reader.close(code); + reader.reset(code); } } } - impl RpcRead for DropCloseRead { + impl RpcRead for DropResetRead { type Error = R::Error; #[track_caller] @@ -116,22 +116,22 @@ mod drop { self.inner.as_mut().unwrap().poll_read(cx) } - fn close(mut self, code: StreamCloseCode) { - Self::close(&mut self, code); + fn reset(mut self, code: ResetCode) { + Self::reset(&mut self, code); } } - impl Drop for DropCloseRead { + impl Drop for DropResetRead { fn drop(&mut self) { - self.close(StreamCloseCode::DROPPED); + self.reset(ResetCode::DROPPED); } } - pub struct DropCloseWrite { + pub struct DropResetWrite { inner: Option, } - impl DropCloseWrite { + impl DropResetWrite { pub fn new(writer: W) -> Self { Self { inner: Some(writer), @@ -139,14 +139,14 @@ mod drop { } #[inline] - pub fn close(&mut self, code: StreamCloseCode) { + pub fn reset(&mut self, code: ResetCode) { if let Some(writer) = self.inner.take() { - writer.close(code); + writer.reset(code); } } } - impl RpcWrite for DropCloseWrite { + impl RpcWrite for DropResetWrite { type Error = W::Error; #[track_caller] @@ -163,14 +163,14 @@ mod drop { self.inner.as_mut().unwrap().poll_finish(cx) } - fn close(mut self, code: StreamCloseCode) { - Self::close(&mut self, code); + fn reset(mut self, code: ResetCode) { + Self::reset(&mut self, code); } } - impl Drop for DropCloseWrite { + impl Drop for DropResetWrite { fn drop(&mut self) { - self.close(StreamCloseCode::DROPPED) + self.reset(ResetCode::DROPPED) } } } diff --git a/ql-runtime/src/command.rs b/ql-runtime/src/command.rs index 1d4857ca..d90f787b 100644 --- a/ql-runtime/src/command.rs +++ b/ql-runtime/src/command.rs @@ -1,6 +1,6 @@ use ql_fsm::{NoSessionError, PairingInvite}; use ql_wire::{ - CloseTarget, PairingToken, PeerBundle, RouteId, ServiceId, SessionCloseCode, StreamCloseCode, + PairingToken, PeerBundle, ResetCode, ResetTarget, RouteId, ServiceId, SessionCloseCode, StreamId, }; @@ -33,10 +33,10 @@ pub enum Command { code: SessionCloseCode, }, Unpair, - CloseStream { + ResetStream { stream_id: StreamId, - target: CloseTarget, - code: StreamCloseCode, + target: ResetTarget, + code: ResetCode, }, } @@ -53,7 +53,7 @@ impl Command { Self::PollStream { .. } => "PollStream", Self::CloseSession { .. } => "CloseSession", Self::Unpair => "Unpair", - Self::CloseStream { .. } => "CloseStream", + Self::ResetStream { .. } => "ResetStream", } } } diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 8b366ad1..a092d578 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -16,7 +16,7 @@ use std::{ use async_channel::Recv; use futures_lite::future::{poll_fn, yield_now}; use ql_fsm::{Event, QlFsm, WriteId}; -use ql_wire::{CloseTarget, StreamCloseCode, StreamCloseOrigin, StreamHeader, StreamId}; +use ql_wire::{ResetCode, ResetOrigin, ResetTarget, StreamHeader, StreamId}; use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; use crate::{ @@ -237,8 +237,8 @@ impl DriverState { log::info!("open stream allocated: service_id={service_id} route_id={route_id} stream_id={stream_id}"); let (reader, writer, reader_io, writer_io) = io::new_stream( stream_id, - CloseTarget::Return, - CloseTarget::Origin, + ResetTarget::Return, + ResetTarget::Origin, RuntimeHandle::new(runtime_tx), ); self.streams.insert( @@ -255,7 +255,7 @@ impl DriverState { stream.inbound_close(); stream.outbound_close(); } - stream_ops.close(CloseTarget::Both, StreamCloseCode::DROPPED); + stream_ops.reset(ResetTarget::Both, ResetCode::DROPPED); drop(stream_ops); return; } @@ -270,26 +270,26 @@ impl DriverState { log::trace!("poll stream requested: stream_id={stream_id}"); self.poll_stream(fsm, stream_id); } - Command::CloseStream { + Command::ResetStream { stream_id, target, code, } => { log::debug!( - "close stream command: stream_id={stream_id} target={target:?} code={code:?}" + "reset stream command: stream_id={stream_id} target={target:?} code={code:?}" ); if let Entry::Occupied(mut entry) = self.streams.entry(stream_id) { let stream = entry.get_mut(); - if target == CloseTarget::Both || target == stream.inbound_target() { + if target == ResetTarget::Both || target == stream.inbound_target() { stream.inbound_close(); } - if target == CloseTarget::Both || target == stream.outbound_target() { + if target == ResetTarget::Both || target == stream.outbound_target() { stream.outbound_close(); } Self::try_reap_stream(entry); } if let Ok(mut stream) = fsm.stream(stream_id) { - stream.close(target, code); + stream.reset(target, code); } } } @@ -341,11 +341,11 @@ impl DriverState { log::info!("outbound finish acknowledged: stream_id={stream_id}"); self.handle_outbound_finished(stream_id); } - Event::Closed(frame) => { - self.handle_closed_stream(&frame); + Event::Reset(frame) => { + self.handle_reset_stream(&frame); } - Event::WritableClosed(frame) => { - self.handle_writable_closed(&frame); + Event::WritableReset(frame) => { + self.handle_writable_reset(&frame); } Event::SessionClosed(close) => { log::info!("session closed: frame={close:?}"); @@ -368,15 +368,15 @@ impl DriverState { "dropping inbound stream because handle channel is unavailable: stream_id={stream_id}" ); if let Ok(mut stream) = fsm.stream(stream_id) { - stream.close(CloseTarget::Both, StreamCloseCode::DISCONNECTED); + stream.reset(ResetTarget::Both, ResetCode::DISCONNECTED); } return; }; let (reader, writer, reader_io, writer_io) = io::new_stream( stream_id, - CloseTarget::Origin, - CloseTarget::Return, + ResetTarget::Origin, + ResetTarget::Return, RuntimeHandle::new(runtime_tx), ); @@ -456,7 +456,7 @@ impl DriverState { stream_ops.commit_read(accepted).unwrap(); } if peer_closed { - stream_ops.close(target, StreamCloseCode::DROPPED); + stream_ops.reset(target, ResetCode::DROPPED); if let Entry::Occupied(entry) = self.streams.entry(stream_id) { Self::try_reap_stream(entry); } @@ -475,9 +475,9 @@ impl DriverState { Self::try_reap_stream(entry); } - fn handle_closed_stream(&mut self, frame: &ql_wire::StreamClose) { + fn handle_reset_stream(&mut self, frame: &ql_wire::StreamReset) { log::info!( - "inbound close frame: stream_id={} target={:?} code={}", + "inbound reset frame: stream_id={} target={:?} code={}", frame.stream_id, frame.target, frame.code @@ -487,24 +487,24 @@ impl DriverState { }; let stream = entry.get_mut(); - if frame.target == CloseTarget::Both || frame.target == stream.inbound_target() { - stream.inbound_fail(QlStreamError::StreamClosed { + if frame.target == ResetTarget::Both || frame.target == stream.inbound_target() { + stream.inbound_fail(QlStreamError::StreamReset { code: frame.code, - origin: StreamCloseOrigin::Peer, + origin: ResetOrigin::Peer, }); } - if frame.target == CloseTarget::Both || frame.target == stream.outbound_target() { - stream.outbound_fail(QlStreamError::StreamClosed { + if frame.target == ResetTarget::Both || frame.target == stream.outbound_target() { + stream.outbound_fail(QlStreamError::StreamReset { code: frame.code, - origin: StreamCloseOrigin::Peer, + origin: ResetOrigin::Peer, }); } Self::try_reap_stream(entry); } - fn handle_writable_closed(&mut self, frame: &ql_wire::StreamClose) { + fn handle_writable_reset(&mut self, frame: &ql_wire::StreamReset) { log::info!( - "writable close frame: stream_id={} target={:?} code={}", + "writable reset frame: stream_id={} target={:?} code={}", frame.stream_id, frame.target, frame.code @@ -513,9 +513,9 @@ impl DriverState { return; }; let stream = entry.get_mut(); - stream.outbound_fail(QlStreamError::StreamClosed { + stream.outbound_fail(QlStreamError::StreamReset { code: frame.code, - origin: StreamCloseOrigin::Peer, + origin: ResetOrigin::Peer, }); Self::try_reap_stream(entry); } diff --git a/ql-runtime/src/driver/state.rs b/ql-runtime/src/driver/state.rs index 0ff8eca8..ccd57765 100644 --- a/ql-runtime/src/driver/state.rs +++ b/ql-runtime/src/driver/state.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use bytes::Bytes; -use ql_wire::{CloseTarget, StreamId}; +use ql_wire::{ResetTarget, StreamId}; use crate::{ command::Command, @@ -34,19 +34,19 @@ impl DriverStreamIo { } } - pub fn inbound_target(&self) -> CloseTarget { + pub fn inbound_target(&self) -> ResetTarget { if self.is_initiator { - CloseTarget::Return + ResetTarget::Return } else { - CloseTarget::Origin + ResetTarget::Origin } } - pub fn outbound_target(&self) -> CloseTarget { + pub fn outbound_target(&self) -> ResetTarget { if self.is_initiator { - CloseTarget::Origin + ResetTarget::Origin } else { - CloseTarget::Return + ResetTarget::Return } } diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs index af4ab63a..878682b3 100644 --- a/ql-runtime/src/driver/test.rs +++ b/ql-runtime/src/driver/test.rs @@ -1,4 +1,4 @@ -use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto, StreamClose, QID}; +use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto, StreamReset, QID}; use super::*; use crate::{ @@ -69,8 +69,8 @@ fn new_inbound_io(capacity: usize) -> InboundIo { let (runtime_tx, _runtime_rx) = async_channel::unbounded(); let stream = io::new_stream( StreamId(99u32.into()), - CloseTarget::Origin, - CloseTarget::Return, + ResetTarget::Origin, + ResetTarget::Return, RuntimeHandle::new(runtime_tx), ); let (_, _, reader_io, _) = stream; @@ -81,8 +81,8 @@ fn new_outbound_io() -> OutboundIo { let (runtime_tx, _runtime_rx) = async_channel::unbounded(); let stream = io::new_stream( StreamId(100u32.into()), - CloseTarget::Return, - CloseTarget::Origin, + ResetTarget::Return, + ResetTarget::Origin, RuntimeHandle::new(runtime_tx), ); let (_, _, _, writer_io) = stream; @@ -90,7 +90,7 @@ fn new_outbound_io() -> OutboundIo { } #[test] -fn handle_inbound_finished_reaps_closed_initiator_stream() { +fn handle_inbound_finished_reaps_reset_initiator_stream() { let (mut state, _fsm) = new_driver_state(); let stream_id = StreamId(1u32.into()); @@ -105,7 +105,7 @@ fn handle_inbound_finished_reaps_closed_initiator_stream() { } #[test] -fn handle_closed_stream_reaps_when_both_halves_close() { +fn handle_reset_stream_reaps_when_both_halves_reset() { let (mut state, _fsm) = new_driver_state(); let stream_id = StreamId(1u32.into()); @@ -114,10 +114,10 @@ fn handle_closed_stream_reaps_when_both_halves_close() { DriverStreamIo::new(false, Some(new_outbound_io()), Some(new_inbound_io(1))), ); - state.handle_closed_stream(&StreamClose { + state.handle_reset_stream(&StreamReset { stream_id, - target: CloseTarget::Both, - code: StreamCloseCode::CANCELLED, + target: ResetTarget::Both, + code: ResetCode::CANCELLED, }); assert!(!state.streams.contains_key(&stream_id)); @@ -130,8 +130,8 @@ fn poll_stream_keeps_outbound_pending_after_local_finish_when_inbound_is_closed( let (runtime_tx, _runtime_rx) = async_channel::unbounded(); let (_, mut writer, _, writer_io) = io::new_stream( stream_id, - CloseTarget::Return, - CloseTarget::Origin, + ResetTarget::Return, + ResetTarget::Origin, RuntimeHandle::new(runtime_tx), ); writer.queue_finish(); @@ -148,14 +148,14 @@ fn poll_stream_keeps_outbound_pending_after_local_finish_when_inbound_is_closed( } #[test] -fn local_close_command_reaps_when_other_half_is_already_closed() { +fn local_reset_command_reaps_when_other_half_is_already_closed() { let (mut state, mut fsm) = new_driver_state(); let stream_id = StreamId(1u32.into()); let (runtime_tx, _runtime_rx) = async_channel::unbounded(); let (_, _, _, writer_io) = io::new_stream( stream_id, - CloseTarget::Return, - CloseTarget::Origin, + ResetTarget::Return, + ResetTarget::Origin, RuntimeHandle::new(runtime_tx), ); @@ -166,10 +166,10 @@ fn local_close_command_reaps_when_other_half_is_already_closed() { state.drive_command( &mut fsm, - Command::CloseStream { + Command::ResetStream { stream_id, - target: CloseTarget::Origin, - code: StreamCloseCode::CANCELLED, + target: ResetTarget::Origin, + code: ResetCode::CANCELLED, }, &NoopCrypto, ); @@ -185,8 +185,8 @@ fn unpaired_status_fails_and_reaps_all_streams() { let (runtime_tx, _runtime_rx) = async_channel::unbounded(); let (_, _, reader_io, writer_io) = io::new_stream( stream_id, - CloseTarget::Origin, - CloseTarget::Return, + ResetTarget::Origin, + ResetTarget::Return, RuntimeHandle::new(runtime_tx), ); diff --git a/ql-runtime/src/error.rs b/ql-runtime/src/error.rs index 55833e88..0cd845ae 100644 --- a/ql-runtime/src/error.rs +++ b/ql-runtime/src/error.rs @@ -1,10 +1,10 @@ -use ql_wire::{StreamCloseCode, StreamCloseOrigin}; +use ql_wire::{ResetCode, ResetOrigin}; #[derive(Debug, Clone, PartialEq, Eq)] pub enum QlStreamError { - StreamClosed { - code: StreamCloseCode, - origin: StreamCloseOrigin, + StreamReset { + code: ResetCode, + origin: ResetOrigin, }, NoSession, } @@ -12,7 +12,7 @@ pub enum QlStreamError { impl std::fmt::Display for QlStreamError { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { - Self::StreamClosed { code, origin } => write!(f, "stream closed {code:?} ({origin:?})"), + Self::StreamReset { code, origin } => write!(f, "stream reset {code:?} ({origin:?})"), Self::NoSession => f.write_str("no session"), } } diff --git a/ql-runtime/src/io/inner.rs b/ql-runtime/src/io/inner.rs index 5d322aad..40866198 100644 --- a/ql-runtime/src/io/inner.rs +++ b/ql-runtime/src/io/inner.rs @@ -292,7 +292,7 @@ mod loom_tests { use bytes::Bytes; use loom::thread; - use ql_wire::{StreamCloseCode, StreamCloseOrigin}; + use ql_wire::{ResetCode, ResetOrigin}; use super::*; use crate::{ @@ -407,9 +407,9 @@ mod loom_tests { let failer = { let shared = shared.clone(); thread::spawn(move || { - shared.rx.fail(QlStreamError::StreamClosed { - code: StreamCloseCode::CANCELLED, - origin: StreamCloseOrigin::Local, + shared.rx.fail(QlStreamError::StreamReset { + code: ResetCode::CANCELLED, + origin: ResetOrigin::Local, }) }) }; @@ -421,14 +421,14 @@ mod loom_tests { (Ok(Item::Chunk(bytes)), None) => { assert_eq!(bytes, Bytes::from_static(b"abc")); match shared.rx.pop() { - Ok(Item::Error(QlStreamError::StreamClosed { code, .. })) => { - assert_eq!(code, StreamCloseCode::CANCELLED); + Ok(Item::Error(QlStreamError::StreamReset { code, .. })) => { + assert_eq!(code, ResetCode::CANCELLED); } _ => panic!("expected terminal reader error"), } } - (Ok(Item::Error(QlStreamError::StreamClosed { code, .. })), Some(bytes)) => { - assert_eq!(code, StreamCloseCode::CANCELLED); + (Ok(Item::Error(QlStreamError::StreamReset { code, .. })), Some(bytes)) => { + assert_eq!(code, ResetCode::CANCELLED); assert_eq!(bytes, Bytes::from_static(b"abc")); assert!(matches!(shared.rx.pop(), Err(PopError))); } @@ -499,9 +499,9 @@ mod loom_tests { let failer = { let shared = shared.clone(); thread::spawn(move || { - let displaced = shared.tx.fail(QlStreamError::StreamClosed { - code: StreamCloseCode::CANCELLED, - origin: StreamCloseOrigin::Local, + let displaced = shared.tx.fail(QlStreamError::StreamReset { + code: ResetCode::CANCELLED, + origin: ResetOrigin::Local, }); assert_eq!(displaced.unwrap(), Some(Bytes::from_static(b"abc"))); }) @@ -512,8 +512,8 @@ mod loom_tests { assert!(TxInner::terminal_ready(shared.tx.load_state())); shared.tx.unregister_waiter(); match shared.tx.pop() { - Ok(Item::Error(QlStreamError::StreamClosed { code, .. })) => { - assert_eq!(code, StreamCloseCode::CANCELLED); + Ok(Item::Error(QlStreamError::StreamReset { code, .. })) => { + assert_eq!(code, ResetCode::CANCELLED); } _ => panic!("expected terminal writer error"), } @@ -570,9 +570,9 @@ mod loom_tests { let failer = { let shared = shared.clone(); thread::spawn(move || { - shared.tx.fail(QlStreamError::StreamClosed { - code: StreamCloseCode::CANCELLED, - origin: StreamCloseOrigin::Local, + shared.tx.fail(QlStreamError::StreamReset { + code: ResetCode::CANCELLED, + origin: ResetOrigin::Local, }) }) }; @@ -597,8 +597,8 @@ mod loom_tests { } match shared.tx.pop() { - Ok(Item::Error(QlStreamError::StreamClosed { code, .. })) => { - assert_eq!(code, StreamCloseCode::CANCELLED); + Ok(Item::Error(QlStreamError::StreamReset { code, .. })) => { + assert_eq!(code, ResetCode::CANCELLED); } _ => panic!("expected terminal writer error"), } @@ -617,9 +617,9 @@ mod loom_tests { let failer = { let shared = shared.clone(); thread::spawn(move || { - shared.tx.fail(QlStreamError::StreamClosed { - code: StreamCloseCode::CANCELLED, - origin: StreamCloseOrigin::Local, + shared.tx.fail(QlStreamError::StreamReset { + code: ResetCode::CANCELLED, + origin: ResetOrigin::Local, }) }) }; @@ -635,8 +635,8 @@ mod loom_tests { Ok(_) => { assert!(!TxInner::terminal_ok(shared.tx.load_state())); match shared.tx.pop() { - Ok(Item::Error(QlStreamError::StreamClosed { code, .. })) => { - assert_eq!(code, StreamCloseCode::CANCELLED); + Ok(Item::Error(QlStreamError::StreamReset { code, .. })) => { + assert_eq!(code, ResetCode::CANCELLED); } _ => panic!("expected terminal writer error"), } diff --git a/ql-runtime/src/io/mod.rs b/ql-runtime/src/io/mod.rs index 2eb7f0f0..2fc4064e 100644 --- a/ql-runtime/src/io/mod.rs +++ b/ql-runtime/src/io/mod.rs @@ -6,7 +6,7 @@ mod writer; use std::ops::Deref; -use ql_wire::{CloseTarget, StreamId}; +use ql_wire::{ResetTarget, StreamId}; pub use self::{reader::StreamReader, slot::PushError, writer::StreamWriter}; use crate::RuntimeHandle; @@ -45,8 +45,8 @@ impl Tx { pub fn new_stream( stream_id: StreamId, - reader_target: CloseTarget, - writer_target: CloseTarget, + reader_target: ResetTarget, + writer_target: ResetTarget, handle: RuntimeHandle, ) -> (StreamReader, StreamWriter, Rx, Tx) { let shared = inner::new(stream_id); diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs index 3f96f2d3..8647d4ee 100644 --- a/ql-runtime/src/io/reader.rs +++ b/ql-runtime/src/io/reader.rs @@ -4,7 +4,7 @@ use std::{ }; use bytes::Bytes; -use ql_wire::{CloseTarget, StreamCloseCode}; +use ql_wire::{ResetCode, ResetTarget}; use super::{ inner::{Item, RxInner}, @@ -15,7 +15,7 @@ use crate::{command::Command, log, QlStreamError, RuntimeHandle}; pub struct StreamReader { rx: Rx, - target: CloseTarget, + target: ResetTarget, terminal: ReaderTerminalState, handle: RuntimeHandle, } @@ -41,7 +41,7 @@ impl std::fmt::Debug for StreamReader { } impl StreamReader { - pub(crate) fn new(shared: Rx, target: CloseTarget, handle: RuntimeHandle) -> Self { + pub(crate) fn new(shared: Rx, target: ResetTarget, handle: RuntimeHandle) -> Self { Self { rx: shared, target, @@ -117,22 +117,22 @@ impl StreamReader { poll_fn(|cx| self.poll_read(cx)).await } - pub fn close(mut self, code: StreamCloseCode) { - self.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); } - fn close_inner(&mut self, code: StreamCloseCode) { + fn reset_inner(&mut self, code: ResetCode) { if matches!(self.terminal, ReaderTerminalState::Delivered) { return; } log::debug!( - "byte reader explicit close: stream_id={:?} target={:?} code={:?}", + "byte reader explicit reset: stream_id={:?} target={:?} code={:?}", self.rx.stream_id(), self.target, code ); self.terminal = ReaderTerminalState::Delivered; - self.handle.try_send(Command::CloseStream { + self.handle.try_send(Command::ResetStream { stream_id: self.rx.stream_id(), target: self.target, code, @@ -146,15 +146,15 @@ impl Drop for StreamReader { return; } log::debug!( - "byte reader drop close: stream_id={:?} target={:?} code={:?}", + "byte reader drop reset: stream_id={:?} target={:?} code={:?}", self.rx.stream_id(), self.target, - StreamCloseCode::DROPPED + ResetCode::DROPPED ); - self.handle.try_send(Command::CloseStream { + self.handle.try_send(Command::ResetStream { stream_id: self.rx.stream_id(), target: self.target, - code: StreamCloseCode::DROPPED, + code: ResetCode::DROPPED, }); } } @@ -165,7 +165,7 @@ mod loom_tests { use bytes::Bytes; use loom::thread; - use ql_wire::CloseTarget; + use ql_wire::ResetTarget; use super::*; use crate::io::sync::loom::*; @@ -174,7 +174,7 @@ mod loom_tests { fn poll_read_observes_chunk_racing_with_registration() { check_model(|| { let inner = shared(); - let mut reader = StreamReader::new(Rx(inner.clone()), CloseTarget::Origin, handle()); + let mut reader = StreamReader::new(Rx(inner.clone()), ResetTarget::Origin, handle()); let mut cx = Context::from_waker(Waker::noop()); let producer = { diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs index f20d0814..4176f46f 100644 --- a/ql-runtime/src/io/writer.rs +++ b/ql-runtime/src/io/writer.rs @@ -4,7 +4,7 @@ use std::{ }; use bytes::Bytes; -use ql_wire::{CloseTarget, StreamCloseCode}; +use ql_wire::{ResetCode, ResetTarget}; use super::{ inner::{Item, TxInner}, @@ -15,7 +15,7 @@ use crate::{command::Command, log, QlStreamError, RuntimeHandle}; pub struct StreamWriter { tx: Tx, - target: CloseTarget, + target: ResetTarget, open: bool, terminal: WriterTerminalState, handle: RuntimeHandle, @@ -39,7 +39,7 @@ impl std::fmt::Debug for StreamWriter { } impl StreamWriter { - pub(crate) fn new(shared: Tx, target: CloseTarget, handle: RuntimeHandle) -> Self { + pub(crate) fn new(shared: Tx, target: ResetTarget, handle: RuntimeHandle) -> Self { Self { tx: shared, target, @@ -139,8 +139,8 @@ impl StreamWriter { self.poll_terminal(cx) } - pub fn close(mut self, code: StreamCloseCode) { - self.close_inner(code); + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); } fn poll_runtime(&self) { @@ -194,18 +194,18 @@ impl StreamWriter { Poll::Pending } - fn close_inner(&mut self, code: StreamCloseCode) { + fn reset_inner(&mut self, code: ResetCode) { if !self.open { return; } self.open = false; log::debug!( - "byte writer close: stream_id={:?} target={:?} code={:?}", + "byte writer reset: stream_id={:?} target={:?} code={:?}", self.tx.stream_id(), self.target, code ); - self.handle.try_send(Command::CloseStream { + self.handle.try_send(Command::ResetStream { stream_id: self.tx.stream_id(), target: self.target, code, @@ -215,7 +215,7 @@ impl StreamWriter { impl Drop for StreamWriter { fn drop(&mut self) { - self.close_inner(StreamCloseCode::DROPPED); + self.reset_inner(ResetCode::DROPPED); } } @@ -225,7 +225,7 @@ mod loom_tests { use bytes::Bytes; use loom::thread; - use ql_wire::CloseTarget; + use ql_wire::ResetTarget; use super::*; use crate::io::sync::loom::*; @@ -236,7 +236,7 @@ mod loom_tests { let inner = shared(); inner.tx.try_write(Bytes::from_static(b"abc")).unwrap(); - let mut writer = StreamWriter::new(Tx(inner.clone()), CloseTarget::Origin, handle()); + let mut writer = StreamWriter::new(Tx(inner.clone()), ResetTarget::Origin, handle()); let mut bytes = Bytes::from_static(b"xyz"); let mut cx = Context::from_waker(Waker::noop()); @@ -267,7 +267,7 @@ mod loom_tests { fn poll_finish_observes_terminal_racing_with_registration() { check_model(|| { let inner = shared(); - let mut writer = StreamWriter::new(Tx(inner.clone()), CloseTarget::Origin, handle()); + let mut writer = StreamWriter::new(Tx(inner.clone()), ResetTarget::Origin, handle()); let mut cx = Context::from_waker(Waker::noop()); writer.queue_finish(); diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index dd6e4813..388bc060 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -1,7 +1,7 @@ use std::task::{Context, Poll}; use bytes::Bytes; -use ql_rpc::{RouteId, RpcRead, RpcStream, RpcWrite, ServiceId, StreamCloseCode}; +use ql_rpc::{ResetCode, RouteId, RpcRead, RpcStream, RpcWrite, ServiceId}; use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; @@ -30,8 +30,8 @@ impl RpcRead for StreamReader { StreamReader::poll_read(self, cx) } - fn close(self, code: StreamCloseCode) { - StreamReader::close(self, code); + fn reset(self, code: ResetCode) { + StreamReader::reset(self, code); } } @@ -50,7 +50,7 @@ impl RpcWrite for StreamWriter { StreamWriter::poll_finish(self, cx) } - fn close(self, code: StreamCloseCode) { - StreamWriter::close(self, code); + fn reset(self, code: ResetCode) { + StreamWriter::reset(self, code); } } diff --git a/ql-runtime/src/rpc/download.rs b/ql-runtime/src/rpc/download.rs index d3b63585..cfcc25f6 100644 --- a/ql-runtime/src/rpc/download.rs +++ b/ql-runtime/src/rpc/download.rs @@ -25,8 +25,8 @@ where Ok((header, DownloadReader { inner })) } - pub fn close(self, code: ql_wire::StreamCloseCode) { - self.inner.close(ql_rpc::StreamCloseCode(code.0)); + pub fn reset(self, code: ql_wire::ResetCode) { + self.inner.reset(ql_rpc::ResetCode(code.0)); } } @@ -48,8 +48,8 @@ where self.inner.complete().await.map_err(RpcError::from) } - pub fn close(self, code: ql_wire::StreamCloseCode) { - self.inner.close(ql_rpc::StreamCloseCode(code.0)); + pub fn reset(self, code: ql_wire::ResetCode) { + self.inner.reset(ql_rpc::ResetCode(code.0)); } } @@ -61,7 +61,7 @@ where Ok(self.inner.read_chunk().await?) } - pub fn close(self, code: ql_wire::StreamCloseCode) { - self.inner.close(ql_rpc::StreamCloseCode(code.0)); + pub fn reset(self, code: ql_wire::ResetCode) { + self.inner.reset(ql_rpc::ResetCode(code.0)); } } diff --git a/ql-runtime/src/rpc/duplex.rs b/ql-runtime/src/rpc/duplex.rs index cdad6670..be040d3e 100644 --- a/ql-runtime/src/rpc/duplex.rs +++ b/ql-runtime/src/rpc/duplex.rs @@ -35,8 +35,8 @@ where self.inner.finish().await } - pub fn close(self, code: ql_wire::StreamCloseCode) { - self.inner.close(ql_rpc::StreamCloseCode(code.0)); + pub fn reset(self, code: ql_wire::ResetCode) { + self.inner.reset(ql_rpc::ResetCode(code.0)); } } @@ -53,7 +53,7 @@ where .await } - pub fn close(self, code: ql_wire::StreamCloseCode) { - self.inner.close(ql_rpc::StreamCloseCode(code.0)); + pub fn reset(self, code: ql_wire::ResetCode) { + self.inner.reset(ql_rpc::ResetCode(code.0)); } } diff --git a/ql-runtime/src/rpc/error.rs b/ql-runtime/src/rpc/error.rs index e1d1f384..965b711a 100644 --- a/ql-runtime/src/rpc/error.rs +++ b/ql-runtime/src/rpc/error.rs @@ -5,9 +5,9 @@ use crate::QlStreamError; #[derive(Debug)] pub enum RpcError { NoSession, - Closed { - code: ql_rpc::StreamCloseCode, - origin: ql_rpc::StreamCloseOrigin, + Reset { + code: ql_rpc::ResetCode, + origin: ql_rpc::ResetOrigin, }, Protocol(ql_rpc::Error), Codec(E), @@ -22,7 +22,7 @@ impl From for RpcError { impl From for RpcError { fn from(error: QlStreamError) -> Self { match error { - QlStreamError::StreamClosed { code, origin } => Self::Closed { code, origin }, + QlStreamError::StreamReset { code, origin } => Self::Reset { code, origin }, QlStreamError::NoSession => Self::NoSession, } } @@ -51,7 +51,7 @@ where fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { Self::NoSession => write!(f, "no session"), - Self::Closed { code, origin } => write!(f, "stream closed {code:?} ({origin:?})"), + Self::Reset { code, origin } => write!(f, "stream reset {code:?} ({origin:?})"), Self::Protocol(error) => write!(f, "{error}"), Self::Codec(error) => write!(f, "{error}"), } @@ -67,7 +67,7 @@ where Self::Protocol(error) => Some(error), Self::Codec(error) => Some(error), RpcError::NoSession => None, - RpcError::Closed { .. } => None, + RpcError::Reset { .. } => None, } } } diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index 828bf0ee..c1d8c7c8 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -38,7 +38,7 @@ impl RpcHandle { .inner .open_stream(open_stream_params(M::SERVICE, M::ROUTE)) .await?; - stream.reader.close(ql_rpc::StreamCloseCode::CANCELLED); + stream.reader.reset(ql_rpc::ResetCode::CANCELLED); stream.writer.write(Bytes::from(payload)).await?; stream.writer.finish().await?; Ok(()) diff --git a/ql-runtime/src/rpc/progress.rs b/ql-runtime/src/rpc/progress.rs index a22da20f..c7e5e77b 100644 --- a/ql-runtime/src/rpc/progress.rs +++ b/ql-runtime/src/rpc/progress.rs @@ -20,8 +20,8 @@ impl ProgressCall where M: Progress, { - pub fn close(self, code: ql_wire::StreamCloseCode) { - self.inner.close(ql_rpc::StreamCloseCode(code.0)); + pub fn reset(self, code: ql_wire::ResetCode) { + self.inner.reset(ql_rpc::ResetCode(code.0)); } } diff --git a/ql-runtime/src/rpc/subscription.rs b/ql-runtime/src/rpc/subscription.rs index 45a08a6b..5652394a 100644 --- a/ql-runtime/src/rpc/subscription.rs +++ b/ql-runtime/src/rpc/subscription.rs @@ -23,8 +23,8 @@ where poll_fn(|cx| Pin::new(&mut *self).poll_next(cx)).await } - pub fn close(self, code: ql_wire::StreamCloseCode) { - self.inner.close(ql_rpc::StreamCloseCode(code.0)); + pub fn reset(self, code: ql_wire::ResetCode) { + self.inner.reset(ql_rpc::ResetCode(code.0)); } } diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index c2c43e1f..fabded16 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -12,9 +12,9 @@ use futures_lite::StreamExt; use ql_rpc::{ DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, NotificationHandlerLocal, ProgressHandlerLocal, ProgressResponder, RequestHandler, - RequestHandlerLocal, Response, RouteId, SendSpawner, ServiceId, Spawner, StreamCloseCode, - StreamCloseOrigin, SubscriptionHandlerLocal, SubscriptionResponder, UploadHandlerLocal, - UploadReader, UploadResponder, + RequestHandlerLocal, ResetCode, ResetOrigin, Response, RouteId, SendSpawner, ServiceId, + Spawner, SubscriptionHandlerLocal, SubscriptionResponder, UploadHandlerLocal, UploadReader, + UploadResponder, }; use super::*; @@ -330,8 +330,8 @@ async fn rpc_router_enforces_max_request_bytes() { let response = rpc.request::(&"hello".to_string()).await; assert!(matches!( response, - Err(RpcError::Closed { code, origin }) - if code == StreamCloseCode::LIMIT && origin == StreamCloseOrigin::Peer + Err(RpcError::Reset { code, origin }) + if code == ResetCode::LIMIT && origin == ResetOrigin::Peer )); tokio::time::timeout(Duration::from_secs(2), responder) diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index fb23665b..7d9e7bb5 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -1,7 +1,7 @@ use std::time::Duration; use bytes::Bytes; -use ql_wire::{StreamCloseCode, StreamCloseOrigin}; +use ql_wire::{ResetCode, ResetOrigin}; use super::*; use crate::QlStreamError; @@ -124,15 +124,15 @@ async fn dropping_responder_closes_initiator_response() { let err = stream.writer.finish().await.unwrap_err(); assert!(matches!( err, - QlStreamError::StreamClosed { code, origin } - if code == StreamCloseCode::DROPPED && origin == StreamCloseOrigin::Peer + QlStreamError::StreamReset { code, origin } + if code == ResetCode::DROPPED && origin == ResetOrigin::Peer )); let err = next_chunk(&mut stream.reader).await.unwrap_err(); assert!(matches!( err, - QlStreamError::StreamClosed { code, origin } - if code == StreamCloseCode::DROPPED && origin == StreamCloseOrigin::Peer + QlStreamError::StreamReset { code, origin } + if code == ResetCode::DROPPED && origin == ResetOrigin::Peer )); tokio::time::timeout(Duration::from_secs(2), responder) @@ -165,8 +165,8 @@ async fn dropping_inbound_reader_cancels_remote_writer() { let err = writer.finish().await.unwrap_err(); assert!(matches!( err, - QlStreamError::StreamClosed { code, origin } - if code == StreamCloseCode::DROPPED && origin == StreamCloseOrigin::Peer + QlStreamError::StreamReset { code, origin } + if code == ResetCode::DROPPED && origin == ResetOrigin::Peer )); }); @@ -213,7 +213,7 @@ async fn closing_initiator_reader_preserves_initiator_writer() { .await .unwrap(); let mut writer = stream.writer; - stream.reader.close(StreamCloseCode::CANCELLED); + stream.reader.reset(ResetCode::CANCELLED); writer.write(Bytes::from_static(&[1, 2])).await.unwrap(); writer.write(Bytes::from_static(&[3, 4])).await.unwrap(); diff --git a/ql-wire/src/encrypted/builder.rs b/ql-wire/src/encrypted/builder.rs index a60f8a57..d64ed2d8 100644 --- a/ql-wire/src/encrypted/builder.rs +++ b/ql-wire/src/encrypted/builder.rs @@ -1,6 +1,6 @@ use bytes::BufMut; -use super::{RecordAck, SessionClose, SessionFrame, StreamClose, StreamData, StreamWindow}; +use super::{RecordAck, SessionClose, SessionFrame, StreamData, StreamReset, StreamWindow}; use crate::{ BufView, ConnectionId, Nonce, QlCrypto, RecordSeq, RecordType, SessionHeader, SessionKey, WireEncode, QL_WIRE_VERSION, @@ -82,8 +82,8 @@ impl SessionRecordBuilder { self.push_frame_payload(super::SessionFrameKind::StreamWindow, frame) } - pub fn push_stream_close(&mut self, frame: &StreamClose) -> bool { - self.push_frame_payload(super::SessionFrameKind::StreamClose, frame) + pub fn push_stream_reset(&mut self, frame: &StreamReset) -> bool { + self.push_frame_payload(super::SessionFrameKind::StreamReset, frame) } pub fn push_close(&mut self, close: &SessionClose) -> bool { @@ -97,7 +97,7 @@ impl SessionRecordBuilder { SessionFrame::Ack(frame) => self.push_ack(frame), SessionFrame::StreamData(frame) => self.push_stream_data(frame), SessionFrame::StreamWindow(frame) => self.push_stream_window(frame), - SessionFrame::StreamClose(frame) => self.push_stream_close(frame), + SessionFrame::StreamReset(frame) => self.push_stream_reset(frame), SessionFrame::Close(close) => self.push_close(close), } } diff --git a/ql-wire/src/encrypted/close.rs b/ql-wire/src/encrypted/close.rs index e0860d7a..c9e0d237 100644 --- a/ql-wire/src/encrypted/close.rs +++ b/ql-wire/src/encrypted/close.rs @@ -1,6 +1,6 @@ use crate::{codec, codec::Reader, ByteSlice, WireEncode, WireError}; -/// closes the whole session immediately with a close code. +/// closes the whole session immediately with a reset code. #[derive(Debug, Clone, PartialEq, Eq)] pub struct SessionClose { pub code: SessionCloseCode, diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index 1f45ac89..93ae23e6 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -6,15 +6,15 @@ use crate::{ mod ack; mod builder; mod close; -mod stream_close; mod stream_data; +mod stream_reset; mod stream_window; pub use ack::*; pub use builder::*; pub use close::*; -pub use stream_close::*; pub use stream_data::*; +pub use stream_reset::*; pub use stream_window::*; varint_wrapper_codec!(RouteId); @@ -29,7 +29,7 @@ pub enum SessionFrame { Ack(RecordAck), StreamData(StreamData), StreamWindow(StreamWindow), - StreamClose(StreamClose), + StreamReset(StreamReset), Close(SessionClose), } @@ -42,7 +42,7 @@ impl WireDecode for SessionFrame { SessionFrameKind::Ack => Self::Ack(reader.decode::()?), SessionFrameKind::StreamData => Self::StreamData(reader.decode::>()?), SessionFrameKind::StreamWindow => Self::StreamWindow(reader.decode::()?), - SessionFrameKind::StreamClose => Self::StreamClose(reader.decode::()?), + SessionFrameKind::StreamReset => Self::StreamReset(reader.decode::()?), SessionFrameKind::Close => Self::Close(reader.decode::()?), }; Ok(frame) @@ -57,7 +57,7 @@ impl SessionFrame { Self::Ack(_) => SessionFrameKind::Ack, Self::StreamData(_) => SessionFrameKind::StreamData, Self::StreamWindow(_) => SessionFrameKind::StreamWindow, - Self::StreamClose(_) => SessionFrameKind::StreamClose, + Self::StreamReset(_) => SessionFrameKind::StreamReset, Self::Close(_) => SessionFrameKind::Close, } } @@ -71,7 +71,7 @@ impl SessionFrame { Self::Ack(frame) => SessionFrame::Ack(frame), Self::StreamData(frame) => SessionFrame::StreamData(frame.into_owned()), Self::StreamWindow(frame) => SessionFrame::StreamWindow(frame), - Self::StreamClose(frame) => SessionFrame::StreamClose(frame), + Self::StreamReset(frame) => SessionFrame::StreamReset(frame), Self::Close(frame) => SessionFrame::Close(frame), } } @@ -84,7 +84,7 @@ impl WireEncode for SessionFrame { Self::Ack(frame) => frame.encoded_len(), Self::StreamData(frame) => frame.encoded_len(), Self::StreamWindow(frame) => frame.encoded_len(), - Self::StreamClose(frame) => frame.encoded_len(), + Self::StreamReset(frame) => frame.encoded_len(), Self::Close(frame) => frame.encoded_len(), } } @@ -96,7 +96,7 @@ impl WireEncode for SessionFrame { Self::Ack(frame) => frame.encode(out), Self::StreamData(frame) => frame.encode(out), Self::StreamWindow(frame) => frame.encode(out), - Self::StreamClose(frame) => frame.encode(out), + Self::StreamReset(frame) => frame.encode(out), Self::Close(frame) => frame.encode(out), } } @@ -109,7 +109,7 @@ pub enum SessionFrameKind { Ack = 2, StreamData = 3, StreamWindow = 4, - StreamClose = 5, + StreamReset = 5, Close = 6, Unpair = 7, } @@ -123,7 +123,7 @@ impl TryFrom for SessionFrameKind { 2 => Ok(Self::Ack), 3 => Ok(Self::StreamData), 4 => Ok(Self::StreamWindow), - 5 => Ok(Self::StreamClose), + 5 => Ok(Self::StreamReset), 6 => Ok(Self::Close), 7 => Ok(Self::Unpair), _ => Err(WireError::InvalidPayload), diff --git a/ql-wire/src/encrypted/stream_close.rs b/ql-wire/src/encrypted/stream_reset.rs similarity index 69% rename from ql-wire/src/encrypted/stream_close.rs rename to ql-wire/src/encrypted/stream_reset.rs index dffbada9..99381e39 100644 --- a/ql-wire/src/encrypted/stream_close.rs +++ b/ql-wire/src/encrypted/stream_reset.rs @@ -1,21 +1,21 @@ use super::StreamId; -use crate::{codec, ByteSlice, StreamCloseCode, WireEncode, WireError}; +use crate::{codec, ByteSlice, ResetCode, WireEncode, WireError}; -/// aborts one or both lanes of a stream with a close code +/// aborts one or both lanes of a stream with a reset code /// /// stream origin is the peer that opened the stream /// origin lane carries bytes sent by the stream origin /// return lane carries bytes sent back toward the stream origin #[derive(Debug, Clone, PartialEq, Eq)] -pub struct StreamClose { +pub struct StreamReset { pub stream_id: StreamId, - pub target: CloseTarget, - pub code: StreamCloseCode, + pub target: ResetTarget, + pub code: ResetCode, } -impl StreamClose {} +impl StreamReset {} -impl WireEncode for StreamClose { +impl WireEncode for StreamReset { fn encoded_len(&self) -> usize { self.stream_id.encoded_len() + self.target.encoded_len() + self.code.encoded_len() } @@ -27,7 +27,7 @@ impl WireEncode for StreamClose { } } -impl codec::WireDecode for StreamClose { +impl codec::WireDecode for StreamReset { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { stream_id: reader.decode()?, @@ -37,25 +37,25 @@ impl codec::WireDecode for StreamClose { } } -/// selects which stream lane a [`StreamClose`] applies to +/// selects which stream lane a [`StreamReset`] applies to #[derive(Debug, Clone, Copy, PartialEq, Eq)] #[repr(u8)] -pub enum CloseTarget { - /// close the lane sent by the stream origin +pub enum ResetTarget { + /// reset the lane sent by the stream origin Origin = 1, - /// close the lane sent back toward the stream origin + /// reset the lane sent back toward the stream origin Return = 2, - /// close both stream lanes + /// reset both stream lanes Both = 3, } -impl CloseTarget { +impl ResetTarget { pub const fn to_wire(self) -> u8 { self as u8 } } -impl WireEncode for CloseTarget { +impl WireEncode for ResetTarget { fn encoded_len(&self) -> usize { size_of::() } @@ -65,7 +65,7 @@ impl WireEncode for CloseTarget { } } -impl TryFrom for CloseTarget { +impl TryFrom for ResetTarget { type Error = WireError; fn try_from(value: u8) -> Result { @@ -78,19 +78,19 @@ impl TryFrom for CloseTarget { } } -impl codec::WireDecode for CloseTarget { +impl codec::WireDecode for ResetTarget { fn decode(reader: &mut codec::Reader) -> Result { reader.decode::()?.try_into() } } -impl codec::WireDecode for StreamCloseCode { +impl codec::WireDecode for ResetCode { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self(reader.decode()?)) } } -impl WireEncode for StreamCloseCode { +impl WireEncode for ResetCode { fn encoded_len(&self) -> usize { size_of::() } diff --git a/ql-wire/src/lib.rs b/ql-wire/src/lib.rs index 2bb1e3e5..0016941b 100644 --- a/ql-wire/src/lib.rs +++ b/ql-wire/src/lib.rs @@ -34,8 +34,7 @@ pub use identity::*; pub use pq::*; pub use qid::*; pub use ql_common::{ - RouteId, ServiceId, StreamCloseCode, StreamCloseOrigin, StreamId, VarInt, VarIntBoundsExceeded, - QID, + ResetCode, ResetOrigin, RouteId, ServiceId, StreamId, VarInt, VarIntBoundsExceeded, QID, }; pub use record::*; #[cfg(any(feature = "test-utils", test))] diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index 6657ea31..b241c4aa 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -739,10 +739,10 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { bytes: b"hello".to_vec(), fin: true, }), - SessionFrame::StreamClose(StreamClose { + SessionFrame::StreamReset(StreamReset { stream_id: stream_id(9), - target: CloseTarget::Both, - code: StreamCloseCode::CANCELLED, + target: ResetTarget::Both, + code: ResetCode::CANCELLED, }), SessionFrame::Close(SessionClose { code: SessionCloseCode::TIMEOUT, From 6820b91b7836a3dbcadd9a795712483bd2608c85 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 23 Jun 2026 17:38:16 -0400 Subject: [PATCH 20/59] ql-rpc: separate poll_finish from finish --- ql-rpc/src/rpc/duplex/client.rs | 7 +++---- ql-rpc/src/rpc/upload/client.rs | 4 ++-- ql-rpc/src/stream.rs | 18 +++++++++++++----- ql-runtime/src/rpc/adapter.rs | 4 ++++ ql-runtime/src/rpc/duplex.rs | 4 ++-- ql-runtime/src/rpc/mod.rs | 2 +- ql-runtime/src/tests/rpc.rs | 4 ++-- 7 files changed, 27 insertions(+), 16 deletions(-) diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index 7b6bb5f0..ccad5e05 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -8,8 +8,7 @@ use bytes::Bytes; use crate::{ duplex::{codec, Duplex, EventReader, ReadStep}, - finish_bytes, write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcCodec, RpcError, - RpcRead, RpcWrite, + write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcCodec, RpcError, RpcRead, RpcWrite, }; pub struct DuplexCall @@ -59,8 +58,8 @@ where write_bytes(writer, Bytes::from(encoded)).await } - pub async fn finish(mut self) -> Result<(), W::Error> { - finish_bytes(&mut self.writer).await + pub fn finish(mut self) { + self.writer.queue_finish(); } pub fn reset(mut self, code: ResetCode) { diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index e16833b3..0cb06510 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -1,7 +1,7 @@ use bytes::{BufMut, Bytes}; use crate::{ - finish_bytes, read_bytes, + read_bytes, rpc::parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, upload::Upload, write_bytes, ChunkQueue, DropResetRead, DropResetWrite, ResetCode, RpcCodec, RpcError, RpcRead, @@ -64,7 +64,7 @@ where write_bytes(writer, Bytes::from(encoded)) .await .map_err(RpcError::Transport)?; - finish_bytes(writer).await.map_err(RpcError::Transport)?; + writer.queue_finish(); let reader = &mut self.reader; let mut bytes = ChunkQueue::default(); diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index d6cffecc..4952e192 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -1,5 +1,5 @@ use std::{ - future::poll_fn, + future::{poll_fn, Future}, task::{Context, Poll}, }; @@ -37,10 +37,13 @@ pub trait RpcWrite { cx: &mut Context<'_>, ) -> Poll>; - /// completes the write side and must be polled until ready without further write or reset calls + /// queues a graceful write-side finish + fn queue_finish(&mut self); + + /// waits for the queued finish to be delivered fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll>; - /// aborts the write side before finish + /// aborts the write side before finish; must not replace a queued finish fn reset(self, code: ResetCode); } @@ -59,11 +62,12 @@ where poll_fn(|cx| writer.poll_write(&mut bytes, cx)).await } -pub async fn finish_bytes(writer: &mut W) -> Result<(), W::Error> +pub fn finish_bytes(writer: &mut W) -> impl Future> + '_ where W: RpcWrite, { - poll_fn(|cx| writer.poll_finish(cx)).await + writer.queue_finish(); + poll_fn(|cx| writer.poll_finish(cx)) } pub fn reset_stream(stream: St, code: ResetCode) @@ -158,6 +162,10 @@ mod drop { self.inner.as_mut().unwrap().poll_write(bytes, cx) } + fn queue_finish(&mut self) { + self.inner.as_mut().unwrap().queue_finish(); + } + #[track_caller] fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll> { self.inner.as_mut().unwrap().poll_finish(cx) diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index 388bc060..93b9f361 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -46,6 +46,10 @@ impl RpcWrite for StreamWriter { StreamWriter::poll_write(self, bytes, cx) } + fn queue_finish(&mut self) { + StreamWriter::queue_finish(self); + } + fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll> { StreamWriter::poll_finish(self, cx) } diff --git a/ql-runtime/src/rpc/duplex.rs b/ql-runtime/src/rpc/duplex.rs index be040d3e..e5c52609 100644 --- a/ql-runtime/src/rpc/duplex.rs +++ b/ql-runtime/src/rpc/duplex.rs @@ -31,8 +31,8 @@ where self.inner.send(event).await } - pub async fn finish(self) -> Result<(), QlStreamError> { - self.inner.finish().await + pub fn finish(self) { + self.inner.finish(); } pub fn reset(self, code: ql_wire::ResetCode) { diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index c1d8c7c8..c1257b4c 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -150,7 +150,7 @@ impl RpcHandle { .open_stream(open_stream_params(service_id, route_id)) .await?; stream.writer.write(Bytes::from(payload)).await?; - stream.writer.finish().await?; + stream.writer.queue_finish(); Ok(stream.reader) } } diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index fabded16..1bd13b22 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -640,7 +640,7 @@ async fn rpc_duplex() { let second = peer.receiver.next_event().await.unwrap().unwrap(); seen.borrow_mut().push(second); - peer.sender.finish().await.unwrap(); + peer.sender.finish(); } } @@ -670,7 +670,7 @@ async fn rpc_duplex() { b"challenge-response".to_vec() ); chat.sender.send(&b"verification".to_vec()).await.unwrap(); - chat.sender.finish().await.unwrap(); + chat.sender.finish(); assert!(chat.receiver.next_event().await.is_none()); assert_eq!( From 80d834e243ac9f68c87684385295c6bfee4da959 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 24 Jun 2026 09:34:54 -0400 Subject: [PATCH 21/59] ql-rpc: stream handler context --- ql-rpc/src/lib.rs | 2 +- ql-rpc/src/router/builder.rs | 360 +++++++++++++++----------- ql-rpc/src/router/mod.rs | 24 +- ql-rpc/src/router/mode.rs | 3 +- ql-rpc/src/rpc/download/server.rs | 15 +- ql-rpc/src/rpc/duplex/server.rs | 10 +- ql-rpc/src/rpc/notification/server.rs | 11 +- ql-rpc/src/rpc/progress/server.rs | 15 +- ql-rpc/src/rpc/request/server.rs | 14 +- ql-rpc/src/rpc/subscription/server.rs | 10 +- ql-rpc/src/rpc/upload/server.rs | 7 +- ql-rpc/src/stream.rs | 8 +- ql-runtime/src/command.rs | 2 +- ql-runtime/src/driver/mod.rs | 12 +- ql-runtime/src/driver/test.rs | 2 +- ql-runtime/src/handle/mod.rs | 16 +- ql-runtime/src/io/reader.rs | 6 +- ql-runtime/src/io/writer.rs | 6 +- ql-runtime/src/platform.rs | 18 +- ql-runtime/src/rpc/adapter.rs | 33 ++- ql-runtime/src/tests/mod.rs | 14 +- ql-runtime/src/tests/rpc.rs | 131 ++++++---- 22 files changed, 436 insertions(+), 283 deletions(-) diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index 90941bc7..09a91ff3 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -14,7 +14,7 @@ pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; use framed_value::*; -pub use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId}; +pub use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId, StreamId, QID}; pub use router::*; pub use rpc::*; pub use stream::*; diff --git a/ql-rpc/src/router/builder.rs b/ql-rpc/src/router/builder.rs index 9a78f812..e655da28 100644 --- a/ql-rpc/src/router/builder.rs +++ b/ql-rpc/src/router/builder.rs @@ -84,17 +84,21 @@ where M: RequestRpc + 'static, S: RequestHandlerLocal + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_request_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_request_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn notification(self) -> Self @@ -102,17 +106,21 @@ where M: NotificationRpc + 'static, S: NotificationHandlerLocal + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_notification_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_notification_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn duplex(self) -> Self @@ -120,16 +128,20 @@ where M: DuplexRpc + 'static, S: DuplexHandlerLocal + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_duplex_inner::( - state, - config, - reader, - writer, - S::handle, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_duplex_inner::( + state, + context, + config, + reader, + writer, + S::handle, + )) + }, + ) } pub fn download(self) -> Self @@ -137,17 +149,21 @@ where M: DownloadRpc + 'static, S: DownloadHandlerLocal + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_download_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_download_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn subscription(self) -> Self @@ -155,17 +171,21 @@ where M: SubscriptionRpc + 'static, S: SubscriptionHandlerLocal + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_subscription_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_subscription_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn progress(self) -> Self @@ -173,17 +193,21 @@ where M: ProgressRpc + 'static, S: ProgressHandlerLocal + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_progress_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_progress_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn upload(self) -> Self @@ -191,17 +215,21 @@ where M: UploadRpc + 'static, S: UploadHandlerLocal + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_upload_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_upload_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } } @@ -218,17 +246,21 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_request_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_request_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn notification(self) -> Self @@ -239,17 +271,21 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_notification_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_notification_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn duplex(self) -> Self @@ -261,16 +297,20 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_duplex_inner::( - state, - config, - reader, - writer, - S::handle, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_duplex_inner::( + state, + context, + config, + reader, + writer, + S::handle, + )) + }, + ) } pub fn download(self) -> Self @@ -281,17 +321,21 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_download_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_download_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn subscription(self) -> Self @@ -302,17 +346,21 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_subscription_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_subscription_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn progress(self) -> Self @@ -323,17 +371,21 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_progress_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_progress_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } pub fn upload(self) -> Self @@ -344,16 +396,20 @@ where St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route(RouteKey::new::(), |spawner, state, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_upload_inner::( - state, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }) + self.add_route( + RouteKey::new::(), + |spawner, state, context, config, stream| { + let (reader, writer) = stream.split(); + spawner.spawn(handle_upload_inner::( + state, + context, + config, + reader, + writer, + S::handle, + S::handle_error, + )) + }, + ) } } diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index 8bf5d541..985a9c5b 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -1,4 +1,4 @@ -use crate::{ResetCode, RouteId, ServiceId}; +use crate::{ResetCode, RouteId, ServiceId, StreamId, QID}; mod builder; mod config; @@ -30,6 +30,12 @@ where routes: Vec>, } +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct Context { + pub qid: QID, + pub stream_id: StreamId, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] pub struct RouteKey { pub service_id: ServiceId, @@ -83,8 +89,12 @@ where } pub fn handle(&self, stream: St) -> Option<(RouteId, Sp::Handle)> { - let service_id = stream.service_id()?; - let route_id = stream.route_id()?; + let service_id = stream.service_id(); + let route_id = stream.route_id(); + let context = Context { + qid: stream.qid(), + stream_id: stream.stream_id(), + }; let key = RouteKey { service_id, route_id, @@ -96,7 +106,13 @@ where let route = self.routes[index].route; Some(( route_id, - route(&self.spawner, self.state.clone(), self.config, stream), + route( + &self.spawner, + self.state.clone(), + context, + self.config, + stream, + ), )) } diff --git a/ql-rpc/src/router/mode.rs b/ql-rpc/src/router/mode.rs index 33b6c06a..3a61ae4b 100644 --- a/ql-rpc/src/router/mode.rs +++ b/ql-rpc/src/router/mode.rs @@ -1,8 +1,9 @@ use std::future::Future; +use super::Context; use crate::RouterConfig; -pub type RouteFn = fn(&Sp, S, RouterConfig, St) -> ::Handle; +pub type RouteFn = fn(&Sp, S, Context, RouterConfig, St) -> ::Handle; pub trait Spawner: Clone + 'static { type Handle; diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs index 95e22822..98b3a4e7 100644 --- a/ql-rpc/src/rpc/download/server.rs +++ b/ql-rpc/src/rpc/download/server.rs @@ -10,7 +10,8 @@ use crate::{ parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, read_eof_request, }, - write_bytes, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, + write_bytes, Context, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, + RpcWrite, }; #[trait_variant::make(DownloadHandler: Send)] @@ -19,7 +20,12 @@ where M: DownloadRpc, St: RpcStream, { - async fn handle(self, message: M::Request, download: DownloadStart); + async fn handle( + self, + context: Context, + message: M::Request, + download: DownloadStart, + ); fn handle_error(&self, _error: &RpcError) {} } @@ -162,6 +168,7 @@ where pub(crate) async fn handle_download_inner( state: S, + context: Context, config: RouterConfig, mut reader: St::Reader, writer: St::Writer, @@ -170,7 +177,7 @@ pub(crate) async fn handle_download_inner( ) where M: DownloadRpc + 'static, St: RpcStream + 'static, - H: FnOnce(S, M::Request, DownloadStart) -> HF, + H: FnOnce(S, Context, M::Request, DownloadStart) -> HF, HF: Future, E: FnOnce(&S, &RpcError), { @@ -187,5 +194,5 @@ pub(crate) async fn handle_download_inner( } }; - handle(state, request, DownloadStart::new(writer)).await; + handle(state, context, request, DownloadStart::new(writer)).await; } diff --git a/ql-rpc/src/rpc/duplex/server.rs b/ql-rpc/src/rpc/duplex/server.rs index bf024335..6dfb94cc 100644 --- a/ql-rpc/src/rpc/duplex/server.rs +++ b/ql-rpc/src/rpc/duplex/server.rs @@ -2,7 +2,7 @@ use std::future::Future; use crate::{ duplex::{Duplex, DuplexReceiver, DuplexSender}, - RpcRead, RpcStream, RpcWrite, + Context, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(DuplexHandler: Send)] @@ -11,7 +11,9 @@ where M: Duplex, St: RpcStream, { - async fn handle(self, peer: DuplexPeer); + async fn handle(self, context: Context, peer: DuplexPeer); + + fn handle_error(&self, _error: &RpcError) {} } pub struct DuplexPeer @@ -26,6 +28,7 @@ where pub(crate) async fn handle_duplex_inner( state: S, + context: Context, _config: crate::RouterConfig, reader: St::Reader, writer: St::Writer, @@ -33,11 +36,12 @@ pub(crate) async fn handle_duplex_inner( ) where M: Duplex + 'static, St: RpcStream + 'static, - H: FnOnce(S, DuplexPeer) -> HF, + H: FnOnce(S, Context, DuplexPeer) -> HF, HF: Future, { handle( state, + context, DuplexPeer { sender: DuplexSender::new(writer), receiver: DuplexReceiver::new(reader), diff --git a/ql-rpc/src/rpc/notification/server.rs b/ql-rpc/src/rpc/notification/server.rs index ae766532..87258b7e 100644 --- a/ql-rpc/src/rpc/notification/server.rs +++ b/ql-rpc/src/rpc/notification/server.rs @@ -1,8 +1,8 @@ use std::future::Future; use crate::{ - notification::Notification as NotificationRpc, rpc::read_eof_request, ResetCode, RouterConfig, - RpcError, RpcRead, RpcStream, RpcWrite, + notification::Notification as NotificationRpc, rpc::read_eof_request, Context, ResetCode, + RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(NotificationHandler: Send)] @@ -11,13 +11,14 @@ where M: NotificationRpc, St: RpcStream, { - async fn handle(self, message: M::Payload); + async fn handle(self, context: Context, message: M::Payload); fn handle_error(&self, _error: &RpcError) {} } pub(crate) async fn handle_notification_inner( state: S, + context: Context, config: RouterConfig, mut reader: St::Reader, writer: St::Writer, @@ -26,7 +27,7 @@ pub(crate) async fn handle_notification_inner( ) where M: NotificationRpc + 'static, St: RpcStream + 'static, - H: FnOnce(S, M::Payload) -> HF, + H: FnOnce(S, Context, M::Payload) -> HF, HF: Future, E: FnOnce(&S, &RpcError), { @@ -44,5 +45,5 @@ pub(crate) async fn handle_notification_inner( }; writer.reset(ResetCode::CANCELLED); - handle(state, notification).await; + handle(state, context, notification).await; } diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index e396d679..eabad230 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -6,7 +6,8 @@ use crate::{ finish_bytes, progress::{encode_progress, encode_response, Progress}, rpc::read_framed_request, - write_bytes, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, + write_bytes, Context, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, + RpcWrite, }; #[trait_variant::make(ProgressHandler: Send)] @@ -15,7 +16,12 @@ where M: Progress, St: RpcStream, { - async fn handle(self, request: M::Request, responder: ProgressResponder); + async fn handle( + self, + context: Context, + request: M::Request, + responder: ProgressResponder, + ); fn handle_error(&self, _error: &RpcError) {} } @@ -62,6 +68,7 @@ where pub(crate) async fn handle_progress_inner( state: S, + context: Context, config: RouterConfig, mut reader: St::Reader, writer: St::Writer, @@ -70,7 +77,7 @@ pub(crate) async fn handle_progress_inner( ) where M: Progress + 'static, St: RpcStream + 'static, - H: FnOnce(S, M::Request, ProgressResponder) -> HF, + H: FnOnce(S, Context, M::Request, ProgressResponder) -> HF, HF: Future, E: FnOnce(&S, &RpcError), { @@ -87,5 +94,5 @@ pub(crate) async fn handle_progress_inner( } }; - handle(state, request, ProgressResponder::new(writer)).await; + handle(state, context, request, ProgressResponder::new(writer)).await; } diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs index ba810e7b..d57365bf 100644 --- a/ql-rpc/src/rpc/request/server.rs +++ b/ql-rpc/src/rpc/request/server.rs @@ -3,7 +3,7 @@ use std::{future::Future, marker::PhantomData}; use bytes::Bytes; use crate::{ - finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, + finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, Context, DropResetWrite, ResetCode, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, }; @@ -13,7 +13,12 @@ where M: RequestRpc, St: RpcStream, { - async fn handle(self, message: M::Request, responder: Response); + async fn handle( + self, + context: Context, + message: M::Request, + responder: Response, + ); fn handle_error(&self, _error: &RpcError) {} } @@ -54,6 +59,7 @@ where pub(crate) async fn handle_request_inner( state: S, + context: Context, config: RouterConfig, mut reader: St::Reader, writer: St::Writer, @@ -62,7 +68,7 @@ pub(crate) async fn handle_request_inner( ) where M: RequestRpc + 'static, St: RpcStream + 'static, - H: FnOnce(S, M::Request, Response) -> HF, + H: FnOnce(S, Context, M::Request, Response) -> HF, HF: Future, E: FnOnce(&S, &RpcError), { @@ -79,5 +85,5 @@ pub(crate) async fn handle_request_inner( } }; - handle(state, request, Response::new(writer)).await; + handle(state, context, request, Response::new(writer)).await; } diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs index f6dcbd54..5f6a2ff4 100644 --- a/ql-rpc/src/rpc/subscription/server.rs +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -4,8 +4,8 @@ use bytes::Bytes; use crate::{ codec, finish_bytes, rpc::read_eof_request, subscription::Subscription as SubscriptionRpc, - write_bytes, DropResetWrite, ResetCode, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, - RpcWrite, + write_bytes, Context, DropResetWrite, ResetCode, RouterConfig, RpcCodec, RpcError, RpcRead, + RpcStream, RpcWrite, }; #[trait_variant::make(SubscriptionHandler: Send)] @@ -16,6 +16,7 @@ where { async fn handle( self, + context: Context, message: M::Request, responder: SubscriptionResponder, ); @@ -62,6 +63,7 @@ where pub(crate) async fn handle_subscription_inner( state: S, + context: Context, config: RouterConfig, mut reader: St::Reader, writer: St::Writer, @@ -70,7 +72,7 @@ pub(crate) async fn handle_subscription_inner( ) where M: SubscriptionRpc + 'static, St: RpcStream + 'static, - H: FnOnce(S, M::Request, SubscriptionResponder) -> HF, + H: FnOnce(S, Context, M::Request, SubscriptionResponder) -> HF, HF: Future, E: FnOnce(&S, &RpcError), { @@ -87,5 +89,5 @@ pub(crate) async fn handle_subscription_inner( } }; - handle(state, request, SubscriptionResponder::new(writer)).await; + handle(state, context, request, SubscriptionResponder::new(writer)).await; } diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index 82245494..8ea9287a 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -8,7 +8,8 @@ use crate::{ parts::{FrameKind, PartFrameReader, PartReadStep}, read_framed_request_prefix, }, - DropResetRead, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, Upload, + Context, DropResetRead, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, + Upload, }; #[trait_variant::make(UploadHandler: Send)] @@ -19,6 +20,7 @@ where { async fn handle( self, + context: Context, request: M::Request, upload: UploadReader, responder: UploadResponder, @@ -184,6 +186,7 @@ where pub(crate) async fn handle_upload_inner( state: S, + context: Context, config: RouterConfig, mut reader: St::Reader, writer: St::Writer, @@ -194,6 +197,7 @@ pub(crate) async fn handle_upload_inner( St: RpcStream + 'static, H: FnOnce( S, + Context, M::Request, UploadReader, UploadResponder, @@ -217,6 +221,7 @@ pub(crate) async fn handle_upload_inner( handle( state, + context, request, UploadReader { stream: DropResetRead::new(reader), diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index 4952e192..37d14d95 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -5,15 +5,17 @@ use std::{ use bytes::Bytes; -use crate::{ResetCode, RouteId, ServiceId}; +use crate::{ResetCode, RouteId, ServiceId, StreamId, QID}; pub trait RpcStream { type Error; type Reader: RpcRead; type Writer: RpcWrite; - fn service_id(&self) -> Option; - fn route_id(&self) -> Option; + fn qid(&self) -> QID; + fn stream_id(&self) -> StreamId; + fn service_id(&self) -> ServiceId; + fn route_id(&self) -> RouteId; fn split(self) -> (Self::Reader, Self::Writer); } diff --git a/ql-runtime/src/command.rs b/ql-runtime/src/command.rs index d90f787b..a27a0888 100644 --- a/ql-runtime/src/command.rs +++ b/ql-runtime/src/command.rs @@ -21,7 +21,7 @@ pub enum Command { OpenStream { service_id: ServiceId, route_id: RouteId, - start: oneshot::Sender>, + start: oneshot::Sender>, }, PollInbound { stream_id: StreamId, diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index a092d578..5841d020 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -21,9 +21,8 @@ use ql_wire::{ResetCode, ResetOrigin, ResetTarget, StreamHeader, StreamId}; use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; use crate::{ command::Command, - handle::QlStream, io, log, - platform::{QlInbound, QlPlatform, QlTimer}, + platform::{QlInbound, QlInboundStream, QlPlatform, QlTimer}, QlStreamError, Runtime, RuntimeHandle, }; @@ -249,14 +248,13 @@ impl DriverState { Some(InboundIo::new(reader_io)), ), ); - if start.send(Ok((stream_id, reader, writer))).is_err() { + if start.send(Ok((reader, writer))).is_err() { log::warn!("open stream cancelled before delivery: stream_id={stream_id}"); if let Some(stream) = self.streams.get_mut(&stream_id) { stream.inbound_close(); stream.outbound_close(); } stream_ops.reset(ResetTarget::Both, ResetCode::DROPPED); - drop(stream_ops); return; } drop(stream_ops); @@ -389,6 +387,7 @@ impl DriverState { ), ); + let qid = fsm.peer().unwrap().qid; let stream = fsm.stream(stream_id).unwrap(); let StreamHeader { service_id, @@ -399,10 +398,11 @@ impl DriverState { "delivering inbound stream to platform: service_id={service_id} route_id={route_id} stream_id={stream_id}", ); - platform.handle_inbound(QlStream { - stream_id, + platform.handle_inbound(QlInboundStream { + qid, route_id, service_id, + stream_id, writer, reader, }); diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs index 878682b3..e3c40987 100644 --- a/ql-runtime/src/driver/test.rs +++ b/ql-runtime/src/driver/test.rs @@ -39,7 +39,7 @@ impl QlPlatform for NoopCrypto { fn handle_peer_status(&self, _peer: Option, _status: ql_fsm::PeerStatus) {} - fn handle_inbound(&self, _event: QlStream) {} + fn handle_inbound(&self, _event: QlInboundStream) {} } impl QlInbound for NoopInbound { diff --git a/ql-runtime/src/handle/mod.rs b/ql-runtime/src/handle/mod.rs index 4275d277..af9c3e0b 100644 --- a/ql-runtime/src/handle/mod.rs +++ b/ql-runtime/src/handle/mod.rs @@ -1,15 +1,11 @@ use ql_fsm::{NoSessionError, OpenStreamParams, PairingInvite}; -use ql_wire::{PairingToken, PeerBundle, RouteId, ServiceId, SessionCloseCode, StreamId}; +use ql_wire::{PairingToken, PeerBundle, SessionCloseCode}; use crate::command::Command; pub use crate::io::{StreamReader, StreamWriter}; #[derive(Debug)] pub struct QlStream { - pub service_id: ServiceId, - pub route_id: RouteId, - pub stream_id: StreamId, - pub writer: StreamWriter, pub reader: StreamReader, } @@ -70,15 +66,9 @@ impl RuntimeHandle { }); // runtime cannot be shutdown while we have a handle - let (stream_id, reader, writer) = start_rx.await.unwrap()?; + let (reader, writer) = start_rx.await.unwrap()?; - Ok(QlStream { - route_id, - service_id, - stream_id, - writer, - reader, - }) + Ok(QlStream { writer, reader }) } #[cfg(feature = "rpc")] diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs index 8647d4ee..c9068fe6 100644 --- a/ql-runtime/src/io/reader.rs +++ b/ql-runtime/src/io/reader.rs @@ -4,7 +4,7 @@ use std::{ }; use bytes::Bytes; -use ql_wire::{ResetCode, ResetTarget}; +use ql_wire::{ResetCode, ResetTarget, StreamId}; use super::{ inner::{Item, RxInner}, @@ -50,6 +50,10 @@ impl StreamReader { } } + pub fn stream_id(&self) -> StreamId { + self.rx.stream_id() + } + pub fn poll_read( &mut self, cx: &mut Context<'_>, diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs index 4176f46f..e985608f 100644 --- a/ql-runtime/src/io/writer.rs +++ b/ql-runtime/src/io/writer.rs @@ -4,7 +4,7 @@ use std::{ }; use bytes::Bytes; -use ql_wire::{ResetCode, ResetTarget}; +use ql_wire::{ResetCode, ResetTarget, StreamId}; use super::{ inner::{Item, TxInner}, @@ -49,6 +49,10 @@ impl StreamWriter { } } + pub fn stream_id(&self) -> StreamId { + self.tx.stream_id() + } + pub fn poll_write( &mut self, bytes: &mut Bytes, diff --git a/ql-runtime/src/platform.rs b/ql-runtime/src/platform.rs index 331bfe7a..24a9108d 100644 --- a/ql-runtime/src/platform.rs +++ b/ql-runtime/src/platform.rs @@ -6,9 +6,19 @@ use std::{ }; use ql_fsm::{PeerStatus, ReceiveError}; -use ql_wire::{PeerBundle, QlCrypto, QID}; - -use crate::QlStream; +use ql_wire::{PeerBundle, QlCrypto, RouteId, ServiceId, StreamId, QID}; + +use crate::{StreamReader, StreamWriter}; + +#[derive(Debug)] +pub struct QlInboundStream { + pub qid: QID, + pub service_id: ServiceId, + pub route_id: RouteId, + pub stream_id: StreamId, + pub writer: StreamWriter, + pub reader: StreamReader, +} pub trait QlTimer { fn set_deadline(self: Pin<&mut Self>, deadline: Option); @@ -38,6 +48,6 @@ pub trait QlPlatform: QlCrypto { fn persist_peer(&self, peer: PeerBundle); fn handle_peer_status(&self, peer: Option, status: PeerStatus); - fn handle_inbound(&self, event: QlStream); + fn handle_inbound(&self, event: QlInboundStream); fn handle_recv_error(&self, _error: ReceiveError) {} } diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index 93b9f361..93a220a5 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -1,21 +1,29 @@ -use std::task::{Context, Poll}; +use std::task::{Context as TaskContext, Poll}; use bytes::Bytes; -use ql_rpc::{ResetCode, RouteId, RpcRead, RpcStream, RpcWrite, ServiceId}; +use ql_rpc::{ResetCode, RouteId, RpcRead, RpcStream, RpcWrite, ServiceId, StreamId, QID}; -use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; +use crate::{QlInboundStream, QlStreamError, StreamReader, StreamWriter}; -impl RpcStream for QlStream { +impl RpcStream for QlInboundStream { type Error = QlStreamError; type Reader = StreamReader; type Writer = StreamWriter; - fn service_id(&self) -> Option { - Some(self.service_id) + fn qid(&self) -> QID { + self.qid } - fn route_id(&self) -> Option { - Some(self.route_id) + fn stream_id(&self) -> StreamId { + self.stream_id + } + + fn service_id(&self) -> ServiceId { + self.service_id + } + + fn route_id(&self) -> RouteId { + self.route_id } fn split(self) -> (Self::Reader, Self::Writer) { @@ -26,7 +34,10 @@ impl RpcStream for QlStream { impl RpcRead for StreamReader { type Error = QlStreamError; - fn poll_read(&mut self, cx: &mut Context<'_>) -> Poll, QlStreamError>> { + fn poll_read( + &mut self, + cx: &mut TaskContext<'_>, + ) -> Poll, QlStreamError>> { StreamReader::poll_read(self, cx) } @@ -41,7 +52,7 @@ impl RpcWrite for StreamWriter { fn poll_write( &mut self, bytes: &mut Bytes, - cx: &mut Context<'_>, + cx: &mut TaskContext<'_>, ) -> Poll> { StreamWriter::poll_write(self, bytes, cx) } @@ -50,7 +61,7 @@ impl RpcWrite for StreamWriter { StreamWriter::queue_finish(self); } - fn poll_finish(&mut self, cx: &mut Context<'_>) -> Poll> { + fn poll_finish(&mut self, cx: &mut TaskContext<'_>) -> Poll> { StreamWriter::poll_finish(self, cx) } diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index 814511d2..19c2ad37 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -20,7 +20,7 @@ use ql_wire::{ use tokio::{task::LocalSet, time::Sleep}; use crate::{ - new_runtime, platform::QlTimer, NoSessionError, PairingInvite, QlFsmConfig, QlStream, + new_runtime, platform::QlTimer, NoSessionError, PairingInvite, QlFsmConfig, QlInboundStream, QlStreamError, RuntimeConfig, RuntimeHandle, }; @@ -97,7 +97,7 @@ struct TestPlatform { _inbound_messages_tx: Sender>, inbound_messages: Option>>, status: Sender, - inbound: Option>, + inbound: Option>, crypto: SoftwareCrypto, encrypted_write_counter: AtomicUsize, fail_encrypted_write_at: Option, @@ -121,7 +121,7 @@ type TestPlatformPartsWithInbound = ( Receiver>, Sender>, Receiver, - Receiver, + Receiver, ); impl TestPlatform { @@ -151,7 +151,7 @@ impl TestPlatform { } fn new_inner( - inbound: Option>, + inbound: Option>, fail_encrypted_write_at: Option, write_delay: Duration, write_stats: Option, @@ -183,7 +183,7 @@ struct TestSide { handle: RuntimeHandle, status: Receiver, peer: QID, - inbound: Receiver, + inbound: Receiver, } struct TestPair { @@ -310,7 +310,7 @@ impl TestPair { .await; } - fn take_inbound(&mut self, side: Side) -> Receiver { + fn take_inbound(&mut self, side: Side) -> Receiver { let replacement = async_channel::unbounded().1; std::mem::replace(&mut self.side_mut(side).inbound, replacement) } @@ -450,7 +450,7 @@ impl crate::platform::QlPlatform for TestPlatform { let _ = self.status.try_send(StatusEvent { peer, status }); } - fn handle_inbound(&self, event: QlStream) { + fn handle_inbound(&self, event: QlInboundStream) { if let Some(tx) = &self.inbound { let _ = tx.try_send(event); } diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index 1bd13b22..00e3d749 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -10,7 +10,7 @@ use std::{ use bytes::Bytes; use futures_lite::StreamExt; use ql_rpc::{ - DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, + Context, DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, NotificationHandlerLocal, ProgressHandlerLocal, ProgressResponder, RequestHandler, RequestHandlerLocal, ResetCode, ResetOrigin, Response, RouteId, SendSpawner, ServiceId, Spawner, SubscriptionHandlerLocal, SubscriptionResponder, UploadHandlerLocal, UploadReader, @@ -18,7 +18,7 @@ use ql_rpc::{ }; use super::*; -use crate::{rpc::RpcError, QlStream, StreamWriter}; +use crate::{rpc::RpcError, QlInboundStream, StreamWriter}; const TEST_SERVICE: ServiceId = ServiceId([7; 16]); @@ -155,8 +155,13 @@ async fn rpc_request() { seen: Arc>>, } - impl RequestHandler for RouterState { - async fn handle(self, request: String, response: Response) { + impl RequestHandler for RouterState { + async fn handle( + self, + _context: Context, + request: String, + response: Response, + ) { let seen = self.seen.clone(); seen.lock().unwrap().push(request); let _ = response.respond("world".into()).await; @@ -170,7 +175,7 @@ async fn rpc_request() { let seen = Arc::new(Mutex::new(Vec::new())); let router = - ql_rpc::Router::<_, QlStream, TokioSendSpawner>::builder_send(TokioSendSpawner) + ql_rpc::Router::<_, QlInboundStream, TokioSendSpawner>::builder_send(TokioSendSpawner) .request::() .build(RouterState { seen: seen.clone() }); @@ -206,8 +211,8 @@ async fn rpc_notification() { seen: Rc>>>, } - impl NotificationHandlerLocal for RouterState { - async fn handle(self, payload: Vec) { + impl NotificationHandlerLocal for RouterState { + async fn handle(self, _context: Context, payload: Vec) { self.seen.borrow_mut().push(payload); } } @@ -218,10 +223,11 @@ async fn rpc_notification() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .notification::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( + TokioLocalSpawner, + ) + .notification::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let inbound = inbound_b.recv().await.unwrap(); @@ -251,9 +257,10 @@ async fn rpc_subscrption() { seen: Rc>>>, } - impl SubscriptionHandlerLocal for RouterState { + impl SubscriptionHandlerLocal for RouterState { async fn handle( self, + _context: Context, request: Vec, mut response: SubscriptionResponder, StreamWriter>, ) { @@ -271,10 +278,11 @@ async fn rpc_subscrption() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .subscription::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( + TokioLocalSpawner, + ) + .subscription::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let inbound = inbound_b.recv().await.unwrap(); @@ -303,8 +311,13 @@ async fn rpc_router_enforces_max_request_bytes() { #[derive(Clone)] struct LimitedState; - impl RequestHandlerLocal for LimitedState { - async fn handle(self, request: String, response: Response) { + impl RequestHandlerLocal for LimitedState { + async fn handle( + self, + _context: Context, + request: String, + response: Response, + ) { let _ = response.respond(request).await; } } @@ -313,11 +326,12 @@ async fn rpc_router_enforces_max_request_bytes() { let mut pair = TestPair::new(default_runtime_config()); pair.connect_and_wait(Side::A).await; let inbound_b = pair.take_inbound(Side::B); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .max_request_bytes(4) - .request::() - .build(LimitedState); + let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( + TokioLocalSpawner, + ) + .max_request_bytes(4) + .request::() + .build(LimitedState); let responder = tokio::task::spawn_local(async move { let inbound = inbound_b.recv().await.unwrap(); @@ -349,9 +363,10 @@ async fn rpc_progress() { seen: Rc>>>, } - impl ProgressHandlerLocal for RouterState { + impl ProgressHandlerLocal for RouterState { async fn handle( self, + _context: Context, request: Vec, mut responder: ProgressResponder, ) { @@ -369,10 +384,11 @@ async fn rpc_progress() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .progress::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( + TokioLocalSpawner, + ) + .progress::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let inbound = inbound_b.recv().await.unwrap(); @@ -405,9 +421,10 @@ async fn rpc_download() { seen: Rc>>>, } - impl DownloadHandlerLocal for RouterState { + impl DownloadHandlerLocal for RouterState { async fn handle( self, + _context: Context, request: Vec, download: DownloadStart, ) { @@ -431,10 +448,11 @@ async fn rpc_download() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .download::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( + TokioLocalSpawner, + ) + .download::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let inbound = inbound_b.recv().await.unwrap(); @@ -490,9 +508,10 @@ async fn rpc_download_complete() { seen: Rc>>>, } - impl DownloadHandlerLocal for RouterState { + impl DownloadHandlerLocal for RouterState { async fn handle( self, + _context: Context, request: Vec, download: DownloadStart, ) { @@ -507,10 +526,11 @@ async fn rpc_download_complete() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .download::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( + TokioLocalSpawner, + ) + .download::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let inbound = inbound_b.recv().await.unwrap(); @@ -545,9 +565,10 @@ async fn rpc_upload() { uploads: Rc>>>, } - impl UploadHandlerLocal for RouterState { + impl UploadHandlerLocal for RouterState { async fn handle( self, + _context: Context, request: Vec, mut upload: UploadReader, responder: UploadResponder, StreamWriter>, @@ -578,13 +599,14 @@ async fn rpc_upload() { let requests = Rc::new(RefCell::new(Vec::new())); let uploads = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .upload::() - .build(RouterState { - requests: requests.clone(), - uploads: uploads.clone(), - }); + let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( + TokioLocalSpawner, + ) + .upload::() + .build(RouterState { + requests: requests.clone(), + uploads: uploads.clone(), + }); let responder = tokio::task::spawn_local(async move { let inbound = inbound_b.recv().await.unwrap(); @@ -626,8 +648,12 @@ async fn rpc_duplex() { seen: Rc>>>, } - impl DuplexHandlerLocal for RouterState { - async fn handle(self, mut peer: DuplexPeer) { + impl DuplexHandlerLocal for RouterState { + async fn handle( + self, + _context: Context, + mut peer: DuplexPeer, + ) { let seen = self.seen.clone(); let first = peer.receiver.next_event().await.unwrap().unwrap(); seen.borrow_mut().push(first); @@ -650,10 +676,11 @@ async fn rpc_duplex() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .duplex::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( + TokioLocalSpawner, + ) + .duplex::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let inbound = inbound_b.recv().await.unwrap(); From d98a9197629a7b4e68c6aef602df096bb187e066 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 24 Jun 2026 10:10:25 -0400 Subject: [PATCH 22/59] ql-runtime: better handle --- ql-runtime/src/driver/mod.rs | 21 +++------------------ ql-runtime/src/driver/state.rs | 2 +- ql-runtime/src/driver/test.rs | 12 ++++++------ ql-runtime/src/handle/mod.rs | 20 +++++++++++++++----- ql-runtime/src/io/mod.rs | 8 ++++---- ql-runtime/src/io/reader.rs | 18 +++++++++++------- ql-runtime/src/io/sync.rs | 6 +++--- ql-runtime/src/io/writer.rs | 16 ++++++++++------ ql-runtime/src/lib.rs | 7 ++++--- 9 files changed, 57 insertions(+), 53 deletions(-) diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 5841d020..125b5ba3 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -23,7 +23,7 @@ use crate::{ command::Command, io, log, platform::{QlInbound, QlInboundStream, QlPlatform, QlTimer}, - QlStreamError, Runtime, RuntimeHandle, + QlStreamError, Runtime, }; impl Runtime

{ @@ -215,11 +215,6 @@ impl DriverState { start, } => { log::info!("open stream requested: route_id={route_id}"); - let Some(runtime_tx) = self.runtime_tx.upgrade() else { - log::warn!("open stream aborted: runtime channel unavailable"); - let _ = start.send(Err(ql_fsm::NoSessionError)); - return; - }; let mut stream_ops = match fsm.open_stream(ql_fsm::OpenStreamParams { service_id, @@ -238,7 +233,7 @@ impl DriverState { stream_id, ResetTarget::Return, ResetTarget::Origin, - RuntimeHandle::new(runtime_tx), + self.runtime_tx.clone(), ); self.streams.insert( stream_id, @@ -361,21 +356,11 @@ impl DriverState { platform: &P, stream_id: StreamId, ) { - let Some(runtime_tx) = self.runtime_tx.upgrade() else { - log::warn!( - "dropping inbound stream because handle channel is unavailable: stream_id={stream_id}" - ); - if let Ok(mut stream) = fsm.stream(stream_id) { - stream.reset(ResetTarget::Both, ResetCode::DISCONNECTED); - } - return; - }; - let (reader, writer, reader_io, writer_io) = io::new_stream( stream_id, ResetTarget::Origin, ResetTarget::Return, - RuntimeHandle::new(runtime_tx), + self.runtime_tx.clone(), ); self.streams.insert( diff --git a/ql-runtime/src/driver/state.rs b/ql-runtime/src/driver/state.rs index ccd57765..b2803343 100644 --- a/ql-runtime/src/driver/state.rs +++ b/ql-runtime/src/driver/state.rs @@ -11,7 +11,7 @@ use crate::{ pub struct DriverState { pub streams: HashMap, - pub runtime_tx: async_channel::WeakSender, + pub runtime_tx: async_channel::Sender, pub max_concurrent_message_writes: usize, } diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs index e3c40987..69bf68e0 100644 --- a/ql-runtime/src/driver/test.rs +++ b/ql-runtime/src/driver/test.rs @@ -53,7 +53,7 @@ fn new_driver_state() -> (DriverState, QlFsm) { ( DriverState { streams: HashMap::new(), - runtime_tx: runtime_tx.downgrade(), + runtime_tx, max_concurrent_message_writes: 1, }, QlFsm::new( @@ -71,7 +71,7 @@ fn new_inbound_io(capacity: usize) -> InboundIo { StreamId(99u32.into()), ResetTarget::Origin, ResetTarget::Return, - RuntimeHandle::new(runtime_tx), + runtime_tx, ); let (_, _, reader_io, _) = stream; InboundIo::new(reader_io) @@ -83,7 +83,7 @@ fn new_outbound_io() -> OutboundIo { StreamId(100u32.into()), ResetTarget::Return, ResetTarget::Origin, - RuntimeHandle::new(runtime_tx), + runtime_tx, ); let (_, _, _, writer_io) = stream; OutboundIo::new(writer_io) @@ -132,7 +132,7 @@ fn poll_stream_keeps_outbound_pending_after_local_finish_when_inbound_is_closed( stream_id, ResetTarget::Return, ResetTarget::Origin, - RuntimeHandle::new(runtime_tx), + runtime_tx, ); writer.queue_finish(); state.streams.insert( @@ -156,7 +156,7 @@ fn local_reset_command_reaps_when_other_half_is_already_closed() { stream_id, ResetTarget::Return, ResetTarget::Origin, - RuntimeHandle::new(runtime_tx), + runtime_tx, ); state.streams.insert( @@ -187,7 +187,7 @@ fn unpaired_status_fails_and_reaps_all_streams() { stream_id, ResetTarget::Origin, ResetTarget::Return, - RuntimeHandle::new(runtime_tx), + runtime_tx, ); state.streams.insert( diff --git a/ql-runtime/src/handle/mod.rs b/ql-runtime/src/handle/mod.rs index af9c3e0b..d9ba3f60 100644 --- a/ql-runtime/src/handle/mod.rs +++ b/ql-runtime/src/handle/mod.rs @@ -1,3 +1,5 @@ +use std::sync::Arc; + use ql_fsm::{NoSessionError, OpenStreamParams, PairingInvite}; use ql_wire::{PairingToken, PeerBundle, SessionCloseCode}; @@ -12,7 +14,7 @@ pub struct QlStream { #[derive(Clone)] pub struct RuntimeHandle { - tx: async_channel::Sender, + inner: Arc, } impl RuntimeHandle { @@ -79,16 +81,24 @@ impl RuntimeHandle { impl RuntimeHandle { pub(crate) fn new(tx: async_channel::Sender) -> Self { - Self { tx } + Self { + inner: Arc::new(Inner { tx }), + } } #[inline] #[track_caller] pub(crate) fn send(&self, cmd: Command) { - self.tx.try_send(cmd).expect("runtime is alive"); + self.inner.tx.try_send(cmd).expect("runtime is alive"); } +} + +struct Inner { + tx: async_channel::Sender, +} - pub(crate) fn try_send(&self, cmd: Command) -> bool { - self.tx.try_send(cmd).is_ok() +impl Drop for Inner { + fn drop(&mut self) { + self.tx.close(); } } diff --git a/ql-runtime/src/io/mod.rs b/ql-runtime/src/io/mod.rs index 2fc4064e..52dba870 100644 --- a/ql-runtime/src/io/mod.rs +++ b/ql-runtime/src/io/mod.rs @@ -9,7 +9,7 @@ use std::ops::Deref; use ql_wire::{ResetTarget, StreamId}; pub use self::{reader::StreamReader, slot::PushError, writer::StreamWriter}; -use crate::RuntimeHandle; +use crate::command::Command; pub struct Rx(sync::Arc); @@ -47,12 +47,12 @@ pub fn new_stream( stream_id: StreamId, reader_target: ResetTarget, writer_target: ResetTarget, - handle: RuntimeHandle, + runtime_tx: async_channel::Sender, ) -> (StreamReader, StreamWriter, Rx, Tx) { let shared = inner::new(stream_id); ( - StreamReader::new(Rx(shared.clone()), reader_target, handle.clone()), - StreamWriter::new(Tx(shared.clone()), writer_target, handle), + StreamReader::new(Rx(shared.clone()), reader_target, runtime_tx.clone()), + StreamWriter::new(Tx(shared.clone()), writer_target, runtime_tx), Rx(shared.clone()), Tx(shared), ) diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs index c9068fe6..591c8cad 100644 --- a/ql-runtime/src/io/reader.rs +++ b/ql-runtime/src/io/reader.rs @@ -11,13 +11,13 @@ use super::{ slot::PopError, Rx, }; -use crate::{command::Command, log, QlStreamError, RuntimeHandle}; +use crate::{command::Command, log, QlStreamError}; pub struct StreamReader { rx: Rx, target: ResetTarget, terminal: ReaderTerminalState, - handle: RuntimeHandle, + runtime_tx: async_channel::Sender, } enum ReaderTerminalState { @@ -41,12 +41,16 @@ impl std::fmt::Debug for StreamReader { } impl StreamReader { - pub(crate) fn new(shared: Rx, target: ResetTarget, handle: RuntimeHandle) -> Self { + pub(crate) fn new( + shared: Rx, + target: ResetTarget, + runtime_tx: async_channel::Sender, + ) -> Self { Self { rx: shared, target, terminal: ReaderTerminalState::Open, - handle, + runtime_tx, } } @@ -87,7 +91,7 @@ impl StreamReader { self.target, bytes.len() ); - self.handle.try_send(Command::PollInbound { + let _ = self.runtime_tx.try_send(Command::PollInbound { stream_id: self.rx.stream_id(), }); Poll::Ready(Ok(Some(bytes))) @@ -136,7 +140,7 @@ impl StreamReader { code ); self.terminal = ReaderTerminalState::Delivered; - self.handle.try_send(Command::ResetStream { + let _ = self.runtime_tx.try_send(Command::ResetStream { stream_id: self.rx.stream_id(), target: self.target, code, @@ -155,7 +159,7 @@ impl Drop for StreamReader { self.target, ResetCode::DROPPED ); - self.handle.try_send(Command::ResetStream { + let _ = self.runtime_tx.try_send(Command::ResetStream { stream_id: self.rx.stream_id(), target: self.target, code: ResetCode::DROPPED, diff --git a/ql-runtime/src/io/sync.rs b/ql-runtime/src/io/sync.rs index c5034076..bc06d474 100644 --- a/ql-runtime/src/io/sync.rs +++ b/ql-runtime/src/io/sync.rs @@ -71,7 +71,7 @@ pub(crate) mod loom { use ql_wire::StreamId; use super::Arc; - use crate::{io::inner::Inner, RuntimeHandle}; + use crate::{command::Command, io::inner::Inner}; pub(crate) fn check_model(f: impl Fn() + Sync + Send + 'static) { let builder = model::Builder::new(); @@ -82,8 +82,8 @@ pub(crate) mod loom { crate::io::inner::new(StreamId(1u32.into())) } - pub(crate) fn handle() -> RuntimeHandle { + pub(crate) fn handle() -> async_channel::Sender { let (tx, _rx) = async_channel::unbounded(); - RuntimeHandle::new(tx) + tx } } diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs index e985608f..c124b017 100644 --- a/ql-runtime/src/io/writer.rs +++ b/ql-runtime/src/io/writer.rs @@ -11,14 +11,14 @@ use super::{ slot::PopError, PushError, Tx, }; -use crate::{command::Command, log, QlStreamError, RuntimeHandle}; +use crate::{command::Command, log, QlStreamError}; pub struct StreamWriter { tx: Tx, target: ResetTarget, open: bool, terminal: WriterTerminalState, - handle: RuntimeHandle, + runtime_tx: async_channel::Sender, } enum WriterTerminalState { @@ -39,13 +39,17 @@ impl std::fmt::Debug for StreamWriter { } impl StreamWriter { - pub(crate) fn new(shared: Tx, target: ResetTarget, handle: RuntimeHandle) -> Self { + pub(crate) fn new( + shared: Tx, + target: ResetTarget, + runtime_tx: async_channel::Sender, + ) -> Self { Self { tx: shared, target, open: true, terminal: WriterTerminalState::Pending, - handle, + runtime_tx, } } @@ -148,7 +152,7 @@ impl StreamWriter { } fn poll_runtime(&self) { - self.handle.try_send(Command::PollStream { + let _ = self.runtime_tx.try_send(Command::PollStream { stream_id: self.tx.stream_id(), }); } @@ -209,7 +213,7 @@ impl StreamWriter { self.target, code ); - self.handle.try_send(Command::ResetStream { + let _ = self.runtime_tx.try_send(Command::ResetStream { stream_id: self.tx.stream_id(), target: self.target, code, diff --git a/ql-runtime/src/lib.rs b/ql-runtime/src/lib.rs index 33783456..40df8dee 100644 --- a/ql-runtime/src/lib.rs +++ b/ql-runtime/src/lib.rs @@ -38,7 +38,7 @@ pub struct Runtime

{ platform: P, config: RuntimeConfig, rx: async_channel::Receiver, - tx: async_channel::WeakSender, + tx: async_channel::Sender, } pub fn new_runtime

( @@ -50,14 +50,15 @@ where P: QlPlatform, { let (tx, rx) = async_channel::unbounded(); + let handle = RuntimeHandle::new(tx.clone()); ( Runtime { identity, platform, config, rx, - tx: tx.downgrade(), + tx, }, - RuntimeHandle::new(tx), + handle, ) } From c6143351b9f5c59676e2df932ec73d8c512de8cf Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 24 Jun 2026 11:57:50 -0400 Subject: [PATCH 23/59] ql-runtime: remove rpc wrappers --- ql-runtime/src/error.rs | 7 ++ ql-runtime/src/rpc/download.rs | 67 ------------ ql-runtime/src/rpc/duplex.rs | 59 ---------- ql-runtime/src/rpc/error.rs | 73 ------------- ql-runtime/src/rpc/mod.rs | 166 +++++++++++++---------------- ql-runtime/src/rpc/progress.rs | 50 --------- ql-runtime/src/rpc/subscription.rs | 43 -------- ql-runtime/src/rpc/upload.rs | 44 -------- ql-runtime/src/tests/rpc.rs | 23 ++-- 9 files changed, 95 insertions(+), 437 deletions(-) delete mode 100644 ql-runtime/src/rpc/download.rs delete mode 100644 ql-runtime/src/rpc/duplex.rs delete mode 100644 ql-runtime/src/rpc/error.rs delete mode 100644 ql-runtime/src/rpc/progress.rs delete mode 100644 ql-runtime/src/rpc/subscription.rs delete mode 100644 ql-runtime/src/rpc/upload.rs diff --git a/ql-runtime/src/error.rs b/ql-runtime/src/error.rs index 0cd845ae..4bdb1d50 100644 --- a/ql-runtime/src/error.rs +++ b/ql-runtime/src/error.rs @@ -1,3 +1,4 @@ +use ql_fsm::NoSessionError; use ql_wire::{ResetCode, ResetOrigin}; #[derive(Debug, Clone, PartialEq, Eq)] @@ -19,3 +20,9 @@ impl std::fmt::Display for QlStreamError { } impl std::error::Error for QlStreamError {} + +impl From for QlStreamError { + fn from(_: NoSessionError) -> Self { + Self::NoSession + } +} diff --git a/ql-runtime/src/rpc/download.rs b/ql-runtime/src/rpc/download.rs deleted file mode 100644 index cfcc25f6..00000000 --- a/ql-runtime/src/rpc/download.rs +++ /dev/null @@ -1,67 +0,0 @@ -use bytes::Bytes; -use ql_rpc::download::Download as DownloadRpc; - -use super::RpcError; -use crate::StreamReader; - -pub struct DownloadCall { - pub(super) inner: ql_rpc::download::DownloadCall, -} - -pub struct DownloadReader { - pub(super) inner: ql_rpc::download::DownloadReader, -} - -pub struct DownloadPart<'a, M: DownloadRpc> { - inner: ql_rpc::download::DownloadPart<'a, M, StreamReader>, -} - -impl DownloadCall -where - M: DownloadRpc, -{ - pub async fn start(self) -> Result<(M::ResponseHeader, DownloadReader), RpcError> { - let (header, inner) = self.inner.start().await?; - Ok((header, DownloadReader { inner })) - } - - pub fn reset(self, code: ql_wire::ResetCode) { - self.inner.reset(ql_rpc::ResetCode(code.0)); - } -} - -impl DownloadReader -where - M: DownloadRpc, -{ - pub async fn next_part( - &mut self, - ) -> Result)>, RpcError> { - Ok(self - .inner - .next_part() - .await? - .map(|(header, inner)| (header, DownloadPart { inner }))) - } - - pub async fn complete(self) -> Result<(), RpcError> { - self.inner.complete().await.map_err(RpcError::from) - } - - pub fn reset(self, code: ql_wire::ResetCode) { - self.inner.reset(ql_rpc::ResetCode(code.0)); - } -} - -impl DownloadPart<'_, M> -where - M: DownloadRpc, -{ - pub async fn read_chunk(&mut self) -> Result, RpcError> { - Ok(self.inner.read_chunk().await?) - } - - pub fn reset(self, code: ql_wire::ResetCode) { - self.inner.reset(ql_rpc::ResetCode(code.0)); - } -} diff --git a/ql-runtime/src/rpc/duplex.rs b/ql-runtime/src/rpc/duplex.rs deleted file mode 100644 index e5c52609..00000000 --- a/ql-runtime/src/rpc/duplex.rs +++ /dev/null @@ -1,59 +0,0 @@ -use futures_lite::future::poll_fn; -use ql_rpc::duplex::Duplex as DuplexRpc; - -use super::RpcError; -use crate::{QlStreamError, StreamReader, StreamWriter}; - -pub struct DuplexCall { - pub sender: DuplexSender, - pub receiver: DuplexReceiver, -} - -pub struct DuplexSender -where - T: ql_rpc::RpcCodec, -{ - pub(super) inner: ql_rpc::duplex::DuplexSender, -} - -pub struct DuplexReceiver -where - T: ql_rpc::RpcCodec, -{ - pub(super) inner: ql_rpc::duplex::DuplexReceiver, -} - -impl DuplexSender -where - T: ql_rpc::RpcCodec, -{ - pub async fn send(&mut self, event: &T) -> Result<(), QlStreamError> { - self.inner.send(event).await - } - - pub fn finish(self) { - self.inner.finish(); - } - - pub fn reset(self, code: ql_wire::ResetCode) { - self.inner.reset(ql_rpc::ResetCode(code.0)); - } -} - -impl DuplexReceiver -where - T: ql_rpc::RpcCodec, -{ - pub async fn next_event(&mut self) -> Option>> { - poll_fn(|cx| { - self.inner - .poll_next_event(cx) - .map(|item| item.map(|result| Ok(result?))) - }) - .await - } - - pub fn reset(self, code: ql_wire::ResetCode) { - self.inner.reset(ql_rpc::ResetCode(code.0)); - } -} diff --git a/ql-runtime/src/rpc/error.rs b/ql-runtime/src/rpc/error.rs deleted file mode 100644 index 965b711a..00000000 --- a/ql-runtime/src/rpc/error.rs +++ /dev/null @@ -1,73 +0,0 @@ -use ql_fsm::NoSessionError; - -use crate::QlStreamError; - -#[derive(Debug)] -pub enum RpcError { - NoSession, - Reset { - code: ql_rpc::ResetCode, - origin: ql_rpc::ResetOrigin, - }, - Protocol(ql_rpc::Error), - Codec(E), -} - -impl From for RpcError { - fn from(_: NoSessionError) -> Self { - Self::NoSession - } -} - -impl From for RpcError { - fn from(error: QlStreamError) -> Self { - match error { - QlStreamError::StreamReset { code, origin } => Self::Reset { code, origin }, - QlStreamError::NoSession => Self::NoSession, - } - } -} - -impl From for RpcError { - fn from(error: ql_rpc::Error) -> Self { - Self::Protocol(error) - } -} - -impl From> for RpcError { - fn from(error: ql_rpc::RpcError) -> Self { - match error { - ql_rpc::RpcError::Protocol(error) => Self::Protocol(error), - ql_rpc::RpcError::Codec(error) => Self::Codec(error), - ql_rpc::RpcError::Transport(error) => error.into(), - } - } -} - -impl std::fmt::Display for RpcError -where - E: std::fmt::Display, -{ - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - match self { - Self::NoSession => write!(f, "no session"), - Self::Reset { code, origin } => write!(f, "stream reset {code:?} ({origin:?})"), - Self::Protocol(error) => write!(f, "{error}"), - Self::Codec(error) => write!(f, "{error}"), - } - } -} - -impl std::error::Error for RpcError -where - E: std::error::Error + 'static, -{ - fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { - match self { - Self::Protocol(error) => Some(error), - Self::Codec(error) => Some(error), - RpcError::NoSession => None, - RpcError::Reset { .. } => None, - } - } -} diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index c1257b4c..b14b0b25 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -1,26 +1,12 @@ -pub use self::{download::*, duplex::*, error::*, progress::*, subscription::*, upload::*}; - mod adapter; -mod download; -mod duplex; -mod error; -mod progress; -mod subscription; -mod upload; use bytes::Bytes; use ql_fsm::OpenStreamParams; -use ql_rpc::{ - download::{self as rpc_download, Download as DownloadRpc}, - duplex::{self as rpc_duplex, Duplex as DuplexRpc}, - notification::{self, Notification}, - progress::{self as rpc_progress, Progress}, - request::{self, Request as RequestRpc}, - subscription::{self as rpc_subscription, Subscription as SubscriptionRpc}, - upload::{self as rpc_upload, Upload as UploadRpc}, -}; +use ql_rpc::{download, duplex, notification, progress, request, subscription, upload}; + +use crate::{QlStream, QlStreamError, RuntimeHandle, StreamReader, StreamWriter}; -use crate::{RuntimeHandle, StreamReader}; +type RpcResult = Result>; #[derive(Clone)] pub struct RpcHandle { @@ -28,108 +14,104 @@ pub struct RpcHandle { } impl RpcHandle { - pub async fn notification(&self, event: &M::Payload) -> Result<(), RpcError> + pub async fn notification(&self, event: &M::Payload) -> RpcResult<(), M::Error> where - M: Notification, + M: notification::Notification, { let mut payload = Vec::new(); notification::encode_notification::(event, &mut payload); - let mut stream = self - .inner - .open_stream(open_stream_params(M::SERVICE, M::ROUTE)) - .await?; + let mut stream = self.open_rpc_stream::().await?; stream.reader.reset(ql_rpc::ResetCode::CANCELLED); - stream.writer.write(Bytes::from(payload)).await?; - stream.writer.finish().await?; + stream + .writer + .write(Bytes::from(payload)) + .await + .map_err(ql_rpc::RpcError::Transport)?; + stream + .writer + .finish() + .await + .map_err(ql_rpc::RpcError::Transport)?; Ok(()) } - pub async fn request(&self, request: &M::Request) -> Result> + pub async fn request(&self, request: &M::Request) -> RpcResult where - M: RequestRpc, + M: request::Request, { let mut payload = Vec::new(); request::encode_request::(request, &mut payload); - let response = self.start_request(M::SERVICE, M::ROUTE, payload).await?; - Ok(request::read_response::(response).await?) + let response = self.start_request::(payload).await?; + request::read_response::(response).await } pub async fn subscribe( &self, request: &M::Request, - ) -> Result, RpcError> + ) -> RpcResult, M::Error> where - M: SubscriptionRpc, + M: subscription::Subscription, { let mut payload = Vec::new(); - rpc_subscription::encode_request::(request, &mut payload); - let response = self.start_request(M::SERVICE, M::ROUTE, payload).await?; - Ok(Subscription { - inner: rpc_subscription::SubscriptionCall::new(response), - }) + subscription::encode_request::(request, &mut payload); + let response = self.start_request::(payload).await?; + Ok(subscription::SubscriptionCall::new(response)) } pub async fn download( &self, request: &M::Request, - ) -> Result, RpcError> + ) -> RpcResult, M::Error> where - M: DownloadRpc, + M: download::Download, { let mut payload = Vec::new(); - rpc_download::encode_request::(request, &mut payload); - let response = self.start_request(M::SERVICE, M::ROUTE, payload).await?; - Ok(DownloadCall { - inner: rpc_download::DownloadCall::new(response), - }) + download::encode_request::(request, &mut payload); + let response = self.start_request::(payload).await?; + Ok(download::DownloadCall::new(response)) } pub async fn progress( &self, request: &M::Request, - ) -> Result, RpcError> + ) -> RpcResult, M::Error> where - M: Progress, + M: progress::Progress, { let mut payload = Vec::new(); - rpc_progress::encode_request::(request, &mut payload); - let response = self.start_request(M::SERVICE, M::ROUTE, payload).await?; - Ok(ProgressCall { - inner: rpc_progress::ProgressCall::new(response), - }) + progress::encode_request::(request, &mut payload); + let response = self.start_request::(payload).await?; + Ok(progress::ProgressCall::new(response)) } - pub async fn upload(&self, request: &M::Request) -> Result, RpcError> + pub async fn upload( + &self, + request: &M::Request, + ) -> RpcResult, M::Error> where - M: UploadRpc, + M: upload::Upload, { let mut payload = Vec::new(); - rpc_upload::encode_request::(request, &mut payload); - let mut stream = self - .inner - .open_stream(open_stream_params(M::SERVICE, M::ROUTE)) - .await?; - stream.writer.write(Bytes::from(payload)).await?; - Ok(UploadCall { - inner: rpc_upload::UploadCall::new(stream.writer, stream.reader), - }) + upload::encode_request::(request, &mut payload); + let mut stream = self.open_rpc_stream::().await?; + stream + .writer + .write(Bytes::from(payload)) + .await + .map_err(ql_rpc::RpcError::Transport)?; + Ok(upload::UploadCall::new(stream.writer, stream.reader)) } - pub async fn duplex(&self) -> Result, RpcError> + pub async fn duplex( + &self, + ) -> RpcResult, M::Error> where - M: DuplexRpc, + M: duplex::Duplex, { - let stream = self - .inner - .open_stream(open_stream_params(M::SERVICE, M::ROUTE)) - .await?; - Ok(DuplexCall { - sender: DuplexSender { - inner: rpc_duplex::DuplexSender::new(stream.writer), - }, - receiver: DuplexReceiver { - inner: rpc_duplex::DuplexReceiver::new(stream.reader), - }, + let stream = self.open_rpc_stream::().await?; + Ok(duplex::DuplexCall { + sender: duplex::DuplexSender::new(stream.writer), + receiver: duplex::DuplexReceiver::new(stream.reader), }) } } @@ -139,28 +121,28 @@ impl RpcHandle { Self { inner } } - async fn start_request( + async fn start_request( &self, - service_id: ql_rpc::ServiceId, - route_id: ql_rpc::RouteId, payload: Vec, - ) -> Result> { - let mut stream = self - .inner - .open_stream(open_stream_params(service_id, route_id)) - .await?; - stream.writer.write(Bytes::from(payload)).await?; + ) -> RpcResult { + let mut stream = self.open_rpc_stream::().await?; + stream + .writer + .write(Bytes::from(payload)) + .await + .map_err(ql_rpc::RpcError::Transport)?; stream.writer.queue_finish(); Ok(stream.reader) } -} -fn open_stream_params( - service_id: ql_rpc::ServiceId, - route_id: ql_rpc::RouteId, -) -> OpenStreamParams { - OpenStreamParams { - service_id, - route_id, + async fn open_rpc_stream(&self) -> RpcResult { + self.inner + .open_stream(OpenStreamParams { + service_id: R::SERVICE, + route_id: R::ROUTE, + }) + .await + .map_err(QlStreamError::from) + .map_err(ql_rpc::RpcError::Transport) } } diff --git a/ql-runtime/src/rpc/progress.rs b/ql-runtime/src/rpc/progress.rs deleted file mode 100644 index c7e5e77b..00000000 --- a/ql-runtime/src/rpc/progress.rs +++ /dev/null @@ -1,50 +0,0 @@ -use std::{ - future::Future, - pin::Pin, - task::{Context, Poll}, -}; - -use futures_lite::Stream; -use ql_rpc::progress::Progress; - -use super::RpcError; -use crate::StreamReader; - -pub struct ProgressCall { - pub(super) inner: ql_rpc::progress::ProgressCall, -} - -impl Unpin for ProgressCall where M: Progress {} - -impl ProgressCall -where - M: Progress, -{ - pub fn reset(self, code: ql_wire::ResetCode) { - self.inner.reset(ql_rpc::ResetCode(code.0)); - } -} - -impl Stream for ProgressCall -where - M: Progress, -{ - type Item = M::Progress; - - fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - self.get_mut().inner.poll_next_progress(cx) - } -} - -impl Future for ProgressCall -where - M: Progress, -{ - type Output = Result>; - - fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { - Pin::new(&mut self.get_mut().inner) - .poll(cx) - .map(|result| result.map_err(RpcError::from)) - } -} diff --git a/ql-runtime/src/rpc/subscription.rs b/ql-runtime/src/rpc/subscription.rs deleted file mode 100644 index 5652394a..00000000 --- a/ql-runtime/src/rpc/subscription.rs +++ /dev/null @@ -1,43 +0,0 @@ -use std::{ - pin::Pin, - task::{Context, Poll}, -}; - -use futures_lite::{future::poll_fn, Stream}; -use ql_rpc::subscription::Subscription as SubscriptionRpc; - -use super::RpcError; -use crate::StreamReader; - -pub struct Subscription { - pub(super) inner: ql_rpc::subscription::SubscriptionCall, -} - -impl Unpin for Subscription where M: SubscriptionRpc {} - -impl Subscription -where - M: SubscriptionRpc, -{ - pub async fn next_event(&mut self) -> Option>> { - poll_fn(|cx| Pin::new(&mut *self).poll_next(cx)).await - } - - pub fn reset(self, code: ql_wire::ResetCode) { - self.inner.reset(ql_rpc::ResetCode(code.0)); - } -} - -impl Stream for Subscription -where - M: SubscriptionRpc, -{ - type Item = Result>; - - fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { - self.get_mut() - .inner - .poll_next_event(cx) - .map(|item| item.map(|result| Ok(result?))) - } -} diff --git a/ql-runtime/src/rpc/upload.rs b/ql-runtime/src/rpc/upload.rs deleted file mode 100644 index 33ee3665..00000000 --- a/ql-runtime/src/rpc/upload.rs +++ /dev/null @@ -1,44 +0,0 @@ -use bytes::Bytes; -use ql_rpc::upload::Upload as UploadRpc; - -use super::RpcError; -use crate::QlStreamError; - -pub struct UploadCall { - pub(super) inner: ql_rpc::upload::UploadCall, -} - -pub struct UploadPartWriter<'a, M: UploadRpc> { - inner: ql_rpc::upload::UploadPartWriter<'a, M, crate::StreamWriter, crate::StreamReader>, -} - -impl UploadCall -where - M: UploadRpc, -{ - pub async fn start_part( - &mut self, - part_header: M::PartHeader, - ) -> Result, QlStreamError> { - Ok(UploadPartWriter { - inner: self.inner.start_part(part_header).await?, - }) - } - - pub async fn finish(self) -> Result> { - self.inner.finish().await.map_err(RpcError::from) - } -} - -impl UploadPartWriter<'_, M> -where - M: UploadRpc, -{ - pub async fn send(&mut self, bytes: Bytes) -> Result<(), QlStreamError> { - self.inner.send(bytes).await - } - - pub async fn finish(self) -> Result<(), QlStreamError> { - self.inner.finish().await - } -} diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index 00e3d749..e570f9b6 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -8,7 +8,6 @@ use std::{ }; use bytes::Bytes; -use futures_lite::StreamExt; use ql_rpc::{ Context, DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, NotificationHandlerLocal, ProgressHandlerLocal, ProgressResponder, RequestHandler, @@ -18,7 +17,7 @@ use ql_rpc::{ }; use super::*; -use crate::{rpc::RpcError, QlInboundStream, StreamWriter}; +use crate::{QlInboundStream, QlStreamError, StreamWriter}; const TEST_SERVICE: ServiceId = ServiceId([7; 16]); @@ -293,9 +292,15 @@ async fn rpc_subscrption() { let rpc = pair.side_mut(Side::A).handle.rpc(); let mut subscription = rpc.subscribe::(&b"watch".to_vec()).await.unwrap(); - assert_eq!(subscription.next().await.unwrap().unwrap(), b"one".to_vec()); - assert_eq!(subscription.next().await.unwrap().unwrap(), b"two".to_vec()); - assert!(subscription.next().await.is_none()); + assert_eq!( + subscription.next_event().await.unwrap().unwrap(), + b"one".to_vec() + ); + assert_eq!( + subscription.next_event().await.unwrap().unwrap(), + b"two".to_vec() + ); + assert!(subscription.next_event().await.is_none()); assert_eq!(seen.borrow().as_slice(), &[b"watch".to_vec()]); tokio::time::timeout(Duration::from_secs(2), responder) @@ -344,7 +349,7 @@ async fn rpc_router_enforces_max_request_bytes() { let response = rpc.request::(&"hello".to_string()).await; assert!(matches!( response, - Err(RpcError::Reset { code, origin }) + Err(ql_rpc::RpcError::Transport(QlStreamError::StreamReset { code, origin })) if code == ResetCode::LIMIT && origin == ResetOrigin::Peer )); @@ -400,9 +405,9 @@ async fn rpc_progress() { let rpc = pair.side_mut(Side::A).handle.rpc(); let mut download = rpc.progress::(&b"logo".to_vec()).await.unwrap(); - assert_eq!(download.next().await, Some(b"10".to_vec())); - assert_eq!(download.next().await, Some(b"90".to_vec())); - assert_eq!(download.next().await, None); + assert_eq!(download.next_progress().await, Some(b"10".to_vec())); + assert_eq!(download.next_progress().await, Some(b"90".to_vec())); + assert_eq!(download.next_progress().await, None); assert_eq!(download.await.unwrap(), b"done".to_vec()); assert_eq!(seen.borrow().as_slice(), &[b"logo".to_vec()]); From 3ff947729194c232a14f6044b305fbbd5b236f52 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 24 Jun 2026 15:06:34 -0400 Subject: [PATCH 24/59] ql-fsm: simplify stream reset --- ql-fsm/src/fsm.rs | 7 +-- ql-fsm/src/lib.rs | 42 +++++++++++++---- ql-fsm/src/session/mod.rs | 33 +++++++++----- ql-fsm/src/session/stream_ops.rs | 17 ++++--- ql-fsm/src/session/tests.rs | 64 +++++++++----------------- ql-fsm/src/tests/proptest.rs | 47 ++++++++++--------- ql-fsm/src/tests/session.rs | 5 +- ql-runtime/src/command.rs | 7 ++- ql-runtime/src/driver/mod.rs | 78 ++++++++------------------------ ql-runtime/src/driver/state.rs | 31 ++----------- ql-runtime/src/driver/test.rs | 55 ++++++---------------- ql-runtime/src/io/mod.rs | 8 ++-- ql-runtime/src/io/reader.rs | 37 ++++++--------- ql-runtime/src/io/writer.rs | 39 ++++++---------- 14 files changed, 190 insertions(+), 280 deletions(-) diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index 6a804af8..804b486b 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -46,11 +46,8 @@ impl session::EventSink for EventSink<'_> { SessionEvent::OutboundFinished(stream_id) => { self.events.push_back(Event::OutboundFinished(stream_id)); } - SessionEvent::Reset(frame) => { - self.events.push_back(Event::Reset(frame)); - } - SessionEvent::WritableReset(frame) => { - self.events.push_back(Event::WritableReset(frame)); + SessionEvent::Reset(reset) => { + self.events.push_back(Event::Reset(reset)); } SessionEvent::SessionClosed(close) => { self.termination = Some(TerminalFrame::Close(close.clone())); diff --git a/ql-fsm/src/lib.rs b/ql-fsm/src/lib.rs index d70fd9f1..8a96f93b 100644 --- a/ql-fsm/src/lib.rs +++ b/ql-fsm/src/lib.rs @@ -36,8 +36,8 @@ pub use bytes::Bytes; pub use error::*; pub use pairing::PairingInvite; use ql_wire::{ - PairingToken, PeerBundle, QlCrypto, QlIdentity, RouteId, ServiceId, SessionClose, - SessionCloseCode, StreamHeader, StreamId, StreamReset, + PairingToken, PeerBundle, QlCrypto, QlIdentity, ResetCode, RouteId, ServiceId, SessionClose, + SessionCloseCode, StreamHeader, StreamId, }; pub use session::{SessionEvent, StreamReadIter, StreamWriter}; @@ -76,10 +76,8 @@ pub enum Event { Finished(StreamId), /// our local FIN was acknowledged by the peer at the session layer OutboundFinished(StreamId), - /// a stream was reset - Reset(StreamReset), - /// local writes on this stream are reset - WritableReset(StreamReset), + /// one or both local stream halves were reset by the peer + Reset(StreamResetEvent), /// the encrypted session was closed /// /// session close is abortive and best-effort. the session ends immediately @@ -88,6 +86,34 @@ pub enum Event { SessionClosed(SessionClose), } +/// stream was reset by remote peer +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct StreamResetEvent { + pub stream_id: StreamId, + pub code: ResetCode, + pub target: StreamResetTarget, +} + +/// local stream halves that can be reset +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum StreamResetTarget { + Reader, + Writer, + Both, +} + +impl StreamResetTarget { + #[inline] + pub fn reader(self) -> bool { + matches!(self, Self::Reader | Self::Both) + } + + #[inline] + pub fn writer(self) -> bool { + matches!(self, Self::Writer | Self::Both) + } +} + /// handle for a session write returned by `QlFsm::take_next_write` #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub struct WriteId(pub(crate) u64); @@ -135,8 +161,8 @@ impl StreamOps<'_> { self.inner.writer() } - /// resets the origin lane, return lane, or both lanes of the stream - pub fn reset(&mut self, target: ql_wire::ResetTarget, code: ql_wire::ResetCode) { + /// resets the local read side, write side, or both sides of the stream + pub fn reset(&mut self, target: StreamResetTarget, code: ql_wire::ResetCode) { self.inner.reset(target, code); } } diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index 05724e89..335623f2 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -30,7 +30,7 @@ use self::{ stream_tx::StreamTxRange, tracked::{TrackedFrame, TrackedRecord, TrackedStreamData}, }; -use crate::{NoSessionError, StreamError}; +use crate::{NoSessionError, StreamError, StreamResetEvent, StreamResetTarget}; #[derive(Debug, Clone, Copy)] pub struct SessionConfig { @@ -72,8 +72,7 @@ pub enum SessionEvent { Writable(StreamId), Finished(StreamId), OutboundFinished(StreamId), - Reset(StreamReset), - WritableReset(StreamReset), + Reset(StreamResetEvent), SessionClosed(SessionClose), Unpaired, } @@ -754,23 +753,35 @@ impl SessionFsm { }, }; - if Self::target_affects_inbound(stream.role, frame.target) + let inbound = Self::target_affects_inbound(stream.role, frame.target) && !matches!( stream.inbound_state, InboundState::Reset(_) | InboundState::Discarding - ) - { + ); + let outbound = Self::target_affects_outbound(stream.role, frame.target) + && !matches!(stream.outbound_state, OutboundState::Closed); + + if inbound { stream.inbound_state = InboundState::Reset(frame.clone()); stream.reset_recv(); - sink.emit(SessionEvent::Reset(frame.clone())); } - if Self::target_affects_outbound(stream.role, frame.target) - && !matches!(stream.outbound_state, OutboundState::Closed) - { + if outbound { stream.outbound_state = OutboundState::Closed; stream.tx.clear(); stream.pending_reset = None; - sink.emit(SessionEvent::WritableReset(frame.clone())); + } + if inbound || outbound { + let target = match (inbound, outbound) { + (true, true) => StreamResetTarget::Both, + (true, false) => StreamResetTarget::Reader, + (false, true) => StreamResetTarget::Writer, + (false, false) => unreachable!(), + }; + sink.emit(SessionEvent::Reset(StreamResetEvent { + stream_id, + code: frame.code, + target, + })); } self.try_reap_stream(frame.stream_id); Ok(()) diff --git a/ql-fsm/src/session/stream_ops.rs b/ql-fsm/src/session/stream_ops.rs index 4b3c36ad..67a1900b 100644 --- a/ql-fsm/src/session/stream_ops.rs +++ b/ql-fsm/src/session/stream_ops.rs @@ -1,11 +1,11 @@ -use ql_wire::{ResetCode, ResetTarget, StreamHeader, StreamId, StreamReset}; +use ql_wire::{ResetCode, StreamHeader, StreamId, StreamReset}; use super::{ state::{InboundState, StreamState}, stream_rx::StreamReadIter, EventSink, SessionEvent, SessionFsm, }; -use crate::CommitReadError; +use crate::{CommitReadError, StreamResetTarget}; pub struct StreamOps<'a, E> { session: &'a mut SessionFsm, @@ -86,14 +86,19 @@ impl<'a, E: EventSink> StreamOps<'a, E> { Some(StreamWriter::new(stream, send_buffer_size)) } - /// resets the origin lane, return lane, or both lanes of the stream - pub fn reset(&mut self, target: ResetTarget, code: ResetCode) { + /// resets the local read side, write side, or both sides of the stream + pub fn reset(&mut self, target: StreamResetTarget, code: ResetCode) { let stream_id = self.stream_id; let stream = self.stream_mut(); - SessionFsm::apply_local_reset_to_stream(stream, target); + let wire_target = match target { + StreamResetTarget::Reader => stream.role.inbound_target(), + StreamResetTarget::Writer => stream.role.outbound_target(), + StreamResetTarget::Both => ql_wire::ResetTarget::Both, + }; + SessionFsm::apply_local_reset_to_stream(stream, wire_target); stream.pending_reset = Some(StreamReset { stream_id, - target, + target: wire_target, code, }); self.reap_on_drop = true; diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index 7a9dc06c..85edec29 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -8,7 +8,7 @@ use ql_wire::{ }; use super::{SessionConfig, SessionEvent, SessionFsm}; -use crate::session::stream_parity::StreamParity; +use crate::{session::stream_parity::StreamParity, StreamResetEvent}; fn seq(value: u64) -> RecordSeq { RecordSeq::from_u64(value).unwrap() @@ -421,7 +421,7 @@ fn remote_stream_reset_is_reliable_and_retried() { fsm.stream(stream_id, |_| {}) .unwrap() - .reset(ResetTarget::Both, ResetCode::CANCELLED); + .reset(crate::StreamResetTarget::Both, ResetCode::CANCELLED); let (write_id, builder) = fsm.take_next_write(now).unwrap(); fsm.complete_write(now, write_id.expect("stream reset should be tracked"), true); @@ -501,10 +501,11 @@ fn duplicate_remote_reset_after_reap_is_ignored() { let first = receive_events(&mut fsm, now, seq(1), &record); assert_eq!( first, - vec![ - SessionEvent::Reset(reset.clone()), - SessionEvent::WritableReset(reset), - ] + vec![SessionEvent::Reset(StreamResetEvent { + stream_id: reset.stream_id, + code: reset.code, + target: crate::StreamResetTarget::Both, + })] ); let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); @@ -532,18 +533,11 @@ fn late_remote_stream_data_after_reset_is_ignored() { let first = receive_events(&mut fsm, now, seq(1), &reset); assert_eq!( first, - vec![ - SessionEvent::Reset(StreamReset { - stream_id, - target: ResetTarget::Both, - code: ResetCode(9), - }), - SessionEvent::WritableReset(StreamReset { - stream_id, - target: ResetTarget::Both, - code: ResetCode(9), - }), - ] + vec![SessionEvent::Reset(StreamResetEvent { + stream_id, + code: ResetCode(9), + target: crate::StreamResetTarget::Both, + })] ); let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &data); @@ -626,35 +620,21 @@ fn out_of_order_remote_stream_first_observations_still_open_once_each() { let first = receive_events(&mut fsm, now, seq(1), &reset3); assert_eq!( first, - vec![ - SessionEvent::Reset(StreamReset { - stream_id: stream_id(3), - target: ResetTarget::Both, - code: REFUSED, - }), - SessionEvent::WritableReset(StreamReset { - stream_id: stream_id(3), - target: ResetTarget::Both, - code: REFUSED, - }), - ] + vec![SessionEvent::Reset(StreamResetEvent { + stream_id: stream_id(3), + code: REFUSED, + target: crate::StreamResetTarget::Both, + })] ); let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &reset1); assert_eq!( second, - vec![ - SessionEvent::Reset(StreamReset { - stream_id: stream_id(1), - target: ResetTarget::Both, - code: TIMEOUT, - }), - SessionEvent::WritableReset(StreamReset { - stream_id: stream_id(1), - target: ResetTarget::Both, - code: TIMEOUT, - }), - ] + vec![SessionEvent::Reset(StreamResetEvent { + stream_id: stream_id(1), + code: TIMEOUT, + target: crate::StreamResetTarget::Both, + })] ); let third = receive_events(&mut fsm, now + Duration::from_millis(2), seq(3), &reset3); diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs index ab68fae8..dcd85bdf 100644 --- a/ql-fsm/src/tests/proptest.rs +++ b/ql-fsm/src/tests/proptest.rs @@ -7,10 +7,12 @@ extern crate proptest as proptest_crate; use bytes::Bytes; use proptest_crate::{collection::vec, prelude::*, test_runner::TestCaseResult}; -use ql_wire::{ResetCode, ResetTarget, RouteId, ServiceId, StreamId, WireError}; +use ql_wire::{ResetCode, RouteId, ServiceId, StreamId, WireError}; use super::*; -use crate::{state::LinkState, Event, OpenStreamParams, PeerStatus, ReceiveError, WriteId}; +use crate::{ + state::LinkState, Event, OpenStreamParams, PeerStatus, ReceiveError, StreamResetTarget, WriteId, +}; const SLOT_COUNT: usize = 4; @@ -337,7 +339,7 @@ impl Runner { .fsm .stream(stream_id) .is_ok_and(|mut stream| { - stream.reset(ResetTarget::Both, ResetCode::CANCELLED); + stream.reset(StreamResetTarget::Both, ResetCode::CANCELLED); true }); if reset { @@ -463,28 +465,28 @@ impl Runner { "side {side:?} emitted duplicate OutboundFinished for {stream_id:?}" ); } - Event::Reset(frame) => { + Event::Reset(reset) => { prop_assert!( - self.known_streams.contains(&frame.stream_id), + self.known_streams.contains(&reset.stream_id), "side {side:?} emitted Reset for unknown stream {:?}", - frame.stream_id - ); - prop_assert!( - self.events[side.idx()].reset.insert(frame.stream_id), - "side {side:?} emitted duplicate Reset for {:?}", - frame.stream_id - ); - } - Event::WritableReset(frame) => { - let stream_id = frame.stream_id; - prop_assert!( - self.known_streams.contains(&stream_id), - "side {side:?} emitted WritableReset for unknown stream {stream_id:?}" - ); - prop_assert!( - self.events[side.idx()].writable_reset.insert(stream_id), - "side {side:?} emitted duplicate WritableReset for {stream_id:?}" + reset.stream_id ); + if reset.target.reader() { + prop_assert!( + self.events[side.idx()].reset.insert(reset.stream_id), + "side {side:?} emitted duplicate inbound Reset for {:?}", + reset.stream_id + ); + } + if reset.target.writer() { + prop_assert!( + self.events[side.idx()] + .writable_reset + .insert(reset.stream_id), + "side {side:?} emitted duplicate outbound Reset for {:?}", + reset.stream_id + ); + } } Event::SessionClosed(_) => { let state = &mut self.events[side.idx()]; @@ -547,6 +549,7 @@ impl Runner { | ReceiveError::InvalidRemoteBundle | ReceiveError::InvalidSessionPayload(WireError::InvalidPayload) | ReceiveError::InvalidSessionPayload(WireError::DecryptFailed) + | ReceiveError::InvalidSessionConnectionId | ReceiveError::InvalidIkHandshake(WireError::InvalidPayload) | ReceiveError::InvalidIkHandshake(WireError::InvalidState) | ReceiveError::InvalidKkHandshake(WireError::InvalidPayload) diff --git a/ql-fsm/src/tests/session.rs b/ql-fsm/src/tests/session.rs index 73aabe57..efd9c8c3 100644 --- a/ql-fsm/src/tests/session.rs +++ b/ql-fsm/src/tests/session.rs @@ -218,7 +218,10 @@ fn disconnected_stream_operations_fail_with_no_session() { ); assert_eq!( harness.a.fsm.stream(missing).map(|mut stream| { - stream.reset(ql_wire::ResetTarget::Both, ql_wire::ResetCode::CANCELLED); + stream.reset( + crate::StreamResetTarget::Both, + ql_wire::ResetCode::CANCELLED, + ); }), Err(StreamError::NoSession) ); diff --git a/ql-runtime/src/command.rs b/ql-runtime/src/command.rs index a27a0888..730ebefa 100644 --- a/ql-runtime/src/command.rs +++ b/ql-runtime/src/command.rs @@ -1,7 +1,6 @@ -use ql_fsm::{NoSessionError, PairingInvite}; +use ql_fsm::{NoSessionError, PairingInvite, StreamResetTarget}; use ql_wire::{ - PairingToken, PeerBundle, ResetCode, ResetTarget, RouteId, ServiceId, SessionCloseCode, - StreamId, + PairingToken, PeerBundle, ResetCode, RouteId, ServiceId, SessionCloseCode, StreamId, }; use crate::{StreamReader, StreamWriter}; @@ -35,7 +34,7 @@ pub enum Command { Unpair, ResetStream { stream_id: StreamId, - target: ResetTarget, + target: StreamResetTarget, code: ResetCode, }, } diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 125b5ba3..7b4b2ca4 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -15,8 +15,8 @@ use std::{ use async_channel::Recv; use futures_lite::future::{poll_fn, yield_now}; -use ql_fsm::{Event, QlFsm, WriteId}; -use ql_wire::{ResetCode, ResetOrigin, ResetTarget, StreamHeader, StreamId}; +use ql_fsm::{Event, QlFsm, StreamResetEvent, StreamResetTarget, WriteId}; +use ql_wire::{ResetCode, ResetOrigin, StreamHeader, StreamId}; use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; use crate::{ @@ -229,16 +229,11 @@ impl DriverState { }; let stream_id = stream_ops.stream_id(); log::info!("open stream allocated: service_id={service_id} route_id={route_id} stream_id={stream_id}"); - let (reader, writer, reader_io, writer_io) = io::new_stream( - stream_id, - ResetTarget::Return, - ResetTarget::Origin, - self.runtime_tx.clone(), - ); + let (reader, writer, reader_io, writer_io) = + io::new_stream(stream_id, self.runtime_tx.clone()); self.streams.insert( stream_id, DriverStreamIo::new( - true, Some(OutboundIo::new(writer_io)), Some(InboundIo::new(reader_io)), ), @@ -249,7 +244,7 @@ impl DriverState { stream.inbound_close(); stream.outbound_close(); } - stream_ops.reset(ResetTarget::Both, ResetCode::DROPPED); + stream_ops.reset(StreamResetTarget::Both, ResetCode::DROPPED); return; } drop(stream_ops); @@ -273,10 +268,10 @@ impl DriverState { ); if let Entry::Occupied(mut entry) = self.streams.entry(stream_id) { let stream = entry.get_mut(); - if target == ResetTarget::Both || target == stream.inbound_target() { + if target.reader() { stream.inbound_close(); } - if target == ResetTarget::Both || target == stream.outbound_target() { + if target.writer() { stream.outbound_close(); } Self::try_reap_stream(entry); @@ -334,11 +329,8 @@ impl DriverState { log::info!("outbound finish acknowledged: stream_id={stream_id}"); self.handle_outbound_finished(stream_id); } - Event::Reset(frame) => { - self.handle_reset_stream(&frame); - } - Event::WritableReset(frame) => { - self.handle_writable_reset(&frame); + Event::Reset(reset) => { + self.handle_stream_reset(reset); } Event::SessionClosed(close) => { log::info!("session closed: frame={close:?}"); @@ -356,17 +348,12 @@ impl DriverState { platform: &P, stream_id: StreamId, ) { - let (reader, writer, reader_io, writer_io) = io::new_stream( - stream_id, - ResetTarget::Origin, - ResetTarget::Return, - self.runtime_tx.clone(), - ); + let (reader, writer, reader_io, writer_io) = + io::new_stream(stream_id, self.runtime_tx.clone()); self.streams.insert( stream_id, DriverStreamIo::new( - false, Some(OutboundIo::new(writer_io)), Some(InboundIo::new(reader_io)), ), @@ -405,12 +392,10 @@ impl DriverState { log::trace!("draining inbound bytes: stream_id={stream_id} readable={readable}"); let mut accepted = 0usize; let mut peer_closed = false; - let target; { let Some(stream) = self.streams.get_mut(&stream_id) else { return; }; - target = stream.inbound_target(); for chunk in stream_ops.read() { if chunk.is_empty() { continue; @@ -427,7 +412,7 @@ impl DriverState { } InboundWriteResult::Closed => { log::warn!( - "inbound consumer closed; sending CANCELLED: stream_id={stream_id} target={target:?}" + "inbound consumer closed; sending CANCELLED: stream_id={stream_id}" ); peer_closed = true; break; @@ -441,7 +426,7 @@ impl DriverState { stream_ops.commit_read(accepted).unwrap(); } if peer_closed { - stream_ops.reset(target, ResetCode::DROPPED); + stream_ops.reset(StreamResetTarget::Reader, ResetCode::DROPPED); if let Entry::Occupied(entry) = self.streams.entry(stream_id) { Self::try_reap_stream(entry); } @@ -460,51 +445,28 @@ impl DriverState { Self::try_reap_stream(entry); } - fn handle_reset_stream(&mut self, frame: &ql_wire::StreamReset) { - log::info!( - "inbound reset frame: stream_id={} target={:?} code={}", - frame.stream_id, - frame.target, - frame.code - ); - let Entry::Occupied(mut entry) = self.streams.entry(frame.stream_id) else { + fn handle_stream_reset(&mut self, reset: StreamResetEvent) { + log::info!("stream reset: {reset:?}",); + let Entry::Occupied(mut entry) = self.streams.entry(reset.stream_id) else { return; }; let stream = entry.get_mut(); - if frame.target == ResetTarget::Both || frame.target == stream.inbound_target() { + if reset.target.reader() { stream.inbound_fail(QlStreamError::StreamReset { - code: frame.code, + code: reset.code, origin: ResetOrigin::Peer, }); } - if frame.target == ResetTarget::Both || frame.target == stream.outbound_target() { + if reset.target.writer() { stream.outbound_fail(QlStreamError::StreamReset { - code: frame.code, + code: reset.code, origin: ResetOrigin::Peer, }); } Self::try_reap_stream(entry); } - fn handle_writable_reset(&mut self, frame: &ql_wire::StreamReset) { - log::info!( - "writable reset frame: stream_id={} target={:?} code={}", - frame.stream_id, - frame.target, - frame.code - ); - let Entry::Occupied(mut entry) = self.streams.entry(frame.stream_id) else { - return; - }; - let stream = entry.get_mut(); - stream.outbound_fail(QlStreamError::StreamReset { - code: frame.code, - origin: ResetOrigin::Peer, - }); - Self::try_reap_stream(entry); - } - fn handle_outbound_finished(&mut self, stream_id: StreamId) { log::info!("outbound finish acknowledged: stream_id={stream_id}"); let Entry::Occupied(mut entry) = self.streams.entry(stream_id) else { diff --git a/ql-runtime/src/driver/state.rs b/ql-runtime/src/driver/state.rs index b2803343..1a85e96c 100644 --- a/ql-runtime/src/driver/state.rs +++ b/ql-runtime/src/driver/state.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use bytes::Bytes; -use ql_wire::{ResetTarget, StreamId}; +use ql_wire::StreamId; use crate::{ command::Command, @@ -16,38 +16,13 @@ pub struct DriverState { } pub struct DriverStreamIo { - is_initiator: bool, outbound: Option, inbound: Option, } impl DriverStreamIo { - pub fn new( - is_initiator: bool, - outbound: Option, - inbound: Option, - ) -> Self { - Self { - is_initiator, - outbound, - inbound, - } - } - - pub fn inbound_target(&self) -> ResetTarget { - if self.is_initiator { - ResetTarget::Return - } else { - ResetTarget::Origin - } - } - - pub fn outbound_target(&self) -> ResetTarget { - if self.is_initiator { - ResetTarget::Origin - } else { - ResetTarget::Return - } + pub fn new(outbound: Option, inbound: Option) -> Self { + Self { outbound, inbound } } pub fn fail_all(&mut self) { diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs index 69bf68e0..58397247 100644 --- a/ql-runtime/src/driver/test.rs +++ b/ql-runtime/src/driver/test.rs @@ -1,4 +1,5 @@ -use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto, StreamReset, QID}; +use ql_fsm::StreamResetEvent; +use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto, QID}; use super::*; use crate::{ @@ -67,24 +68,14 @@ fn new_driver_state() -> (DriverState, QlFsm) { fn new_inbound_io(capacity: usize) -> InboundIo { let _ = capacity; let (runtime_tx, _runtime_rx) = async_channel::unbounded(); - let stream = io::new_stream( - StreamId(99u32.into()), - ResetTarget::Origin, - ResetTarget::Return, - runtime_tx, - ); + let stream = io::new_stream(StreamId(99u32.into()), runtime_tx); let (_, _, reader_io, _) = stream; InboundIo::new(reader_io) } fn new_outbound_io() -> OutboundIo { let (runtime_tx, _runtime_rx) = async_channel::unbounded(); - let stream = io::new_stream( - StreamId(100u32.into()), - ResetTarget::Return, - ResetTarget::Origin, - runtime_tx, - ); + let stream = io::new_stream(StreamId(100u32.into()), runtime_tx); let (_, _, _, writer_io) = stream; OutboundIo::new(writer_io) } @@ -96,7 +87,7 @@ fn handle_inbound_finished_reaps_reset_initiator_stream() { state.streams.insert( stream_id, - DriverStreamIo::new(true, None, Some(new_inbound_io(1))), + DriverStreamIo::new(None, Some(new_inbound_io(1))), ); state.handle_inbound_finished(stream_id); @@ -105,19 +96,19 @@ fn handle_inbound_finished_reaps_reset_initiator_stream() { } #[test] -fn handle_reset_stream_reaps_when_both_halves_reset() { +fn handle_stream_reset_reaps_when_both_halves_reset() { let (mut state, _fsm) = new_driver_state(); let stream_id = StreamId(1u32.into()); state.streams.insert( stream_id, - DriverStreamIo::new(false, Some(new_outbound_io()), Some(new_inbound_io(1))), + DriverStreamIo::new(Some(new_outbound_io()), Some(new_inbound_io(1))), ); - state.handle_reset_stream(&StreamReset { + state.handle_stream_reset(StreamResetEvent { stream_id, - target: ResetTarget::Both, code: ResetCode::CANCELLED, + target: StreamResetTarget::Both, }); assert!(!state.streams.contains_key(&stream_id)); @@ -128,16 +119,11 @@ fn poll_stream_keeps_outbound_pending_after_local_finish_when_inbound_is_closed( let (mut state, mut fsm) = new_driver_state(); let stream_id = StreamId(1u32.into()); let (runtime_tx, _runtime_rx) = async_channel::unbounded(); - let (_, mut writer, _, writer_io) = io::new_stream( - stream_id, - ResetTarget::Return, - ResetTarget::Origin, - runtime_tx, - ); + let (_, mut writer, _, writer_io) = io::new_stream(stream_id, runtime_tx); writer.queue_finish(); state.streams.insert( stream_id, - DriverStreamIo::new(true, Some(OutboundIo::new(writer_io)), None), + DriverStreamIo::new(Some(OutboundIo::new(writer_io)), None), ); state.poll_stream(&mut fsm, stream_id); @@ -152,23 +138,18 @@ fn local_reset_command_reaps_when_other_half_is_already_closed() { let (mut state, mut fsm) = new_driver_state(); let stream_id = StreamId(1u32.into()); let (runtime_tx, _runtime_rx) = async_channel::unbounded(); - let (_, _, _, writer_io) = io::new_stream( - stream_id, - ResetTarget::Return, - ResetTarget::Origin, - runtime_tx, - ); + let (_, _, _, writer_io) = io::new_stream(stream_id, runtime_tx); state.streams.insert( stream_id, - DriverStreamIo::new(true, Some(OutboundIo::new(writer_io)), None), + DriverStreamIo::new(Some(OutboundIo::new(writer_io)), None), ); state.drive_command( &mut fsm, Command::ResetStream { stream_id, - target: ResetTarget::Origin, + target: StreamResetTarget::Writer, code: ResetCode::CANCELLED, }, &NoopCrypto, @@ -183,17 +164,11 @@ fn unpaired_status_fails_and_reaps_all_streams() { let peer = generate_identity(&SoftwareCrypto, "peer").unwrap().bundle(); let stream_id = StreamId(1u32.into()); let (runtime_tx, _runtime_rx) = async_channel::unbounded(); - let (_, _, reader_io, writer_io) = io::new_stream( - stream_id, - ResetTarget::Origin, - ResetTarget::Return, - runtime_tx, - ); + let (_, _, reader_io, writer_io) = io::new_stream(stream_id, runtime_tx); state.streams.insert( stream_id, DriverStreamIo::new( - false, Some(OutboundIo::new(writer_io)), Some(InboundIo::new(reader_io)), ), diff --git a/ql-runtime/src/io/mod.rs b/ql-runtime/src/io/mod.rs index 52dba870..039363e3 100644 --- a/ql-runtime/src/io/mod.rs +++ b/ql-runtime/src/io/mod.rs @@ -6,7 +6,7 @@ mod writer; use std::ops::Deref; -use ql_wire::{ResetTarget, StreamId}; +use ql_wire::StreamId; pub use self::{reader::StreamReader, slot::PushError, writer::StreamWriter}; use crate::command::Command; @@ -45,14 +45,12 @@ impl Tx { pub fn new_stream( stream_id: StreamId, - reader_target: ResetTarget, - writer_target: ResetTarget, runtime_tx: async_channel::Sender, ) -> (StreamReader, StreamWriter, Rx, Tx) { let shared = inner::new(stream_id); ( - StreamReader::new(Rx(shared.clone()), reader_target, runtime_tx.clone()), - StreamWriter::new(Tx(shared.clone()), writer_target, runtime_tx), + StreamReader::new(Rx(shared.clone()), runtime_tx.clone()), + StreamWriter::new(Tx(shared.clone()), runtime_tx), Rx(shared.clone()), Tx(shared), ) diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs index 591c8cad..9482ca4a 100644 --- a/ql-runtime/src/io/reader.rs +++ b/ql-runtime/src/io/reader.rs @@ -4,7 +4,8 @@ use std::{ }; use bytes::Bytes; -use ql_wire::{ResetCode, ResetTarget, StreamId}; +use ql_fsm::StreamResetTarget; +use ql_wire::{ResetCode, StreamId}; use super::{ inner::{Item, RxInner}, @@ -15,7 +16,6 @@ use crate::{command::Command, log, QlStreamError}; pub struct StreamReader { rx: Rx, - target: ResetTarget, terminal: ReaderTerminalState, runtime_tx: async_channel::Sender, } @@ -31,7 +31,6 @@ impl std::fmt::Debug for StreamReader { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("StreamReader") .field("stream_id", &self.rx.stream_id()) - .field("target", &self.target) .field( "terminal", &matches!(self.terminal, ReaderTerminalState::Delivered), @@ -41,14 +40,9 @@ impl std::fmt::Debug for StreamReader { } impl StreamReader { - pub(crate) fn new( - shared: Rx, - target: ResetTarget, - runtime_tx: async_channel::Sender, - ) -> Self { + pub(crate) fn new(shared: Rx, runtime_tx: async_channel::Sender) -> Self { Self { rx: shared, - target, terminal: ReaderTerminalState::Open, runtime_tx, } @@ -86,9 +80,8 @@ impl StreamReader { match self.rx.pop() { Ok(Item::Chunk(bytes)) => { log::trace!( - "byte reader received chunk: stream_id={} target={:?} len={}", + "byte reader received chunk: stream_id={} len={}", self.rx.stream_id(), - self.target, bytes.len() ); let _ = self.runtime_tx.try_send(Command::PollInbound { @@ -98,9 +91,8 @@ impl StreamReader { } Ok(Item::Error(error)) => { log::debug!( - "byte reader delivered terminal error: stream_id={} target={:?} error={:?}", + "byte reader delivered terminal error: stream_id={} error={:?}", self.rx.stream_id(), - self.target, error ); self.terminal = ReaderTerminalState::Delivered; @@ -109,9 +101,8 @@ impl StreamReader { Err(PopError) => { if RxInner::is_finished(self.rx.load_state()) { log::debug!( - "byte reader delivered clean eof: stream_id={} target={:?}", - self.rx.stream_id(), - self.target + "byte reader delivered clean eof: stream_id={}", + self.rx.stream_id() ); self.terminal = ReaderTerminalState::Delivered; return Poll::Ready(Ok(None)); @@ -134,15 +125,14 @@ impl StreamReader { return; } log::debug!( - "byte reader explicit reset: stream_id={:?} target={:?} code={:?}", + "byte reader explicit reset: stream_id={:?} code={:?}", self.rx.stream_id(), - self.target, code ); self.terminal = ReaderTerminalState::Delivered; let _ = self.runtime_tx.try_send(Command::ResetStream { stream_id: self.rx.stream_id(), - target: self.target, + target: StreamResetTarget::Reader, code, }); } @@ -154,14 +144,13 @@ impl Drop for StreamReader { return; } log::debug!( - "byte reader drop reset: stream_id={:?} target={:?} code={:?}", + "byte reader drop reset: stream_id={:?} code={:?}", self.rx.stream_id(), - self.target, ResetCode::DROPPED ); let _ = self.runtime_tx.try_send(Command::ResetStream { stream_id: self.rx.stream_id(), - target: self.target, + target: StreamResetTarget::Reader, code: ResetCode::DROPPED, }); } @@ -173,7 +162,7 @@ mod loom_tests { use bytes::Bytes; use loom::thread; - use ql_wire::ResetTarget; + use ql_fsm::StreamResetTarget; use super::*; use crate::io::sync::loom::*; @@ -182,7 +171,7 @@ mod loom_tests { fn poll_read_observes_chunk_racing_with_registration() { check_model(|| { let inner = shared(); - let mut reader = StreamReader::new(Rx(inner.clone()), ResetTarget::Origin, handle()); + let mut reader = StreamReader::new(Rx(inner.clone()), handle()); let mut cx = Context::from_waker(Waker::noop()); let producer = { diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs index c124b017..89f637cc 100644 --- a/ql-runtime/src/io/writer.rs +++ b/ql-runtime/src/io/writer.rs @@ -4,7 +4,8 @@ use std::{ }; use bytes::Bytes; -use ql_wire::{ResetCode, ResetTarget, StreamId}; +use ql_fsm::StreamResetTarget; +use ql_wire::{ResetCode, StreamId}; use super::{ inner::{Item, TxInner}, @@ -15,7 +16,6 @@ use crate::{command::Command, log, QlStreamError}; pub struct StreamWriter { tx: Tx, - target: ResetTarget, open: bool, terminal: WriterTerminalState, runtime_tx: async_channel::Sender, @@ -32,21 +32,15 @@ impl std::fmt::Debug for StreamWriter { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("StreamWriter") .field("stream_id", &self.tx.stream_id()) - .field("target", &self.target) .field("closed", &!self.open) .finish_non_exhaustive() } } impl StreamWriter { - pub(crate) fn new( - shared: Tx, - target: ResetTarget, - runtime_tx: async_channel::Sender, - ) -> Self { + pub(crate) fn new(shared: Tx, runtime_tx: async_channel::Sender) -> Self { Self { tx: shared, - target, open: true, terminal: WriterTerminalState::Pending, runtime_tx, @@ -73,9 +67,8 @@ impl StreamWriter { match self.tx.try_write(std::mem::take(bytes)) { Ok(()) => { log::trace!( - "byte writer accepted chunk: stream_id={} target={:?}", - self.tx.stream_id(), - self.target + "byte writer accepted chunk: stream_id={}", + self.tx.stream_id() ); self.poll_runtime(); return Poll::Ready(Ok(())); @@ -96,9 +89,8 @@ impl StreamWriter { Ok(()) => { self.tx.unregister_waiter(); log::trace!( - "byte writer accepted chunk: stream_id={} target={:?}", - self.tx.stream_id(), - self.target + "byte writer accepted chunk: stream_id={}", + self.tx.stream_id() ); self.poll_runtime(); Poll::Ready(Ok(())) @@ -125,11 +117,7 @@ impl StreamWriter { if !self.open { return; } - log::debug!( - "byte writer finish: stream_id={} target={:?}", - self.tx.stream_id(), - self.target - ); + log::debug!("byte writer finish: stream_id={}", self.tx.stream_id()); self.open = false; self.tx.request_finish(); self.poll_runtime(); @@ -208,14 +196,13 @@ impl StreamWriter { } self.open = false; log::debug!( - "byte writer reset: stream_id={:?} target={:?} code={:?}", + "byte writer reset: stream_id={:?} code={:?}", self.tx.stream_id(), - self.target, code ); let _ = self.runtime_tx.try_send(Command::ResetStream { stream_id: self.tx.stream_id(), - target: self.target, + target: StreamResetTarget::Writer, code, }); } @@ -233,7 +220,7 @@ mod loom_tests { use bytes::Bytes; use loom::thread; - use ql_wire::ResetTarget; + use ql_fsm::StreamResetTarget; use super::*; use crate::io::sync::loom::*; @@ -244,7 +231,7 @@ mod loom_tests { let inner = shared(); inner.tx.try_write(Bytes::from_static(b"abc")).unwrap(); - let mut writer = StreamWriter::new(Tx(inner.clone()), ResetTarget::Origin, handle()); + let mut writer = StreamWriter::new(Tx(inner.clone()), handle()); let mut bytes = Bytes::from_static(b"xyz"); let mut cx = Context::from_waker(Waker::noop()); @@ -275,7 +262,7 @@ mod loom_tests { fn poll_finish_observes_terminal_racing_with_registration() { check_model(|| { let inner = shared(); - let mut writer = StreamWriter::new(Tx(inner.clone()), ResetTarget::Origin, handle()); + let mut writer = StreamWriter::new(Tx(inner.clone()), handle()); let mut cx = Context::from_waker(Waker::noop()); writer.queue_finish(); From e4e1580a77bf91841d94a0aa72b1822d37815781 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 24 Jun 2026 18:05:08 -0400 Subject: [PATCH 25/59] ql-rpc: cleanup --- ql-rpc/src/codec.rs | 9 +++ ql-rpc/src/rpc/download/client.rs | 33 ++++++---- ql-rpc/src/rpc/download/mod.rs | 7 +-- ql-rpc/src/rpc/duplex/client.rs | 60 ++++++++++++++++++- ql-rpc/src/rpc/duplex/codec.rs | 86 --------------------------- ql-rpc/src/rpc/duplex/mod.rs | 4 +- ql-rpc/src/rpc/notification/client.rs | 17 +++--- ql-rpc/src/rpc/notification/mod.rs | 2 +- ql-rpc/src/rpc/parts.rs | 15 +---- ql-rpc/src/rpc/progress/client.rs | 29 ++++++++- ql-rpc/src/rpc/progress/codec.rs | 86 +-------------------------- ql-rpc/src/rpc/progress/mod.rs | 3 +- ql-rpc/src/rpc/progress/server.rs | 28 +++++---- ql-rpc/src/rpc/request/client.rs | 40 +++++-------- ql-rpc/src/rpc/request/mod.rs | 2 +- ql-rpc/src/rpc/subscription/client.rs | 24 +++++++- ql-rpc/src/rpc/subscription/codec.rs | 13 +--- ql-rpc/src/rpc/subscription/mod.rs | 3 +- ql-rpc/src/rpc/upload/client.rs | 44 +++++++------- ql-rpc/src/rpc/upload/mod.rs | 2 +- ql-rpc/src/rpc/upload/server.rs | 29 +-------- ql-rpc/src/rpc/utils.rs | 41 +++++++++++-- ql-rpc/src/stream.rs | 1 + ql-runtime/src/rpc/mod.rs | 78 ++++++------------------ 24 files changed, 270 insertions(+), 386 deletions(-) delete mode 100644 ql-rpc/src/rpc/duplex/codec.rs diff --git a/ql-rpc/src/codec.rs b/ql-rpc/src/codec.rs index 51da527b..57f091f8 100644 --- a/ql-rpc/src/codec.rs +++ b/ql-rpc/src/codec.rs @@ -67,6 +67,15 @@ pub fn encode_value_part>(value: &T, out: & backpatch_length(out, payload_start); } +pub fn encode_tagged_value_part>( + tag: u8, + value: &T, + out: &mut B, +) { + out.put_u8(tag); + encode_value_part(value, out); +} + /// reads one length-delimited rpc value from buffered byte chunks pub fn reserve_length>(out: &mut B) -> usize { let start = out.as_mut().len(); diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index b586ca15..eecdaaee 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -1,13 +1,30 @@ use std::future::poll_fn; -use bytes::{BufMut, Bytes}; +use bytes::Bytes; use crate::{ - download::{Download, PartReadStep}, - rpc::parts::FrameKind, - DropResetRead, FramedPrefixStep, FramedReader, ResetCode, RpcCodec, RpcError, RpcRead, + download::Download, + parts::{PartFrameReader, PartReadStep}, + rpc::{parts::FrameKind, write_eof_value}, + DropResetRead, FramedPrefixStep, FramedReader, ResetCode, RpcError, RpcRead, RpcWrite, }; +pub async fn start( + reader: R, + mut writer: W, + request: &M::Request, +) -> Result, RpcError> +where + M: Download, + R: RpcRead, + W: RpcWrite, +{ + write_eof_value(&mut writer, request) + .await + .map_err(RpcError::Transport)?; + Ok(DownloadCall::new(reader)) +} + pub struct DownloadCall where M: Download, @@ -32,7 +49,7 @@ where R: RpcRead, { stream: DropResetRead, - reader: crate::download::PartFrameReader, + reader: PartFrameReader, } impl DownloadCall @@ -58,7 +75,7 @@ where value, DownloadReader { stream: self.stream, - reader: crate::download::PartFrameReader::::new(bytes), + reader: PartFrameReader::::new(bytes), }, )); } @@ -210,7 +227,3 @@ where } } } - -pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { - request.encode_value(out) -} diff --git a/ql-rpc/src/rpc/download/mod.rs b/ql-rpc/src/rpc/download/mod.rs index 5ed34aed..cfd4e374 100644 --- a/ql-rpc/src/rpc/download/mod.rs +++ b/ql-rpc/src/rpc/download/mod.rs @@ -4,16 +4,11 @@ use crate::RpcCodec; pub(crate) mod client; pub(crate) mod server; -pub use client::{encode_request, DownloadCall, DownloadPart, DownloadReader}; +pub use client::{start, DownloadCall, DownloadPart, DownloadReader}; pub use server::{ DownloadHandler, DownloadHandlerLocal, DownloadPartWriter, DownloadStart, DownloadWriter, }; -pub use crate::rpc::parts::{ - encode_body_chunk, encode_end_part, encode_finish, encode_part_header, PartFrameReader, - PartReadStep, -}; - /// rpc where the responder returns metadata first and then zero or more byte parts /// /// the typed portion of the response ends at [`Self::ResponseHeader`] diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index ccad5e05..424e1772 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -7,10 +7,22 @@ use std::{ use bytes::Bytes; use crate::{ - duplex::{codec, Duplex, EventReader, ReadStep}, - write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcCodec, RpcError, RpcRead, RpcWrite, + codec, duplex::Duplex, write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcCodec, + RpcError, RpcRead, RpcWrite, }; +pub fn start(writer: W, reader: R) -> DuplexCall +where + M: Duplex, + W: RpcWrite, + R: RpcRead, +{ + DuplexCall { + sender: DuplexSender::new(writer), + receiver: DuplexReceiver::new(reader), + } +} + pub struct DuplexCall where M: Duplex, @@ -54,7 +66,7 @@ where pub async fn send(&mut self, event: &T) -> Result<(), W::Error> { let writer = &mut self.writer; let mut encoded = Vec::new(); - codec::encode_event(event, &mut encoded); + codec::encode_value_part(event, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await } @@ -132,3 +144,45 @@ where DropResetRead::reset(&mut self.stream, code); } } + +enum ReadStep { + NeedMore, + Event(T), +} + +struct EventReader { + bytes: codec::ChunkQueue, + marker: PhantomData T>, +} + +impl Default for EventReader { + fn default() -> Self { + Self { + bytes: codec::ChunkQueue::default(), + marker: PhantomData, + } + } +} + +impl EventReader { + fn push(&mut self, chunk: Bytes) { + self.bytes.push(chunk); + } + + fn is_empty(&self) -> bool { + self.bytes.remaining() == 0 + } + + fn advance(&mut self) -> Result, RpcError> { + let Some(mut body) = self.bytes.try_take_part().map_err(RpcError::Protocol)? else { + return Ok(ReadStep::NeedMore); + }; + + let value = { + let value = T::decode_value(&mut body).map_err(RpcError::Codec)?; + drop(body); + value + }; + Ok(ReadStep::Event(value)) + } +} diff --git a/ql-rpc/src/rpc/duplex/codec.rs b/ql-rpc/src/rpc/duplex/codec.rs deleted file mode 100644 index 0392dfd9..00000000 --- a/ql-rpc/src/rpc/duplex/codec.rs +++ /dev/null @@ -1,86 +0,0 @@ -use std::marker::PhantomData; - -use bytes::{BufMut, Bytes}; - -use crate::{codec, RpcCodec, RpcError}; - -pub fn encode_event(event: &T, out: &mut (impl BufMut + AsMut<[u8]>)) -where - T: RpcCodec, -{ - codec::encode_value_part(event, out) -} - -pub enum ReadStep { - NeedMore, - Event(T), -} - -pub struct EventReader { - bytes: codec::ChunkQueue, - marker: PhantomData T>, -} - -impl Default for EventReader { - fn default() -> Self { - Self { - bytes: codec::ChunkQueue::default(), - marker: PhantomData, - } - } -} - -impl EventReader { - pub fn push(&mut self, chunk: Bytes) { - self.bytes.push(chunk); - } - - pub fn is_empty(&self) -> bool { - self.bytes.remaining() == 0 - } - - pub fn advance(&mut self) -> Result, RpcError> { - let Some(mut body) = self.bytes.try_take_part().map_err(RpcError::Protocol)? else { - return Ok(ReadStep::NeedMore); - }; - - let value = { - let value = T::decode_value(&mut body).map_err(RpcError::Codec)?; - drop(body); - value - }; - Ok(ReadStep::Event(value)) - } -} - -#[cfg(test)] -mod tests { - use bytes::Bytes; - - use super::{encode_event, EventReader, ReadStep}; - - #[test] - fn event_reader_emits_multiple_events() { - let mut encoded = Vec::new(); - encode_event(&b"one".to_vec(), &mut encoded); - encode_event(&b"two".to_vec(), &mut encoded); - - let mut reader = EventReader::>::default(); - reader.push(Bytes::from(encoded)); - - match reader.advance::().unwrap() { - ReadStep::Event(value) => { - assert_eq!(value, b"one".to_vec()); - } - _ => unreachable!(), - }; - - match reader.advance::().unwrap() { - ReadStep::Event(value) => { - assert_eq!(value, b"two".to_vec()); - assert!(reader.is_empty()); - } - _ => unreachable!(), - } - } -} diff --git a/ql-rpc/src/rpc/duplex/mod.rs b/ql-rpc/src/rpc/duplex/mod.rs index a9622029..11e15b6f 100644 --- a/ql-rpc/src/rpc/duplex/mod.rs +++ b/ql-rpc/src/rpc/duplex/mod.rs @@ -2,11 +2,9 @@ use super::Route; use crate::RpcCodec; pub(crate) mod client; -pub(crate) mod codec; pub(crate) mod server; -pub use client::{DuplexCall, DuplexReceiver, DuplexSender}; -pub use codec::{encode_event, EventReader, ReadStep}; +pub use client::{start, DuplexCall, DuplexReceiver, DuplexSender}; pub use server::{DuplexHandler, DuplexHandlerLocal, DuplexPeer}; /// rpc where both sides exchange typed events on the same stream diff --git a/ql-rpc/src/rpc/notification/client.rs b/ql-rpc/src/rpc/notification/client.rs index 72b6900a..ae9e1020 100644 --- a/ql-rpc/src/rpc/notification/client.rs +++ b/ql-rpc/src/rpc/notification/client.rs @@ -1,10 +1,11 @@ -use bytes::BufMut; +use crate::{notification::Notification, rpc::write_eof_value, ResetCode, RpcRead, RpcWrite}; -use crate::{notification::Notification, RpcCodec}; - -pub fn encode_notification( - payload: &M::Payload, - out: &mut (impl BufMut + AsMut<[u8]>), -) { - payload.encode_value(out) +pub async fn send(reader: R, mut writer: W, payload: &M::Payload) -> Result<(), W::Error> +where + M: Notification, + R: RpcRead, + W: RpcWrite, +{ + reader.reset(ResetCode::CANCELLED); + write_eof_value(&mut writer, payload).await } diff --git a/ql-rpc/src/rpc/notification/mod.rs b/ql-rpc/src/rpc/notification/mod.rs index 4740a64f..3da03299 100644 --- a/ql-rpc/src/rpc/notification/mod.rs +++ b/ql-rpc/src/rpc/notification/mod.rs @@ -4,7 +4,7 @@ use crate::RpcCodec; pub(crate) mod client; pub(crate) mod server; -pub use client::encode_notification; +pub use client::send; pub use server::{NotificationHandler, NotificationHandlerLocal}; /// one-way rpc that carries a single typed payload and no typed response diff --git a/ql-rpc/src/rpc/parts.rs b/ql-rpc/src/rpc/parts.rs index f5821680..a1832159 100644 --- a/ql-rpc/src/rpc/parts.rs +++ b/ql-rpc/src/rpc/parts.rs @@ -111,11 +111,11 @@ impl PartFrameReader { } pub fn encode_part_header(part_header: &H, out: &mut (impl BufMut + AsMut<[u8]>)) { - encode_tagged_value_part(FrameKind::PartHeader, part_header, out) + codec::encode_tagged_value_part(FrameKind::PartHeader.tag(), part_header, out) } pub fn encode_body_chunk(bytes: &Bytes, out: &mut (impl BufMut + AsMut<[u8]>)) { - encode_tagged_value_part(FrameKind::BodyChunk, bytes, out) + codec::encode_tagged_value_part(FrameKind::BodyChunk.tag(), bytes, out) } pub fn encode_end_part(out: &mut (impl BufMut + AsMut<[u8]>)) { @@ -155,17 +155,6 @@ impl TryFrom for FrameKind { } } -fn encode_tagged_value_part>( - kind: FrameKind, - value: &T, - out: &mut B, -) { - out.put_u8(kind.tag()); - let payload_start = codec::reserve_length(out); - value.encode_value(out); - codec::backpatch_length(out, payload_start); -} - fn encode_tagged_empty_part>(kind: FrameKind, out: &mut B) { out.put_u8(kind.tag()); let payload_start = codec::reserve_length(out); diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs index 755653c8..7be24bec 100644 --- a/ql-rpc/src/rpc/progress/client.rs +++ b/ql-rpc/src/rpc/progress/client.rs @@ -4,11 +4,36 @@ use std::{ task::{Context, Poll}, }; +use bytes::Bytes; + use crate::{ - progress::{Progress, ReadStep, ResponseReader}, - DropResetRead, Error, ResetCode, RpcError, RpcRead, + codec, finish_bytes, + progress::Progress, + rpc::progress::codec::{ReadStep, ResponseReader}, + write_bytes, DropResetRead, Error, ResetCode, RpcError, RpcRead, RpcWrite, }; +pub async fn start( + reader: R, + mut writer: W, + request: &M::Request, +) -> Result, RpcError> +where + M: Progress, + R: RpcRead, + W: RpcWrite, +{ + let mut payload = Vec::new(); + codec::encode_value_part(request, &mut payload); + write_bytes(&mut writer, Bytes::from(payload)) + .await + .map_err(RpcError::Transport)?; + finish_bytes(&mut writer) + .await + .map_err(RpcError::Transport)?; + Ok(ProgressCall::new(reader)) +} + pub struct ProgressCall where M: Progress, diff --git a/ql-rpc/src/rpc/progress/codec.rs b/ql-rpc/src/rpc/progress/codec.rs index b354e472..e98d381d 100644 --- a/ql-rpc/src/rpc/progress/codec.rs +++ b/ql-rpc/src/rpc/progress/codec.rs @@ -1,6 +1,6 @@ use std::marker::PhantomData; -use bytes::{BufMut, Bytes}; +use bytes::Bytes; use crate::{codec, progress::Progress, Error, RpcCodec, RpcError}; @@ -61,91 +61,9 @@ impl ResponseReader { } } -pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { - codec::encode_value_part(request, out) -} - -pub fn encode_progress(progress: &M::Progress, out: &mut (impl BufMut + AsMut<[u8]>)) { - encode_tagged_value_part(FrameKind::Progress, progress, out) -} - -pub fn encode_response(response: &M::Response, out: &mut (impl BufMut + AsMut<[u8]>)) { - encode_tagged_value_part(FrameKind::Response, response, out) -} - #[derive(Debug, Clone, Copy, PartialEq, Eq)] #[repr(u8)] -enum FrameKind { +pub enum FrameKind { Progress = 1, Response = 2, } - -fn encode_tagged_value_part>( - kind: FrameKind, - value: &T, - out: &mut B, -) { - out.put_u8(kind as u8); - let payload_start = codec::reserve_length(out); - value.encode_value(out); - codec::backpatch_length(out, payload_start); -} - -#[cfg(test)] -mod tests { - use bytes::Bytes; - - use super::{encode_progress, encode_response, ReadStep, ResponseReader}; - use crate::{progress::Progress, Route, RouteId, ServiceId}; - - const TEST_SERVICE: ServiceId = ServiceId([7; 16]); - - struct Watch; - - impl Route for Watch { - const SERVICE: ServiceId = TEST_SERVICE; - const ROUTE: RouteId = RouteId::from_u32(11); - } - - impl Progress for Watch { - type Error = core::convert::Infallible; - type Request = Vec; - type Progress = Vec; - type Response = Vec; - } - - #[test] - fn response_reader_emits_progress_then_response() { - let mut encoded = Vec::new(); - encode_progress::(&b"10%".to_vec(), &mut encoded); - encode_response::(&b"done".to_vec(), &mut encoded); - - let mut reader = ResponseReader::::default(); - reader.push(Bytes::from(encoded)); - - match reader.advance::().unwrap() { - ReadStep::Progress(value) => { - assert_eq!(value, b"10%".to_vec()); - } - _ => unreachable!(), - }; - match reader.advance::().unwrap() { - ReadStep::Response(value) => assert_eq!(value, b"done".to_vec()), - _ => unreachable!(), - } - } - - #[test] - fn response_reader_handles_response_only() { - let mut encoded = Vec::new(); - encode_response::(&b"done".to_vec(), &mut encoded); - - let mut reader = ResponseReader::::default(); - reader.push(Bytes::from(encoded)); - - match reader.advance::().unwrap() { - ReadStep::Response(value) => assert_eq!(value, b"done".to_vec()), - _ => unreachable!(), - } - } -} diff --git a/ql-rpc/src/rpc/progress/mod.rs b/ql-rpc/src/rpc/progress/mod.rs index b21c826d..bb93884d 100644 --- a/ql-rpc/src/rpc/progress/mod.rs +++ b/ql-rpc/src/rpc/progress/mod.rs @@ -5,8 +5,7 @@ pub(crate) mod client; pub(crate) mod codec; pub(crate) mod server; -pub use client::ProgressCall; -pub use codec::{encode_progress, encode_request, encode_response, ReadStep, ResponseReader}; +pub use client::{start, ProgressCall}; pub use server::{ProgressHandler, ProgressHandlerLocal, ProgressResponder}; /// rpc where the responder streams progress values before a final response diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index eabad230..8ea1add5 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -3,9 +3,9 @@ use std::{future::Future, marker::PhantomData}; use bytes::Bytes; use crate::{ - finish_bytes, - progress::{encode_progress, encode_response, Progress}, - rpc::read_framed_request, + codec, finish_bytes, + progress::Progress, + rpc::{progress::codec::FrameKind, read_framed_request}, write_bytes, Context, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, }; @@ -40,23 +40,16 @@ where M: Progress, W: RpcWrite, { - pub(crate) fn new(writer: W) -> Self { - Self { - writer: DropResetWrite::new(writer), - marker: PhantomData, - } - } - pub async fn send(&mut self, progress: M::Progress) -> Result<(), W::Error> { let writer = &mut self.writer; let mut encoded = Vec::new(); - encode_progress::(&progress, &mut encoded); + codec::encode_tagged_value_part(FrameKind::Progress as u8, &progress, &mut encoded); write_bytes(writer, Bytes::from(encoded)).await } pub async fn finish(mut self, response: M::Response) -> Result<(), W::Error> { let mut encoded = Vec::new(); - encode_response::(&response, &mut encoded); + codec::encode_tagged_value_part(FrameKind::Response as u8, &response, &mut encoded); write_bytes(&mut self.writer, Bytes::from(encoded)).await?; finish_bytes(&mut self.writer).await } @@ -94,5 +87,14 @@ pub(crate) async fn handle_progress_inner( } }; - handle(state, context, request, ProgressResponder::new(writer)).await; + handle( + state, + context, + request, + ProgressResponder { + writer: DropResetWrite::new(writer), + marker: PhantomData, + }, + ) + .await; } diff --git a/ql-rpc/src/rpc/request/client.rs b/ql-rpc/src/rpc/request/client.rs index 0d91263b..4520604a 100644 --- a/ql-rpc/src/rpc/request/client.rs +++ b/ql-rpc/src/rpc/request/client.rs @@ -1,29 +1,21 @@ -use bytes::BufMut; +use crate::{ + request::Request, + rpc::{read_eof_value, write_eof_value}, + RpcError, RpcRead, RpcWrite, +}; -use crate::{read_bytes, request::Request, ChunkQueue, RpcCodec, RpcError, RpcRead}; - -pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { - request.encode_value(out) -} - -pub fn encode_response(response: &M::Response, out: &mut (impl BufMut + AsMut<[u8]>)) { - response.encode_value(out) -} - -pub async fn read_response(mut reader: R) -> Result> +pub async fn call( + mut reader: R, + mut writer: W, + request: &M::Request, +) -> Result> where M: Request, - R: RpcRead, + R: RpcRead, + W: RpcWrite, { - let mut bytes = ChunkQueue::default(); - - while let Some(chunk) = read_bytes(&mut reader).await.map_err(RpcError::Transport)? { - bytes.push(chunk); - } - - let value = M::Response::decode_value(&mut bytes).map_err(RpcError::Codec)?; - if bytes.remaining() > 0 { - return Err(crate::Error::TrailingBytes.into()); - } - Ok(value) + write_eof_value(&mut writer, request) + .await + .map_err(RpcError::Transport)?; + read_eof_value::(&mut reader).await } diff --git a/ql-rpc/src/rpc/request/mod.rs b/ql-rpc/src/rpc/request/mod.rs index adf32597..14f15953 100644 --- a/ql-rpc/src/rpc/request/mod.rs +++ b/ql-rpc/src/rpc/request/mod.rs @@ -4,7 +4,7 @@ use crate::RpcCodec; pub(crate) mod client; pub(crate) mod server; -pub use client::{encode_request, encode_response, read_response}; +pub use client::call; pub use server::{RequestHandler, RequestHandlerLocal, Response}; /// request-response rpc with exactly one typed value in each direction diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index 95259406..ac44ac17 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -4,10 +4,30 @@ use std::{ }; use crate::{ - subscription::{ReadStep, ResponseReader, Subscription}, - DropResetRead, ResetCode, RpcError, RpcRead, + rpc::{ + subscription::codec::{ReadStep, ResponseReader}, + write_eof_value, + }, + subscription::Subscription, + DropResetRead, ResetCode, RpcError, RpcRead, RpcWrite, }; +pub async fn start( + reader: R, + mut writer: W, + request: &M::Request, +) -> Result, RpcError> +where + M: Subscription, + R: RpcRead, + W: RpcWrite, +{ + write_eof_value(&mut writer, request) + .await + .map_err(RpcError::Transport)?; + Ok(SubscriptionCall::new(reader)) +} + pub struct SubscriptionCall where M: Subscription, diff --git a/ql-rpc/src/rpc/subscription/codec.rs b/ql-rpc/src/rpc/subscription/codec.rs index 2025d4d2..234e99b7 100644 --- a/ql-rpc/src/rpc/subscription/codec.rs +++ b/ql-rpc/src/rpc/subscription/codec.rs @@ -1,20 +1,9 @@ use std::marker::PhantomData; -use bytes::{BufMut, Bytes}; +use bytes::Bytes; use crate::{codec, subscription::Subscription, RpcCodec, RpcError}; -pub fn encode_request( - request: &M::Request, - out: &mut (impl BufMut + AsMut<[u8]>), -) { - request.encode_value(out) -} - -pub fn encode_item(item: &M::Event, out: &mut (impl BufMut + AsMut<[u8]>)) { - codec::encode_value_part(item, out) -} - pub enum ReadStep { NeedMore, Item(M::Event), diff --git a/ql-rpc/src/rpc/subscription/mod.rs b/ql-rpc/src/rpc/subscription/mod.rs index 672eb9bc..e83b9751 100644 --- a/ql-rpc/src/rpc/subscription/mod.rs +++ b/ql-rpc/src/rpc/subscription/mod.rs @@ -5,8 +5,7 @@ pub(crate) mod client; pub(crate) mod codec; pub(crate) mod server; -pub use client::SubscriptionCall; -pub use codec::{encode_item, encode_request, ReadStep, ResponseReader}; +pub use client::{start, SubscriptionCall}; pub use server::{SubscriptionHandler, SubscriptionHandlerLocal, SubscriptionResponder}; /// rpc where one request opens a stream of typed events diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index 0cb06510..cd3f5fe6 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -1,13 +1,30 @@ -use bytes::{BufMut, Bytes}; +use bytes::Bytes; use crate::{ - read_bytes, - rpc::parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, + rpc::{ + parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, + read_eof_value, + }, upload::Upload, - write_bytes, ChunkQueue, DropResetRead, DropResetWrite, ResetCode, RpcCodec, RpcError, RpcRead, - RpcWrite, + write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcError, RpcRead, RpcWrite, }; +pub async fn start( + mut writer: W, + reader: R, + request: &M::Request, +) -> Result, W::Error> +where + M: Upload, + W: RpcWrite, + R: RpcRead, +{ + let mut payload = Vec::new(); + crate::codec::encode_value_part(request, &mut payload); + write_bytes(&mut writer, Bytes::from(payload)).await?; + Ok(UploadCall::new(writer, reader)) +} + pub struct UploadCall where M: Upload, @@ -66,18 +83,7 @@ where .map_err(RpcError::Transport)?; writer.queue_finish(); - let reader = &mut self.reader; - let mut bytes = ChunkQueue::default(); - - while let Some(chunk) = read_bytes(reader).await.map_err(RpcError::Transport)? { - bytes.push(chunk); - } - - let value = M::Response::decode_value(&mut bytes).map_err(RpcError::Codec)?; - if bytes.remaining() > 0 { - return Err(crate::Error::TrailingBytes.into()); - } - Ok(value) + read_eof_value::(&mut self.reader).await } fn reset(&mut self, code: ResetCode) { @@ -121,7 +127,3 @@ where } } } - -pub fn encode_request(request: &M::Request, out: &mut (impl BufMut + AsMut<[u8]>)) { - crate::codec::encode_value_part(request, out) -} diff --git a/ql-rpc/src/rpc/upload/mod.rs b/ql-rpc/src/rpc/upload/mod.rs index 9f96a824..433eeafb 100644 --- a/ql-rpc/src/rpc/upload/mod.rs +++ b/ql-rpc/src/rpc/upload/mod.rs @@ -4,7 +4,7 @@ use crate::RpcCodec; pub(crate) mod client; pub(crate) mod server; -pub use client::{encode_request, UploadCall, UploadPartWriter}; +pub use client::{start, UploadCall, UploadPartWriter}; pub use server::{UploadHandler, UploadHandlerLocal, UploadPart, UploadReader, UploadResponder}; /// rpc where the caller uploads zero or more byte parts after a typed request diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index 8ea9287a..4e32fd64 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -47,12 +47,7 @@ where finished: bool, } -pub struct UploadResponder -where - W: RpcWrite, -{ - inner: Response, -} +pub type UploadResponder = Response; impl UploadReader where @@ -164,26 +159,6 @@ where } } -impl UploadResponder -where - T: crate::RpcCodec, - W: RpcWrite, -{ - pub(crate) fn new(writer: W) -> Self { - Self { - inner: Response::new(writer), - } - } - - pub async fn respond(self, response: T) -> Result<(), W::Error> { - self.inner.respond(response).await - } - - pub fn reset(self, code: ResetCode) { - self.inner.reset(code); - } -} - pub(crate) async fn handle_upload_inner( state: S, context: Context, @@ -227,7 +202,7 @@ pub(crate) async fn handle_upload_inner( stream: DropResetRead::new(reader), reader: PartFrameReader::new(buffered), }, - UploadResponder::new(writer), + Response::new(writer), ) .await; } diff --git a/ql-rpc/src/rpc/utils.rs b/ql-rpc/src/rpc/utils.rs index 457fe982..5f210b3f 100644 --- a/ql-rpc/src/rpc/utils.rs +++ b/ql-rpc/src/rpc/utils.rs @@ -1,10 +1,41 @@ +use bytes::Bytes; + use crate::{ - read_bytes, ChunkQueue, Error, FramedPrefixStep, FramedReadStep, FramedReader, RouterConfig, - RpcCodec, RpcError, RpcRead, + finish_bytes, read_bytes, write_bytes, ChunkQueue, Error, FramedPrefixStep, FramedReadStep, + FramedReader, RouterConfig, RpcCodec, RpcError, RpcRead, RpcWrite, }; +pub async fn write_eof_value(writer: &mut W, value: &T) -> Result<(), W::Error> +where + T: RpcCodec, + W: RpcWrite, +{ + let mut encoded = Vec::new(); + value.encode_value(&mut encoded); + write_bytes(writer, Bytes::from(encoded)).await?; + finish_bytes(writer).await +} + +pub async fn read_eof_value(reader: &mut R) -> Result> +where + T: RpcCodec, + R: RpcRead, +{ + let mut bytes = ChunkQueue::default(); + + while let Some(chunk) = read_bytes(reader).await.map_err(RpcError::Transport)? { + bytes.push(chunk); + } + + let value = T::decode_value(&mut bytes).map_err(RpcError::Codec)?; + if bytes.remaining() > 0 { + return Err(RpcError::Protocol(Error::TrailingBytes)); + } + Ok(value) +} + /// reads one length-delimited value and rejects trailing bytes -pub(crate) async fn read_framed_request( +pub async fn read_framed_request( reader: &mut R, config: RouterConfig, ) -> Result> @@ -38,7 +69,7 @@ where } /// reads one length-delimited value and returns any bytes already buffered -pub(crate) async fn read_framed_request_prefix( +pub async fn read_framed_request_prefix( reader: &mut R, config: RouterConfig, ) -> Result<(T, ChunkQueue), RpcError> @@ -66,7 +97,7 @@ where } /// reads one eof-delimited value up to the configured request limit -pub(crate) async fn read_eof_request( +pub async fn read_eof_request( reader: &mut R, config: RouterConfig, ) -> Result> diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index 37d14d95..e9b09509 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -164,6 +164,7 @@ mod drop { self.inner.as_mut().unwrap().poll_write(bytes, cx) } + #[track_caller] fn queue_finish(&mut self) { self.inner.as_mut().unwrap().queue_finish(); } diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index b14b0b25..44a1d07c 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -1,8 +1,7 @@ mod adapter; -use bytes::Bytes; use ql_fsm::OpenStreamParams; -use ql_rpc::{download, duplex, notification, progress, request, subscription, upload}; +use ql_rpc::{download, duplex, notification, progress, request, subscription, upload, Route}; use crate::{QlStream, QlStreamError, RuntimeHandle, StreamReader, StreamWriter}; @@ -18,31 +17,18 @@ impl RpcHandle { where M: notification::Notification, { - let mut payload = Vec::new(); - notification::encode_notification::(event, &mut payload); - let mut stream = self.open_rpc_stream::().await?; - stream.reader.reset(ql_rpc::ResetCode::CANCELLED); - stream - .writer - .write(Bytes::from(payload)) + let stream = self.open_rpc_stream::().await?; + notification::send::(stream.reader, stream.writer, event) .await - .map_err(ql_rpc::RpcError::Transport)?; - stream - .writer - .finish() - .await - .map_err(ql_rpc::RpcError::Transport)?; - Ok(()) + .map_err(ql_rpc::RpcError::Transport) } pub async fn request(&self, request: &M::Request) -> RpcResult where M: request::Request, { - let mut payload = Vec::new(); - request::encode_request::(request, &mut payload); - let response = self.start_request::(payload).await?; - request::read_response::(response).await + let stream = self.open_rpc_stream::().await?; + request::call::(stream.reader, stream.writer, request).await } pub async fn subscribe( @@ -52,10 +38,8 @@ impl RpcHandle { where M: subscription::Subscription, { - let mut payload = Vec::new(); - subscription::encode_request::(request, &mut payload); - let response = self.start_request::(payload).await?; - Ok(subscription::SubscriptionCall::new(response)) + let stream = self.open_rpc_stream::().await?; + subscription::start::(stream.reader, stream.writer, request).await } pub async fn download( @@ -65,10 +49,8 @@ impl RpcHandle { where M: download::Download, { - let mut payload = Vec::new(); - download::encode_request::(request, &mut payload); - let response = self.start_request::(payload).await?; - Ok(download::DownloadCall::new(response)) + let stream = self.open_rpc_stream::().await?; + download::start::(stream.reader, stream.writer, request).await } pub async fn progress( @@ -78,10 +60,8 @@ impl RpcHandle { where M: progress::Progress, { - let mut payload = Vec::new(); - progress::encode_request::(request, &mut payload); - let response = self.start_request::(payload).await?; - Ok(progress::ProgressCall::new(response)) + let stream = self.open_rpc_stream::().await?; + progress::start::(stream.reader, stream.writer, request).await } pub async fn upload( @@ -91,15 +71,10 @@ impl RpcHandle { where M: upload::Upload, { - let mut payload = Vec::new(); - upload::encode_request::(request, &mut payload); - let mut stream = self.open_rpc_stream::().await?; - stream - .writer - .write(Bytes::from(payload)) + let stream = self.open_rpc_stream::().await?; + upload::start::(stream.writer, stream.reader, request) .await - .map_err(ql_rpc::RpcError::Transport)?; - Ok(upload::UploadCall::new(stream.writer, stream.reader)) + .map_err(ql_rpc::RpcError::Transport) } pub async fn duplex( @@ -108,11 +83,8 @@ impl RpcHandle { where M: duplex::Duplex, { - let stream = self.open_rpc_stream::().await?; - Ok(duplex::DuplexCall { - sender: duplex::DuplexSender::new(stream.writer), - receiver: duplex::DuplexReceiver::new(stream.reader), - }) + let stream = self.open_rpc_stream::().await?; + Ok(duplex::start::(stream.writer, stream.reader)) } } @@ -121,21 +93,7 @@ impl RpcHandle { Self { inner } } - async fn start_request( - &self, - payload: Vec, - ) -> RpcResult { - let mut stream = self.open_rpc_stream::().await?; - stream - .writer - .write(Bytes::from(payload)) - .await - .map_err(ql_rpc::RpcError::Transport)?; - stream.writer.queue_finish(); - Ok(stream.reader) - } - - async fn open_rpc_stream(&self) -> RpcResult { + async fn open_rpc_stream(&self) -> RpcResult { self.inner .open_stream(OpenStreamParams { service_id: R::SERVICE, From f10e8f8862ff9c6d9011360dbf7857290bc2d85a Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 24 Jun 2026 22:47:03 -0400 Subject: [PATCH 26/59] ql-rpc: separate streaminfo --- Cargo.lock | 1 + ql-common/src/lib.rs | 8 ++ ql-rpc/src/lib.rs | 2 +- ql-rpc/src/router/mod.rs | 17 ++-- ql-rpc/src/rpc/download/client.rs | 13 ++- ql-rpc/src/rpc/duplex/client.rs | 8 +- ql-rpc/src/rpc/notification/client.rs | 8 +- ql-rpc/src/rpc/progress/client.rs | 13 ++- ql-rpc/src/rpc/request/client.rs | 13 ++- ql-rpc/src/rpc/subscription/client.rs | 13 ++- ql-rpc/src/rpc/upload/client.rs | 13 ++- ql-rpc/src/stream.rs | 6 +- ql-runtime/Cargo.toml | 1 + ql-runtime/src/driver/mod.rs | 19 ++-- ql-runtime/src/driver/test.rs | 5 +- ql-runtime/src/io/reader.rs | 5 +- ql-runtime/src/io/writer.rs | 5 +- ql-runtime/src/platform.rs | 17 +--- ql-runtime/src/rpc/adapter.rs | 22 +--- ql-runtime/src/rpc/mod.rs | 14 +-- ql-runtime/src/tests/handshake.rs | 2 +- ql-runtime/src/tests/mod.rs | 22 ++-- ql-runtime/src/tests/rpc.rs | 138 ++++++++++++-------------- ql-runtime/src/tests/session.rs | 6 +- ql-runtime/src/tests/stream.rs | 20 ++-- 25 files changed, 178 insertions(+), 213 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 008d5fed..54ffae1f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1086,6 +1086,7 @@ dependencies = [ "log", "loom", "oneshot", + "ql-common", "ql-fsm", "ql-rpc", "ql-wire", diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index 5dfd6851..3904b4b3 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -84,3 +84,11 @@ impl std::fmt::Display for ServiceId { varint_wrapper!(RouteId); varint_wrapper!(StreamId); + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct StreamInfo { + pub qid: QID, + pub stream_id: StreamId, + pub service_id: ServiceId, + pub route_id: RouteId, +} diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index 09a91ff3..40dd5a1a 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -14,7 +14,7 @@ pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; use framed_value::*; -pub use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId, StreamId, QID}; +pub use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId, StreamId, StreamInfo, QID}; pub use router::*; pub use rpc::*; pub use stream::*; diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index 985a9c5b..5d32b529 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -1,4 +1,4 @@ -use crate::{ResetCode, RouteId, ServiceId, StreamId, QID}; +use crate::{ResetCode, RouteId, ServiceId, StreamId, StreamInfo, QID}; mod builder; mod config; @@ -88,13 +88,14 @@ where RouterBuilder::::new(spawner) } - pub fn handle(&self, stream: St) -> Option<(RouteId, Sp::Handle)> { - let service_id = stream.service_id(); - let route_id = stream.route_id(); - let context = Context { - qid: stream.qid(), - stream_id: stream.stream_id(), - }; + pub fn handle(&self, info: StreamInfo, stream: St) -> Option<(RouteId, Sp::Handle)> { + let StreamInfo { + qid, + stream_id, + service_id, + route_id, + } = info; + let context = Context { qid, stream_id }; let key = RouteKey { service_id, route_id, diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index eecdaaee..9001e15c 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -6,19 +6,18 @@ use crate::{ download::Download, parts::{PartFrameReader, PartReadStep}, rpc::{parts::FrameKind, write_eof_value}, - DropResetRead, FramedPrefixStep, FramedReader, ResetCode, RpcError, RpcRead, RpcWrite, + DropResetRead, FramedPrefixStep, FramedReader, ResetCode, RpcError, RpcRead, RpcStream, }; -pub async fn start( - reader: R, - mut writer: W, +pub async fn start( + stream: St, request: &M::Request, -) -> Result, RpcError> +) -> Result, RpcError> where M: Download, - R: RpcRead, - W: RpcWrite, + St: RpcStream, { + let (reader, mut writer) = stream.split(); write_eof_value(&mut writer, request) .await .map_err(RpcError::Transport)?; diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index 424e1772..1f9ea080 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -8,15 +8,15 @@ use bytes::Bytes; use crate::{ codec, duplex::Duplex, write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcCodec, - RpcError, RpcRead, RpcWrite, + RpcError, RpcRead, RpcStream, RpcWrite, }; -pub fn start(writer: W, reader: R) -> DuplexCall +pub fn start(stream: St) -> DuplexCall where M: Duplex, - W: RpcWrite, - R: RpcRead, + St: RpcStream, { + let (reader, writer) = stream.split(); DuplexCall { sender: DuplexSender::new(writer), receiver: DuplexReceiver::new(reader), diff --git a/ql-rpc/src/rpc/notification/client.rs b/ql-rpc/src/rpc/notification/client.rs index ae9e1020..27bd23b4 100644 --- a/ql-rpc/src/rpc/notification/client.rs +++ b/ql-rpc/src/rpc/notification/client.rs @@ -1,11 +1,11 @@ -use crate::{notification::Notification, rpc::write_eof_value, ResetCode, RpcRead, RpcWrite}; +use crate::{notification::Notification, rpc::write_eof_value, ResetCode, RpcRead, RpcStream}; -pub async fn send(reader: R, mut writer: W, payload: &M::Payload) -> Result<(), W::Error> +pub async fn send(stream: St, payload: &M::Payload) -> Result<(), St::Error> where M: Notification, - R: RpcRead, - W: RpcWrite, + St: RpcStream, { + let (reader, mut writer) = stream.split(); reader.reset(ResetCode::CANCELLED); write_eof_value(&mut writer, payload).await } diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs index 7be24bec..61d3fee0 100644 --- a/ql-rpc/src/rpc/progress/client.rs +++ b/ql-rpc/src/rpc/progress/client.rs @@ -10,19 +10,18 @@ use crate::{ codec, finish_bytes, progress::Progress, rpc::progress::codec::{ReadStep, ResponseReader}, - write_bytes, DropResetRead, Error, ResetCode, RpcError, RpcRead, RpcWrite, + write_bytes, DropResetRead, Error, ResetCode, RpcError, RpcRead, RpcStream, }; -pub async fn start( - reader: R, - mut writer: W, +pub async fn start( + stream: St, request: &M::Request, -) -> Result, RpcError> +) -> Result, RpcError> where M: Progress, - R: RpcRead, - W: RpcWrite, + St: RpcStream, { + let (reader, mut writer) = stream.split(); let mut payload = Vec::new(); codec::encode_value_part(request, &mut payload); write_bytes(&mut writer, Bytes::from(payload)) diff --git a/ql-rpc/src/rpc/request/client.rs b/ql-rpc/src/rpc/request/client.rs index 4520604a..2bd08ef6 100644 --- a/ql-rpc/src/rpc/request/client.rs +++ b/ql-rpc/src/rpc/request/client.rs @@ -1,19 +1,18 @@ use crate::{ request::Request, rpc::{read_eof_value, write_eof_value}, - RpcError, RpcRead, RpcWrite, + RpcError, RpcStream, }; -pub async fn call( - mut reader: R, - mut writer: W, +pub async fn call( + stream: St, request: &M::Request, -) -> Result> +) -> Result> where M: Request, - R: RpcRead, - W: RpcWrite, + St: RpcStream, { + let (mut reader, mut writer) = stream.split(); write_eof_value(&mut writer, request) .await .map_err(RpcError::Transport)?; diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index ac44ac17..24d22aa6 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -9,19 +9,18 @@ use crate::{ write_eof_value, }, subscription::Subscription, - DropResetRead, ResetCode, RpcError, RpcRead, RpcWrite, + DropResetRead, ResetCode, RpcError, RpcRead, RpcStream, }; -pub async fn start( - reader: R, - mut writer: W, +pub async fn start( + stream: St, request: &M::Request, -) -> Result, RpcError> +) -> Result, RpcError> where M: Subscription, - R: RpcRead, - W: RpcWrite, + St: RpcStream, { + let (reader, mut writer) = stream.split(); write_eof_value(&mut writer, request) .await .map_err(RpcError::Transport)?; diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index cd3f5fe6..b56400c9 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -6,19 +6,18 @@ use crate::{ read_eof_value, }, upload::Upload, - write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcError, RpcRead, RpcWrite, + write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcError, RpcRead, RpcStream, RpcWrite, }; -pub async fn start( - mut writer: W, - reader: R, +pub async fn start( + stream: St, request: &M::Request, -) -> Result, W::Error> +) -> Result, St::Error> where M: Upload, - W: RpcWrite, - R: RpcRead, + St: RpcStream, { + let (reader, mut writer) = stream.split(); let mut payload = Vec::new(); crate::codec::encode_value_part(request, &mut payload); write_bytes(&mut writer, Bytes::from(payload)).await?; diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index e9b09509..043e1673 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -5,17 +5,13 @@ use std::{ use bytes::Bytes; -use crate::{ResetCode, RouteId, ServiceId, StreamId, QID}; +use crate::ResetCode; pub trait RpcStream { type Error; type Reader: RpcRead; type Writer: RpcWrite; - fn qid(&self) -> QID; - fn stream_id(&self) -> StreamId; - fn service_id(&self) -> ServiceId; - fn route_id(&self) -> RouteId; fn split(self) -> (Self::Reader, Self::Writer); } diff --git a/ql-runtime/Cargo.toml b/ql-runtime/Cargo.toml index 564cee1c..1fc53189 100644 --- a/ql-runtime/Cargo.toml +++ b/ql-runtime/Cargo.toml @@ -18,6 +18,7 @@ futures-lite = { version = "2.5" } log = { version = "0.4", optional = true } oneshot = { version = "0.1.11" } ql-fsm = { workspace = true } +ql-common = { workspace = true } ql-rpc = { workspace = true, optional = true } ql-wire = { workspace = true } diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 7b4b2ca4..5d4a5b98 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -22,7 +22,7 @@ use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, Ou use crate::{ command::Command, io, log, - platform::{QlInbound, QlInboundStream, QlPlatform, QlTimer}, + platform::{QlInbound, QlPlatform, QlTimer, StreamInfo}, QlStreamError, Runtime, }; @@ -370,14 +370,15 @@ impl DriverState { "delivering inbound stream to platform: service_id={service_id} route_id={route_id} stream_id={stream_id}", ); - platform.handle_inbound(QlInboundStream { - qid, - route_id, - service_id, - stream_id, - writer, - reader, - }); + platform.handle_inbound( + StreamInfo { + qid, + stream_id, + service_id, + route_id, + }, + crate::QlStream { writer, reader }, + ); } fn handle_inbound_readable(&mut self, fsm: &mut QlFsm, stream_id: StreamId) { diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs index 58397247..d891647b 100644 --- a/ql-runtime/src/driver/test.rs +++ b/ql-runtime/src/driver/test.rs @@ -4,8 +4,7 @@ use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto, QID}; use super::*; use crate::{ driver::state::{InboundIo, OutboundIo}, - io, - platform::QlInbound, + io, StreamInfo, }; pub struct NoopTimer; @@ -40,7 +39,7 @@ impl QlPlatform for NoopCrypto { fn handle_peer_status(&self, _peer: Option, _status: ql_fsm::PeerStatus) {} - fn handle_inbound(&self, _event: QlInboundStream) {} + fn handle_inbound(&self, _info: StreamInfo, _stream: crate::QlStream) {} } impl QlInbound for NoopInbound { diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs index 9482ca4a..28b8f154 100644 --- a/ql-runtime/src/io/reader.rs +++ b/ql-runtime/src/io/reader.rs @@ -4,6 +4,7 @@ use std::{ }; use bytes::Bytes; +use ql_common::ResetCode; use ql_fsm::StreamResetTarget; use ql_wire::{ResetCode, StreamId}; @@ -48,10 +49,6 @@ impl StreamReader { } } - pub fn stream_id(&self) -> StreamId { - self.rx.stream_id() - } - pub fn poll_read( &mut self, cx: &mut Context<'_>, diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs index 89f637cc..ea67b7e1 100644 --- a/ql-runtime/src/io/writer.rs +++ b/ql-runtime/src/io/writer.rs @@ -4,6 +4,7 @@ use std::{ }; use bytes::Bytes; +use ql_common::ResetCode; use ql_fsm::StreamResetTarget; use ql_wire::{ResetCode, StreamId}; @@ -47,10 +48,6 @@ impl StreamWriter { } } - pub fn stream_id(&self) -> StreamId { - self.tx.stream_id() - } - pub fn poll_write( &mut self, bytes: &mut Bytes, diff --git a/ql-runtime/src/platform.rs b/ql-runtime/src/platform.rs index 24a9108d..cf45aedb 100644 --- a/ql-runtime/src/platform.rs +++ b/ql-runtime/src/platform.rs @@ -5,20 +5,9 @@ use std::{ time::Instant, }; +pub use ql_common::StreamInfo; use ql_fsm::{PeerStatus, ReceiveError}; -use ql_wire::{PeerBundle, QlCrypto, RouteId, ServiceId, StreamId, QID}; - -use crate::{StreamReader, StreamWriter}; - -#[derive(Debug)] -pub struct QlInboundStream { - pub qid: QID, - pub service_id: ServiceId, - pub route_id: RouteId, - pub stream_id: StreamId, - pub writer: StreamWriter, - pub reader: StreamReader, -} +use ql_wire::{PeerBundle, QlCrypto, QID}; pub trait QlTimer { fn set_deadline(self: Pin<&mut Self>, deadline: Option); @@ -48,6 +37,6 @@ pub trait QlPlatform: QlCrypto { fn persist_peer(&self, peer: PeerBundle); fn handle_peer_status(&self, peer: Option, status: PeerStatus); - fn handle_inbound(&self, event: QlInboundStream); + fn handle_inbound(&self, info: StreamInfo, stream: crate::QlStream); fn handle_recv_error(&self, _error: ReceiveError) {} } diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index 93a220a5..857bf0bc 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -1,31 +1,15 @@ use std::task::{Context as TaskContext, Poll}; use bytes::Bytes; -use ql_rpc::{ResetCode, RouteId, RpcRead, RpcStream, RpcWrite, ServiceId, StreamId, QID}; +use ql_rpc::{ResetCode, RpcRead, RpcStream, RpcWrite}; -use crate::{QlInboundStream, QlStreamError, StreamReader, StreamWriter}; +use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; -impl RpcStream for QlInboundStream { +impl RpcStream for QlStream { type Error = QlStreamError; type Reader = StreamReader; type Writer = StreamWriter; - fn qid(&self) -> QID { - self.qid - } - - fn stream_id(&self) -> StreamId { - self.stream_id - } - - fn service_id(&self) -> ServiceId { - self.service_id - } - - fn route_id(&self) -> RouteId { - self.route_id - } - fn split(self) -> (Self::Reader, Self::Writer) { (self.reader, self.writer) } diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index 44a1d07c..d58efa17 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -18,7 +18,7 @@ impl RpcHandle { M: notification::Notification, { let stream = self.open_rpc_stream::().await?; - notification::send::(stream.reader, stream.writer, event) + notification::send::(stream, event) .await .map_err(ql_rpc::RpcError::Transport) } @@ -28,7 +28,7 @@ impl RpcHandle { M: request::Request, { let stream = self.open_rpc_stream::().await?; - request::call::(stream.reader, stream.writer, request).await + request::call::(stream, request).await } pub async fn subscribe( @@ -39,7 +39,7 @@ impl RpcHandle { M: subscription::Subscription, { let stream = self.open_rpc_stream::().await?; - subscription::start::(stream.reader, stream.writer, request).await + subscription::start::(stream, request).await } pub async fn download( @@ -50,7 +50,7 @@ impl RpcHandle { M: download::Download, { let stream = self.open_rpc_stream::().await?; - download::start::(stream.reader, stream.writer, request).await + download::start::(stream, request).await } pub async fn progress( @@ -61,7 +61,7 @@ impl RpcHandle { M: progress::Progress, { let stream = self.open_rpc_stream::().await?; - progress::start::(stream.reader, stream.writer, request).await + progress::start::(stream, request).await } pub async fn upload( @@ -72,7 +72,7 @@ impl RpcHandle { M: upload::Upload, { let stream = self.open_rpc_stream::().await?; - upload::start::(stream.writer, stream.reader, request) + upload::start::(stream, request) .await .map_err(ql_rpc::RpcError::Transport) } @@ -84,7 +84,7 @@ impl RpcHandle { M: duplex::Duplex, { let stream = self.open_rpc_stream::().await?; - Ok(duplex::start::(stream.writer, stream.reader)) + Ok(duplex::start::(stream)) } } diff --git a/ql-runtime/src/tests/handshake.rs b/ql-runtime/src/tests/handshake.rs index b2666f02..fb915d4d 100644 --- a/ql-runtime/src/tests/handshake.rs +++ b/ql-runtime/src/tests/handshake.rs @@ -82,7 +82,7 @@ async fn rejected_session_write_is_reissued() { await_status(&status_b, Some(identity_a.qid), PeerStatus::Connected).await; let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let request = read_all(stream.reader).await.unwrap(); stream.writer.finish().await.unwrap(); request diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index 19c2ad37..00ac8a74 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -20,10 +20,14 @@ use ql_wire::{ use tokio::{task::LocalSet, time::Sleep}; use crate::{ - new_runtime, platform::QlTimer, NoSessionError, PairingInvite, QlFsmConfig, QlInboundStream, - QlStreamError, RuntimeConfig, RuntimeHandle, + new_runtime, + platform::{QlTimer, StreamInfo}, + NoSessionError, PairingInvite, QlFsmConfig, QlStream, QlStreamError, RuntimeConfig, + RuntimeHandle, }; +type InboundStream = (StreamInfo, QlStream); + mod handshake; #[cfg(feature = "rpc")] mod rpc; @@ -97,7 +101,7 @@ struct TestPlatform { _inbound_messages_tx: Sender>, inbound_messages: Option>>, status: Sender, - inbound: Option>, + inbound: Option>, crypto: SoftwareCrypto, encrypted_write_counter: AtomicUsize, fail_encrypted_write_at: Option, @@ -121,7 +125,7 @@ type TestPlatformPartsWithInbound = ( Receiver>, Sender>, Receiver, - Receiver, + Receiver, ); impl TestPlatform { @@ -151,7 +155,7 @@ impl TestPlatform { } fn new_inner( - inbound: Option>, + inbound: Option>, fail_encrypted_write_at: Option, write_delay: Duration, write_stats: Option, @@ -183,7 +187,7 @@ struct TestSide { handle: RuntimeHandle, status: Receiver, peer: QID, - inbound: Receiver, + inbound: Receiver, } struct TestPair { @@ -310,7 +314,7 @@ impl TestPair { .await; } - fn take_inbound(&mut self, side: Side) -> Receiver { + fn take_inbound(&mut self, side: Side) -> Receiver { let replacement = async_channel::unbounded().1; std::mem::replace(&mut self.side_mut(side).inbound, replacement) } @@ -450,9 +454,9 @@ impl crate::platform::QlPlatform for TestPlatform { let _ = self.status.try_send(StatusEvent { peer, status }); } - fn handle_inbound(&self, event: QlInboundStream) { + fn handle_inbound(&self, info: StreamInfo, stream: QlStream) { if let Some(tx) = &self.inbound { - let _ = tx.try_send(event); + let _ = tx.try_send((info, stream)); } } } diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index e570f9b6..c26f5c03 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -17,7 +17,7 @@ use ql_rpc::{ }; use super::*; -use crate::{QlInboundStream, QlStreamError, StreamWriter}; +use crate::{QlStream, QlStreamError, StreamWriter}; const TEST_SERVICE: ServiceId = ServiceId([7; 16]); @@ -154,7 +154,7 @@ async fn rpc_request() { seen: Arc>>, } - impl RequestHandler for RouterState { + impl RequestHandler for RouterState { async fn handle( self, _context: Context, @@ -174,13 +174,13 @@ async fn rpc_request() { let seen = Arc::new(Mutex::new(Vec::new())); let router = - ql_rpc::Router::<_, QlInboundStream, TokioSendSpawner>::builder_send(TokioSendSpawner) + ql_rpc::Router::<_, QlStream, TokioSendSpawner>::builder_send(TokioSendSpawner) .request::() .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { let fut = assert_send(fut); fut.await.unwrap(); } @@ -210,7 +210,7 @@ async fn rpc_notification() { seen: Rc>>>, } - impl NotificationHandlerLocal for RouterState { + impl NotificationHandlerLocal for RouterState { async fn handle(self, _context: Context, payload: Vec) { self.seen.borrow_mut().push(payload); } @@ -222,15 +222,14 @@ async fn rpc_notification() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( - TokioLocalSpawner, - ) - .notification::() - .build(RouterState { seen: seen.clone() }); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .notification::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -256,7 +255,7 @@ async fn rpc_subscrption() { seen: Rc>>>, } - impl SubscriptionHandlerLocal for RouterState { + impl SubscriptionHandlerLocal for RouterState { async fn handle( self, _context: Context, @@ -277,15 +276,14 @@ async fn rpc_subscrption() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( - TokioLocalSpawner, - ) - .subscription::() - .build(RouterState { seen: seen.clone() }); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .subscription::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -316,7 +314,7 @@ async fn rpc_router_enforces_max_request_bytes() { #[derive(Clone)] struct LimitedState; - impl RequestHandlerLocal for LimitedState { + impl RequestHandlerLocal for LimitedState { async fn handle( self, _context: Context, @@ -331,16 +329,15 @@ async fn rpc_router_enforces_max_request_bytes() { let mut pair = TestPair::new(default_runtime_config()); pair.connect_and_wait(Side::A).await; let inbound_b = pair.take_inbound(Side::B); - let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( - TokioLocalSpawner, - ) - .max_request_bytes(4) - .request::() - .build(LimitedState); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .max_request_bytes(4) + .request::() + .build(LimitedState); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -368,7 +365,7 @@ async fn rpc_progress() { seen: Rc>>>, } - impl ProgressHandlerLocal for RouterState { + impl ProgressHandlerLocal for RouterState { async fn handle( self, _context: Context, @@ -389,15 +386,14 @@ async fn rpc_progress() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( - TokioLocalSpawner, - ) - .progress::() - .build(RouterState { seen: seen.clone() }); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .progress::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -426,7 +422,7 @@ async fn rpc_download() { seen: Rc>>>, } - impl DownloadHandlerLocal for RouterState { + impl DownloadHandlerLocal for RouterState { async fn handle( self, _context: Context, @@ -453,15 +449,14 @@ async fn rpc_download() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( - TokioLocalSpawner, - ) - .download::() - .build(RouterState { seen: seen.clone() }); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .download::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -513,7 +508,7 @@ async fn rpc_download_complete() { seen: Rc>>>, } - impl DownloadHandlerLocal for RouterState { + impl DownloadHandlerLocal for RouterState { async fn handle( self, _context: Context, @@ -531,15 +526,14 @@ async fn rpc_download_complete() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( - TokioLocalSpawner, - ) - .download::() - .build(RouterState { seen: seen.clone() }); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .download::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -570,7 +564,7 @@ async fn rpc_upload() { uploads: Rc>>>, } - impl UploadHandlerLocal for RouterState { + impl UploadHandlerLocal for RouterState { async fn handle( self, _context: Context, @@ -604,18 +598,17 @@ async fn rpc_upload() { let requests = Rc::new(RefCell::new(Vec::new())); let uploads = Rc::new(RefCell::new(Vec::new())); - let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( - TokioLocalSpawner, - ) - .upload::() - .build(RouterState { - requests: requests.clone(), - uploads: uploads.clone(), - }); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .upload::() + .build(RouterState { + requests: requests.clone(), + uploads: uploads.clone(), + }); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -653,7 +646,7 @@ async fn rpc_duplex() { seen: Rc>>>, } - impl DuplexHandlerLocal for RouterState { + impl DuplexHandlerLocal for RouterState { async fn handle( self, _context: Context, @@ -681,15 +674,14 @@ async fn rpc_duplex() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = ql_rpc::Router::<_, QlInboundStream, TokioLocalSpawner>::builder_local( - TokioLocalSpawner, - ) - .duplex::() - .build(RouterState { seen: seen.clone() }); + let router = + ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) + .duplex::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(inbound) { + let (info, stream) = inbound_b.recv().await.unwrap(); + if let Some((_, fut)) = router.handle(info, stream) { fut.await.unwrap(); } }); diff --git a/ql-runtime/src/tests/session.rs b/ql-runtime/src/tests/session.rs index ae58d933..89066e00 100644 --- a/ql-runtime/src/tests/session.rs +++ b/ql-runtime/src/tests/session.rs @@ -21,7 +21,7 @@ async fn close_session_aborts_active_streams_and_allows_reconnect() { pair.connect_and_wait(Side::A).await; let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let mut reader = stream.reader; assert_eq!( @@ -86,7 +86,7 @@ async fn unpair_aborts_active_streams_and_prevents_reconnect() { pair.connect_and_wait(Side::A).await; let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let mut reader = stream.reader; assert_eq!( @@ -193,7 +193,7 @@ async fn session_timeout_disconnects_and_fails_pending_open() { await_status(&status_b, Some(identity_a.qid), PeerStatus::Connected).await; let responder_task = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let _ = read_all(stream.reader).await; let err = stream.writer.finish().await.unwrap_err(); assert!(matches!(err, QlStreamError::NoSession)); diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index 7d9e7bb5..83baf7a2 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -14,7 +14,7 @@ async fn open_stream_duplex_happy_path() { let inbound_b = pair.take_inbound(Side::B); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); + let (_, inbound) = inbound_b.recv().await.unwrap(); let mut writer = inbound.writer; let mut reader = inbound.reader; @@ -69,7 +69,7 @@ async fn large_stream_payload_round_trips() { let inbound_b = pair.take_inbound(Side::B); let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let request_data = read_all(stream.reader).await.unwrap(); stream.writer.finish().await.unwrap(); done_tx.send(request_data).await.unwrap(); @@ -111,7 +111,7 @@ async fn dropping_responder_closes_initiator_response() { let inbound_b = pair.take_inbound(Side::B); let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); drop(stream.reader); }); @@ -152,7 +152,7 @@ async fn dropping_inbound_reader_cancels_remote_writer() { pair.connect_and_wait(Side::A).await; let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let mut writer = stream.writer; let mut reader = stream.reader; assert_eq!(next_chunk(&mut reader).await.unwrap(), None); @@ -201,7 +201,7 @@ async fn closing_initiator_reader_preserves_initiator_writer() { let (done_tx, done_rx) = async_channel::bounded(1); let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let request = read_all(stream.reader).await.unwrap(); done_tx.send(request).await.unwrap(); }); @@ -264,7 +264,7 @@ async fn max_concurrent_message_writes_is_respected() { let responder = tokio::task::spawn_local(async move { for _ in 0..4 { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let _ = read_all(stream.reader).await; let mut writer = stream.writer; writer.queue_finish(); @@ -338,7 +338,7 @@ async fn stream_round_trip_survives_encrypted_packet_drops() { await_status(&status_b, Some(identity_a.qid), PeerStatus::Connected).await; let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let received_request = read_all(stream.reader).await.unwrap(); let mut writer = stream.writer; writer @@ -421,7 +421,7 @@ async fn multi_megabyte_stream_survives_asymmetric_loss_and_delay() { let inbound_b = pair.take_inbound(Side::B); let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); eprintln!("responder accepted inbound stream"); let mut reader = stream.reader; let mut received = Vec::new(); @@ -540,7 +540,7 @@ async fn reproducer_writer_stalls_after_reverse_path_impairment() { let inbound_b = pair.take_inbound(Side::B); let responder = tokio::task::spawn_local(async move { - let stream = inbound_b.recv().await.unwrap(); + let (_, stream) = inbound_b.recv().await.unwrap(); let mut reader = stream.reader; while next_chunk(&mut reader).await.unwrap().is_some() {} }); @@ -591,7 +591,7 @@ async fn responder_drains_multiple_local_chunks_per_writable_wake() { let inbound_b = pair.take_inbound(Side::B); let responder = tokio::task::spawn_local(async move { - let inbound = inbound_b.recv().await.unwrap(); + let (_, inbound) = inbound_b.recv().await.unwrap(); let _ = read_all(inbound.reader).await.unwrap(); let mut writer = inbound.writer; From 6e742dec7bd4828f89af4112c081866104a394ad Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 24 Jun 2026 23:17:30 -0400 Subject: [PATCH 27/59] ql: refer to ql-common directly --- Cargo.lock | 1 + ql-fsm/Cargo.toml | 1 + ql-fsm/src/fsm.rs | 3 ++- ql-fsm/src/handshake/mod.rs | 3 ++- ql-fsm/src/handshake/xx.rs | 3 ++- ql-fsm/src/lib.rs | 6 +++--- ql-fsm/src/pairing.rs | 3 ++- ql-fsm/src/session/mod.rs | 6 +++--- ql-fsm/src/session/remote_stream_history.rs | 2 +- ql-fsm/src/session/state.rs | 3 ++- ql-fsm/src/session/stream_ops.rs | 3 ++- ql-fsm/src/session/stream_parity.rs | 4 ++-- ql-fsm/src/session/tests.rs | 6 +++--- ql-fsm/src/session/tracked.rs | 3 ++- ql-fsm/src/tests/handshake.rs | 2 +- ql-fsm/src/tests/mod.rs | 3 ++- ql-fsm/src/tests/proptest.rs | 3 ++- ql-fsm/src/tests/session.rs | 5 +++-- ql-rpc/src/error.rs | 2 +- ql-rpc/src/lib.rs | 1 - ql-rpc/src/router/mod.rs | 2 +- ql-rpc/src/rpc/download/client.rs | 3 ++- ql-rpc/src/rpc/download/server.rs | 4 ++-- ql-rpc/src/rpc/duplex/client.rs | 5 +++-- ql-rpc/src/rpc/mod.rs | 2 +- ql-rpc/src/rpc/notification/client.rs | 4 +++- ql-rpc/src/rpc/notification/server.rs | 6 ++++-- ql-rpc/src/rpc/progress/client.rs | 3 ++- ql-rpc/src/rpc/progress/server.rs | 4 ++-- ql-rpc/src/rpc/request/server.rs | 3 ++- ql-rpc/src/rpc/subscription/client.rs | 4 +++- ql-rpc/src/rpc/subscription/server.rs | 5 +++-- ql-rpc/src/rpc/upload/client.rs | 3 ++- ql-rpc/src/rpc/upload/server.rs | 4 ++-- ql-rpc/src/stream.rs | 3 +-- ql-runtime/src/command.rs | 5 ++--- ql-runtime/src/driver/mod.rs | 5 +++-- ql-runtime/src/driver/state.rs | 2 +- ql-runtime/src/driver/test.rs | 5 +++-- ql-runtime/src/error.rs | 2 +- ql-runtime/src/io/inner.rs | 4 ++-- ql-runtime/src/io/mod.rs | 5 ++--- ql-runtime/src/io/reader.rs | 1 - ql-runtime/src/io/sync.rs | 2 +- ql-runtime/src/io/writer.rs | 1 - ql-runtime/src/platform.rs | 4 ++-- ql-runtime/src/rpc/adapter.rs | 3 ++- ql-runtime/src/tests/mod.rs | 9 ++++----- ql-runtime/src/tests/rpc.rs | 6 +++--- ql-runtime/src/tests/stream.rs | 2 +- ql-wire/src/encrypted/ack.rs | 8 ++++++-- ql-wire/src/encrypted/mod.rs | 4 +++- ql-wire/src/encrypted/stream_data.rs | 6 ++---- ql-wire/src/encrypted/stream_reset.rs | 4 +++- ql-wire/src/encrypted/stream_window.rs | 4 +++- ql-wire/src/handshake/mod.rs | 4 +++- ql-wire/src/handshake/xx.rs | 4 +++- ql-wire/src/identity.rs | 4 +++- ql-wire/src/lib.rs | 3 --- ql-wire/src/macros.rs | 2 +- ql-wire/src/qid.rs | 4 +++- ql-wire/src/tests.rs | 2 ++ 62 files changed, 129 insertions(+), 94 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 54ffae1f..00b789c1 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1061,6 +1061,7 @@ dependencies = [ "bytes", "indexmap", "proptest", + "ql-common", "ql-wire", ] diff --git a/ql-fsm/Cargo.toml b/ql-fsm/Cargo.toml index be937f01..93006d35 100644 --- a/ql-fsm/Cargo.toml +++ b/ql-fsm/Cargo.toml @@ -8,6 +8,7 @@ license = "Proprietary" [dependencies] bytes = { workspace = true } indexmap = "2" +ql-common = { workspace = true } ql-wire = { workspace = true } [dev-dependencies] diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index 804b486b..6e4be4a0 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -1,7 +1,8 @@ use std::{collections::VecDeque, time::Instant}; use bytes::Bytes; -use ql_wire::{self as wire, QlCrypto, SessionCloseCode, StreamId, WireDecode}; +use ql_common::StreamId; +use ql_wire::{self as wire, QlCrypto, SessionCloseCode, WireDecode}; use crate::{ handshake, diff --git a/ql-fsm/src/handshake/mod.rs b/ql-fsm/src/handshake/mod.rs index 2e36ee65..dd2bfc5c 100644 --- a/ql-fsm/src/handshake/mod.rs +++ b/ql-fsm/src/handshake/mod.rs @@ -2,8 +2,9 @@ mod ik; mod kk; mod xx; +use ql_common::QID; use ql_wire::{ - self as wire, EphemeralPublicKey, HandshakeId, HandshakeMeta, QlCrypto, QlHandshakeRecord, QID, + self as wire, EphemeralPublicKey, HandshakeId, HandshakeMeta, QlCrypto, QlHandshakeRecord, }; use crate::{ diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs index 4aa634ef..cf325f8f 100644 --- a/ql-fsm/src/handshake/xx.rs +++ b/ql-fsm/src/handshake/xx.rs @@ -1,4 +1,5 @@ -use ql_wire::{self as wire, PairingToken, QlCrypto, QlHandshakeRecord, Xx1, Xx2, Xx3, Xx4, QID}; +use ql_common::QID; +use ql_wire::{self as wire, PairingToken, QlCrypto, QlHandshakeRecord, Xx1, Xx2, Xx3, Xx4}; use super::{ emit_peer_status, enqueue_handshake, establish_session, reset_connected_session_if_needed, diff --git a/ql-fsm/src/lib.rs b/ql-fsm/src/lib.rs index 8a96f93b..d78d5f16 100644 --- a/ql-fsm/src/lib.rs +++ b/ql-fsm/src/lib.rs @@ -35,9 +35,9 @@ use std::{ pub use bytes::Bytes; pub use error::*; pub use pairing::PairingInvite; +use ql_common::{ResetCode, RouteId, ServiceId, StreamId}; use ql_wire::{ - PairingToken, PeerBundle, QlCrypto, QlIdentity, ResetCode, RouteId, ServiceId, SessionClose, - SessionCloseCode, StreamHeader, StreamId, + PairingToken, PeerBundle, QlCrypto, QlIdentity, SessionClose, SessionCloseCode, StreamHeader, }; pub use session::{SessionEvent, StreamReadIter, StreamWriter}; @@ -162,7 +162,7 @@ impl StreamOps<'_> { } /// resets the local read side, write side, or both sides of the stream - pub fn reset(&mut self, target: StreamResetTarget, code: ql_wire::ResetCode) { + pub fn reset(&mut self, target: StreamResetTarget, code: ResetCode) { self.inner.reset(target, code); } } diff --git a/ql-fsm/src/pairing.rs b/ql-fsm/src/pairing.rs index 4b8361b8..c1f239a1 100644 --- a/ql-fsm/src/pairing.rs +++ b/ql-fsm/src/pairing.rs @@ -1,4 +1,5 @@ -use ql_wire::{ByteSlice, PairingToken, Reader, WireDecode, WireEncode, WireError, QID}; +use ql_common::QID; +use ql_wire::{ByteSlice, PairingToken, Reader, WireDecode, WireEncode, WireError}; /// Out-of-band invite consumed by the initiator of an XX pairing #[derive(Debug, Clone, Copy, PartialEq, Eq)] diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index 335623f2..1ab9772f 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -17,10 +17,10 @@ use std::time::{Duration, Instant}; use bytes::Bytes; use indexmap::IndexMap; +use ql_common::{RouteId, ServiceId, StreamId, VarInt}; use ql_wire::{ - RecordAck, RecordSeq, ResetTarget, RouteId, ServiceId, SessionClose, SessionCloseCode, - SessionFrame, SessionRecordBuilder, StreamData, StreamHeader, StreamId, StreamReset, - StreamWindow, VarInt, WireError, + RecordAck, RecordSeq, ResetTarget, SessionClose, SessionCloseCode, SessionFrame, + SessionRecordBuilder, StreamData, StreamHeader, StreamReset, StreamWindow, WireError, }; use self::{ diff --git a/ql-fsm/src/session/remote_stream_history.rs b/ql-fsm/src/session/remote_stream_history.rs index 0d126b18..e9b0f5e0 100644 --- a/ql-fsm/src/session/remote_stream_history.rs +++ b/ql-fsm/src/session/remote_stream_history.rs @@ -1,4 +1,4 @@ -use ql_wire::StreamId; +use ql_common::StreamId; use super::{range_set::RangeSet, stream_parity::StreamParity}; diff --git a/ql-fsm/src/session/state.rs b/ql-fsm/src/session/state.rs index 310d7f0f..a71ec5fc 100644 --- a/ql-fsm/src/session/state.rs +++ b/ql-fsm/src/session/state.rs @@ -1,7 +1,8 @@ use std::time::Instant; use indexmap::IndexMap; -use ql_wire::{RecordSeq, ResetTarget, SessionClose, StreamHeader, StreamId, StreamReset}; +use ql_common::StreamId; +use ql_wire::{RecordSeq, ResetTarget, SessionClose, StreamHeader, StreamReset}; use super::{ ack_tracker::AckTracker, remote_stream_history::RemoteStreamHistory, stream_rx::StreamRx, diff --git a/ql-fsm/src/session/stream_ops.rs b/ql-fsm/src/session/stream_ops.rs index 67a1900b..5261edef 100644 --- a/ql-fsm/src/session/stream_ops.rs +++ b/ql-fsm/src/session/stream_ops.rs @@ -1,4 +1,5 @@ -use ql_wire::{ResetCode, StreamHeader, StreamId, StreamReset}; +use ql_common::{ResetCode, StreamId}; +use ql_wire::{StreamHeader, StreamReset}; use super::{ state::{InboundState, StreamState}, diff --git a/ql-fsm/src/session/stream_parity.rs b/ql-fsm/src/session/stream_parity.rs index 26846d79..8d2d74b2 100644 --- a/ql-fsm/src/session/stream_parity.rs +++ b/ql-fsm/src/session/stream_parity.rs @@ -1,4 +1,4 @@ -use ql_wire::{StreamId, QID}; +use ql_common::{StreamId, VarInt, QID}; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum StreamParity { @@ -36,7 +36,7 @@ impl StreamParity { } pub fn make_stream_id(self, ordinal: u32) -> StreamId { - StreamId(ql_wire::VarInt::from_u32( + StreamId(VarInt::from_u32( self.first_stream_id() .saturating_add(ordinal.saturating_mul(2)), )) diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index 85edec29..84386785 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -1,10 +1,10 @@ use std::time::{Duration, Instant}; use bytes::Bytes; +use ql_common::{ResetCode, RouteId, ServiceId, StreamId, VarInt, QID}; use ql_wire::{ - decode_session_frames, parse_session_frames, RecordAck, RecordSeq, ResetCode, ResetTarget, - RouteId, ServiceId, SessionFrame, SessionRecordBuilder, StreamData, StreamHeader, StreamId, - StreamReset, VarInt, QID, + decode_session_frames, parse_session_frames, RecordAck, RecordSeq, ResetTarget, SessionFrame, + SessionRecordBuilder, StreamData, StreamHeader, StreamReset, }; use super::{SessionConfig, SessionEvent, SessionFsm}; diff --git a/ql-fsm/src/session/tracked.rs b/ql-fsm/src/session/tracked.rs index 72439875..e74b2928 100644 --- a/ql-fsm/src/session/tracked.rs +++ b/ql-fsm/src/session/tracked.rs @@ -2,7 +2,8 @@ use std::time::Instant; -use ql_wire::{RecordAck, RecordSeq, StreamId, StreamReset}; +use ql_common::StreamId; +use ql_wire::{RecordAck, RecordSeq, StreamReset}; #[derive(Debug, Clone)] pub struct TrackedRecord { diff --git a/ql-fsm/src/tests/handshake.rs b/ql-fsm/src/tests/handshake.rs index 4de4f06e..d44e8789 100644 --- a/ql-fsm/src/tests/handshake.rs +++ b/ql-fsm/src/tests/handshake.rs @@ -106,7 +106,7 @@ fn connect_methods_require_bound_peer() { fsm.connect_xx( time, PairingInvite { - qid: ql_wire::QID([2; ql_wire::QID::SIZE]), + qid: ql_common::QID([2; ql_common::QID::SIZE]), token: pairing_token(2), }, &crypto, diff --git a/ql-fsm/src/tests/mod.rs b/ql-fsm/src/tests/mod.rs index 2df7d14b..dbc6bcd1 100644 --- a/ql-fsm/src/tests/mod.rs +++ b/ql-fsm/src/tests/mod.rs @@ -4,9 +4,10 @@ mod session; use std::time::{Duration, Instant}; +use ql_common::QID; use ql_wire::{ self, generate_identity, test_identities, ConnectionId, HandshakeId, PairingToken, QlCrypto, - SessionKey, SoftwareCrypto, TransportParams, QID, + SessionKey, SoftwareCrypto, TransportParams, }; use crate::{ diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs index dcd85bdf..6f2dc07d 100644 --- a/ql-fsm/src/tests/proptest.rs +++ b/ql-fsm/src/tests/proptest.rs @@ -7,7 +7,8 @@ extern crate proptest as proptest_crate; use bytes::Bytes; use proptest_crate::{collection::vec, prelude::*, test_runner::TestCaseResult}; -use ql_wire::{ResetCode, RouteId, ServiceId, StreamId, WireError}; +use ql_common::{ResetCode, RouteId, ServiceId, StreamId}; +use ql_wire::WireError; use super::*; use crate::{ diff --git a/ql-fsm/src/tests/session.rs b/ql-fsm/src/tests/session.rs index efd9c8c3..cc240a0d 100644 --- a/ql-fsm/src/tests/session.rs +++ b/ql-fsm/src/tests/session.rs @@ -1,7 +1,8 @@ use std::time::Duration; use bytes::Bytes; -use ql_wire::{RouteId, ServiceId, SessionClose, StreamId, VarInt}; +use ql_common::{RouteId, ServiceId, StreamId, VarInt}; +use ql_wire::SessionClose; use super::*; use crate::{ @@ -220,7 +221,7 @@ fn disconnected_stream_operations_fail_with_no_session() { harness.a.fsm.stream(missing).map(|mut stream| { stream.reset( crate::StreamResetTarget::Both, - ql_wire::ResetCode::CANCELLED, + ql_common::ResetCode::CANCELLED, ); }), Err(StreamError::NoSession) diff --git a/ql-rpc/src/error.rs b/ql-rpc/src/error.rs index 8c47362d..6c176a7f 100644 --- a/ql-rpc/src/error.rs +++ b/ql-rpc/src/error.rs @@ -1,4 +1,4 @@ -use crate::ResetCode; +use ql_common::ResetCode; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum Error { diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index 40dd5a1a..7dd45047 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -14,7 +14,6 @@ pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; use framed_value::*; -pub use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId, StreamId, StreamInfo, QID}; pub use router::*; pub use rpc::*; pub use stream::*; diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index 5d32b529..9bb691f5 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -1,4 +1,4 @@ -use crate::{ResetCode, RouteId, ServiceId, StreamId, StreamInfo, QID}; +use ql_common::{ResetCode, RouteId, ServiceId, StreamId, StreamInfo, QID}; mod builder; mod config; diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index 9001e15c..b759f10a 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -1,12 +1,13 @@ use std::future::poll_fn; use bytes::Bytes; +use ql_common::ResetCode; use crate::{ download::Download, parts::{PartFrameReader, PartReadStep}, rpc::{parts::FrameKind, write_eof_value}, - DropResetRead, FramedPrefixStep, FramedReader, ResetCode, RpcError, RpcRead, RpcStream, + DropResetRead, FramedPrefixStep, FramedReader, RpcError, RpcRead, RpcStream, }; pub async fn start( diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs index 98b3a4e7..a5ebbccd 100644 --- a/ql-rpc/src/rpc/download/server.rs +++ b/ql-rpc/src/rpc/download/server.rs @@ -1,6 +1,7 @@ use std::{future::Future, marker::PhantomData}; use bytes::Bytes; +use ql_common::ResetCode; use crate::{ codec, @@ -10,8 +11,7 @@ use crate::{ parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, read_eof_request, }, - write_bytes, Context, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, - RpcWrite, + write_bytes, Context, DropResetWrite, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(DownloadHandler: Send)] diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index 1f9ea080..6a2cbe9a 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -5,10 +5,11 @@ use std::{ }; use bytes::Bytes; +use ql_common::ResetCode; use crate::{ - codec, duplex::Duplex, write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcCodec, - RpcError, RpcRead, RpcStream, RpcWrite, + codec, duplex::Duplex, write_bytes, DropResetRead, DropResetWrite, RpcCodec, RpcError, RpcRead, + RpcStream, RpcWrite, }; pub fn start(stream: St) -> DuplexCall diff --git a/ql-rpc/src/rpc/mod.rs b/ql-rpc/src/rpc/mod.rs index 0c7d9788..a8d08bcb 100644 --- a/ql-rpc/src/rpc/mod.rs +++ b/ql-rpc/src/rpc/mod.rs @@ -5,7 +5,7 @@ //! route dispatch uses [`crate::RouteId`] and the submodules provide the matching //! client and server helpers for encoding, decoding, and handler glue -use crate::{RouteId, ServiceId}; +use ql_common::{RouteId, ServiceId}; pub mod download; pub mod duplex; diff --git a/ql-rpc/src/rpc/notification/client.rs b/ql-rpc/src/rpc/notification/client.rs index 27bd23b4..9d5d6bf2 100644 --- a/ql-rpc/src/rpc/notification/client.rs +++ b/ql-rpc/src/rpc/notification/client.rs @@ -1,4 +1,6 @@ -use crate::{notification::Notification, rpc::write_eof_value, ResetCode, RpcRead, RpcStream}; +use ql_common::ResetCode; + +use crate::{notification::Notification, rpc::write_eof_value, RpcRead, RpcStream}; pub async fn send(stream: St, payload: &M::Payload) -> Result<(), St::Error> where diff --git a/ql-rpc/src/rpc/notification/server.rs b/ql-rpc/src/rpc/notification/server.rs index 87258b7e..282d6659 100644 --- a/ql-rpc/src/rpc/notification/server.rs +++ b/ql-rpc/src/rpc/notification/server.rs @@ -1,8 +1,10 @@ use std::future::Future; +use ql_common::ResetCode; + use crate::{ - notification::Notification as NotificationRpc, rpc::read_eof_request, Context, ResetCode, - RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, + notification::Notification as NotificationRpc, rpc::read_eof_request, Context, RouterConfig, + RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(NotificationHandler: Send)] diff --git a/ql-rpc/src/rpc/progress/client.rs b/ql-rpc/src/rpc/progress/client.rs index 61d3fee0..0998ae7b 100644 --- a/ql-rpc/src/rpc/progress/client.rs +++ b/ql-rpc/src/rpc/progress/client.rs @@ -5,12 +5,13 @@ use std::{ }; use bytes::Bytes; +use ql_common::ResetCode; use crate::{ codec, finish_bytes, progress::Progress, rpc::progress::codec::{ReadStep, ResponseReader}, - write_bytes, DropResetRead, Error, ResetCode, RpcError, RpcRead, RpcStream, + write_bytes, DropResetRead, Error, RpcError, RpcRead, RpcStream, }; pub async fn start( diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index 8ea1add5..12ae6c7b 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -1,13 +1,13 @@ use std::{future::Future, marker::PhantomData}; use bytes::Bytes; +use ql_common::ResetCode; use crate::{ codec, finish_bytes, progress::Progress, rpc::{progress::codec::FrameKind, read_framed_request}, - write_bytes, Context, DropResetWrite, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, - RpcWrite, + write_bytes, Context, DropResetWrite, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(ProgressHandler: Send)] diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs index d57365bf..68fdf862 100644 --- a/ql-rpc/src/rpc/request/server.rs +++ b/ql-rpc/src/rpc/request/server.rs @@ -1,10 +1,11 @@ use std::{future::Future, marker::PhantomData}; use bytes::Bytes; +use ql_common::ResetCode; use crate::{ finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, Context, - DropResetWrite, ResetCode, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, + DropResetWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(RequestHandler: Send)] diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index 24d22aa6..9767211a 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -3,13 +3,15 @@ use std::{ task::{Context, Poll}, }; +use ql_common::ResetCode; + use crate::{ rpc::{ subscription::codec::{ReadStep, ResponseReader}, write_eof_value, }, subscription::Subscription, - DropResetRead, ResetCode, RpcError, RpcRead, RpcStream, + DropResetRead, RpcError, RpcRead, RpcStream, }; pub async fn start( diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs index 5f6a2ff4..df2a4ec8 100644 --- a/ql-rpc/src/rpc/subscription/server.rs +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -1,11 +1,12 @@ use std::{future::Future, marker::PhantomData}; use bytes::Bytes; +use ql_common::ResetCode; use crate::{ codec, finish_bytes, rpc::read_eof_request, subscription::Subscription as SubscriptionRpc, - write_bytes, Context, DropResetWrite, ResetCode, RouterConfig, RpcCodec, RpcError, RpcRead, - RpcStream, RpcWrite, + write_bytes, Context, DropResetWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, + RpcWrite, }; #[trait_variant::make(SubscriptionHandler: Send)] diff --git a/ql-rpc/src/rpc/upload/client.rs b/ql-rpc/src/rpc/upload/client.rs index b56400c9..b5120599 100644 --- a/ql-rpc/src/rpc/upload/client.rs +++ b/ql-rpc/src/rpc/upload/client.rs @@ -1,4 +1,5 @@ use bytes::Bytes; +use ql_common::ResetCode; use crate::{ rpc::{ @@ -6,7 +7,7 @@ use crate::{ read_eof_value, }, upload::Upload, - write_bytes, DropResetRead, DropResetWrite, ResetCode, RpcError, RpcRead, RpcStream, RpcWrite, + write_bytes, DropResetRead, DropResetWrite, RpcError, RpcRead, RpcStream, RpcWrite, }; pub async fn start( diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index 4e32fd64..299aa5a3 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -1,6 +1,7 @@ use std::future::{poll_fn, Future}; use bytes::Bytes; +use ql_common::ResetCode; use crate::{ request::Response, @@ -8,8 +9,7 @@ use crate::{ parts::{FrameKind, PartFrameReader, PartReadStep}, read_framed_request_prefix, }, - Context, DropResetRead, ResetCode, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, - Upload, + Context, DropResetRead, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, Upload, }; #[trait_variant::make(UploadHandler: Send)] diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index 043e1673..e6b1f73e 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -4,8 +4,7 @@ use std::{ }; use bytes::Bytes; - -use crate::ResetCode; +use ql_common::ResetCode; pub trait RpcStream { type Error; diff --git a/ql-runtime/src/command.rs b/ql-runtime/src/command.rs index 730ebefa..72138c71 100644 --- a/ql-runtime/src/command.rs +++ b/ql-runtime/src/command.rs @@ -1,7 +1,6 @@ +use ql_common::{ResetCode, RouteId, ServiceId, StreamId}; use ql_fsm::{NoSessionError, PairingInvite, StreamResetTarget}; -use ql_wire::{ - PairingToken, PeerBundle, ResetCode, RouteId, ServiceId, SessionCloseCode, StreamId, -}; +use ql_wire::{PairingToken, PeerBundle, SessionCloseCode}; use crate::{StreamReader, StreamWriter}; diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 5d4a5b98..e2feb2d7 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -15,14 +15,15 @@ use std::{ use async_channel::Recv; use futures_lite::future::{poll_fn, yield_now}; +use ql_common::{ResetCode, ResetOrigin, StreamId, StreamInfo}; use ql_fsm::{Event, QlFsm, StreamResetEvent, StreamResetTarget, WriteId}; -use ql_wire::{ResetCode, ResetOrigin, StreamHeader, StreamId}; +use ql_wire::StreamHeader; use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; use crate::{ command::Command, io, log, - platform::{QlInbound, QlPlatform, QlTimer, StreamInfo}, + platform::{QlInbound, QlPlatform, QlTimer}, QlStreamError, Runtime, }; diff --git a/ql-runtime/src/driver/state.rs b/ql-runtime/src/driver/state.rs index 1a85e96c..df725838 100644 --- a/ql-runtime/src/driver/state.rs +++ b/ql-runtime/src/driver/state.rs @@ -1,7 +1,7 @@ use std::collections::HashMap; use bytes::Bytes; -use ql_wire::StreamId; +use ql_common::StreamId; use crate::{ command::Command, diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs index d891647b..f3e7b528 100644 --- a/ql-runtime/src/driver/test.rs +++ b/ql-runtime/src/driver/test.rs @@ -1,10 +1,11 @@ +use ql_common::{StreamInfo, QID}; use ql_fsm::StreamResetEvent; -use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto, QID}; +use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto}; use super::*; use crate::{ driver::state::{InboundIo, OutboundIo}, - io, StreamInfo, + io, }; pub struct NoopTimer; diff --git a/ql-runtime/src/error.rs b/ql-runtime/src/error.rs index 4bdb1d50..d947871d 100644 --- a/ql-runtime/src/error.rs +++ b/ql-runtime/src/error.rs @@ -1,5 +1,5 @@ +use ql_common::{ResetCode, ResetOrigin}; use ql_fsm::NoSessionError; -use ql_wire::{ResetCode, ResetOrigin}; #[derive(Debug, Clone, PartialEq, Eq)] pub enum QlStreamError { diff --git a/ql-runtime/src/io/inner.rs b/ql-runtime/src/io/inner.rs index 40866198..5b2e8ab5 100644 --- a/ql-runtime/src/io/inner.rs +++ b/ql-runtime/src/io/inner.rs @@ -6,7 +6,7 @@ use std::task::Waker; use bytes::Bytes; use diatomic_waker::DiatomicWaker; -use ql_wire::StreamId; +use ql_common::StreamId; use super::{ slot::{PopError, PushError, Slot}, @@ -292,7 +292,7 @@ mod loom_tests { use bytes::Bytes; use loom::thread; - use ql_wire::{ResetCode, ResetOrigin}; + use ql_common::{ResetCode, ResetOrigin}; use super::*; use crate::{ diff --git a/ql-runtime/src/io/mod.rs b/ql-runtime/src/io/mod.rs index 039363e3..7575cbea 100644 --- a/ql-runtime/src/io/mod.rs +++ b/ql-runtime/src/io/mod.rs @@ -6,10 +6,9 @@ mod writer; use std::ops::Deref; -use ql_wire::StreamId; +use ql_common::StreamId; pub use self::{reader::StreamReader, slot::PushError, writer::StreamWriter}; -use crate::command::Command; pub struct Rx(sync::Arc); @@ -45,7 +44,7 @@ impl Tx { pub fn new_stream( stream_id: StreamId, - runtime_tx: async_channel::Sender, + runtime_tx: async_channel::Sender, ) -> (StreamReader, StreamWriter, Rx, Tx) { let shared = inner::new(stream_id); ( diff --git a/ql-runtime/src/io/reader.rs b/ql-runtime/src/io/reader.rs index 28b8f154..f8e16cea 100644 --- a/ql-runtime/src/io/reader.rs +++ b/ql-runtime/src/io/reader.rs @@ -6,7 +6,6 @@ use std::{ use bytes::Bytes; use ql_common::ResetCode; use ql_fsm::StreamResetTarget; -use ql_wire::{ResetCode, StreamId}; use super::{ inner::{Item, RxInner}, diff --git a/ql-runtime/src/io/sync.rs b/ql-runtime/src/io/sync.rs index bc06d474..4863478d 100644 --- a/ql-runtime/src/io/sync.rs +++ b/ql-runtime/src/io/sync.rs @@ -68,7 +68,7 @@ pub use inner::*; #[cfg(all(test, loom))] pub(crate) mod loom { use loom::model; - use ql_wire::StreamId; + use ql_common::StreamId; use super::Arc; use crate::{command::Command, io::inner::Inner}; diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs index ea67b7e1..be354c03 100644 --- a/ql-runtime/src/io/writer.rs +++ b/ql-runtime/src/io/writer.rs @@ -6,7 +6,6 @@ use std::{ use bytes::Bytes; use ql_common::ResetCode; use ql_fsm::StreamResetTarget; -use ql_wire::{ResetCode, StreamId}; use super::{ inner::{Item, TxInner}, diff --git a/ql-runtime/src/platform.rs b/ql-runtime/src/platform.rs index cf45aedb..18ad6672 100644 --- a/ql-runtime/src/platform.rs +++ b/ql-runtime/src/platform.rs @@ -5,9 +5,9 @@ use std::{ time::Instant, }; -pub use ql_common::StreamInfo; +use ql_common::{StreamInfo, QID}; use ql_fsm::{PeerStatus, ReceiveError}; -use ql_wire::{PeerBundle, QlCrypto, QID}; +use ql_wire::{PeerBundle, QlCrypto}; pub trait QlTimer { fn set_deadline(self: Pin<&mut Self>, deadline: Option); diff --git a/ql-runtime/src/rpc/adapter.rs b/ql-runtime/src/rpc/adapter.rs index 857bf0bc..293caa3d 100644 --- a/ql-runtime/src/rpc/adapter.rs +++ b/ql-runtime/src/rpc/adapter.rs @@ -1,7 +1,8 @@ use std::task::{Context as TaskContext, Poll}; use bytes::Bytes; -use ql_rpc::{ResetCode, RpcRead, RpcStream, RpcWrite}; +use ql_common::ResetCode; +use ql_rpc::{RpcRead, RpcStream, RpcWrite}; use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index 00ac8a74..d0f28942 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -11,19 +11,18 @@ use std::{ use async_channel::{Receiver, Sender}; use futures_lite::Stream; +use ql_common::{RouteId, ServiceId, StreamInfo, QID}; use ql_fsm::{OpenStreamParams, PeerStatus}; use ql_wire::{ generate_identity, test_identities, MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, Nonce, PairingToken, PeerBundle, QlAead, QlHash, QlIdentity, QlKem, QlRandom, - RecordHeader, RecordType, RouteId, ServiceId, SessionKey, SoftwareCrypto, WireDecode, QID, + RecordHeader, RecordType, SessionKey, SoftwareCrypto, WireDecode, }; use tokio::{task::LocalSet, time::Sleep}; use crate::{ - new_runtime, - platform::{QlTimer, StreamInfo}, - NoSessionError, PairingInvite, QlFsmConfig, QlStream, QlStreamError, RuntimeConfig, - RuntimeHandle, + new_runtime, platform::QlTimer, NoSessionError, PairingInvite, QlFsmConfig, QlStream, + QlStreamError, RuntimeConfig, RuntimeHandle, }; type InboundStream = (StreamInfo, QlStream); diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index c26f5c03..c72425b9 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -8,12 +8,12 @@ use std::{ }; use bytes::Bytes; +use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId}; use ql_rpc::{ Context, DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, NotificationHandlerLocal, ProgressHandlerLocal, ProgressResponder, RequestHandler, - RequestHandlerLocal, ResetCode, ResetOrigin, Response, RouteId, SendSpawner, ServiceId, - Spawner, SubscriptionHandlerLocal, SubscriptionResponder, UploadHandlerLocal, UploadReader, - UploadResponder, + RequestHandlerLocal, Response, SendSpawner, Spawner, SubscriptionHandlerLocal, + SubscriptionResponder, UploadHandlerLocal, UploadReader, UploadResponder, }; use super::*; diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index 83baf7a2..151e6ed4 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -1,7 +1,7 @@ use std::time::Duration; use bytes::Bytes; -use ql_wire::{ResetCode, ResetOrigin}; +use ql_common::{ResetCode, ResetOrigin}; use super::*; use crate::QlStreamError; diff --git a/ql-wire/src/encrypted/ack.rs b/ql-wire/src/encrypted/ack.rs index eaaa8bb3..099fe42f 100644 --- a/ql-wire/src/encrypted/ack.rs +++ b/ql-wire/src/encrypted/ack.rs @@ -1,6 +1,8 @@ use std::{fmt, ops::RangeInclusive}; -use crate::{codec, ByteSlice, RecordSeq, VarInt, WireEncode, WireError}; +use ql_common::VarInt; + +use crate::{codec, ByteSlice, RecordSeq, WireEncode, WireError}; #[derive(Debug, Clone, PartialEq, Eq)] pub struct RecordAck { @@ -272,8 +274,10 @@ impl RecordAckBuilder { mod tests { use std::ops::RangeInclusive; + use ql_common::VarInt; + use super::{RecordAck, RecordAckBlock, RecordAckBuilder, RecordAckRangeError}; - use crate::{RecordSeq, VarInt, WireDecode, WireEncode, WireError}; + use crate::{RecordSeq, WireDecode, WireEncode, WireError}; fn seq(value: u64) -> RecordSeq { RecordSeq::from_u64(value).unwrap() diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index 93ae23e6..c079fba1 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -1,6 +1,8 @@ +use ql_common::{RouteId, ServiceId, StreamId}; + use crate::{ codec, encrypted_message::EncryptedMessage, BufView, ByteSlice, Nonce, QlCrypto, Reader, - RouteId, ServiceId, SessionHeader, SessionKey, StreamId, WireDecode, WireEncode, WireError, + SessionHeader, SessionKey, WireDecode, WireEncode, WireError, }; mod ack; diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index f789ad09..3cc9eb50 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -1,9 +1,7 @@ use bytes::Buf; +use ql_common::{RouteId, ServiceId, StreamId, VarInt}; -use crate::{ - codec, BufView, ByteSlice, RouteId, ServiceId, StreamId, VarInt, WireDecode, WireEncode, - WireError, -}; +use crate::{codec, BufView, ByteSlice, WireDecode, WireEncode, WireError}; /// carries bytes for a stream and may finish that sending direction. #[derive(Debug, Clone, PartialEq, Eq)] diff --git a/ql-wire/src/encrypted/stream_reset.rs b/ql-wire/src/encrypted/stream_reset.rs index 99381e39..6ee73722 100644 --- a/ql-wire/src/encrypted/stream_reset.rs +++ b/ql-wire/src/encrypted/stream_reset.rs @@ -1,5 +1,7 @@ +use ql_common::ResetCode; + use super::StreamId; -use crate::{codec, ByteSlice, ResetCode, WireEncode, WireError}; +use crate::{codec, ByteSlice, WireEncode, WireError}; /// aborts one or both lanes of a stream with a reset code /// diff --git a/ql-wire/src/encrypted/stream_window.rs b/ql-wire/src/encrypted/stream_window.rs index 6a2274f9..f932a2e6 100644 --- a/ql-wire/src/encrypted/stream_window.rs +++ b/ql-wire/src/encrypted/stream_window.rs @@ -1,5 +1,7 @@ +use ql_common::VarInt; + use super::StreamId; -use crate::{codec, ByteSlice, VarInt, WireEncode, WireError}; +use crate::{codec, ByteSlice, WireEncode, WireError}; /// advertises the highest byte offset the peer may send on a stream. #[derive(Debug, Clone, PartialEq, Eq)] diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index b8f06248..b8bfe204 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -1,7 +1,9 @@ +use ql_common::QID; + use crate::{ codec, derive_qid, ByteSlice, ConnectionId, HandshakeKind, MlKemCiphertext, MlKemKeyPair, MlKemPublicKey, Nonce, PeerBundle, QlCrypto, SessionKey, WireDecode, WireEncode, WireError, - ENCRYPTED_MESSAGE_AUTH_SIZE, QID, + ENCRYPTED_MESSAGE_AUTH_SIZE, }; mod ik; diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs index 0b6452d4..f812744a 100644 --- a/ql-wire/src/handshake/xx.rs +++ b/ql-wire/src/handshake/xx.rs @@ -1,3 +1,5 @@ +use ql_common::QID; + use super::{ decrypt_mlkem_ciphertext, decrypt_peer_bundle, encrypt_mlkem_ciphertext, encrypt_peer_bundle, finalize_handshake, generate_ephemeral_keypair, init_xx_symmetric, initialize_handshake_meta, @@ -8,7 +10,7 @@ use super::{ }; use crate::{ codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PairingId, PairingToken, - PeerBundle, QlCrypto, QlIdentity, WireEncode, WireError, QID, + PeerBundle, QlCrypto, QlIdentity, WireEncode, WireError, }; #[derive(Debug, Clone, PartialEq, Eq)] diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index e8c30c90..57602178 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -1,8 +1,10 @@ use std::ops::Deref; +use ql_common::{VarInt, QID}; + use crate::{ codec, derive_qid, ByteSlice, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash, - VarInt, WireEncode, WireError, QID, + WireEncode, WireError, }; #[derive(Debug, Clone, PartialEq, Eq)] diff --git a/ql-wire/src/lib.rs b/ql-wire/src/lib.rs index 0016941b..b55c3dd9 100644 --- a/ql-wire/src/lib.rs +++ b/ql-wire/src/lib.rs @@ -33,9 +33,6 @@ pub use header::*; pub use identity::*; pub use pq::*; pub use qid::*; -pub use ql_common::{ - ResetCode, ResetOrigin, RouteId, ServiceId, StreamId, VarInt, VarIntBoundsExceeded, QID, -}; pub use record::*; #[cfg(any(feature = "test-utils", test))] pub use testing::*; diff --git a/ql-wire/src/macros.rs b/ql-wire/src/macros.rs index 8cd15a5a..29fc4e7d 100644 --- a/ql-wire/src/macros.rs +++ b/ql-wire/src/macros.rs @@ -12,7 +12,7 @@ macro_rules! varint_wrapper_codec { impl $crate::WireDecode for $name { fn decode(reader: &mut $crate::Reader) -> Result { - Ok(<$name>::from(reader.decode::<$crate::VarInt>()?)) + Ok(<$name>::from(reader.decode::<::ql_common::VarInt>()?)) } } }; diff --git a/ql-wire/src/qid.rs b/ql-wire/src/qid.rs index b5149bda..38672045 100644 --- a/ql-wire/src/qid.rs +++ b/ql-wire/src/qid.rs @@ -1,4 +1,6 @@ -use crate::{MlKemPublicKey, QlHash, ML_KEM_SUITE_TAG, QID}; +use ql_common::QID; + +use crate::{MlKemPublicKey, QlHash, ML_KEM_SUITE_TAG}; array_wrapper_codec!(QID); diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index b241c4aa..af6aa05e 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -1,5 +1,7 @@ use std::ops::RangeInclusive; +use ql_common::{ResetCode, StreamId, VarInt, QID}; + use super::*; fn decode_handshake_record(bytes: &[u8]) -> QlHandshakeRecord { From 032f92df7e0f2d0de07ca6ade0d34b17a1992f60 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 25 Jun 2026 09:06:55 -0400 Subject: [PATCH 28/59] ql-rpc: more internal cleanup --- ql-rpc/src/codec.rs | 2 - ql-rpc/src/router/builder.rs | 269 ++++---------------------- ql-rpc/src/router/mod.rs | 38 ++-- ql-rpc/src/rpc/download/mod.rs | 9 +- ql-rpc/src/rpc/download/server.rs | 56 +++--- ql-rpc/src/rpc/duplex/client.rs | 8 +- ql-rpc/src/rpc/duplex/mod.rs | 7 +- ql-rpc/src/rpc/duplex/server.rs | 34 ++-- ql-rpc/src/rpc/mod.rs | 12 +- ql-rpc/src/rpc/notification/mod.rs | 7 +- ql-rpc/src/rpc/notification/server.rs | 50 ++--- ql-rpc/src/rpc/parts.rs | 2 +- ql-rpc/src/rpc/progress/codec.rs | 6 +- ql-rpc/src/rpc/progress/mod.rs | 7 +- ql-rpc/src/rpc/progress/server.rs | 54 +++--- ql-rpc/src/rpc/request/mod.rs | 7 +- ql-rpc/src/rpc/request/server.rs | 49 ++--- ql-rpc/src/rpc/subscription/codec.rs | 6 +- ql-rpc/src/rpc/subscription/mod.rs | 7 +- ql-rpc/src/rpc/subscription/server.rs | 50 ++--- ql-rpc/src/rpc/upload/mod.rs | 9 +- ql-rpc/src/rpc/upload/server.rs | 60 +++--- ql-rpc/src/stream.rs | 9 - ql-runtime/src/tests/rpc.rs | 30 +-- 24 files changed, 300 insertions(+), 488 deletions(-) diff --git a/ql-rpc/src/codec.rs b/ql-rpc/src/codec.rs index 57f091f8..b2718c01 100644 --- a/ql-rpc/src/codec.rs +++ b/ql-rpc/src/codec.rs @@ -2,8 +2,6 @@ use std::{convert::Infallible, str::Utf8Error}; use bytes::{Buf, BufMut, Bytes}; -pub use crate::chunk_queue::ChunkQueue; - pub trait RpcCodec: Sized { type Error; diff --git a/ql-rpc/src/router/builder.rs b/ql-rpc/src/router/builder.rs index e655da28..05ff6e39 100644 --- a/ql-rpc/src/router/builder.rs +++ b/ql-rpc/src/router/builder.rs @@ -1,16 +1,8 @@ use std::marker::PhantomData; -use super::{ - LocalSpawner, RouteEntry, RouteFn, Router, RouterConfig, RpcStream, SendSpawner, Spawner, -}; +use super::*; use crate::{ - download::{server::*, Download as DownloadRpc}, - duplex::{server::*, Duplex as DuplexRpc}, - notification::{server::*, Notification as NotificationRpc}, - progress::{server::*, Progress as ProgressRpc}, - request::{server::*, Request as RequestRpc}, - subscription::{server::*, Subscription as SubscriptionRpc}, - upload::{server::*, Upload as UploadRpc}, + download::*, duplex::*, notification::*, progress::*, request::*, subscription::*, upload::*, RouteKey, }; @@ -81,155 +73,58 @@ where { pub fn request(self) -> Self where - M: RequestRpc + 'static, + M: Request + 'static, S: RequestHandlerLocal + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_request_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_request, S::handle, S::handle_error) } pub fn notification(self) -> Self where - M: NotificationRpc + 'static, + M: Notification + 'static, S: NotificationHandlerLocal + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_notification_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_notification, S::handle, S::handle_error) } pub fn duplex(self) -> Self where - M: DuplexRpc + 'static, + M: Duplex + 'static, S: DuplexHandlerLocal + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_duplex_inner::( - state, - context, - config, - reader, - writer, - S::handle, - )) - }, - ) + add_route!(self, M, handle_duplex, S::handle) } pub fn download(self) -> Self where - M: DownloadRpc + 'static, + M: Download + 'static, S: DownloadHandlerLocal + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_download_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_download, S::handle, S::handle_error) } pub fn subscription(self) -> Self where - M: SubscriptionRpc + 'static, + M: Subscription + 'static, S: SubscriptionHandlerLocal + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_subscription_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_subscription, S::handle, S::handle_error) } pub fn progress(self) -> Self where - M: ProgressRpc + 'static, + M: Progress + 'static, S: ProgressHandlerLocal + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_progress_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_progress, S::handle, S::handle_error) } pub fn upload(self) -> Self where - M: UploadRpc + 'static, + M: Upload + 'static, S: UploadHandlerLocal + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_upload_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_upload, S::handle, S::handle_error) } } @@ -240,176 +135,98 @@ where { pub fn request(self) -> Self where - M: RequestRpc + 'static, + M: Request + 'static, M::Request: Send + 'static, S: RequestHandler + Send + 'static, St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_request_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_request, S::handle, S::handle_error) } pub fn notification(self) -> Self where - M: NotificationRpc + 'static, + M: Notification + 'static, M::Payload: Send + 'static, S: NotificationHandler + Send + 'static, St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_notification_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_notification, S::handle, S::handle_error) } pub fn duplex(self) -> Self where - M: DuplexRpc + 'static, + M: Duplex + 'static, M::InitiatorEvent: Send + 'static, M::ResponderEvent: Send + 'static, S: DuplexHandler + Send + 'static, St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_duplex_inner::( - state, - context, - config, - reader, - writer, - S::handle, - )) - }, - ) + add_route!(self, M, handle_duplex, S::handle) } pub fn download(self) -> Self where - M: DownloadRpc + 'static, + M: Download + 'static, M::Request: Send + 'static, S: DownloadHandler + Send + 'static, St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_download_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_download, S::handle, S::handle_error) } pub fn subscription(self) -> Self where - M: SubscriptionRpc + 'static, + M: Subscription + 'static, M::Request: Send + 'static, S: SubscriptionHandler + Send + 'static, St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_subscription_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_subscription, S::handle, S::handle_error) } pub fn progress(self) -> Self where - M: ProgressRpc + 'static, + M: Progress + 'static, M::Request: Send + 'static, S: ProgressHandler + Send + 'static, St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route( - RouteKey::new::(), - |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_progress_inner::( - state, - context, - config, - reader, - writer, - S::handle, - S::handle_error, - )) - }, - ) + add_route!(self, M, handle_progress, S::handle, S::handle_error) } pub fn upload(self) -> Self where - M: UploadRpc + 'static, + M: Upload + 'static, M::Request: Send + 'static, S: UploadHandler + Send + 'static, St::Reader: Send + 'static, St::Writer: Send + 'static, { - self.add_route( - RouteKey::new::(), + add_route!(self, M, handle_upload, S::handle, S::handle_error) + } +} + +macro_rules! add_route { + ($builder:expr, $rpc:ty, $handler:ident, $($arg:path),+ $(,)?) => { + $builder.add_route( + RouteKey::new::<$rpc>(), |spawner, state, context, config, stream| { - let (reader, writer) = stream.split(); - spawner.spawn(handle_upload_inner::( + spawner.spawn($handler( state, context, config, - reader, - writer, - S::handle, - S::handle_error, + stream, + $($arg),+ )) }, ) - } + }; } + +use add_route; diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index 9bb691f5..568e5e3d 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -4,21 +4,8 @@ mod builder; mod config; mod mode; -pub use self::{ - builder::{LocalRoutes, RouterBuilder, SendRoutes}, - config::RouterConfig, - mode::*, -}; -pub use crate::{ - download::{DownloadHandler, DownloadHandlerLocal, DownloadStart, DownloadWriter}, - duplex::{DuplexHandler, DuplexHandlerLocal, DuplexPeer}, - notification::{NotificationHandler, NotificationHandlerLocal}, - progress::{ProgressHandler, ProgressHandlerLocal, ProgressResponder}, - request::{RequestHandler, RequestHandlerLocal, Response}, - subscription::{SubscriptionHandler, SubscriptionHandlerLocal, SubscriptionResponder}, - upload::{UploadHandler, UploadHandlerLocal, UploadReader, UploadResponder}, -}; -use crate::{reset_stream, RpcStream}; +pub use self::{builder::*, config::*, mode::*}; +use crate::{RpcRead, RpcStream, RpcWrite}; pub struct Router where @@ -88,7 +75,7 @@ where RouterBuilder::::new(spawner) } - pub fn handle(&self, info: StreamInfo, stream: St) -> Option<(RouteId, Sp::Handle)> { + pub fn handle(&self, info: StreamInfo, stream: St) -> Option { let StreamInfo { qid, stream_id, @@ -101,19 +88,18 @@ where route_id, }; let Ok(index) = self.routes.binary_search_by_key(&key, |entry| entry.key) else { - reset_stream(stream, ResetCode::UNKNOWN_ROUTE); + let (reader, writer) = stream.split(); + reader.reset(ResetCode::UNKNOWN_ROUTE); + writer.reset(ResetCode::UNKNOWN_ROUTE); return None; }; let route = self.routes[index].route; - Some(( - route_id, - route( - &self.spawner, - self.state.clone(), - context, - self.config, - stream, - ), + Some(route( + &self.spawner, + self.state.clone(), + context, + self.config, + stream, )) } diff --git a/ql-rpc/src/rpc/download/mod.rs b/ql-rpc/src/rpc/download/mod.rs index cfd4e374..3e9a613e 100644 --- a/ql-rpc/src/rpc/download/mod.rs +++ b/ql-rpc/src/rpc/download/mod.rs @@ -1,13 +1,10 @@ use super::Route; use crate::RpcCodec; -pub(crate) mod client; -pub(crate) mod server; +mod client; +mod server; -pub use client::{start, DownloadCall, DownloadPart, DownloadReader}; -pub use server::{ - DownloadHandler, DownloadHandlerLocal, DownloadPartWriter, DownloadStart, DownloadWriter, -}; +pub use self::{client::*, server::*}; /// rpc where the responder returns metadata first and then zero or more byte parts /// diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs index a5ebbccd..84fa3c08 100644 --- a/ql-rpc/src/rpc/download/server.rs +++ b/ql-rpc/src/rpc/download/server.rs @@ -5,7 +5,7 @@ use ql_common::ResetCode; use crate::{ codec, - download::Download as DownloadRpc, + download::Download, finish_bytes, rpc::{ parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, @@ -17,7 +17,7 @@ use crate::{ #[trait_variant::make(DownloadHandler: Send)] pub trait DownloadHandlerLocal where - M: DownloadRpc, + M: Download, St: RpcStream, { async fn handle( @@ -32,7 +32,7 @@ where pub struct DownloadStart where - M: DownloadRpc, + M: Download, W: RpcWrite, { writer: DropResetWrite, @@ -41,7 +41,7 @@ where pub struct DownloadWriter where - M: DownloadRpc, + M: Download, W: RpcWrite, { writer: DropResetWrite, @@ -50,7 +50,7 @@ where pub struct DownloadPartWriter<'a, M, W> where - M: DownloadRpc, + M: Download, W: RpcWrite, { parent: &'a mut DownloadWriter, @@ -59,7 +59,7 @@ where impl DownloadStart where - M: DownloadRpc, + M: Download, W: RpcWrite, { pub(crate) fn new(writer: W) -> Self { @@ -102,7 +102,7 @@ where impl DownloadWriter where - M: DownloadRpc, + M: Download, W: RpcWrite, { pub async fn start_part( @@ -134,7 +134,7 @@ where impl DownloadPartWriter<'_, M, W> where - M: DownloadRpc, + M: Download, W: RpcWrite, { pub async fn send(&mut self, bytes: Bytes) -> Result<(), W::Error> { @@ -156,7 +156,7 @@ where impl Drop for DownloadPartWriter<'_, M, W> where - M: DownloadRpc, + M: Download, W: RpcWrite, { fn drop(&mut self) { @@ -166,33 +166,37 @@ where } } -pub(crate) async fn handle_download_inner( +pub(crate) fn handle_download( state: S, context: Context, config: RouterConfig, - mut reader: St::Reader, - writer: St::Writer, + stream: St, handle: H, handle_error: E, -) where - M: DownloadRpc + 'static, +) -> impl Future +where + M: Download + 'static, St: RpcStream + 'static, H: FnOnce(S, Context, M::Request, DownloadStart) -> HF, HF: Future, E: FnOnce(&S, &RpcError), { - let request = match read_eof_request::(&mut reader, config).await { - Ok(request) => request, - Err(error) => { - let code = error.reset_code(); - handle_error(&state, &error); - if let Some(code) = code { - reader.reset(code); - writer.reset(code); + let (mut reader, writer) = stream.split(); + + async move { + let request = match read_eof_request::(&mut reader, config).await { + Ok(request) => request, + Err(error) => { + let code = error.reset_code(); + handle_error(&state, &error); + if let Some(code) = code { + reader.reset(code); + writer.reset(code); + } + return; } - return; - } - }; + }; - handle(state, context, request, DownloadStart::new(writer)).await; + handle(state, context, request, DownloadStart::new(writer)).await; + } } diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index 6a2cbe9a..346c08cd 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -8,8 +8,8 @@ use bytes::Bytes; use ql_common::ResetCode; use crate::{ - codec, duplex::Duplex, write_bytes, DropResetRead, DropResetWrite, RpcCodec, RpcError, RpcRead, - RpcStream, RpcWrite, + codec, duplex::Duplex, write_bytes, ChunkQueue, DropResetRead, DropResetWrite, RpcCodec, + RpcError, RpcRead, RpcStream, RpcWrite, }; pub fn start(stream: St) -> DuplexCall @@ -152,14 +152,14 @@ enum ReadStep { } struct EventReader { - bytes: codec::ChunkQueue, + bytes: ChunkQueue, marker: PhantomData T>, } impl Default for EventReader { fn default() -> Self { Self { - bytes: codec::ChunkQueue::default(), + bytes: ChunkQueue::default(), marker: PhantomData, } } diff --git a/ql-rpc/src/rpc/duplex/mod.rs b/ql-rpc/src/rpc/duplex/mod.rs index 11e15b6f..8eb385bd 100644 --- a/ql-rpc/src/rpc/duplex/mod.rs +++ b/ql-rpc/src/rpc/duplex/mod.rs @@ -1,11 +1,10 @@ use super::Route; use crate::RpcCodec; -pub(crate) mod client; -pub(crate) mod server; +mod client; +mod server; -pub use client::{start, DuplexCall, DuplexReceiver, DuplexSender}; -pub use server::{DuplexHandler, DuplexHandlerLocal, DuplexPeer}; +pub use self::{client::*, server::*}; /// rpc where both sides exchange typed events on the same stream /// diff --git a/ql-rpc/src/rpc/duplex/server.rs b/ql-rpc/src/rpc/duplex/server.rs index 6dfb94cc..fd2f941d 100644 --- a/ql-rpc/src/rpc/duplex/server.rs +++ b/ql-rpc/src/rpc/duplex/server.rs @@ -2,7 +2,7 @@ use std::future::Future; use crate::{ duplex::{Duplex, DuplexReceiver, DuplexSender}, - Context, RpcError, RpcRead, RpcStream, RpcWrite, + Context, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(DuplexHandler: Send)] @@ -12,8 +12,6 @@ where St: RpcStream, { async fn handle(self, context: Context, peer: DuplexPeer); - - fn handle_error(&self, _error: &RpcError) {} } pub struct DuplexPeer @@ -26,26 +24,30 @@ where pub receiver: DuplexReceiver, } -pub(crate) async fn handle_duplex_inner( +pub(crate) fn handle_duplex( state: S, context: Context, _config: crate::RouterConfig, - reader: St::Reader, - writer: St::Writer, + stream: St, handle: H, -) where +) -> impl Future +where M: Duplex + 'static, St: RpcStream + 'static, H: FnOnce(S, Context, DuplexPeer) -> HF, HF: Future, { - handle( - state, - context, - DuplexPeer { - sender: DuplexSender::new(writer), - receiver: DuplexReceiver::new(reader), - }, - ) - .await; + let (reader, writer) = stream.split(); + + async move { + handle( + state, + context, + DuplexPeer { + sender: DuplexSender::new(writer), + receiver: DuplexReceiver::new(reader), + }, + ) + .await; + } } diff --git a/ql-rpc/src/rpc/mod.rs b/ql-rpc/src/rpc/mod.rs index a8d08bcb..f51803aa 100644 --- a/ql-rpc/src/rpc/mod.rs +++ b/ql-rpc/src/rpc/mod.rs @@ -25,11 +25,9 @@ pub trait Route { const ROUTE: RouteId; } -pub use download::Download; -pub use duplex::Duplex; -pub use notification::Notification; -pub use progress::Progress; -pub use request::Request; -pub use subscription::Subscription; -pub use upload::Upload; use utils::*; + +pub use self::{ + download::Download, duplex::Duplex, notification::Notification, progress::Progress, + request::Request, subscription::Subscription, upload::Upload, +}; diff --git a/ql-rpc/src/rpc/notification/mod.rs b/ql-rpc/src/rpc/notification/mod.rs index 3da03299..2a2e9691 100644 --- a/ql-rpc/src/rpc/notification/mod.rs +++ b/ql-rpc/src/rpc/notification/mod.rs @@ -1,11 +1,10 @@ use super::Route; use crate::RpcCodec; -pub(crate) mod client; -pub(crate) mod server; +mod client; +mod server; -pub use client::send; -pub use server::{NotificationHandler, NotificationHandlerLocal}; +pub use self::{client::*, server::*}; /// one-way rpc that carries a single typed payload and no typed response /// diff --git a/ql-rpc/src/rpc/notification/server.rs b/ql-rpc/src/rpc/notification/server.rs index 282d6659..ed39ca66 100644 --- a/ql-rpc/src/rpc/notification/server.rs +++ b/ql-rpc/src/rpc/notification/server.rs @@ -3,14 +3,14 @@ use std::future::Future; use ql_common::ResetCode; use crate::{ - notification::Notification as NotificationRpc, rpc::read_eof_request, Context, RouterConfig, - RpcError, RpcRead, RpcStream, RpcWrite, + notification::Notification, rpc::read_eof_request, Context, RouterConfig, RpcCodec, RpcError, + RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(NotificationHandler: Send)] pub trait NotificationHandlerLocal where - M: NotificationRpc, + M: Notification, St: RpcStream, { async fn handle(self, context: Context, message: M::Payload); @@ -18,34 +18,38 @@ where fn handle_error(&self, _error: &RpcError) {} } -pub(crate) async fn handle_notification_inner( +pub(crate) fn handle_notification( state: S, context: Context, config: RouterConfig, - mut reader: St::Reader, - writer: St::Writer, + stream: St, handle: H, handle_error: E, -) where - M: NotificationRpc + 'static, +) -> impl Future +where + Payload: RpcCodec + 'static, St: RpcStream + 'static, - H: FnOnce(S, Context, M::Payload) -> HF, + H: FnOnce(S, Context, Payload) -> HF, HF: Future, - E: FnOnce(&S, &RpcError), + E: FnOnce(&S, &RpcError), { - let notification = match read_eof_request::(&mut reader, config).await { - Ok(notification) => notification, - Err(error) => { - let code = error.reset_code(); - handle_error(&state, &error); - if let Some(code) = code { - reader.reset(code); - writer.reset(code); + let (mut reader, writer) = stream.split(); + + async move { + let notification = match read_eof_request::(&mut reader, config).await { + Ok(notification) => notification, + Err(error) => { + let code = error.reset_code(); + handle_error(&state, &error); + if let Some(code) = code { + reader.reset(code); + writer.reset(code); + } + return; } - return; - } - }; + }; - writer.reset(ResetCode::CANCELLED); - handle(state, context, notification).await; + writer.reset(ResetCode::CANCELLED); + handle(state, context, notification).await; + } } diff --git a/ql-rpc/src/rpc/parts.rs b/ql-rpc/src/rpc/parts.rs index a1832159..f5032eb6 100644 --- a/ql-rpc/src/rpc/parts.rs +++ b/ql-rpc/src/rpc/parts.rs @@ -13,7 +13,7 @@ pub enum PartReadStep { } pub struct PartFrameReader { - bytes: codec::ChunkQueue, + bytes: ChunkQueue, pending_frame: PendingFrame, marker: PhantomData H>, } diff --git a/ql-rpc/src/rpc/progress/codec.rs b/ql-rpc/src/rpc/progress/codec.rs index e98d381d..a31c98c3 100644 --- a/ql-rpc/src/rpc/progress/codec.rs +++ b/ql-rpc/src/rpc/progress/codec.rs @@ -2,7 +2,7 @@ use std::marker::PhantomData; use bytes::Bytes; -use crate::{codec, progress::Progress, Error, RpcCodec, RpcError}; +use crate::{progress::Progress, ChunkQueue, Error, RpcCodec, RpcError}; pub enum ReadStep { NeedMore, @@ -11,14 +11,14 @@ pub enum ReadStep { } pub struct ResponseReader { - bytes: codec::ChunkQueue, + bytes: ChunkQueue, marker: PhantomData M>, } impl Default for ResponseReader { fn default() -> Self { Self { - bytes: codec::ChunkQueue::default(), + bytes: ChunkQueue::default(), marker: PhantomData, } } diff --git a/ql-rpc/src/rpc/progress/mod.rs b/ql-rpc/src/rpc/progress/mod.rs index bb93884d..327a35fb 100644 --- a/ql-rpc/src/rpc/progress/mod.rs +++ b/ql-rpc/src/rpc/progress/mod.rs @@ -1,12 +1,11 @@ use super::Route; use crate::RpcCodec; -pub(crate) mod client; +mod client; pub(crate) mod codec; -pub(crate) mod server; +mod server; -pub use client::{start, ProgressCall}; -pub use server::{ProgressHandler, ProgressHandlerLocal, ProgressResponder}; +pub use self::{client::*, server::*}; /// rpc where the responder streams progress values before a final response /// diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index 12ae6c7b..42875152 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -59,42 +59,46 @@ where } } -pub(crate) async fn handle_progress_inner( +pub(crate) fn handle_progress( state: S, context: Context, config: RouterConfig, - mut reader: St::Reader, - writer: St::Writer, + stream: St, handle: H, handle_error: E, -) where +) -> impl Future +where M: Progress + 'static, St: RpcStream + 'static, H: FnOnce(S, Context, M::Request, ProgressResponder) -> HF, HF: Future, E: FnOnce(&S, &RpcError), { - let request = match read_framed_request::(&mut reader, config).await { - Ok(request) => request, - Err(error) => { - let code = error.reset_code(); - handle_error(&state, &error); - if let Some(code) = code { - reader.reset(code); - writer.reset(code); + let (mut reader, writer) = stream.split(); + + async move { + let request = match read_framed_request::(&mut reader, config).await { + Ok(request) => request, + Err(error) => { + let code = error.reset_code(); + handle_error(&state, &error); + if let Some(code) = code { + reader.reset(code); + writer.reset(code); + } + return; } - return; - } - }; + }; - handle( - state, - context, - request, - ProgressResponder { - writer: DropResetWrite::new(writer), - marker: PhantomData, - }, - ) - .await; + handle( + state, + context, + request, + ProgressResponder { + writer: DropResetWrite::new(writer), + marker: PhantomData, + }, + ) + .await; + } } diff --git a/ql-rpc/src/rpc/request/mod.rs b/ql-rpc/src/rpc/request/mod.rs index 14f15953..3ad45db7 100644 --- a/ql-rpc/src/rpc/request/mod.rs +++ b/ql-rpc/src/rpc/request/mod.rs @@ -1,11 +1,10 @@ use super::Route; use crate::RpcCodec; -pub(crate) mod client; -pub(crate) mod server; +mod client; +mod server; -pub use client::call; -pub use server::{RequestHandler, RequestHandlerLocal, Response}; +pub use self::{client::*, server::*}; /// request-response rpc with exactly one typed value in each direction /// diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs index 68fdf862..ba434c4d 100644 --- a/ql-rpc/src/rpc/request/server.rs +++ b/ql-rpc/src/rpc/request/server.rs @@ -4,14 +4,14 @@ use bytes::Bytes; use ql_common::ResetCode; use crate::{ - finish_bytes, request::Request as RequestRpc, rpc::read_eof_request, write_bytes, Context, - DropResetWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, + finish_bytes, request::Request, rpc::read_eof_request, write_bytes, Context, DropResetWrite, + RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(RequestHandler: Send)] pub trait RequestHandlerLocal where - M: RequestRpc, + M: Request, St: RpcStream, { async fn handle( @@ -58,33 +58,38 @@ where } } -pub(crate) async fn handle_request_inner( +pub(crate) fn handle_request( state: S, context: Context, config: RouterConfig, - mut reader: St::Reader, - writer: St::Writer, + stream: St, handle: H, handle_error: E, -) where - M: RequestRpc + 'static, +) -> impl Future +where + Req: RpcCodec + 'static, + Res: RpcCodec + 'static, St: RpcStream + 'static, - H: FnOnce(S, Context, M::Request, Response) -> HF, + H: FnOnce(S, Context, Req, Response) -> HF, HF: Future, - E: FnOnce(&S, &RpcError), + E: FnOnce(&S, &RpcError), { - let request = match read_eof_request::(&mut reader, config).await { - Ok(request) => request, - Err(error) => { - let code = error.reset_code(); - handle_error(&state, &error); - if let Some(code) = code { - reader.reset(code); - writer.reset(code); + let (mut reader, writer) = stream.split(); + + async move { + let request = match read_eof_request::(&mut reader, config).await { + Ok(request) => request, + Err(error) => { + let code = error.reset_code(); + handle_error(&state, &error); + if let Some(code) = code { + reader.reset(code); + writer.reset(code); + } + return; } - return; - } - }; + }; - handle(state, context, request, Response::new(writer)).await; + handle(state, context, request, Response::new(writer)).await; + } } diff --git a/ql-rpc/src/rpc/subscription/codec.rs b/ql-rpc/src/rpc/subscription/codec.rs index 234e99b7..9563df0f 100644 --- a/ql-rpc/src/rpc/subscription/codec.rs +++ b/ql-rpc/src/rpc/subscription/codec.rs @@ -2,7 +2,7 @@ use std::marker::PhantomData; use bytes::Bytes; -use crate::{codec, subscription::Subscription, RpcCodec, RpcError}; +use crate::{subscription::Subscription, ChunkQueue, RpcCodec, RpcError}; pub enum ReadStep { NeedMore, @@ -10,14 +10,14 @@ pub enum ReadStep { } pub struct ResponseReader { - bytes: codec::ChunkQueue, + bytes: ChunkQueue, marker: PhantomData M>, } impl Default for ResponseReader { fn default() -> Self { Self { - bytes: codec::ChunkQueue::default(), + bytes: ChunkQueue::default(), marker: PhantomData, } } diff --git a/ql-rpc/src/rpc/subscription/mod.rs b/ql-rpc/src/rpc/subscription/mod.rs index e83b9751..d1c4346f 100644 --- a/ql-rpc/src/rpc/subscription/mod.rs +++ b/ql-rpc/src/rpc/subscription/mod.rs @@ -1,12 +1,11 @@ use super::Route; use crate::RpcCodec; -pub(crate) mod client; +mod client; pub(crate) mod codec; -pub(crate) mod server; +mod server; -pub use client::{start, SubscriptionCall}; -pub use server::{SubscriptionHandler, SubscriptionHandlerLocal, SubscriptionResponder}; +pub use self::{client::*, server::*}; /// rpc where one request opens a stream of typed events /// diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs index df2a4ec8..c0766067 100644 --- a/ql-rpc/src/rpc/subscription/server.rs +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -4,15 +4,14 @@ use bytes::Bytes; use ql_common::ResetCode; use crate::{ - codec, finish_bytes, rpc::read_eof_request, subscription::Subscription as SubscriptionRpc, - write_bytes, Context, DropResetWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, - RpcWrite, + codec, finish_bytes, rpc::read_eof_request, subscription::Subscription, write_bytes, Context, + DropResetWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, }; #[trait_variant::make(SubscriptionHandler: Send)] pub trait SubscriptionHandlerLocal where - M: SubscriptionRpc, + M: Subscription, St: RpcStream, { async fn handle( @@ -62,33 +61,38 @@ where } } -pub(crate) async fn handle_subscription_inner( +pub(crate) fn handle_subscription( state: S, context: Context, config: RouterConfig, - mut reader: St::Reader, - writer: St::Writer, + stream: St, handle: H, handle_error: E, -) where - M: SubscriptionRpc + 'static, +) -> impl Future +where + Req: RpcCodec + 'static, + Event: RpcCodec + 'static, St: RpcStream + 'static, - H: FnOnce(S, Context, M::Request, SubscriptionResponder) -> HF, + H: FnOnce(S, Context, Req, SubscriptionResponder) -> HF, HF: Future, - E: FnOnce(&S, &RpcError), + E: FnOnce(&S, &RpcError), { - let request = match read_eof_request::(&mut reader, config).await { - Ok(request) => request, - Err(error) => { - let code = error.reset_code(); - handle_error(&state, &error); - if let Some(code) = code { - reader.reset(code); - writer.reset(code); + let (mut reader, writer) = stream.split(); + + async move { + let request = match read_eof_request::(&mut reader, config).await { + Ok(request) => request, + Err(error) => { + let code = error.reset_code(); + handle_error(&state, &error); + if let Some(code) = code { + reader.reset(code); + writer.reset(code); + } + return; } - return; - } - }; + }; - handle(state, context, request, SubscriptionResponder::new(writer)).await; + handle(state, context, request, SubscriptionResponder::new(writer)).await; + } } diff --git a/ql-rpc/src/rpc/upload/mod.rs b/ql-rpc/src/rpc/upload/mod.rs index 433eeafb..de73630b 100644 --- a/ql-rpc/src/rpc/upload/mod.rs +++ b/ql-rpc/src/rpc/upload/mod.rs @@ -1,11 +1,10 @@ -use super::Route; +use super::*; use crate::RpcCodec; -pub(crate) mod client; -pub(crate) mod server; +mod client; +mod server; -pub use client::{start, UploadCall, UploadPartWriter}; -pub use server::{UploadHandler, UploadHandlerLocal, UploadPart, UploadReader, UploadResponder}; +pub use self::{client::*, server::*}; /// rpc where the caller uploads zero or more byte parts after a typed request /// diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index 299aa5a3..32676154 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -159,15 +159,15 @@ where } } -pub(crate) async fn handle_upload_inner( +pub(crate) fn handle_upload( state: S, context: Context, config: RouterConfig, - mut reader: St::Reader, - writer: St::Writer, + stream: St, handle: H, handle_error: E, -) where +) -> impl Future +where M: Upload + 'static, St: RpcStream + 'static, H: FnOnce( @@ -180,29 +180,33 @@ pub(crate) async fn handle_upload_inner( HF: Future, E: FnOnce(&S, &RpcError), { - let (request, buffered) = - match read_framed_request_prefix::(&mut reader, config).await { - Ok(value) => value, - Err(error) => { - let code = error.reset_code(); - handle_error(&state, &error); - if let Some(code) = code { - reader.reset(code); - writer.reset(code); + let (mut reader, writer) = stream.split(); + + async move { + let (request, buffered) = + match read_framed_request_prefix::(&mut reader, config).await { + Ok(value) => value, + Err(error) => { + let code = error.reset_code(); + handle_error(&state, &error); + if let Some(code) = code { + reader.reset(code); + writer.reset(code); + } + return; } - return; - } - }; - - handle( - state, - context, - request, - UploadReader { - stream: DropResetRead::new(reader), - reader: PartFrameReader::new(buffered), - }, - Response::new(writer), - ) - .await; + }; + + handle( + state, + context, + request, + UploadReader { + stream: DropResetRead::new(reader), + reader: PartFrameReader::new(buffered), + }, + Response::new(writer), + ) + .await; + } } diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index e6b1f73e..82e31994 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -67,15 +67,6 @@ where poll_fn(|cx| writer.poll_finish(cx)) } -pub fn reset_stream(stream: St, code: ResetCode) -where - St: RpcStream, -{ - let (reader, writer) = stream.split(); - reader.reset(code); - writer.reset(code); -} - pub(crate) use drop::*; mod drop { use super::*; diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index c72425b9..9c6b608a 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -10,10 +10,14 @@ use std::{ use bytes::Bytes; use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId}; use ql_rpc::{ - Context, DownloadHandlerLocal, DownloadStart, DuplexHandlerLocal, DuplexPeer, LocalSpawner, - NotificationHandlerLocal, ProgressHandlerLocal, ProgressResponder, RequestHandler, - RequestHandlerLocal, Response, SendSpawner, Spawner, SubscriptionHandlerLocal, - SubscriptionResponder, UploadHandlerLocal, UploadReader, UploadResponder, + download::{DownloadHandlerLocal, DownloadStart}, + duplex::{DuplexHandlerLocal, DuplexPeer}, + notification::NotificationHandlerLocal, + progress::{ProgressHandlerLocal, ProgressResponder}, + request::{RequestHandler, RequestHandlerLocal, Response}, + subscription::{SubscriptionHandlerLocal, SubscriptionResponder}, + upload::{UploadHandlerLocal, UploadReader, UploadResponder}, + Context, LocalSpawner, SendSpawner, Spawner, }; use super::*; @@ -180,7 +184,7 @@ async fn rpc_request() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { let fut = assert_send(fut); fut.await.unwrap(); } @@ -229,7 +233,7 @@ async fn rpc_notification() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -283,7 +287,7 @@ async fn rpc_subscrption() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -337,7 +341,7 @@ async fn rpc_router_enforces_max_request_bytes() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -393,7 +397,7 @@ async fn rpc_progress() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -456,7 +460,7 @@ async fn rpc_download() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -533,7 +537,7 @@ async fn rpc_download_complete() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -608,7 +612,7 @@ async fn rpc_upload() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { fut.await.unwrap(); } }); @@ -681,7 +685,7 @@ async fn rpc_duplex() { let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); - if let Some((_, fut)) = router.handle(info, stream) { + if let Some(fut) = router.handle(info, stream) { fut.await.unwrap(); } }); From a6f6c6464a151f60b61c12c66643d6c055427b21 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 25 Jun 2026 10:19:07 -0400 Subject: [PATCH 29/59] ql-fsm: remove unused error variant --- ql-fsm/src/error.rs | 2 -- ql-fsm/src/tests/session.rs | 4 +--- 2 files changed, 1 insertion(+), 5 deletions(-) diff --git a/ql-fsm/src/error.rs b/ql-fsm/src/error.rs index 9bf2a915..c6382371 100644 --- a/ql-fsm/src/error.rs +++ b/ql-fsm/src/error.rs @@ -89,7 +89,6 @@ impl Error for NoSessionError {} #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum StreamError { MissingStream, - NotWritable, NoSession, } @@ -97,7 +96,6 @@ impl Display for StreamError { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { let message = match self { Self::MissingStream => "missing stream", - Self::NotWritable => "stream is not writable", Self::NoSession => "no session", }; f.write_str(message) diff --git a/ql-fsm/src/tests/session.rs b/ql-fsm/src/tests/session.rs index cc240a0d..a70fd091 100644 --- a/ql-fsm/src/tests/session.rs +++ b/ql-fsm/src/tests/session.rs @@ -30,9 +30,7 @@ fn write_stream_bytes( ) -> Result { let mut bytes = Bytes::copy_from_slice(bytes); let mut stream = fsm.stream(stream_id)?; - let Some(mut writer) = stream.writer() else { - return Err(StreamError::NotWritable); - }; + let mut writer = stream.writer().expect("stream is not writable"); Ok(writer.write(&mut bytes)) } From 73cc6eaae406703e358bc748b436202ac6a820d1 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 25 Jun 2026 10:59:44 -0400 Subject: [PATCH 30/59] ql-wire: always have QID in the header --- ql-fsm/src/fsm.rs | 21 ++- ql-fsm/src/handshake/ik.rs | 37 ++-- ql-fsm/src/handshake/kk.rs | 34 ++-- ql-fsm/src/handshake/mod.rs | 28 +-- ql-fsm/src/handshake/xx.rs | 73 +++++--- ql-fsm/src/state.rs | 7 +- ql-fsm/src/tests/mod.rs | 5 +- ql-fsm/src/tests/proptest.rs | 2 +- ql-wire/src/encrypted/builder.rs | 21 ++- ql-wire/src/encrypted/mod.rs | 3 +- ql-wire/src/error.rs | 4 +- ql-wire/src/handshake/ik.rs | 52 +++--- ql-wire/src/handshake/kk.rs | 48 ++--- ql-wire/src/handshake/mod.rs | 40 +--- ql-wire/src/handshake/xx.rs | 82 ++++----- ql-wire/src/header.rs | 35 +++- ql-wire/src/record.rs | 27 ++- ql-wire/src/tests.rs | 305 ++++++++++++++++++++----------- 18 files changed, 499 insertions(+), 325 deletions(-) diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index 6e4be4a0..d52be08e 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -111,17 +111,23 @@ pub fn receive( if header.version != wire::QL_WIRE_VERSION { return Err(ReceiveError::InvalidRecordVersion); } + if header.route.recipient != fsm.identity.qid { + return Err(ReceiveError::InvalidQid); + } match header.record_type { wire::RecordType::Handshake => { let record = wire::QlHandshakeRecord::decode(&mut reader) .map_err(ReceiveError::InvalidHandshakeRecord)?; - handshake::handle_handshake_record(fsm, crypto, &record) + handshake::handle_handshake_record(fsm, crypto, header.route, &record) } wire::RecordType::Session => { let termination = { let QlFsm { state, events, .. } = fsm; let conn = state.link.connected_mut_or_err()?; + if header.route.sender != conn.transport.remote_qid { + return Err(ReceiveError::InvalidQid); + } let (decrypt_len, seq) = { let record = wire::QlSessionRecord::decode(&mut reader) .map_err(ReceiveError::InvalidSessionRecord)?; @@ -130,6 +136,7 @@ pub fn receive( } let payload = wire::decrypt_record( crypto, + &header, &record.header, record.payload, &conn.transport.rx_key, @@ -186,8 +193,11 @@ pub fn next_deadline(fsm: &QlFsm) -> Option { } pub fn take_next_write(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Option { - if let Some(record) = fsm.state.handshake.take() { - let record = wire::encode_record_vec(ql_wire::RecordType::Handshake, &record); + if let Some((route, record)) = fsm.state.handshake.take() { + let record = wire::encode_record_vec( + wire::RecordHeader::new(route, ql_wire::RecordType::Handshake), + &record, + ); return Some(OutboundWrite { record, write_id: None, @@ -196,10 +206,15 @@ pub fn take_next_write(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Option Result<(), ReceiveError> { - if should_ignore_inbound(fsm, message) { + if should_ignore_inbound(fsm, route, message) { return Ok(()); } - if message.header.recipient != fsm.identity.qid { - return Err(ReceiveError::InvalidQid); - } if let Some(peer) = fsm.state.peer.as_ref() { - if message.header.sender != peer.qid { + if route.sender != peer.qid { return Err(ReceiveError::InvalidQid); } } @@ -54,7 +56,7 @@ pub fn handle_ik1( super::local_transport_params(fsm), ); handshake - .read_1(crypto, message) + .read_1(crypto, route, message) .map_err(ReceiveError::InvalidIkHandshake)?; let outbound = handshake .write_2(crypto, message.meta) @@ -67,13 +69,21 @@ pub fn handle_ik1( .map_err(ReceiveError::InvalidIkHandshake)?, )?; fsm.state.handshake = None; - enqueue_handshake(fsm, QlHandshakeRecord::Ik2(outbound)); + enqueue_handshake( + fsm, + RouteHeader { + sender: fsm.identity.qid, + recipient: route.sender, + }, + QlHandshakeRecord::Ik2(outbound), + ); Ok(()) } pub fn handle_ik2( fsm: &mut QlFsm, crypto: &impl QlCrypto, + route: RouteHeader, message: &Ik2, ) -> Result<(), ReceiveError> { { @@ -87,7 +97,7 @@ pub fn handle_ik2( state .handshake - .read_2(crypto, message) + .read_2(crypto, route, message) .map_err(ReceiveError::InvalidIkHandshake)?; } @@ -104,7 +114,7 @@ pub fn handle_ik2( ) } -pub fn should_ignore_inbound(fsm: &QlFsm, message: &Ik1) -> bool { +pub fn should_ignore_inbound(fsm: &QlFsm, route: RouteHeader, message: &Ik1) -> bool { match &fsm.state.link { LinkState::Idle | LinkState::KkInitiator(_) @@ -113,11 +123,10 @@ pub fn should_ignore_inbound(fsm: &QlFsm, message: &Ik1) -> bool { LinkState::Connected(_) => super::is_connected_replay( fsm, message.meta.handshake_id, - message.header.sender, - message.header.recipient, + route.sender, ), LinkState::IkInitiator(state) => { - if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(message.header.sender) { + if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(route.sender) { return false; } super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) diff --git a/ql-fsm/src/handshake/kk.rs b/ql-fsm/src/handshake/kk.rs index 21eaa787..0ece2434 100644 --- a/ql-fsm/src/handshake/kk.rs +++ b/ql-fsm/src/handshake/kk.rs @@ -1,4 +1,4 @@ -use ql_wire::{self as wire, Kk1, Kk2, PeerBundle, QlCrypto, QlHandshakeRecord}; +use ql_wire::{self as wire, Kk1, Kk2, PeerBundle, QlCrypto, QlHandshakeRecord, RouteHeader}; use super::{ emit_peer_status, enqueue_handshake, establish_session, reset_connected_session_if_needed, @@ -10,6 +10,10 @@ use crate::{ pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle) { let meta = super::next_handshake_meta(fsm); + let route = RouteHeader { + sender: fsm.identity.qid, + recipient: peer.qid, + }; let mut handshake = wire::KkHandshake::new_initiator( crypto, fsm.identity.clone(), @@ -24,23 +28,24 @@ pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle handshake, deadline: fsm.state.now + fsm.config.handshake_timeout, }); - enqueue_handshake(fsm, QlHandshakeRecord::Kk1(message)); + enqueue_handshake(fsm, route, QlHandshakeRecord::Kk1(message)); emit_peer_status(fsm, fsm.state.link.status()); } pub fn handle_kk1( fsm: &mut QlFsm, crypto: &impl QlCrypto, + route: RouteHeader, message: &Kk1, ) -> Result<(), ReceiveError> { - if should_ignore_inbound(fsm, message) { + if should_ignore_inbound(fsm, route, message) { return Ok(()); } let Some(peer) = fsm.state.peer.clone() else { return Err(ReceiveError::NoPeer); }; - if message.header.recipient != fsm.identity.qid || message.header.sender != peer.qid { + if route.sender != peer.qid { return Err(ReceiveError::InvalidQid); } @@ -53,7 +58,7 @@ pub fn handle_kk1( super::local_transport_params(fsm), ); handshake - .read_1(crypto, message) + .read_1(crypto, route, message) .map_err(ReceiveError::InvalidKkHandshake)?; let outbound = handshake .write_2(crypto, message.meta) @@ -66,13 +71,21 @@ pub fn handle_kk1( .map_err(ReceiveError::InvalidKkHandshake)?, )?; fsm.state.handshake = None; - enqueue_handshake(fsm, QlHandshakeRecord::Kk2(outbound)); + enqueue_handshake( + fsm, + RouteHeader { + sender: fsm.identity.qid, + recipient: route.sender, + }, + QlHandshakeRecord::Kk2(outbound), + ); Ok(()) } pub fn handle_kk2( fsm: &mut QlFsm, crypto: &impl QlCrypto, + route: RouteHeader, message: &Kk2, ) -> Result<(), ReceiveError> { { @@ -86,7 +99,7 @@ pub fn handle_kk2( state .handshake - .read_2(crypto, message) + .read_2(crypto, route, message) .map_err(ReceiveError::InvalidKkHandshake)?; } @@ -103,18 +116,17 @@ pub fn handle_kk2( ) } -pub fn should_ignore_inbound(fsm: &QlFsm, message: &Kk1) -> bool { +pub fn should_ignore_inbound(fsm: &QlFsm, route: RouteHeader, message: &Kk1) -> bool { match &fsm.state.link { LinkState::Idle | LinkState::XxInitiator(_) | LinkState::XxResponder(_) => false, LinkState::Connected(_) => super::is_connected_replay( fsm, message.meta.handshake_id, - message.header.sender, - message.header.recipient, + route.sender, ), LinkState::IkInitiator(_) => true, LinkState::KkInitiator(state) => { - if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(message.header.sender) { + if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(route.sender) { return false; } super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) diff --git a/ql-fsm/src/handshake/mod.rs b/ql-fsm/src/handshake/mod.rs index dd2bfc5c..70ac7dc7 100644 --- a/ql-fsm/src/handshake/mod.rs +++ b/ql-fsm/src/handshake/mod.rs @@ -5,6 +5,7 @@ mod xx; use ql_common::QID; use ql_wire::{ self as wire, EphemeralPublicKey, HandshakeId, HandshakeMeta, QlCrypto, QlHandshakeRecord, + RouteHeader, }; use crate::{ @@ -39,9 +40,9 @@ pub fn next_handshake_meta(fsm: &mut QlFsm) -> HandshakeMeta { HandshakeMeta { handshake_id } } -pub fn enqueue_handshake(fsm: &mut QlFsm, record: QlHandshakeRecord) { +pub fn enqueue_handshake(fsm: &mut QlFsm, route: RouteHeader, record: QlHandshakeRecord) { debug_assert!(fsm.state.handshake.is_none()); - fsm.state.handshake = Some(record); + fsm.state.handshake = Some((route, record)); } pub fn handle_disarm_pairing(fsm: &mut QlFsm) { @@ -62,17 +63,18 @@ pub fn prepare_for_outbound_connect(fsm: &mut QlFsm) { pub fn handle_handshake_record( fsm: &mut QlFsm, crypto: &impl QlCrypto, + route: RouteHeader, record: &QlHandshakeRecord, ) -> Result<(), ReceiveError> { match record { - QlHandshakeRecord::Ik1(message) => ik::handle_ik1(fsm, crypto, message), - QlHandshakeRecord::Ik2(message) => ik::handle_ik2(fsm, crypto, message), - QlHandshakeRecord::Kk1(message) => kk::handle_kk1(fsm, crypto, message), - QlHandshakeRecord::Kk2(message) => kk::handle_kk2(fsm, crypto, message), - QlHandshakeRecord::Xx1(message) => xx::handle_xx1(fsm, crypto, message), - QlHandshakeRecord::Xx2(message) => xx::handle_xx2(fsm, crypto, message), - QlHandshakeRecord::Xx3(message) => xx::handle_xx3(fsm, crypto, message), - QlHandshakeRecord::Xx4(message) => xx::handle_xx4(fsm, crypto, message), + QlHandshakeRecord::Ik1(message) => ik::handle_ik1(fsm, crypto, route, message), + QlHandshakeRecord::Ik2(message) => ik::handle_ik2(fsm, crypto, route, message), + QlHandshakeRecord::Kk1(message) => kk::handle_kk1(fsm, crypto, route, message), + QlHandshakeRecord::Kk2(message) => kk::handle_kk2(fsm, crypto, route, message), + QlHandshakeRecord::Xx1(message) => xx::handle_xx1(fsm, crypto, route, message), + QlHandshakeRecord::Xx2(message) => xx::handle_xx2(fsm, crypto, route, message), + QlHandshakeRecord::Xx3(message) => xx::handle_xx3(fsm, crypto, route, message), + QlHandshakeRecord::Xx4(message) => xx::handle_xx4(fsm, crypto, route, message), } } @@ -143,6 +145,7 @@ pub fn establish_session( finalized: wire::FinalizedHandshake, ) -> Result<(), ReceiveError> { let transport = SessionTransport { + remote_qid: finalized.remote_bundle.qid, tx_key: finalized.tx_key, rx_key: finalized.rx_key, tx_connection_id: finalized.tx_connection_id, @@ -166,13 +169,10 @@ fn is_connected_replay( fsm: &QlFsm, handshake_id: HandshakeId, sender: QID, - recipient: QID, ) -> bool { let LinkState::Connected(connected) = &fsm.state.link else { return false; }; - connected.handshake_id == handshake_id - && recipient == fsm.identity.qid - && fsm.state.peer.as_ref().map(|peer| peer.qid) == Some(sender) + connected.handshake_id == handshake_id && fsm.state.peer.as_ref().map(|peer| peer.qid) == Some(sender) } diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs index cf325f8f..4f9952dc 100644 --- a/ql-fsm/src/handshake/xx.rs +++ b/ql-fsm/src/handshake/xx.rs @@ -1,5 +1,7 @@ use ql_common::QID; -use ql_wire::{self as wire, PairingToken, QlCrypto, QlHandshakeRecord, Xx1, Xx2, Xx3, Xx4}; +use ql_wire::{ + self as wire, PairingToken, QlCrypto, QlHandshakeRecord, RouteHeader, Xx1, Xx2, Xx3, Xx4, +}; use super::{ emit_peer_status, enqueue_handshake, establish_session, reset_connected_session_if_needed, @@ -16,6 +18,10 @@ pub fn start_initiator( remote_qid: QID, ) { let meta = super::next_handshake_meta(fsm); + let route = RouteHeader { + sender: fsm.identity.qid, + recipient: remote_qid, + }; let mut handshake = wire::XxHandshake::new_initiator( crypto, fsm.identity.clone(), @@ -31,16 +37,17 @@ pub fn start_initiator( handshake, deadline: fsm.state.now + fsm.config.handshake_timeout, }); - enqueue_handshake(fsm, QlHandshakeRecord::Xx1(message)); + enqueue_handshake(fsm, route, QlHandshakeRecord::Xx1(message)); emit_peer_status(fsm, fsm.state.link.status()); } pub fn handle_xx1( fsm: &mut QlFsm, crypto: &impl QlCrypto, + route: RouteHeader, message: &Xx1, ) -> Result<(), ReceiveError> { - if should_ignore_inbound(fsm, crypto, message) { + if should_ignore_inbound(fsm, crypto, route, message) { return Ok(()); } match fsm.state.armed_pairing_token { @@ -50,24 +57,18 @@ pub fn handle_xx1( actual: message.pairing_id, }) } - Some(_) - if message.header.recipient != fsm.identity.qid - || message.header.sender == fsm.identity.qid => - { - Err(ReceiveError::InvalidQid) - } Some(token) => { reset_connected_session_if_needed(fsm); let mut handshake = wire::XxHandshake::new_responder( crypto, fsm.identity.clone(), - message.header.sender, + route.sender, token, super::local_transport_params(fsm), ); handshake - .read_1(crypto, message) + .read_1(crypto, route, message) .map_err(ReceiveError::InvalidXxHandshake)?; let outbound = handshake .write_2(crypto, message.meta) @@ -78,7 +79,14 @@ pub fn handle_xx1( deadline: fsm.state.now + fsm.config.handshake_timeout, }); fsm.state.handshake = None; - enqueue_handshake(fsm, QlHandshakeRecord::Xx2(outbound)); + enqueue_handshake( + fsm, + RouteHeader { + sender: fsm.identity.qid, + recipient: route.sender, + }, + QlHandshakeRecord::Xx2(outbound), + ); Ok(()) } None => Err(ReceiveError::NotPairingMode), @@ -88,6 +96,7 @@ pub fn handle_xx1( pub fn handle_xx2( fsm: &mut QlFsm, crypto: &impl QlCrypto, + route: RouteHeader, message: &Xx2, ) -> Result<(), ReceiveError> { { @@ -101,14 +110,21 @@ pub fn handle_xx2( state .handshake - .read_2(crypto, message) + .read_2(crypto, route, message) .map_err(ReceiveError::InvalidXxHandshake)?; let outbound = state .handshake .write_3(crypto, message.meta) .map_err(ReceiveError::InvalidXxHandshake)?; fsm.state.handshake = None; - enqueue_handshake(fsm, QlHandshakeRecord::Xx3(outbound)); + enqueue_handshake( + fsm, + RouteHeader { + sender: fsm.identity.qid, + recipient: route.sender, + }, + QlHandshakeRecord::Xx3(outbound), + ); } Ok(()) @@ -117,6 +133,7 @@ pub fn handle_xx2( pub fn handle_xx3( fsm: &mut QlFsm, crypto: &impl QlCrypto, + route: RouteHeader, message: &Xx3, ) -> Result<(), ReceiveError> { let LinkState::XxResponder(state) = &mut fsm.state.link else { @@ -129,7 +146,7 @@ pub fn handle_xx3( state .handshake - .read_3(crypto, message) + .read_3(crypto, route, message) .map_err(ReceiveError::InvalidXxHandshake)?; let handshake_meta = state.handshake_meta; let LinkState::XxResponder(mut state) = fsm.state.link.take() else { @@ -140,7 +157,14 @@ pub fn handle_xx3( .write_4(crypto, handshake_meta) .map_err(ReceiveError::InvalidXxHandshake)?; fsm.state.handshake = None; - enqueue_handshake(fsm, QlHandshakeRecord::Xx4(outbound)); + enqueue_handshake( + fsm, + RouteHeader { + sender: fsm.identity.qid, + recipient: route.sender, + }, + QlHandshakeRecord::Xx4(outbound), + ); establish_session( fsm, message.meta.handshake_id, @@ -154,6 +178,7 @@ pub fn handle_xx3( pub fn handle_xx4( fsm: &mut QlFsm, crypto: &impl QlCrypto, + route: RouteHeader, message: &Xx4, ) -> Result<(), ReceiveError> { { @@ -167,7 +192,7 @@ pub fn handle_xx4( state .handshake - .read_4(crypto, message) + .read_4(crypto, route, message) .map_err(ReceiveError::InvalidXxHandshake)?; } @@ -191,23 +216,25 @@ pub fn disarm_pairing(fsm: &mut QlFsm) { } } -pub fn should_ignore_inbound(fsm: &QlFsm, crypto: &impl QlCrypto, message: &Xx1) -> bool { +pub fn should_ignore_inbound( + fsm: &QlFsm, + crypto: &impl QlCrypto, + route: RouteHeader, + message: &Xx1, +) -> bool { match &fsm.state.link { LinkState::Idle => false, LinkState::Connected(_) => super::is_connected_replay( fsm, message.meta.handshake_id, - message.header.sender, - message.header.recipient, + route.sender, ), LinkState::IkInitiator(_) | LinkState::KkInitiator(_) | LinkState::XxResponder(_) => true, LinkState::XxInitiator(state) => { if state.handshake.pairing_id(crypto) != message.pairing_id { return false; } - if message.header.recipient != fsm.identity.qid - || message.header.sender != state.handshake.remote_qid() - { + if route.sender != state.handshake.remote_qid() { return false; } super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) diff --git a/ql-fsm/src/state.rs b/ql-fsm/src/state.rs index d57661bd..6e406ae7 100644 --- a/ql-fsm/src/state.rs +++ b/ql-fsm/src/state.rs @@ -1,8 +1,10 @@ use std::time::Instant; +use ql_common::QID; use ql_wire::{ ConnectionId, EphemeralPublicKey, HandshakeId, HandshakeMeta, IkHandshake, KkHandshake, - PairingToken, PeerBundle, QlHandshakeRecord, SessionKey, TransportParams, XxHandshake, + PairingToken, PeerBundle, QlHandshakeRecord, RouteHeader, SessionKey, TransportParams, + XxHandshake, }; use crate::{session::SessionFsm, NoSessionError, PeerStatus}; @@ -11,13 +13,14 @@ pub struct QlFsmState { pub next_control_id: u32, pub peer: Option, pub armed_pairing_token: Option, - pub handshake: Option, + pub handshake: Option<(RouteHeader, QlHandshakeRecord)>, pub link: LinkState, pub now: Instant, } #[derive(Debug, Clone, PartialEq, Eq)] pub struct SessionTransport { + pub remote_qid: QID, pub tx_key: SessionKey, pub rx_key: SessionKey, pub tx_connection_id: ConnectionId, diff --git a/ql-fsm/src/tests/mod.rs b/ql-fsm/src/tests/mod.rs index dbc6bcd1..1bd84652 100644 --- a/ql-fsm/src/tests/mod.rs +++ b/ql-fsm/src/tests/mod.rs @@ -105,6 +105,7 @@ impl Harness { harness.a.fsm.state.link = LinkState::Connected(ConnectedState { handshake_id: HandshakeId(0), transport: SessionTransport { + remote_qid: harness.b.fsm.identity.qid, tx_key: a_to_b_key.clone(), rx_key: b_to_a_key.clone(), tx_connection_id: a_to_b_conn, @@ -122,6 +123,7 @@ impl Harness { harness.b.fsm.state.link = LinkState::Connected(ConnectedState { handshake_id: HandshakeId(0), transport: SessionTransport { + remote_qid: harness.a.fsm.identity.qid, tx_key: b_to_a_key, rx_key: a_to_b_key, tx_connection_id: b_to_a_conn, @@ -338,10 +340,11 @@ fn decrypt_record( record: &[u8], session_key: &SessionKey, ) -> (ql_wire::SessionHeader, Vec>>) { - let (_header, record) = + let (header, record) = ql_wire::decode_record::, _>(record).unwrap(); let plaintext = ql_wire::decrypt_record( crypto, + &header, &record.header, record.payload.into_owned(), session_key, diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs index 6f2dc07d..655204c7 100644 --- a/ql-fsm/src/tests/proptest.rs +++ b/ql-fsm/src/tests/proptest.rs @@ -968,7 +968,7 @@ proptest_crate::proptest! { let config = QlFsmConfig { session_record_ack_delay: Duration::from_millis(1), session_record_retransmit_timeout: Duration::from_millis(10), - session_record_max_size: 96, + session_record_max_size: ql_wire::SessionRecordBuilder::MIN_CAPACITY + 94, session_pending_ack_range_limit: 512, ..QlFsmConfig::default() }; diff --git a/ql-wire/src/encrypted/builder.rs b/ql-wire/src/encrypted/builder.rs index d64ed2d8..95b977c2 100644 --- a/ql-wire/src/encrypted/builder.rs +++ b/ql-wire/src/encrypted/builder.rs @@ -2,8 +2,8 @@ use bytes::BufMut; use super::{RecordAck, SessionClose, SessionFrame, StreamData, StreamReset, StreamWindow}; use crate::{ - BufView, ConnectionId, Nonce, QlCrypto, RecordSeq, RecordType, SessionHeader, SessionKey, - WireEncode, QL_WIRE_VERSION, + BufView, ConnectionId, Nonce, QlCrypto, RecordHeader, RecordSeq, RecordType, RouteHeader, + SessionHeader, SessionKey, WireEncode, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -15,15 +15,16 @@ pub struct SessionRecordBuilder { } impl SessionRecordBuilder { - pub const MIN_CAPACITY: usize = 1 - + 1 + pub const MIN_CAPACITY: usize = RecordHeader::WIRE_SIZE + ConnectionId::SIZE + RecordSeq::MAX_ENCODED_LEN + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; pub fn new(seq: RecordSeq, max_capacity: usize) -> Self { - let prefix_len = - 1 + 1 + ConnectionId::SIZE + seq.encoded_len() + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; + let prefix_len = RecordHeader::WIRE_SIZE + + ConnectionId::SIZE + + seq.encoded_len() + + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; assert!(max_capacity >= prefix_len); Self { seq, @@ -105,15 +106,17 @@ impl SessionRecordBuilder { pub fn encrypt( mut self, crypto: &impl QlCrypto, + route: RouteHeader, connection_id: ConnectionId, session_key: &SessionKey, ) -> Vec { self.ensure_prefix_capacity(0); + let record_header = RecordHeader::new(route, RecordType::Session); let header = SessionHeader { connection_id, seq: self.seq, }; - let aad = header.aad(); + let aad = header.aad(route); let nonce = Nonce::from_counter(self.seq.0.into_inner()); let auth = crypto.aes256_gcm_encrypt( session_key, @@ -123,9 +126,7 @@ impl SessionRecordBuilder { ); let mut prefix = &mut self.bytes[..self.prefix_len]; - prefix[0] = QL_WIRE_VERSION; - prefix[1] = RecordType::Session as u8; - prefix = &mut prefix[2..]; + record_header.encode(&mut prefix); header.encode(&mut prefix); auth.encode(&mut prefix); debug_assert!(prefix.is_empty()); diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index c079fba1..ca12e9e8 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -170,11 +170,12 @@ impl Iterator for SessionFrameIter { pub fn decrypt_record>( crypto: &impl QlCrypto, + record_header: &crate::RecordHeader, header: &SessionHeader, encrypted: EncryptedMessage, session_key: &SessionKey, ) -> Result { - let aad = header.aad(); + let aad = header.aad(record_header.route); let nonce = Nonce::from_counter(header.seq.0.into_inner()); let mut ciphertext = encrypted.ciphertext; if !crypto.aes256_gcm_decrypt( diff --git a/ql-wire/src/error.rs b/ql-wire/src/error.rs index 8da1eec0..c3c039a5 100644 --- a/ql-wire/src/error.rs +++ b/ql-wire/src/error.rs @@ -3,7 +3,7 @@ use core::fmt; #[derive(Debug, Clone, PartialEq, Eq)] pub enum WireError { InvalidPayload, - InvalidHandshakeHeader, + InvalidRouteHeader, InvalidHandshakeMeta, InvalidPairingId, InvalidRemoteBundle, @@ -17,7 +17,7 @@ impl fmt::Display for WireError { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { let message = match self { Self::InvalidPayload => "invalid payload", - Self::InvalidHandshakeHeader => "invalid handshake header", + Self::InvalidRouteHeader => "invalid route header", Self::InvalidHandshakeMeta => "invalid handshake meta", Self::InvalidPairingId => "invalid pairing id", Self::InvalidRemoteBundle => "invalid remote bundle", diff --git a/ql-wire/src/handshake/ik.rs b/ql-wire/src/handshake/ik.rs index 628e30e7..113caf1d 100644 --- a/ql-wire/src/handshake/ik.rs +++ b/ql-wire/src/handshake/ik.rs @@ -3,7 +3,7 @@ use super::{ finalize_handshake, generate_ephemeral_keypair, init_ik_symmetric, initialize_handshake_meta, mix_hash_ephemeral, mix_hash_routed_handshake, require_handshake_meta, EncryptedMlKemCiphertext, EncryptedPeerBundle, EphemeralKeyPair, EphemeralPublicKey, - FinalizedHandshake, HandshakeHeader, Role, SymmetricState, TransportParams, + FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, }; use crate::{ codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, @@ -12,7 +12,6 @@ use crate::{ #[derive(Debug, Clone, PartialEq, Eq)] pub struct Ik1 { - pub header: HandshakeHeader, pub meta: HandshakeMeta, pub transport_params: TransportParams, pub skem_ciphertext: MlKemCiphertext, @@ -23,7 +22,6 @@ pub struct Ik1 { impl codec::WireDecode for Ik1 { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - header: reader.decode()?, meta: reader.decode()?, transport_params: reader.decode()?, skem_ciphertext: reader.decode()?, @@ -35,8 +33,7 @@ impl codec::WireDecode for Ik1 { impl WireEncode for Ik1 { fn encoded_len(&self) -> usize { - HandshakeHeader::WIRE_SIZE - + HandshakeMeta::WIRE_SIZE + HandshakeMeta::WIRE_SIZE + TransportParams::WIRE_SIZE + MlKemCiphertext::SIZE + EphemeralPublicKey::WIRE_SIZE @@ -44,7 +41,6 @@ impl WireEncode for Ik1 { } fn encode(&self, out: &mut W) { - self.header.encode(out); self.meta.encode(out); self.transport_params.encode(out); self.skem_ciphertext.encode(out); @@ -55,7 +51,6 @@ impl WireEncode for Ik1 { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Ik2 { - pub header: HandshakeHeader, pub meta: HandshakeMeta, pub transport_params: TransportParams, pub ekem_ciphertext: MlKemCiphertext, @@ -63,8 +58,7 @@ pub struct Ik2 { } impl Ik2 { - pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE - + HandshakeMeta::WIRE_SIZE + pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE + TransportParams::WIRE_SIZE + MlKemCiphertext::SIZE + EncryptedMlKemCiphertext::WIRE_SIZE; @@ -73,7 +67,6 @@ impl Ik2 { impl codec::WireDecode for Ik2 { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - header: reader.decode()?, meta: reader.decode()?, transport_params: reader.decode()?, ekem_ciphertext: reader.decode()?, @@ -88,7 +81,6 @@ impl WireEncode for Ik2 { } fn encode(&self, out: &mut W) { - self.header.encode(out); self.meta.encode(out); self.transport_params.encode(out); self.ekem_ciphertext.encode(out); @@ -166,15 +158,15 @@ impl IkHandshake { self.step == IkStep::Done } - fn outbound_header(&self) -> Result { + fn outbound_header(&self) -> Result { let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; - Ok(HandshakeHeader { + Ok(RouteHeader { sender: self.local.qid, recipient: remote_bundle.qid, }) } - fn ensure_inbound_recipient(&self, header: HandshakeHeader) -> Result<(), WireError> { + fn ensure_inbound_recipient(&self, header: RouteHeader) -> Result<(), WireError> { if header.recipient == self.local.qid { Ok(()) } else { @@ -182,7 +174,7 @@ impl IkHandshake { } } - fn ensure_known_remote_sender(&self, header: HandshakeHeader) -> Result<(), WireError> { + fn ensure_known_remote_sender(&self, header: RouteHeader) -> Result<(), WireError> { if let Some(remote_bundle) = self.remote_bundle.as_ref() { if header.sender != remote_bundle.qid { return Err(WireError::InvalidPayload); @@ -225,7 +217,6 @@ impl IkHandshake { self.local_ephemeral = Some(local_ephemeral); self.step = IkStep::Recv2; Ok(Ik1 { - header, meta, transport_params: self.local_transport_params, skem_ciphertext, @@ -271,7 +262,6 @@ impl IkHandshake { self.step = IkStep::Done; Ok(Ik2 { - header, meta, transport_params: self.local_transport_params, ekem_ciphertext, @@ -279,17 +269,22 @@ impl IkHandshake { }) } - pub fn read_1(&mut self, crypto: &impl QlCrypto, message: &Ik1) -> Result<(), WireError> { + pub fn read_1( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Ik1, + ) -> Result<(), WireError> { if self.step != IkStep::Recv1 { return Err(WireError::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; - self.ensure_inbound_recipient(message.header)?; - self.ensure_known_remote_sender(message.header)?; + self.ensure_inbound_recipient(header)?; + self.ensure_known_remote_sender(header)?; mix_hash_routed_handshake( &mut self.symmetric, crypto, - message.header, + header, HandshakeKind::Ik1, message.meta, message.transport_params, @@ -306,7 +301,7 @@ impl IkHandshake { let remote_bundle = decrypt_peer_bundle(crypto, &mut self.symmetric, &message.static_bundle)?; - if remote_bundle.qid != message.header.sender { + if remote_bundle.qid != header.sender { return Err(WireError::InvalidPayload); } match self.remote_bundle.as_ref() { @@ -321,17 +316,22 @@ impl IkHandshake { Ok(()) } - pub fn read_2(&mut self, crypto: &impl QlCrypto, message: &Ik2) -> Result<(), WireError> { + pub fn read_2( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Ik2, + ) -> Result<(), WireError> { if self.step != IkStep::Recv2 { return Err(WireError::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_recipient(message.header)?; - self.ensure_known_remote_sender(message.header)?; + self.ensure_inbound_recipient(header)?; + self.ensure_known_remote_sender(header)?; mix_hash_routed_handshake( &mut self.symmetric, crypto, - message.header, + header, HandshakeKind::Ik2, message.meta, message.transport_params, diff --git a/ql-wire/src/handshake/kk.rs b/ql-wire/src/handshake/kk.rs index 2ad5ee2a..e71775b3 100644 --- a/ql-wire/src/handshake/kk.rs +++ b/ql-wire/src/handshake/kk.rs @@ -2,7 +2,7 @@ use super::{ decrypt_mlkem_ciphertext, encrypt_mlkem_ciphertext, finalize_handshake, generate_ephemeral_keypair, init_kk_symmetric, initialize_handshake_meta, mix_hash_ephemeral, mix_hash_routed_handshake, require_handshake_meta, EncryptedMlKemCiphertext, EphemeralKeyPair, - EphemeralPublicKey, FinalizedHandshake, HandshakeHeader, Role, SymmetricState, TransportParams, + EphemeralPublicKey, FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, }; use crate::{ codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, @@ -11,7 +11,6 @@ use crate::{ #[derive(Debug, Clone, PartialEq, Eq)] pub struct Kk1 { - pub header: HandshakeHeader, pub meta: HandshakeMeta, pub transport_params: TransportParams, pub skem_ciphertext: MlKemCiphertext, @@ -19,8 +18,7 @@ pub struct Kk1 { } impl Kk1 { - pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE - + HandshakeMeta::WIRE_SIZE + pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE + TransportParams::WIRE_SIZE + MlKemCiphertext::SIZE + EphemeralPublicKey::WIRE_SIZE; @@ -29,7 +27,6 @@ impl Kk1 { impl codec::WireDecode for Kk1 { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - header: reader.decode()?, meta: reader.decode()?, transport_params: reader.decode()?, skem_ciphertext: reader.decode()?, @@ -44,7 +41,6 @@ impl WireEncode for Kk1 { } fn encode(&self, out: &mut W) { - self.header.encode(out); self.meta.encode(out); self.transport_params.encode(out); self.skem_ciphertext.encode(out); @@ -54,7 +50,6 @@ impl WireEncode for Kk1 { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Kk2 { - pub header: HandshakeHeader, pub meta: HandshakeMeta, pub transport_params: TransportParams, pub ekem_ciphertext: MlKemCiphertext, @@ -62,8 +57,7 @@ pub struct Kk2 { } impl Kk2 { - pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE - + HandshakeMeta::WIRE_SIZE + pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE + TransportParams::WIRE_SIZE + MlKemCiphertext::SIZE + EncryptedMlKemCiphertext::WIRE_SIZE; @@ -72,7 +66,6 @@ impl Kk2 { impl codec::WireDecode for Kk2 { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - header: reader.decode()?, meta: reader.decode()?, transport_params: reader.decode()?, ekem_ciphertext: reader.decode()?, @@ -87,7 +80,6 @@ impl WireEncode for Kk2 { } fn encode(&self, out: &mut W) { - self.header.encode(out); self.meta.encode(out); self.transport_params.encode(out); self.ekem_ciphertext.encode(out); @@ -165,21 +157,21 @@ impl KkHandshake { self.step == KkStep::Done } - fn outbound_header(&self) -> HandshakeHeader { - HandshakeHeader { + fn outbound_header(&self) -> RouteHeader { + RouteHeader { sender: self.local.qid, recipient: self.remote_bundle.qid, } } - fn inbound_header(&self) -> HandshakeHeader { - HandshakeHeader { + fn inbound_header(&self) -> RouteHeader { + RouteHeader { sender: self.remote_bundle.qid, recipient: self.local.qid, } } - fn ensure_inbound_header(&self, header: HandshakeHeader) -> Result<(), WireError> { + fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), WireError> { if header == self.inbound_header() { Ok(()) } else { @@ -219,7 +211,6 @@ impl KkHandshake { self.local_ephemeral = Some(local_ephemeral); self.step = KkStep::Recv2; Ok(Kk1 { - header, meta, transport_params: self.local_transport_params, skem_ciphertext, @@ -263,7 +254,6 @@ impl KkHandshake { self.step = KkStep::Done; Ok(Kk2 { - header, meta, transport_params: self.local_transport_params, ekem_ciphertext, @@ -271,16 +261,21 @@ impl KkHandshake { }) } - pub fn read_1(&mut self, crypto: &impl QlCrypto, message: &Kk1) -> Result<(), WireError> { + pub fn read_1( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Kk1, + ) -> Result<(), WireError> { if self.step != KkStep::Recv1 { return Err(WireError::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; - self.ensure_inbound_header(message.header)?; + self.ensure_inbound_header(header)?; mix_hash_routed_handshake( &mut self.symmetric, crypto, - message.header, + header, HandshakeKind::Kk1, message.meta, message.transport_params, @@ -299,16 +294,21 @@ impl KkHandshake { Ok(()) } - pub fn read_2(&mut self, crypto: &impl QlCrypto, message: &Kk2) -> Result<(), WireError> { + pub fn read_2( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Kk2, + ) -> Result<(), WireError> { if self.step != KkStep::Recv2 { return Err(WireError::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_header(message.header)?; + self.ensure_inbound_header(header)?; mix_hash_routed_handshake( &mut self.symmetric, crypto, - message.header, + header, HandshakeKind::Kk2, message.meta, message.transport_params, diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index b8bfe204..34897b27 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -1,9 +1,7 @@ -use ql_common::QID; - use crate::{ codec, derive_qid, ByteSlice, ConnectionId, HandshakeKind, MlKemCiphertext, MlKemKeyPair, - MlKemPublicKey, Nonce, PeerBundle, QlCrypto, SessionKey, WireDecode, WireEncode, WireError, - ENCRYPTED_MESSAGE_AUTH_SIZE, + MlKemPublicKey, Nonce, PeerBundle, QlCrypto, RouteHeader, SessionKey, WireDecode, WireEncode, + WireError, ENCRYPTED_MESSAGE_AUTH_SIZE, }; mod ik; @@ -27,36 +25,6 @@ const PROTOCOL_XX: &[u8] = b"ql-wire:pq-xx:v1"; const CONNECTION_ID_DOMAIN: &[u8] = b"ql-wire:conn-id:v1"; const HANDSHAKE_PREAMBLE_DOMAIN: &[u8] = b"ql-wire:handshake-preamble:v1"; -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct HandshakeHeader { - pub sender: QID, - pub recipient: QID, -} - -impl HandshakeHeader { - pub const WIRE_SIZE: usize = QID::SIZE * 2; -} - -impl WireEncode for HandshakeHeader { - fn encoded_len(&self) -> usize { - Self::WIRE_SIZE - } - - fn encode(&self, out: &mut W) { - self.sender.encode(out); - self.recipient.encode(out); - } -} - -impl codec::WireDecode for HandshakeHeader { - fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self { - sender: reader.decode()?, - recipient: reader.decode()?, - }) - } -} - #[derive(Debug, Clone, PartialEq, Eq)] pub struct EphemeralPublicKey { pub mlkem_public_key: MlKemPublicKey, @@ -361,7 +329,7 @@ fn mix_hash_ephemeral( fn mix_hash_routed_handshake( symmetric: &mut SymmetricState, crypto: &impl QlCrypto, - header: HandshakeHeader, + header: RouteHeader, kind: HandshakeKind, meta: HandshakeMeta, transport_params: TransportParams, @@ -379,7 +347,7 @@ fn mix_hash_routed_handshake( fn mix_hash_pairing_handshake( symmetric: &mut SymmetricState, crypto: &impl QlCrypto, - header: HandshakeHeader, + header: RouteHeader, kind: HandshakeKind, meta: HandshakeMeta, pairing_id: PairingId, diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs index f812744a..f859f2e3 100644 --- a/ql-wire/src/handshake/xx.rs +++ b/ql-wire/src/handshake/xx.rs @@ -6,7 +6,7 @@ use super::{ initialize_transport_params, mix_hash_ephemeral, mix_hash_pairing_handshake, mix_psk_pairing_token, require_handshake_meta, require_transport_params, EncryptedMlKemCiphertext, EncryptedPeerBundle, EphemeralKeyPair, EphemeralPublicKey, - FinalizedHandshake, HandshakeHeader, Role, SymmetricState, TransportParams, + FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, }; use crate::{ codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PairingId, PairingToken, @@ -15,7 +15,6 @@ use crate::{ #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx1 { - pub header: HandshakeHeader, pub meta: HandshakeMeta, pub pairing_id: PairingId, pub transport_params: TransportParams, @@ -23,8 +22,7 @@ pub struct Xx1 { } impl Xx1 { - pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE - + HandshakeMeta::WIRE_SIZE + pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE + PairingId::SIZE + TransportParams::WIRE_SIZE + EphemeralPublicKey::WIRE_SIZE; @@ -33,7 +31,6 @@ impl Xx1 { impl codec::WireDecode for Xx1 { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - header: reader.decode()?, meta: reader.decode()?, pairing_id: reader.decode()?, transport_params: reader.decode()?, @@ -48,7 +45,6 @@ impl WireEncode for Xx1 { } fn encode(&self, out: &mut W) { - self.header.encode(out); self.meta.encode(out); self.pairing_id.encode(out); self.transport_params.encode(out); @@ -58,7 +54,6 @@ impl WireEncode for Xx1 { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx2 { - pub header: HandshakeHeader, pub meta: HandshakeMeta, pub pairing_id: PairingId, pub transport_params: TransportParams, @@ -69,7 +64,6 @@ pub struct Xx2 { impl codec::WireDecode for Xx2 { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - header: reader.decode()?, meta: reader.decode()?, pairing_id: reader.decode()?, transport_params: reader.decode()?, @@ -81,8 +75,7 @@ impl codec::WireDecode for Xx2 { impl WireEncode for Xx2 { fn encoded_len(&self) -> usize { - HandshakeHeader::WIRE_SIZE - + HandshakeMeta::WIRE_SIZE + HandshakeMeta::WIRE_SIZE + PairingId::SIZE + TransportParams::WIRE_SIZE + MlKemCiphertext::SIZE @@ -90,7 +83,6 @@ impl WireEncode for Xx2 { } fn encode(&self, out: &mut W) { - self.header.encode(out); self.meta.encode(out); self.pairing_id.encode(out); self.transport_params.encode(out); @@ -101,7 +93,6 @@ impl WireEncode for Xx2 { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx3 { - pub header: HandshakeHeader, pub meta: HandshakeMeta, pub pairing_id: PairingId, pub transport_params: TransportParams, @@ -112,7 +103,6 @@ pub struct Xx3 { impl codec::WireDecode for Xx3 { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - header: reader.decode()?, meta: reader.decode()?, pairing_id: reader.decode()?, transport_params: reader.decode()?, @@ -124,8 +114,7 @@ impl codec::WireDecode for Xx3 { impl WireEncode for Xx3 { fn encoded_len(&self) -> usize { - HandshakeHeader::WIRE_SIZE - + HandshakeMeta::WIRE_SIZE + HandshakeMeta::WIRE_SIZE + PairingId::SIZE + TransportParams::WIRE_SIZE + EncryptedMlKemCiphertext::WIRE_SIZE @@ -133,7 +122,6 @@ impl WireEncode for Xx3 { } fn encode(&self, out: &mut W) { - self.header.encode(out); self.meta.encode(out); self.pairing_id.encode(out); self.transport_params.encode(out); @@ -144,7 +132,6 @@ impl WireEncode for Xx3 { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx4 { - pub header: HandshakeHeader, pub meta: HandshakeMeta, pub pairing_id: PairingId, pub transport_params: TransportParams, @@ -152,8 +139,7 @@ pub struct Xx4 { } impl Xx4 { - pub const WIRE_SIZE: usize = HandshakeHeader::WIRE_SIZE - + HandshakeMeta::WIRE_SIZE + pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE + PairingId::SIZE + TransportParams::WIRE_SIZE + EncryptedMlKemCiphertext::WIRE_SIZE; @@ -162,7 +148,6 @@ impl Xx4 { impl codec::WireDecode for Xx4 { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - header: reader.decode()?, meta: reader.decode()?, pairing_id: reader.decode()?, transport_params: reader.decode()?, @@ -177,7 +162,6 @@ impl WireEncode for Xx4 { } fn encode(&self, out: &mut W) { - self.header.encode(out); self.meta.encode(out); self.pairing_id.encode(out); self.transport_params.encode(out); @@ -281,8 +265,8 @@ impl XxHandshake { self.remote_bundle.as_ref() } - fn header(&self) -> HandshakeHeader { - HandshakeHeader { + fn header(&self) -> RouteHeader { + RouteHeader { sender: self.local.qid, recipient: self.remote_qid, } @@ -291,11 +275,11 @@ impl XxHandshake { fn ensure_inbound_header( &self, crypto: &impl QlCrypto, - header: HandshakeHeader, + header: RouteHeader, pairing_id: PairingId, ) -> Result<(), WireError> { if header.sender != self.remote_qid || header.recipient != self.local.qid { - return Err(WireError::InvalidHandshakeHeader); + return Err(WireError::InvalidRouteHeader); } if pairing_id != self.pairing_token.id(crypto) { return Err(WireError::InvalidPairingId); @@ -340,7 +324,6 @@ impl XxHandshake { self.local_ephemeral = Some(local_ephemeral); self.step = XxStep::Recv2; Ok(Xx1 { - header, meta, pairing_id, transport_params: self.local_transport_params, @@ -348,16 +331,21 @@ impl XxHandshake { }) } - pub fn read_1(&mut self, crypto: &impl QlCrypto, message: &Xx1) -> Result<(), WireError> { + pub fn read_1( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Xx1, + ) -> Result<(), WireError> { if self.step != XxStep::Recv1 { return Err(WireError::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; - self.ensure_inbound_header(crypto, message.header, message.pairing_id)?; + self.ensure_inbound_header(crypto, header, message.pairing_id)?; mix_hash_pairing_handshake( &mut self.symmetric, crypto, - message.header, + header, HandshakeKind::Xx1, message.meta, message.pairing_id, @@ -406,7 +394,6 @@ impl XxHandshake { self.step = XxStep::Recv3; Ok(Xx2 { - header, meta, pairing_id, transport_params: self.local_transport_params, @@ -415,16 +402,21 @@ impl XxHandshake { }) } - pub fn read_2(&mut self, crypto: &impl QlCrypto, message: &Xx2) -> Result<(), WireError> { + pub fn read_2( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Xx2, + ) -> Result<(), WireError> { if self.step != XxStep::Recv2 { return Err(WireError::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_header(crypto, message.header, message.pairing_id)?; + self.ensure_inbound_header(crypto, header, message.pairing_id)?; mix_hash_pairing_handshake( &mut self.symmetric, crypto, - message.header, + header, HandshakeKind::Xx2, message.meta, message.pairing_id, @@ -483,7 +475,6 @@ impl XxHandshake { self.step = XxStep::Recv4; Ok(Xx3 { - header, meta, pairing_id, transport_params: self.local_transport_params, @@ -492,12 +483,17 @@ impl XxHandshake { }) } - pub fn read_3(&mut self, crypto: &impl QlCrypto, message: &Xx3) -> Result<(), WireError> { + pub fn read_3( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Xx3, + ) -> Result<(), WireError> { if self.step != XxStep::Recv3 { return Err(WireError::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_header(crypto, message.header, message.pairing_id)?; + self.ensure_inbound_header(crypto, header, message.pairing_id)?; require_transport_params( self.remote_transport_params.as_ref(), message.transport_params, @@ -505,7 +501,7 @@ impl XxHandshake { mix_hash_pairing_handshake( &mut self.symmetric, crypto, - message.header, + header, HandshakeKind::Xx3, message.meta, message.pairing_id, @@ -557,7 +553,6 @@ impl XxHandshake { self.step = XxStep::Done; Ok(Xx4 { - header, meta, pairing_id, transport_params: self.local_transport_params, @@ -565,12 +560,17 @@ impl XxHandshake { }) } - pub fn read_4(&mut self, crypto: &impl QlCrypto, message: &Xx4) -> Result<(), WireError> { + pub fn read_4( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Xx4, + ) -> Result<(), WireError> { if self.step != XxStep::Recv4 { return Err(WireError::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_header(crypto, message.header, message.pairing_id)?; + self.ensure_inbound_header(crypto, header, message.pairing_id)?; require_transport_params( self.remote_transport_params.as_ref(), message.transport_params, @@ -578,7 +578,7 @@ impl XxHandshake { mix_hash_pairing_handshake( &mut self.symmetric, crypto, - message.header, + header, HandshakeKind::Xx4, message.meta, message.pairing_id, diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index bfe496f3..0d232903 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -1,7 +1,38 @@ use ::bytes::BufMut; +use ql_common::QID; use crate::{codec, ByteSlice, WireEncode, WireError, QL_WIRE_VERSION}; +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct RouteHeader { + pub sender: QID, + pub recipient: QID, +} + +impl RouteHeader { + pub const WIRE_SIZE: usize = QID::SIZE * 2; +} + +impl WireEncode for RouteHeader { + fn encoded_len(&self) -> usize { + Self::WIRE_SIZE + } + + fn encode(&self, out: &mut W) { + self.sender.encode(out); + self.recipient.encode(out); + } +} + +impl codec::WireDecode for RouteHeader { + fn decode(reader: &mut codec::Reader) -> Result { + Ok(Self { + sender: reader.decode()?, + recipient: reader.decode()?, + }) + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct SessionHeader { pub connection_id: ConnectionId, @@ -19,16 +50,18 @@ impl SessionHeader { const AAD_DOMAIN: &[u8] = b"ql-wire:session-aad:v1"; const AAD_RECORD_KIND_SESSION: u8 = 1; - pub fn aad(&self) -> Vec { + pub fn aad(&self, route: RouteHeader) -> Vec { let aad_len = Self::AAD_DOMAIN.len() + size_of::() + size_of::() + + RouteHeader::WIRE_SIZE + ConnectionId::SIZE + self.seq.encoded_len(); let mut aad = Vec::with_capacity(aad_len); aad.put_slice(Self::AAD_DOMAIN); aad.put_u8(QL_WIRE_VERSION); aad.put_u8(Self::AAD_RECORD_KIND_SESSION); + route.encode(&mut aad); self.connection_id.encode(&mut aad); self.seq.encode(&mut aad); debug_assert_eq!(aad.len(), aad_len); diff --git a/ql-wire/src/record.rs b/ql-wire/src/record.rs index 163a1bff..51a29085 100644 --- a/ql-wire/src/record.rs +++ b/ql-wire/src/record.rs @@ -2,25 +2,21 @@ use crate::{ codec, encrypted_message::EncryptedMessage, handshake::{Ik1, Ik2, Kk1, Kk2, Xx1, Xx2, Xx3, Xx4}, - ByteSlice, SessionHeader, WireDecode, WireEncode, WireError, QL_WIRE_VERSION, + ByteSlice, RouteHeader, SessionHeader, WireDecode, WireEncode, WireError, QL_WIRE_VERSION, }; -pub fn encode_record(out: &mut W, record_type: RecordType, body: &T) +pub fn encode_record(out: &mut W, header: RecordHeader, body: &T) where W: bytes::BufMut + ?Sized, T: WireEncode + ?Sized, { - RecordHeader { - version: QL_WIRE_VERSION, - record_type, - } - .encode(out); + header.encode(out); body.encode(out); } -pub fn encode_record_vec(record_type: RecordType, body: &T) -> Vec { +pub fn encode_record_vec(header: RecordHeader, body: &T) -> Vec { let mut out = Vec::with_capacity(RecordHeader::WIRE_SIZE + body.encoded_len()); - encode_record(&mut out, record_type, body); + encode_record(&mut out, header, body); out } @@ -36,17 +32,27 @@ where #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RecordHeader { pub version: u8, + pub route: RouteHeader, pub record_type: RecordType, } impl RecordHeader { - pub const WIRE_SIZE: usize = size_of::() + size_of::(); + pub const WIRE_SIZE: usize = size_of::() + RouteHeader::WIRE_SIZE + size_of::(); + + pub fn new(route: RouteHeader, record_type: RecordType) -> Self { + Self { + version: QL_WIRE_VERSION, + route, + record_type, + } + } } impl WireDecode for RecordHeader { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { version: reader.decode()?, + route: reader.decode()?, record_type: reader.decode()?, }) } @@ -59,6 +65,7 @@ impl WireEncode for RecordHeader { fn encode(&self, out: &mut W) { out.put_u8(self.version); + self.route.encode(out); self.record_type.encode(out); } } diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index af6aa05e..370f6097 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -13,24 +13,12 @@ fn decode_session_record(bytes: &[u8]) -> QlSessionRecord> { record.into_owned() } -fn qid(byte: u8) -> QID { - QID([byte; QID::SIZE]) -} - fn varint(value: u64) -> VarInt { VarInt::from_u64(value).unwrap() } -fn record_seq(value: u64) -> RecordSeq { - RecordSeq(varint(value)) -} - fn record_ack_range(start: u64, end: u64) -> RangeInclusive { - record_seq(start)..=record_seq(end) -} - -fn stream_id(value: u64) -> StreamId { - StreamId(varint(value)) + RecordSeq(varint(start))..=RecordSeq(varint(end)) } fn handshake_meta(id: u32) -> HandshakeMeta { @@ -45,30 +33,29 @@ fn handshake_transport_params(window: u32) -> TransportParams { } } -fn handshake_header(sender: u8, recipient: u8) -> HandshakeHeader { - HandshakeHeader { - sender: qid(sender), - recipient: qid(recipient), +fn route(sender: u8, recipient: u8) -> RouteHeader { + RouteHeader { + sender: QID([sender; QID::SIZE]), + recipient: QID([recipient; QID::SIZE]), } } -fn pairing_token(byte: u8) -> PairingToken { - PairingToken([byte; PairingToken::SIZE]) -} - -fn pairing_id(byte: u8) -> PairingId { - PairingId([byte; PairingId::SIZE]) -} - -fn xx_header(sender: u8, recipient: u8) -> HandshakeHeader { - HandshakeHeader { - sender: qid(sender), - recipient: qid(recipient), - } +fn identity_routes(a: &QlIdentity, b: &QlIdentity) -> (RouteHeader, RouteHeader) { + ( + RouteHeader { + sender: a.qid, + recipient: b.qid, + }, + RouteHeader { + sender: b.qid, + recipient: a.qid, + }, + ) } fn encrypt_record( crypto: &impl QlCrypto, + route: RouteHeader, header: SessionHeader, session_key: &SessionKey, body: &[SessionFrame>], @@ -80,7 +67,7 @@ fn encrypt_record( } decode_session_record( builder - .encrypt(crypto, header.connection_id, session_key) + .encrypt(crypto, route, header.connection_id, session_key) .as_slice(), ) } @@ -115,7 +102,6 @@ fn identity_name_validation() { #[test] fn handshake_record_round_trip_supports_ik_kk_and_xx() { let ik = QlHandshakeRecord::Ik1(Ik1 { - header: handshake_header(1, 2), meta: handshake_meta(1), transport_params: handshake_transport_params(65_536), skem_ciphertext: MlKemCiphertext::new(Box::new([7; MlKemCiphertext::SIZE])), @@ -124,18 +110,19 @@ fn handshake_record_round_trip_supports_ik_kk_and_xx() { }, static_bundle: EncryptedPeerBundle(vec![13; 64].into_boxed_slice()), }); - let ik_encoded = encode_record_vec(RecordType::Handshake, &ik); + let ik_route = route(1, 2); + let ik_encoded = encode_record_vec(RecordHeader::new(ik_route, RecordType::Handshake), &ik); assert_eq!( RecordHeader::decode_bytes(ik_encoded.as_slice()).unwrap(), RecordHeader { version: QL_WIRE_VERSION, + route: ik_route, record_type: RecordType::Handshake, } ); assert_eq!(decode_handshake_record(ik_encoded.as_slice()), ik); let kk = QlHandshakeRecord::Kk1(Kk1 { - header: handshake_header(1, 2), meta: handshake_meta(2), transport_params: handshake_transport_params(131_072), skem_ciphertext: MlKemCiphertext::new(Box::new([11; MlKemCiphertext::SIZE])), @@ -143,30 +130,33 @@ fn handshake_record_round_trip_supports_ik_kk_and_xx() { mlkem_public_key: MlKemPublicKey::new(Box::new([15; MlKemPublicKey::SIZE])), }, }); - let kk_encoded = encode_record_vec(RecordType::Handshake, &kk); + let kk_route = route(1, 2); + let kk_encoded = encode_record_vec(RecordHeader::new(kk_route, RecordType::Handshake), &kk); assert_eq!( RecordHeader::decode_bytes(kk_encoded.as_slice()).unwrap(), RecordHeader { version: QL_WIRE_VERSION, + route: kk_route, record_type: RecordType::Handshake, } ); assert_eq!(decode_handshake_record(kk_encoded.as_slice()), kk); let xx = QlHandshakeRecord::Xx1(Xx1 { - header: xx_header(1, 2), meta: handshake_meta(3), - pairing_id: pairing_id(3), + pairing_id: PairingId([3; PairingId::SIZE]), transport_params: handshake_transport_params(196_608), ephemeral: EphemeralPublicKey { mlkem_public_key: MlKemPublicKey::new(Box::new([17; MlKemPublicKey::SIZE])), }, }); - let xx_encoded = encode_record_vec(RecordType::Handshake, &xx); + let xx_route = route(1, 2); + let xx_encoded = encode_record_vec(RecordHeader::new(xx_route, RecordType::Handshake), &xx); assert_eq!( RecordHeader::decode_bytes(xx_encoded.as_slice()).unwrap(), RecordHeader { version: QL_WIRE_VERSION, + route: xx_route, record_type: RecordType::Handshake, } ); @@ -177,6 +167,7 @@ fn handshake_record_round_trip_supports_ik_kk_and_xx() { fn ik_handshake_rejects_tampered_handshake_meta() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let mut initiator_state = IkHandshake::new_initiator( &crypto, @@ -190,7 +181,9 @@ fn ik_handshake_rejects_tampered_handshake_meta() { let m1 = initiator_state .write_1(&crypto, handshake_meta(77)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); let mut m2 = responder_state .write_2(&crypto, handshake_meta(77)) @@ -198,7 +191,7 @@ fn ik_handshake_rejects_tampered_handshake_meta() { m2.meta.handshake_id = HandshakeId(78); assert_eq!( - initiator_state.read_2(&crypto, &m2), + initiator_state.read_2(&crypto, responder_to_initiator, &m2), Err(WireError::InvalidHandshakeMeta) ); } @@ -207,6 +200,7 @@ fn ik_handshake_rejects_tampered_handshake_meta() { fn kk_handshake_rejects_tampered_handshake_header() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (initiator_to_responder, _) = identity_routes(&initiator, &responder); let mut initiator_state = KkHandshake::new_initiator( &crypto, @@ -224,15 +218,17 @@ fn kk_handshake_rejects_tampered_handshake_header() { let m1 = initiator_state .write_1(&crypto, handshake_meta(88)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); - let mut m2 = responder_state + let m2 = responder_state .write_2(&crypto, handshake_meta(88)) .unwrap(); - m2.header = handshake_header(9, 1); + let tampered_route = route(9, 1); assert_eq!( - initiator_state.read_2(&crypto, &m2), + initiator_state.read_2(&crypto, tampered_route, &m2), Err(WireError::InvalidPayload) ); } @@ -241,6 +237,7 @@ fn kk_handshake_rejects_tampered_handshake_header() { fn ik_handshake_rejects_tampered_transport_params() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let mut initiator_state = IkHandshake::new_initiator( &crypto, @@ -254,7 +251,9 @@ fn ik_handshake_rejects_tampered_transport_params() { let m1 = initiator_state .write_1(&crypto, handshake_meta(89)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); let mut m2 = responder_state .write_2(&crypto, handshake_meta(89)) @@ -262,7 +261,7 @@ fn ik_handshake_rejects_tampered_transport_params() { m2.transport_params.initial_stream_receive_window += 1; assert_eq!( - initiator_state.read_2(&crypto, &m2), + initiator_state.read_2(&crypto, responder_to_initiator, &m2), Err(WireError::DecryptFailed) ); } @@ -271,6 +270,7 @@ fn ik_handshake_rejects_tampered_transport_params() { fn ik_handshake_rejects_tampered_handshake_header() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (mut initiator_to_responder, _) = identity_routes(&initiator, &responder); let mut initiator_state = IkHandshake::new_initiator( &crypto, @@ -281,13 +281,13 @@ fn ik_handshake_rejects_tampered_handshake_header() { let mut responder_state = IkHandshake::new_responder(&crypto, responder, None, TransportParams::default()); - let mut m1 = initiator_state + let m1 = initiator_state .write_1(&crypto, handshake_meta(90)) .unwrap(); - m1.header.sender = qid(9); + initiator_to_responder.sender = QID([9; QID::SIZE]); assert_eq!( - responder_state.read_1(&crypto, &m1), + responder_state.read_1(&crypto, initiator_to_responder, &m1), Err(WireError::DecryptFailed) ); } @@ -297,6 +297,7 @@ fn ik_handshake_rejects_bound_remote_bundle_mismatch() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); let bogus = generate_identity(&crypto, "bogus").unwrap(); + let (initiator_to_responder, _) = identity_routes(&initiator, &responder); let mut initiator_state = IkHandshake::new_initiator( &crypto, @@ -316,7 +317,7 @@ fn ik_handshake_rejects_bound_remote_bundle_mismatch() { .unwrap(); assert_eq!( - responder_state.read_1(&crypto, &m1), + responder_state.read_1(&crypto, initiator_to_responder, &m1), Err(WireError::InvalidPayload) ); } @@ -325,6 +326,7 @@ fn ik_handshake_rejects_bound_remote_bundle_mismatch() { fn ik_handshake_round_trip_derives_matching_transport_and_learns_remote() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let initiator_params = handshake_transport_params(4096); let responder_params = handshake_transport_params(8192); @@ -340,12 +342,16 @@ fn ik_handshake_round_trip_derives_matching_transport_and_learns_remote() { let m1 = initiator_state .write_1(&crypto, handshake_meta(11)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); let m2 = responder_state .write_2(&crypto, handshake_meta(11)) .unwrap(); - initiator_state.read_2(&crypto, &m2).unwrap(); + initiator_state + .read_2(&crypto, responder_to_initiator, &m2) + .unwrap(); let initiator_final = initiator_state.finalize(&crypto).unwrap(); let responder_final = responder_state.finalize(&crypto).unwrap(); @@ -374,6 +380,7 @@ fn ik_handshake_round_trip_derives_matching_transport_and_learns_remote() { fn ik_handshake_round_trip_derives_matching_transport_with_bound_responder() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let initiator_params = handshake_transport_params(16_384); let responder_params = handshake_transport_params(32_768); @@ -393,12 +400,16 @@ fn ik_handshake_round_trip_derives_matching_transport_with_bound_responder() { let m1 = initiator_state .write_1(&crypto, handshake_meta(12)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); let m2 = responder_state .write_2(&crypto, handshake_meta(12)) .unwrap(); - initiator_state.read_2(&crypto, &m2).unwrap(); + initiator_state + .read_2(&crypto, responder_to_initiator, &m2) + .unwrap(); let initiator_final = initiator_state.finalize(&crypto).unwrap(); let responder_final = responder_state.finalize(&crypto).unwrap(); @@ -427,6 +438,7 @@ fn ik_handshake_round_trip_derives_matching_transport_with_bound_responder() { fn kk_handshake_round_trip_derives_matching_transport() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let initiator_params = handshake_transport_params(24_576); let responder_params = handshake_transport_params(49_152); @@ -446,12 +458,16 @@ fn kk_handshake_round_trip_derives_matching_transport() { let m1 = initiator_state .write_1(&crypto, handshake_meta(21)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); let m2 = responder_state .write_2(&crypto, handshake_meta(21)) .unwrap(); - initiator_state.read_2(&crypto, &m2).unwrap(); + initiator_state + .read_2(&crypto, responder_to_initiator, &m2) + .unwrap(); let initiator_final = initiator_state.finalize(&crypto).unwrap(); let responder_final = responder_state.finalize(&crypto).unwrap(); @@ -480,6 +496,7 @@ fn kk_handshake_round_trip_derives_matching_transport() { fn kk_handshake_rejects_tampered_transport_params() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let mut initiator_state = KkHandshake::new_initiator( &crypto, @@ -497,7 +514,9 @@ fn kk_handshake_rejects_tampered_transport_params() { let m1 = initiator_state .write_1(&crypto, handshake_meta(22)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); let mut m2 = responder_state .write_2(&crypto, handshake_meta(22)) @@ -505,7 +524,7 @@ fn kk_handshake_rejects_tampered_transport_params() { m2.transport_params.initial_stream_receive_window += 1; assert_eq!( - initiator_state.read_2(&crypto, &m2), + initiator_state.read_2(&crypto, responder_to_initiator, &m2), Err(WireError::DecryptFailed) ); } @@ -514,7 +533,8 @@ fn kk_handshake_rejects_tampered_transport_params() { fn xx_handshake_rejects_tampered_pairing_id() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); - let token = pairing_token(7); + let token = PairingToken([7; PairingToken::SIZE]); + let (initiator_to_responder, _) = identity_routes(&initiator, &responder); let mut initiator_state = XxHandshake::new_initiator( &crypto, @@ -534,10 +554,10 @@ fn xx_handshake_rejects_tampered_pairing_id() { let mut m1 = initiator_state .write_1(&crypto, handshake_meta(31)) .unwrap(); - m1.pairing_id = pairing_id(8); + m1.pairing_id = PairingId([8; PairingId::SIZE]); assert_eq!( - responder_state.read_1(&crypto, &m1), + responder_state.read_1(&crypto, initiator_to_responder, &m1), Err(WireError::InvalidPairingId) ); } @@ -546,7 +566,7 @@ fn xx_handshake_rejects_tampered_pairing_id() { fn xx_handshake_rejects_tampered_sender_or_recipient() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); - let token = pairing_token(7); + let token = PairingToken([7; PairingToken::SIZE]); let mut initiator_state = XxHandshake::new_initiator( &crypto, @@ -563,14 +583,15 @@ fn xx_handshake_rejects_tampered_sender_or_recipient() { TransportParams::default(), ); - let mut m1 = initiator_state + let m1 = initiator_state .write_1(&crypto, handshake_meta(31)) .unwrap(); - m1.header.sender = responder.qid; + let (mut route, _) = identity_routes(&initiator, &responder); + route.sender = responder.qid; assert_eq!( - responder_state.read_1(&crypto, &m1), - Err(WireError::InvalidHandshakeHeader) + responder_state.read_1(&crypto, route, &m1), + Err(WireError::InvalidRouteHeader) ); let mut initiator_state = XxHandshake::new_initiator( @@ -588,14 +609,15 @@ fn xx_handshake_rejects_tampered_sender_or_recipient() { TransportParams::default(), ); - let mut m1 = initiator_state + let m1 = initiator_state .write_1(&crypto, handshake_meta(31)) .unwrap(); - m1.header.recipient = initiator.qid; + let (mut route, _) = identity_routes(&initiator, &responder); + route.recipient = initiator.qid; assert_eq!( - responder_state.read_1(&crypto, &m1), - Err(WireError::InvalidHandshakeHeader) + responder_state.read_1(&crypto, route, &m1), + Err(WireError::InvalidRouteHeader) ); } @@ -603,7 +625,8 @@ fn xx_handshake_rejects_tampered_sender_or_recipient() { fn xx_handshake_rejects_repeated_transport_param_change() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); - let token = pairing_token(9); + let token = PairingToken([9; PairingToken::SIZE]); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let mut initiator_state = XxHandshake::new_initiator( &crypto, @@ -623,12 +646,16 @@ fn xx_handshake_rejects_repeated_transport_param_change() { let m1 = initiator_state .write_1(&crypto, handshake_meta(32)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); let m2 = responder_state .write_2(&crypto, handshake_meta(32)) .unwrap(); - initiator_state.read_2(&crypto, &m2).unwrap(); + initiator_state + .read_2(&crypto, responder_to_initiator, &m2) + .unwrap(); let mut m3 = initiator_state .write_3(&crypto, handshake_meta(32)) @@ -636,7 +663,7 @@ fn xx_handshake_rejects_repeated_transport_param_change() { m3.transport_params.initial_stream_receive_window += 1; assert_eq!( - responder_state.read_3(&crypto, &m3), + responder_state.read_3(&crypto, initiator_to_responder, &m3), Err(WireError::InvalidTransportParams) ); } @@ -645,7 +672,8 @@ fn xx_handshake_rejects_repeated_transport_param_change() { fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); - let token = pairing_token(10); + let token = PairingToken([10; PairingToken::SIZE]); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let initiator_params = handshake_transport_params(28_672); let responder_params = handshake_transport_params(57_344); @@ -674,25 +702,33 @@ fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { let m1 = initiator_state .write_1(&crypto, handshake_meta(33)) .unwrap(); - responder_state.read_1(&crypto, &m1).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); let m2 = responder_state .write_2(&crypto, handshake_meta(33)) .unwrap(); - initiator_state.read_2(&crypto, &m2).unwrap(); + initiator_state + .read_2(&crypto, responder_to_initiator, &m2) + .unwrap(); assert_eq!(initiator_state.remote_bundle(), Some(&responder.bundle())); assert!(responder_state.remote_bundle().is_none()); let m3 = initiator_state .write_3(&crypto, handshake_meta(33)) .unwrap(); - responder_state.read_3(&crypto, &m3).unwrap(); + responder_state + .read_3(&crypto, initiator_to_responder, &m3) + .unwrap(); assert_eq!(responder_state.remote_bundle(), Some(&initiator.bundle())); let m4 = responder_state .write_4(&crypto, handshake_meta(33)) .unwrap(); - initiator_state.read_4(&crypto, &m4).unwrap(); + initiator_state + .read_4(&crypto, responder_to_initiator, &m4) + .unwrap(); let initiator_final = initiator_state.finalize(&crypto).unwrap(); let responder_final = responder_state.finalize(&crypto).unwrap(); @@ -722,8 +758,9 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { let crypto = SoftwareCrypto; let header = SessionHeader { connection_id: ConnectionId([0x44; ConnectionId::SIZE]), - seq: record_seq(11), + seq: RecordSeq(varint(11)), }; + let session_route = route(1, 2); let body = vec![ SessionFrame::Ping, SessionFrame::Unpair, @@ -731,18 +768,18 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { RecordAck::from_ranges([record_ack_range(20, 23), record_ack_range(12, 13)]).unwrap(), ), SessionFrame::StreamWindow(StreamWindow { - stream_id: stream_id(9), + stream_id: StreamId(varint(9)), maximum_offset: varint(65_536), }), SessionFrame::StreamData(StreamData { - stream_id: stream_id(9), + stream_id: StreamId(varint(9)), offset: varint(1024), header: None, bytes: b"hello".to_vec(), fin: true, }), SessionFrame::StreamReset(StreamReset { - stream_id: stream_id(9), + stream_id: StreamId(varint(9)), target: ResetTarget::Both, code: ResetCode::CANCELLED, }), @@ -751,13 +788,15 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { }), ]; let session_key = SessionKey([7; SessionKey::SIZE]); - let record = encrypt_record(&crypto, header, &session_key, &body); + let record = encrypt_record(&crypto, session_route, header, &session_key, &body); - let bytes = encode_record_vec(RecordType::Session, &record); + let record_header = RecordHeader::new(session_route, RecordType::Session); + let bytes = encode_record_vec(record_header, &record); assert_eq!( RecordHeader::decode_bytes(bytes.as_slice()).unwrap(), RecordHeader { version: QL_WIRE_VERSION, + route: session_route, record_type: RecordType::Session, } ); @@ -765,8 +804,14 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { assert_eq!(decoded.header, header); let encrypted = decoded.payload; - let decrypted = - encrypted::decrypt_record(&crypto, &header, encrypted.clone(), &session_key).unwrap(); + let decrypted = encrypted::decrypt_record( + &crypto, + &record_header, + &header, + encrypted.clone(), + &session_key, + ) + .unwrap(); assert_eq!(decode_session_frames(&decrypted).unwrap(), body); let wrong_header = SessionHeader { @@ -774,16 +819,40 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { seq: header.seq, }; assert_eq!( - encrypted::decrypt_record(&crypto, &wrong_header, encrypted.clone(), &session_key), + encrypted::decrypt_record( + &crypto, + &record_header, + &wrong_header, + encrypted.clone(), + &session_key, + ), + Err(WireError::DecryptFailed) + ); + + let wrong_record_header = RecordHeader::new(route(2, 1), RecordType::Session); + assert_eq!( + encrypted::decrypt_record( + &crypto, + &wrong_record_header, + &header, + encrypted.clone(), + &session_key, + ), Err(WireError::DecryptFailed) ); let wrong_seq_header = SessionHeader { connection_id: header.connection_id, - seq: record_seq(header.seq.0.into_inner() + 1), + seq: RecordSeq(varint(header.seq.0.into_inner() + 1)), }; assert_eq!( - encrypted::decrypt_record(&crypto, &wrong_seq_header, encrypted, &session_key), + encrypted::decrypt_record( + &crypto, + &record_header, + &wrong_seq_header, + encrypted, + &session_key, + ), Err(WireError::DecryptFailed) ); } @@ -796,6 +865,7 @@ fn protocol_record_size_breakdown() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); + let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); let mut ik_initiator = IkHandshake::new_initiator( &crypto, @@ -807,10 +877,14 @@ fn protocol_record_size_breakdown() { IkHandshake::new_responder(&crypto, responder.clone(), None, TransportParams::default()); let ik1 = ik_initiator.write_1(&crypto, handshake_meta(101)).unwrap(); - ik_responder.read_1(&crypto, &ik1).unwrap(); + ik_responder + .read_1(&crypto, initiator_to_responder, &ik1) + .unwrap(); let ik2 = ik_responder.write_2(&crypto, handshake_meta(101)).unwrap(); - ik_initiator.read_2(&crypto, &ik2).unwrap(); + ik_initiator + .read_2(&crypto, responder_to_initiator, &ik2) + .unwrap(); let ik1 = QlHandshakeRecord::Ik1(ik1); let ik2 = QlHandshakeRecord::Ik2(ik2); @@ -829,15 +903,19 @@ fn protocol_record_size_breakdown() { ); let kk1 = kk_initiator.write_1(&crypto, handshake_meta(201)).unwrap(); - kk_responder.read_1(&crypto, &kk1).unwrap(); + kk_responder + .read_1(&crypto, initiator_to_responder, &kk1) + .unwrap(); let kk2 = kk_responder.write_2(&crypto, handshake_meta(201)).unwrap(); - kk_initiator.read_2(&crypto, &kk2).unwrap(); + kk_initiator + .read_2(&crypto, responder_to_initiator, &kk2) + .unwrap(); let kk1 = QlHandshakeRecord::Kk1(kk1); let kk2 = QlHandshakeRecord::Kk2(kk2); - let token = pairing_token(0x42); + let token = PairingToken([0x42; PairingToken::SIZE]); let mut xx_initiator = XxHandshake::new_initiator( &crypto, initiator.clone(), @@ -854,16 +932,24 @@ fn protocol_record_size_breakdown() { ); let xx1 = xx_initiator.write_1(&crypto, handshake_meta(301)).unwrap(); - xx_responder.read_1(&crypto, &xx1).unwrap(); + xx_responder + .read_1(&crypto, initiator_to_responder, &xx1) + .unwrap(); let xx2 = xx_responder.write_2(&crypto, handshake_meta(301)).unwrap(); - xx_initiator.read_2(&crypto, &xx2).unwrap(); + xx_initiator + .read_2(&crypto, responder_to_initiator, &xx2) + .unwrap(); let xx3 = xx_initiator.write_3(&crypto, handshake_meta(301)).unwrap(); - xx_responder.read_3(&crypto, &xx3).unwrap(); + xx_responder + .read_3(&crypto, initiator_to_responder, &xx3) + .unwrap(); let xx4 = xx_responder.write_4(&crypto, handshake_meta(301)).unwrap(); - xx_initiator.read_4(&crypto, &xx4).unwrap(); + xx_initiator + .read_4(&crypto, responder_to_initiator, &xx4) + .unwrap(); let xx1 = QlHandshakeRecord::Xx1(xx1); let xx2 = QlHandshakeRecord::Xx2(xx2); @@ -871,20 +957,26 @@ fn protocol_record_size_breakdown() { let xx4 = QlHandshakeRecord::Xx4(xx4); let session = ik_initiator.finalize(&crypto).unwrap(); + let session_route = RouteHeader { + sender: initiator.qid, + recipient: responder.qid, + }; let session_ping = encrypt_record( &crypto, + session_route, SessionHeader { connection_id: session.tx_connection_id, - seq: record_seq(1), + seq: RecordSeq(varint(1)), }, &session.tx_key, &[SessionFrame::Ping], ); let session_ack = encrypt_record( &crypto, + session_route, SessionHeader { connection_id: session.tx_connection_id, - seq: record_seq(2), + seq: RecordSeq(varint(2)), }, &session.tx_key, &[SessionFrame::Ack( @@ -893,22 +985,24 @@ fn protocol_record_size_breakdown() { ); let session_unpair = encrypt_record( &crypto, + session_route, SessionHeader { connection_id: session.tx_connection_id, - seq: record_seq(3), + seq: RecordSeq(varint(3)), }, &session.tx_key, &[SessionFrame::Unpair], ); let session_stream_empty = encrypt_record( &crypto, + session_route, SessionHeader { connection_id: session.tx_connection_id, - seq: record_seq(4), + seq: RecordSeq(varint(4)), }, &session.tx_key, &[SessionFrame::StreamData(StreamData { - stream_id: stream_id(1), + stream_id: StreamId(varint(1)), offset: varint(0), header: None, fin: false, @@ -917,9 +1011,10 @@ fn protocol_record_size_breakdown() { ); let session_close = encrypt_record( &crypto, + session_route, SessionHeader { connection_id: session.tx_connection_id, - seq: record_seq(5), + seq: RecordSeq(varint(5)), }, &session.tx_key, &[SessionFrame::Close(SessionClose { From 4e71eeada5abd5c8c7545fef5ba2544a7fa0a3b6 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 25 Jun 2026 11:32:26 -0400 Subject: [PATCH 31/59] ql-wire: remove connectionid --- ql-fsm/src/error.rs | 2 -- ql-fsm/src/fsm.rs | 10 +----- ql-fsm/src/handshake/ik.rs | 8 ++--- ql-fsm/src/handshake/kk.rs | 8 ++--- ql-fsm/src/handshake/mod.rs | 11 ++---- ql-fsm/src/handshake/xx.rs | 8 ++--- ql-fsm/src/state.rs | 7 ++-- ql-fsm/src/tests/mod.rs | 10 ++---- ql-fsm/src/tests/proptest.rs | 1 - ql-wire/src/encrypted/builder.rs | 22 ++++-------- ql-wire/src/handshake/mod.rs | 29 ++------------- ql-wire/src/header.rs | 12 ++----- ql-wire/src/tests.rs | 62 ++------------------------------ 13 files changed, 31 insertions(+), 159 deletions(-) diff --git a/ql-fsm/src/error.rs b/ql-fsm/src/error.rs index c6382371..94448778 100644 --- a/ql-fsm/src/error.rs +++ b/ql-fsm/src/error.rs @@ -11,7 +11,6 @@ pub enum ReceiveError { InvalidRecordVersion, InvalidHandshakeRecord(WireError), InvalidSessionRecord(WireError), - InvalidSessionConnectionId, InvalidSessionPayload(WireError), InvalidIkHandshake(WireError), InvalidKkHandshake(WireError), @@ -36,7 +35,6 @@ impl Display for ReceiveError { write!(f, "invalid handshake record: {error}") } Self::InvalidSessionRecord(error) => write!(f, "invalid session record: {error}"), - Self::InvalidSessionConnectionId => f.write_str("invalid session connection id"), Self::InvalidSessionPayload(error) => write!(f, "invalid session payload: {error}"), Self::InvalidIkHandshake(error) => write!(f, "invalid ik handshake: {error}"), Self::InvalidKkHandshake(error) => write!(f, "invalid kk handshake: {error}"), diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index d52be08e..2e2f42a9 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -131,9 +131,6 @@ pub fn receive( let (decrypt_len, seq) = { let record = wire::QlSessionRecord::decode(&mut reader) .map_err(ReceiveError::InvalidSessionRecord)?; - if record.header.connection_id != conn.transport.rx_connection_id { - return Err(ReceiveError::InvalidSessionConnectionId); - } let payload = wire::decrypt_record( crypto, &header, @@ -212,12 +209,7 @@ pub fn take_next_write(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Option | LinkState::KkInitiator(_) | LinkState::XxInitiator(_) | LinkState::XxResponder(_) => false, - LinkState::Connected(_) => super::is_connected_replay( - fsm, - message.meta.handshake_id, - route.sender, - ), + LinkState::Connected(_) => { + super::is_connected_replay(fsm, message.meta.handshake_id, route.sender) + } LinkState::IkInitiator(state) => { if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(route.sender) { return false; diff --git a/ql-fsm/src/handshake/kk.rs b/ql-fsm/src/handshake/kk.rs index 0ece2434..243766ae 100644 --- a/ql-fsm/src/handshake/kk.rs +++ b/ql-fsm/src/handshake/kk.rs @@ -119,11 +119,9 @@ pub fn handle_kk2( pub fn should_ignore_inbound(fsm: &QlFsm, route: RouteHeader, message: &Kk1) -> bool { match &fsm.state.link { LinkState::Idle | LinkState::XxInitiator(_) | LinkState::XxResponder(_) => false, - LinkState::Connected(_) => super::is_connected_replay( - fsm, - message.meta.handshake_id, - route.sender, - ), + LinkState::Connected(_) => { + super::is_connected_replay(fsm, message.meta.handshake_id, route.sender) + } LinkState::IkInitiator(_) => true, LinkState::KkInitiator(state) => { if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(route.sender) { diff --git a/ql-fsm/src/handshake/mod.rs b/ql-fsm/src/handshake/mod.rs index 70ac7dc7..7add5539 100644 --- a/ql-fsm/src/handshake/mod.rs +++ b/ql-fsm/src/handshake/mod.rs @@ -148,8 +148,6 @@ pub fn establish_session( remote_qid: finalized.remote_bundle.qid, tx_key: finalized.tx_key, rx_key: finalized.rx_key, - tx_connection_id: finalized.tx_connection_id, - rx_connection_id: finalized.rx_connection_id, remote_transport_params: finalized.remote_transport_params, }; finish_handshake(fsm, handshake_id, transport, finalized.remote_bundle) @@ -165,14 +163,11 @@ fn local_start_wins(local: &EphemeralPublicKey, inbound: &EphemeralPublicKey) -> local.mlkem_public_key.as_bytes() <= inbound.mlkem_public_key.as_bytes() } -fn is_connected_replay( - fsm: &QlFsm, - handshake_id: HandshakeId, - sender: QID, -) -> bool { +fn is_connected_replay(fsm: &QlFsm, handshake_id: HandshakeId, sender: QID) -> bool { let LinkState::Connected(connected) = &fsm.state.link else { return false; }; - connected.handshake_id == handshake_id && fsm.state.peer.as_ref().map(|peer| peer.qid) == Some(sender) + connected.handshake_id == handshake_id + && fsm.state.peer.as_ref().map(|peer| peer.qid) == Some(sender) } diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs index 4f9952dc..7ef17f99 100644 --- a/ql-fsm/src/handshake/xx.rs +++ b/ql-fsm/src/handshake/xx.rs @@ -224,11 +224,9 @@ pub fn should_ignore_inbound( ) -> bool { match &fsm.state.link { LinkState::Idle => false, - LinkState::Connected(_) => super::is_connected_replay( - fsm, - message.meta.handshake_id, - route.sender, - ), + LinkState::Connected(_) => { + super::is_connected_replay(fsm, message.meta.handshake_id, route.sender) + } LinkState::IkInitiator(_) | LinkState::KkInitiator(_) | LinkState::XxResponder(_) => true, LinkState::XxInitiator(state) => { if state.handshake.pairing_id(crypto) != message.pairing_id { diff --git a/ql-fsm/src/state.rs b/ql-fsm/src/state.rs index 6e406ae7..093bd4b0 100644 --- a/ql-fsm/src/state.rs +++ b/ql-fsm/src/state.rs @@ -2,9 +2,8 @@ use std::time::Instant; use ql_common::QID; use ql_wire::{ - ConnectionId, EphemeralPublicKey, HandshakeId, HandshakeMeta, IkHandshake, KkHandshake, - PairingToken, PeerBundle, QlHandshakeRecord, RouteHeader, SessionKey, TransportParams, - XxHandshake, + EphemeralPublicKey, HandshakeId, HandshakeMeta, IkHandshake, KkHandshake, PairingToken, + PeerBundle, QlHandshakeRecord, RouteHeader, SessionKey, TransportParams, XxHandshake, }; use crate::{session::SessionFsm, NoSessionError, PeerStatus}; @@ -23,8 +22,6 @@ pub struct SessionTransport { pub remote_qid: QID, pub tx_key: SessionKey, pub rx_key: SessionKey, - pub tx_connection_id: ConnectionId, - pub rx_connection_id: ConnectionId, pub remote_transport_params: TransportParams, } diff --git a/ql-fsm/src/tests/mod.rs b/ql-fsm/src/tests/mod.rs index 1bd84652..d7003faa 100644 --- a/ql-fsm/src/tests/mod.rs +++ b/ql-fsm/src/tests/mod.rs @@ -6,8 +6,8 @@ use std::time::{Duration, Instant}; use ql_common::QID; use ql_wire::{ - self, generate_identity, test_identities, ConnectionId, HandshakeId, PairingToken, QlCrypto, - SessionKey, SoftwareCrypto, TransportParams, + self, generate_identity, test_identities, HandshakeId, PairingToken, QlCrypto, SessionKey, + SoftwareCrypto, TransportParams, }; use crate::{ @@ -99,8 +99,6 @@ impl Harness { let mut harness = Self::paired_known(config); let a_to_b_key = SessionKey([7; SessionKey::SIZE]); let b_to_a_key = SessionKey([9; SessionKey::SIZE]); - let a_to_b_conn = ConnectionId([0xA1; ConnectionId::SIZE]); - let b_to_a_conn = ConnectionId([0xB2; ConnectionId::SIZE]); harness.a.fsm.state.link = LinkState::Connected(ConnectedState { handshake_id: HandshakeId(0), @@ -108,8 +106,6 @@ impl Harness { remote_qid: harness.b.fsm.identity.qid, tx_key: a_to_b_key.clone(), rx_key: b_to_a_key.clone(), - tx_connection_id: a_to_b_conn, - rx_connection_id: b_to_a_conn, remote_transport_params: TransportParams { initial_stream_receive_window: harness .b @@ -126,8 +122,6 @@ impl Harness { remote_qid: harness.a.fsm.identity.qid, tx_key: b_to_a_key, rx_key: a_to_b_key, - tx_connection_id: b_to_a_conn, - rx_connection_id: a_to_b_conn, remote_transport_params: TransportParams { initial_stream_receive_window: harness .a diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs index 655204c7..ac1fd440 100644 --- a/ql-fsm/src/tests/proptest.rs +++ b/ql-fsm/src/tests/proptest.rs @@ -550,7 +550,6 @@ impl Runner { | ReceiveError::InvalidRemoteBundle | ReceiveError::InvalidSessionPayload(WireError::InvalidPayload) | ReceiveError::InvalidSessionPayload(WireError::DecryptFailed) - | ReceiveError::InvalidSessionConnectionId | ReceiveError::InvalidIkHandshake(WireError::InvalidPayload) | ReceiveError::InvalidIkHandshake(WireError::InvalidState) | ReceiveError::InvalidKkHandshake(WireError::InvalidPayload) diff --git a/ql-wire/src/encrypted/builder.rs b/ql-wire/src/encrypted/builder.rs index 95b977c2..f25f6859 100644 --- a/ql-wire/src/encrypted/builder.rs +++ b/ql-wire/src/encrypted/builder.rs @@ -2,8 +2,8 @@ use bytes::BufMut; use super::{RecordAck, SessionClose, SessionFrame, StreamData, StreamReset, StreamWindow}; use crate::{ - BufView, ConnectionId, Nonce, QlCrypto, RecordHeader, RecordSeq, RecordType, RouteHeader, - SessionHeader, SessionKey, WireEncode, + BufView, Nonce, QlCrypto, RecordHeader, RecordSeq, RecordType, RouteHeader, SessionHeader, + SessionKey, WireEncode, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -15,16 +15,12 @@ pub struct SessionRecordBuilder { } impl SessionRecordBuilder { - pub const MIN_CAPACITY: usize = RecordHeader::WIRE_SIZE - + ConnectionId::SIZE - + RecordSeq::MAX_ENCODED_LEN - + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; + pub const MIN_CAPACITY: usize = + RecordHeader::WIRE_SIZE + RecordSeq::MAX_ENCODED_LEN + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; pub fn new(seq: RecordSeq, max_capacity: usize) -> Self { - let prefix_len = RecordHeader::WIRE_SIZE - + ConnectionId::SIZE - + seq.encoded_len() - + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; + let prefix_len = + RecordHeader::WIRE_SIZE + seq.encoded_len() + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; assert!(max_capacity >= prefix_len); Self { seq, @@ -107,15 +103,11 @@ impl SessionRecordBuilder { mut self, crypto: &impl QlCrypto, route: RouteHeader, - connection_id: ConnectionId, session_key: &SessionKey, ) -> Vec { self.ensure_prefix_capacity(0); let record_header = RecordHeader::new(route, RecordType::Session); - let header = SessionHeader { - connection_id, - seq: self.seq, - }; + let header = SessionHeader { seq: self.seq }; let aad = header.aad(route); let nonce = Nonce::from_counter(self.seq.0.into_inner()); let auth = crypto.aes256_gcm_encrypt( diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index 34897b27..0a581132 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -1,7 +1,7 @@ use crate::{ - codec, derive_qid, ByteSlice, ConnectionId, HandshakeKind, MlKemCiphertext, MlKemKeyPair, - MlKemPublicKey, Nonce, PeerBundle, QlCrypto, RouteHeader, SessionKey, WireDecode, WireEncode, - WireError, ENCRYPTED_MESSAGE_AUTH_SIZE, + codec, derive_qid, ByteSlice, HandshakeKind, MlKemCiphertext, MlKemKeyPair, MlKemPublicKey, + Nonce, PeerBundle, QlCrypto, RouteHeader, SessionKey, WireDecode, WireEncode, WireError, + ENCRYPTED_MESSAGE_AUTH_SIZE, }; mod ik; @@ -22,7 +22,6 @@ const SHA256_BLOCK_LEN: usize = 64; const PROTOCOL_IK: &[u8] = b"ql-wire:pq-ik:v1"; const PROTOCOL_KK: &[u8] = b"ql-wire:pq-kk:v1"; const PROTOCOL_XX: &[u8] = b"ql-wire:pq-xx:v1"; -const CONNECTION_ID_DOMAIN: &[u8] = b"ql-wire:conn-id:v1"; const HANDSHAKE_PREAMBLE_DOMAIN: &[u8] = b"ql-wire:handshake-preamble:v1"; #[derive(Debug, Clone, PartialEq, Eq)] @@ -113,8 +112,6 @@ impl codec::WireDecode for EncryptedPeerBundle { pub struct FinalizedHandshake { pub tx_key: SessionKey, pub rx_key: SessionKey, - pub tx_connection_id: ConnectionId, - pub rx_connection_id: ConnectionId, pub handshake_hash: [u8; 32], pub remote_bundle: PeerBundle, /// Transport parameters advertised by the remote peer @@ -476,35 +473,15 @@ fn finalize_handshake( ) -> FinalizedHandshake { let handshake_hash = symmetric.handshake_hash; let (tx_key, rx_key) = symmetric.split_for_role(crypto, role); - let (initiator_rx, responder_rx) = derive_connection_ids(crypto, &handshake_hash); - let (tx_connection_id, rx_connection_id) = match role { - Role::Initiator => (responder_rx, initiator_rx), - Role::Responder => (initiator_rx, responder_rx), - }; FinalizedHandshake { tx_key, rx_key, - tx_connection_id, - rx_connection_id, handshake_hash, remote_bundle, remote_transport_params, } } -fn derive_connection_ids( - crypto: &impl QlCrypto, - handshake_hash: &[u8; 32], -) -> (ConnectionId, ConnectionId) { - let initiator = crypto.sha256(&[CONNECTION_ID_DOMAIN, handshake_hash, b"initiator-rx"]); - let responder = crypto.sha256(&[CONNECTION_ID_DOMAIN, handshake_hash, b"responder-rx"]); - let mut initiator_rx = [0u8; ConnectionId::SIZE]; - let mut responder_rx = [0u8; ConnectionId::SIZE]; - initiator_rx.copy_from_slice(&initiator[..ConnectionId::SIZE]); - responder_rx.copy_from_slice(&responder[..ConnectionId::SIZE]); - (ConnectionId(initiator_rx), ConnectionId(responder_rx)) -} - fn hkdf2( crypto: &impl QlCrypto, chaining_key: &[u8; 32], diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index 0d232903..79897c57 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -35,18 +35,14 @@ impl codec::WireDecode for RouteHeader { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct SessionHeader { - pub connection_id: ConnectionId, pub seq: RecordSeq, } ql_common::varint_wrapper!(RecordSeq); - varint_wrapper_codec!(RecordSeq); -array_wrapper!(ConnectionId, 8); - impl SessionHeader { - pub const MAX_ENCODED_LEN: usize = ConnectionId::SIZE + RecordSeq::MAX_ENCODED_LEN; + pub const MAX_ENCODED_LEN: usize = RecordSeq::MAX_ENCODED_LEN; const AAD_DOMAIN: &[u8] = b"ql-wire:session-aad:v1"; const AAD_RECORD_KIND_SESSION: u8 = 1; @@ -55,14 +51,12 @@ impl SessionHeader { + size_of::() + size_of::() + RouteHeader::WIRE_SIZE - + ConnectionId::SIZE + self.seq.encoded_len(); let mut aad = Vec::with_capacity(aad_len); aad.put_slice(Self::AAD_DOMAIN); aad.put_u8(QL_WIRE_VERSION); aad.put_u8(Self::AAD_RECORD_KIND_SESSION); route.encode(&mut aad); - self.connection_id.encode(&mut aad); self.seq.encode(&mut aad); debug_assert_eq!(aad.len(), aad_len); aad @@ -71,11 +65,10 @@ impl SessionHeader { impl WireEncode for SessionHeader { fn encoded_len(&self) -> usize { - ConnectionId::SIZE + self.seq.encoded_len() + self.seq.encoded_len() } fn encode(&self, out: &mut W) { - self.connection_id.encode(out); self.seq.encode(out); } } @@ -83,7 +76,6 @@ impl WireEncode for SessionHeader { impl codec::WireDecode for SessionHeader { fn decode(reader: &mut codec::Reader) -> Result { Ok(Self { - connection_id: reader.decode()?, seq: reader.decode()?, }) } diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index 370f6097..a73cd743 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -65,11 +65,7 @@ fn encrypt_record( let pushed = builder.push_frame(frame); debug_assert!(pushed); } - decode_session_record( - builder - .encrypt(crypto, route, header.connection_id, session_key) - .as_slice(), - ) + decode_session_record(builder.encrypt(crypto, route, session_key).as_slice()) } #[test] @@ -362,14 +358,6 @@ fn ik_handshake_round_trip_derives_matching_transport_and_learns_remote() { ); assert_eq!(initiator_final.tx_key, responder_final.rx_key); assert_eq!(initiator_final.rx_key, responder_final.tx_key); - assert_eq!( - initiator_final.tx_connection_id, - responder_final.rx_connection_id - ); - assert_eq!( - initiator_final.rx_connection_id, - responder_final.tx_connection_id - ); assert_eq!(initiator_final.remote_bundle, responder.bundle()); assert_eq!(responder_final.remote_bundle, initiator.bundle()); assert_eq!(initiator_final.remote_transport_params, responder_params); @@ -420,14 +408,6 @@ fn ik_handshake_round_trip_derives_matching_transport_with_bound_responder() { ); assert_eq!(initiator_final.tx_key, responder_final.rx_key); assert_eq!(initiator_final.rx_key, responder_final.tx_key); - assert_eq!( - initiator_final.tx_connection_id, - responder_final.rx_connection_id - ); - assert_eq!( - initiator_final.rx_connection_id, - responder_final.tx_connection_id - ); assert_eq!(initiator_final.remote_bundle, responder.bundle()); assert_eq!(responder_final.remote_bundle, initiator.bundle()); assert_eq!(initiator_final.remote_transport_params, responder_params); @@ -478,14 +458,6 @@ fn kk_handshake_round_trip_derives_matching_transport() { ); assert_eq!(initiator_final.tx_key, responder_final.rx_key); assert_eq!(initiator_final.rx_key, responder_final.tx_key); - assert_eq!( - initiator_final.tx_connection_id, - responder_final.rx_connection_id - ); - assert_eq!( - initiator_final.rx_connection_id, - responder_final.tx_connection_id - ); assert_eq!(initiator_final.remote_bundle, responder.bundle()); assert_eq!(responder_final.remote_bundle, initiator.bundle()); assert_eq!(initiator_final.remote_transport_params, responder_params); @@ -739,14 +711,6 @@ fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { ); assert_eq!(initiator_final.tx_key, responder_final.rx_key); assert_eq!(initiator_final.rx_key, responder_final.tx_key); - assert_eq!( - initiator_final.tx_connection_id, - responder_final.rx_connection_id - ); - assert_eq!( - initiator_final.rx_connection_id, - responder_final.tx_connection_id - ); assert_eq!(initiator_final.remote_bundle, responder.bundle()); assert_eq!(responder_final.remote_bundle, initiator.bundle()); assert_eq!(initiator_final.remote_transport_params, responder_params); @@ -754,10 +718,9 @@ fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { } #[test] -fn encrypted_session_record_round_trip_uses_connection_id_header() { +fn encrypted_session_record_round_trip_authenticates_header() { let crypto = SoftwareCrypto; let header = SessionHeader { - connection_id: ConnectionId([0x44; ConnectionId::SIZE]), seq: RecordSeq(varint(11)), }; let session_route = route(1, 2); @@ -814,21 +777,6 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { .unwrap(); assert_eq!(decode_session_frames(&decrypted).unwrap(), body); - let wrong_header = SessionHeader { - connection_id: ConnectionId([0x99; ConnectionId::SIZE]), - seq: header.seq, - }; - assert_eq!( - encrypted::decrypt_record( - &crypto, - &record_header, - &wrong_header, - encrypted.clone(), - &session_key, - ), - Err(WireError::DecryptFailed) - ); - let wrong_record_header = RecordHeader::new(route(2, 1), RecordType::Session); assert_eq!( encrypted::decrypt_record( @@ -842,7 +790,6 @@ fn encrypted_session_record_round_trip_uses_connection_id_header() { ); let wrong_seq_header = SessionHeader { - connection_id: header.connection_id, seq: RecordSeq(varint(header.seq.0.into_inner() + 1)), }; assert_eq!( @@ -965,7 +912,6 @@ fn protocol_record_size_breakdown() { &crypto, session_route, SessionHeader { - connection_id: session.tx_connection_id, seq: RecordSeq(varint(1)), }, &session.tx_key, @@ -975,7 +921,6 @@ fn protocol_record_size_breakdown() { &crypto, session_route, SessionHeader { - connection_id: session.tx_connection_id, seq: RecordSeq(varint(2)), }, &session.tx_key, @@ -987,7 +932,6 @@ fn protocol_record_size_breakdown() { &crypto, session_route, SessionHeader { - connection_id: session.tx_connection_id, seq: RecordSeq(varint(3)), }, &session.tx_key, @@ -997,7 +941,6 @@ fn protocol_record_size_breakdown() { &crypto, session_route, SessionHeader { - connection_id: session.tx_connection_id, seq: RecordSeq(varint(4)), }, &session.tx_key, @@ -1013,7 +956,6 @@ fn protocol_record_size_breakdown() { &crypto, session_route, SessionHeader { - connection_id: session.tx_connection_id, seq: RecordSeq(varint(5)), }, &session.tx_key, From 0877b2bc8a856a373afd5ae86b5151bc35c8e59e Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 29 Jun 2026 10:55:34 -0400 Subject: [PATCH 32/59] ql-wire: LenBytes helper --- ql-wire/src/codec.rs | 43 ++++++++++++++++++++++++++-- ql-wire/src/encrypted/stream_data.rs | 24 ++++++---------- ql-wire/src/identity.rs | 15 ++++------ 3 files changed, 54 insertions(+), 28 deletions(-) diff --git a/ql-wire/src/codec.rs b/ql-wire/src/codec.rs index 0245ef6d..a8a164ec 100644 --- a/ql-wire/src/codec.rs +++ b/ql-wire/src/codec.rs @@ -1,6 +1,7 @@ -use bytes::BufMut; +use bytes::{Buf, BufMut}; +use ql_common::VarInt; -use crate::{ByteSlice, WireError}; +use crate::{BufView, ByteSlice, WireError}; pub trait WireEncode { fn encoded_len(&self) -> usize; @@ -243,3 +244,41 @@ impl Reader { T::decode(self) } } + +/// bytes encoded with a VarInt len prefix +pub struct LenBytes(pub B); + +impl WireEncode for LenBytes +where + B: BufView, +{ + fn encoded_len(&self) -> usize { + let len = self.0.buf().remaining(); + let len_prefix = VarInt::try_from(len).unwrap(); + len_prefix.encoded_len() + len + } + + fn encode(&self, out: &mut W) { + let mut bytes = self.0.buf(); + let len = bytes.remaining(); + let len_prefix = VarInt::try_from(len).unwrap(); + len_prefix.encode(out); + while bytes.has_remaining() { + let chunk = bytes.chunk(); + out.put_slice(chunk); + bytes.advance(chunk.len()); + } + } +} + +impl WireDecode for LenBytes +where + B: ByteSlice, +{ + fn decode(reader: &mut Reader) -> Result { + let len = reader.decode::()?.into_inner(); + let len = usize::try_from(len).map_err(|_| WireError::InvalidPayload)?; + let bytes = reader.take_bytes(len)?; + Ok(Self(bytes)) + } +} diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index 3cc9eb50..d160d47e 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -1,7 +1,9 @@ -use bytes::Buf; use ql_common::{RouteId, ServiceId, StreamId, VarInt}; -use crate::{codec, BufView, ByteSlice, WireDecode, WireEncode, WireError}; +use crate::{ + codec::{self, LenBytes}, + BufView, ByteSlice, WireDecode, WireEncode, WireError, +}; /// carries bytes for a stream and may finish that sending direction. #[derive(Debug, Clone, PartialEq, Eq)] @@ -33,15 +35,14 @@ impl WireDecode for StreamData { } else { None }; - let bytes_len = usize::try_from(reader.decode::()?.into_inner()) - .map_err(|_| WireError::InvalidPayload)?; + let bytes = reader.decode::>()?.0; Ok(Self { stream_id, offset, header, fin, - bytes: reader.take_bytes(bytes_len)?, + bytes, }) } } @@ -63,14 +64,11 @@ impl StreamData { impl WireEncode for StreamData { fn encoded_len(&self) -> usize { - let bytes = self.bytes.buf(); - let bytes_len = bytes.remaining(); self.stream_id.encoded_len() + self.offset.encoded_len() + size_of::() + self.header.as_ref().map_or(0, WireEncode::encoded_len) - + VarInt::try_from(bytes_len).unwrap().encoded_len() - + bytes_len + + LenBytes(&self.bytes).encoded_len() } fn encode(&self, out: &mut W) { @@ -92,13 +90,7 @@ impl WireEncode for StreamData { if let Some(header) = &self.header { header.encode(out); } - let mut bytes = self.bytes.buf(); - VarInt::try_from(bytes.remaining()).unwrap().encode(out); - while bytes.has_remaining() { - let chunk = bytes.chunk(); - out.put_slice(chunk); - bytes.advance(chunk.len()); - } + LenBytes(&self.bytes).encode(out); } } diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index 57602178..65a2901e 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -1,5 +1,6 @@ use std::ops::Deref; +use codec::LenBytes; use ql_common::{VarInt, QID}; use crate::{ @@ -149,26 +150,20 @@ impl Deref for QlName { impl WireEncode for QlName { fn encoded_len(&self) -> usize { - let len = VarInt::try_from(self.0.len()).unwrap(); - len.encoded_len() + self.0.len() + LenBytes(self.0.as_bytes()).encoded_len() } fn encode(&self, out: &mut W) { - VarInt::try_from(self.0.len()) - .expect("identity name length fits in varint") - .encode(out); - self.0.as_bytes().encode(out); + LenBytes(self.0.as_bytes()).encode(out) } } impl codec::WireDecode for QlName { fn decode(reader: &mut codec::Reader) -> Result { - let len = usize::try_from(reader.decode::()?.into_inner()) - .map_err(|_| WireError::InvalidPayload)?; - if len == 0 || len > Self::MAX_LEN { + let bytes = reader.decode::>()?.0; + if bytes.is_empty() || bytes.len() > Self::MAX_LEN { return Err(WireError::InvalidPayload); } - let bytes = reader.take_bytes(len)?; let name = std::str::from_utf8(&bytes).map_err(|_| WireError::InvalidPayload)?; QlName::new(name) } From 985b0f4002578e0d2b4512d9187e9cf326400d6a Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 29 Jun 2026 17:39:43 -0400 Subject: [PATCH 33/59] ql-wire: remove name validation --- ql-fsm/src/tests/handshake.rs | 4 ++-- ql-runtime/src/driver/test.rs | 4 ++-- ql-runtime/src/tests/mod.rs | 4 ++-- ql-wire/src/identity.rs | 35 ++++++++--------------------------- ql-wire/src/testing.rs | 4 ++-- ql-wire/src/tests.rs | 17 ++--------------- 6 files changed, 18 insertions(+), 50 deletions(-) diff --git a/ql-fsm/src/tests/handshake.rs b/ql-fsm/src/tests/handshake.rs index d44e8789..922a046e 100644 --- a/ql-fsm/src/tests/handshake.rs +++ b/ql-fsm/src/tests/handshake.rs @@ -96,7 +96,7 @@ fn ik_connect_learns_remote_initial_stream_receive_window() { #[test] fn connect_methods_require_bound_peer() { let time = Harness::paired_known(QlFsmConfig::default()).time(); - let identity = generate_identity(&SoftwareCrypto, "identity").unwrap(); + let identity = generate_identity(&SoftwareCrypto, "identity"); let mut fsm = QlFsm::new(QlFsmConfig::default(), identity, time); let crypto = SoftwareCrypto; @@ -343,7 +343,7 @@ fn bind_peer_clears_queued_handshake_output() { harness .a .fsm - .bind_peer(generate_identity(&SoftwareCrypto, "peer").unwrap().bundle()); + .bind_peer(generate_identity(&SoftwareCrypto, "peer").bundle()); assert!(harness.drain_events(Side::A).is_empty()); assert!(harness.next_outbound(Side::A).is_none()); diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs index f3e7b528..dab61c0c 100644 --- a/ql-runtime/src/driver/test.rs +++ b/ql-runtime/src/driver/test.rs @@ -59,7 +59,7 @@ fn new_driver_state() -> (DriverState, QlFsm) { }, QlFsm::new( ql_fsm::QlFsmConfig::default(), - generate_identity(&SoftwareCrypto, "driver").unwrap(), + generate_identity(&SoftwareCrypto, "driver"), Instant::now(), ), ) @@ -161,7 +161,7 @@ fn local_reset_command_reaps_when_other_half_is_already_closed() { #[test] fn unpaired_status_fails_and_reaps_all_streams() { let (mut state, mut fsm) = new_driver_state(); - let peer = generate_identity(&SoftwareCrypto, "peer").unwrap().bundle(); + let peer = generate_identity(&SoftwareCrypto, "peer").bundle(); let stream_id = StreamId(1u32.into()); let (runtime_tx, _runtime_rx) = async_channel::unbounded(); let (_, _, reader_io, writer_io) = io::new_stream(stream_id, runtime_tx); diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index d0f28942..db78ba45 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -682,7 +682,7 @@ fn default_runtime_config() -> RuntimeConfig { #[test] fn runtime_is_send() { let config = default_runtime_config(); - let identity = generate_identity(&SoftwareCrypto, "runtime").unwrap(); + let identity = generate_identity(&SoftwareCrypto, "runtime"); let (platform, _, _, _) = TestPlatform::new(); let (runtime, _handle) = new_runtime(identity, platform, config); let _run: Box + Send> = Box::new(runtime.run()); @@ -691,7 +691,7 @@ fn runtime_is_send() { #[test] fn runtime_exits_when_last_handle_drops() { let config = default_runtime_config(); - let identity = generate_identity(&SoftwareCrypto, "runtime").unwrap(); + let identity = generate_identity(&SoftwareCrypto, "runtime"); let (platform, _, _, _) = TestPlatform::new(); let (runtime, handle) = new_runtime(identity, platform, config); let (done_tx, done_rx) = oneshot::channel(); diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index 65a2901e..90486375 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -1,7 +1,7 @@ use std::ops::Deref; use codec::LenBytes; -use ql_common::{VarInt, QID}; +use ql_common::QID; use crate::{ codec, derive_qid, ByteSlice, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash, @@ -61,23 +61,22 @@ pub struct QlIdentity { impl QlIdentity { pub const FIXED_WIRE_SIZE: usize = QID::SIZE + MlKemPrivateKey::SIZE + MlKemPublicKey::SIZE + size_of::(); - pub const MAX_WIRE_SIZE: usize = Self::FIXED_WIRE_SIZE + VarInt::MAX_SIZE + QlName::MAX_LEN; pub fn new( crypto: &impl QlHash, mlkem_private_key: MlKemPrivateKey, mlkem_public_key: MlKemPublicKey, name: impl Into, - ) -> Result { - let name = QlName::new(name)?; + ) -> Self { + let name = QlName(name.into()); let qid = derive_qid(crypto, &mlkem_public_key); - Ok(Self { + Self { qid, mlkem_private_key, mlkem_public_key, capabilities: 0, name, - }) + } } pub fn bundle(&self) -> PeerBundle { @@ -117,28 +116,13 @@ impl codec::WireDecode for QlIdentity { } } -pub fn generate_identity( - crypto: &impl QlCrypto, - name: impl Into, -) -> Result { +pub fn generate_identity(crypto: &impl QlCrypto, name: impl Into) -> QlIdentity { let MlKemKeyPair { private, public } = crypto.mlkem_generate_keypair(); QlIdentity::new(crypto, private, public, name) } #[derive(Debug, Clone, PartialEq, Eq)] -pub struct QlName(String); - -impl QlName { - pub const MAX_LEN: usize = 256; - - pub fn new(name: impl Into) -> Result { - let name = name.into(); - if name.is_empty() || name.len() > Self::MAX_LEN { - return Err(WireError::InvalidPayload); - } - Ok(Self(name)) - } -} +pub struct QlName(pub String); impl Deref for QlName { type Target = str; @@ -161,10 +145,7 @@ impl WireEncode for QlName { impl codec::WireDecode for QlName { fn decode(reader: &mut codec::Reader) -> Result { let bytes = reader.decode::>()?.0; - if bytes.is_empty() || bytes.len() > Self::MAX_LEN { - return Err(WireError::InvalidPayload); - } let name = std::str::from_utf8(&bytes).map_err(|_| WireError::InvalidPayload)?; - QlName::new(name) + Ok(QlName(name.into())) } } diff --git a/ql-wire/src/testing.rs b/ql-wire/src/testing.rs index adc2e093..1ce88b55 100644 --- a/ql-wire/src/testing.rs +++ b/ql-wire/src/testing.rs @@ -15,8 +15,8 @@ pub struct NoopCrypto; pub fn test_identities(crypto: &impl QlCrypto) -> (QlIdentity, QlIdentity) { ( - crate::generate_identity(crypto, "alice").unwrap(), - crate::generate_identity(crypto, "bob").unwrap(), + crate::generate_identity(crypto, "alice"), + crate::generate_identity(crypto, "bob"), ) } diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index a73cd743..53e5bdcd 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -71,7 +71,7 @@ fn encrypt_record( #[test] fn peer_bundle_round_trip() { let crypto = SoftwareCrypto; - let mut identity = generate_identity(&crypto, "alice").unwrap(); + let mut identity = generate_identity(&crypto, "alice"); identity.capabilities = 1231; let bundle = identity.bundle(); @@ -82,19 +82,6 @@ fn peer_bundle_round_trip() { assert_eq!(&*decoded.name, "alice"); } -#[test] -fn identity_name_validation() { - assert_eq!( - QlName::new("a".repeat(QlName::MAX_LEN)).unwrap().len(), - QlName::MAX_LEN - ); - assert!(matches!(QlName::new(""), Err(WireError::InvalidPayload))); - assert!(matches!( - QlName::new("a".repeat(QlName::MAX_LEN + 1)), - Err(WireError::InvalidPayload) - )); -} - #[test] fn handshake_record_round_trip_supports_ik_kk_and_xx() { let ik = QlHandshakeRecord::Ik1(Ik1 { @@ -292,7 +279,7 @@ fn ik_handshake_rejects_tampered_handshake_header() { fn ik_handshake_rejects_bound_remote_bundle_mismatch() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); - let bogus = generate_identity(&crypto, "bogus").unwrap(); + let bogus = generate_identity(&crypto, "bogus"); let (initiator_to_responder, _) = identity_routes(&initiator, &responder); let mut initiator_state = IkHandshake::new_initiator( From 8199cada1cacb8508ed8eb9c8ce82eec2efc5c89 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 1 Jul 2026 12:19:30 -0400 Subject: [PATCH 34/59] ql: arbitrary stream header --- ql-common/src/lib.rs | 5 +-- ql-common/src/varint.rs | 59 ++++++++++++++++++++++++++++ ql-fsm/src/fsm.rs | 9 ++--- ql-fsm/src/lib.rs | 20 +++------- ql-fsm/src/session/mod.rs | 20 ++++------ ql-fsm/src/session/state.rs | 9 +++-- ql-fsm/src/session/stream_ops.rs | 4 +- ql-fsm/src/session/tests.rs | 23 ++++------- ql-fsm/src/tests/proptest.rs | 11 ++---- ql-fsm/src/tests/session.rs | 34 ++++------------ ql-rpc/src/lib.rs | 2 + ql-rpc/src/metadata.rs | 25 ++++++++++++ ql-rpc/src/router/mod.rs | 13 +++--- ql-runtime/src/command.rs | 5 +-- ql-runtime/src/driver/mod.rs | 30 ++++---------- ql-runtime/src/handle/mod.rs | 12 ++---- ql-runtime/src/rpc/mod.rs | 6 +-- ql-runtime/src/tests/mod.rs | 15 ++----- ql-wire/src/codec.rs | 1 + ql-wire/src/encrypted/builder.rs | 2 +- ql-wire/src/encrypted/mod.rs | 4 +- ql-wire/src/encrypted/stream_data.rs | 47 +++++----------------- ql-wire/src/varint.rs | 41 +++---------------- 23 files changed, 171 insertions(+), 226 deletions(-) create mode 100644 ql-rpc/src/metadata.rs diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index 3904b4b3..373cc509 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -85,10 +85,9 @@ impl std::fmt::Display for ServiceId { varint_wrapper!(RouteId); varint_wrapper!(StreamId); -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct StreamInfo { pub qid: QID, pub stream_id: StreamId, - pub service_id: ServiceId, - pub route_id: RouteId, + pub header: Box<[u8]>, } diff --git a/ql-common/src/varint.rs b/ql-common/src/varint.rs index 8ac64028..a88d7765 100644 --- a/ql-common/src/varint.rs +++ b/ql-common/src/varint.rs @@ -52,6 +52,65 @@ impl VarInt { 8 } } + + /// Return the encoded length from the first encoded byte. + pub const fn encoded_len_from_first_byte(first: u8) -> usize { + 1usize << (first >> 6) + } + + /// Encode this value by writing its encoded bytes to `write`. + #[allow(clippy::cast_possible_truncation)] + pub fn write_bytes(self, mut write: impl FnMut(&[u8])) { + let value = self.0; + match self.size() { + 1 => write(&[value as u8]), + 2 => write(&((value as u16) | 0x4000).to_be_bytes()), + 4 => write(&((value as u32) | 0x8000_0000).to_be_bytes()), + 8 => write(&(value | 0xC000_0000_0000_0000).to_be_bytes()), + _ => unreachable!(), + } + } + + /// decode a value from the start of `bytes`, returning the value and remaining bytes + pub fn decode_bytes(bytes: &[u8]) -> Option<(Self, &[u8])> { + let first = *bytes.first()?; + let len = Self::encoded_len_from_first_byte(first); + Some(( + Self::decode_with_first_byte(first, bytes.get(1..len)?)?, + &bytes[len..], + )) + } + + /// decode a value after the first encoded byte has already been consumed + pub fn decode_with_first_byte(first: u8, tail: &[u8]) -> Option { + let len = Self::encoded_len_from_first_byte(first); + if tail.len() != len - 1 { + return None; + } + let value = match len { + 1 => u64::from(first & 0x3f), + 2 => u64::from(u16::from_be_bytes([first & 0x3f, tail[0]])), + 4 => { + let bytes = [first & 0x3f, tail[0], tail[1], tail[2]]; + u64::from(u32::from_be_bytes(bytes)) + } + 8 => { + let bytes = [ + first & 0x3f, + tail[0], + tail[1], + tail[2], + tail[3], + tail[4], + tail[5], + tail[6], + ]; + u64::from_be_bytes(bytes) + } + _ => unreachable!(), + }; + Self::from_u64(value).ok() + } } impl From for u64 { diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index 2e2f42a9..df74501e 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -8,8 +8,7 @@ use crate::{ handshake, session::{self, SessionEvent, TerminalFrame}, state::LinkState, - Event, NoPeerError, NoSessionError, OpenStreamParams, OutboundWrite, QlFsm, ReceiveError, - StreamError, WriteId, + Event, NoPeerError, NoSessionError, OutboundWrite, QlFsm, ReceiveError, StreamError, WriteId, }; pub struct EventSink<'a> { @@ -238,13 +237,11 @@ pub fn close_session(fsm: &mut QlFsm, code: SessionCloseCode) { pub fn open_stream( fsm: &mut QlFsm, - params: OpenStreamParams, + header: Box<[u8]>, ) -> Result, NoSessionError> { let QlFsm { state, events, .. } = fsm; let conn = state.link.connected_mut_or_err()?; - let inner = - conn.session - .open_stream(params.service_id, params.route_id, EventSink::new(events))?; + let inner = conn.session.open_stream(header, EventSink::new(events))?; Ok(crate::StreamOps { inner }) } diff --git a/ql-fsm/src/lib.rs b/ql-fsm/src/lib.rs index d78d5f16..a10b1f11 100644 --- a/ql-fsm/src/lib.rs +++ b/ql-fsm/src/lib.rs @@ -35,10 +35,8 @@ use std::{ pub use bytes::Bytes; pub use error::*; pub use pairing::PairingInvite; -use ql_common::{ResetCode, RouteId, ServiceId, StreamId}; -use ql_wire::{ - PairingToken, PeerBundle, QlCrypto, QlIdentity, SessionClose, SessionCloseCode, StreamHeader, -}; +use ql_common::{ResetCode, StreamId}; +use ql_wire::{PairingToken, PeerBundle, QlCrypto, QlIdentity, SessionClose, SessionCloseCode}; pub use session::{SessionEvent, StreamReadIter, StreamWriter}; use crate::state::{LinkState, QlFsmState}; @@ -137,7 +135,7 @@ impl StreamOps<'_> { self.inner.stream_id() } - pub fn header(&self) -> &StreamHeader { + pub fn header(&self) -> &[u8] { self.inner.header() } @@ -210,11 +208,6 @@ impl Default for QlFsmConfig { } } -pub struct OpenStreamParams { - pub service_id: ServiceId, - pub route_id: RouteId, -} - /// synchronous driver for peer binding, handshake, and encrypted streams pub struct QlFsm { config: QlFsmConfig, @@ -350,11 +343,8 @@ impl QlFsm { } /// opens a new outgoing stream - pub fn open_stream( - &mut self, - params: OpenStreamParams, - ) -> Result, NoSessionError> { - fsm::open_stream(self, params) + pub fn open_stream(&mut self, header: Box<[u8]>) -> Result, NoSessionError> { + fsm::open_stream(self, header) } /// returns a facade for an open stream diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index 1ab9772f..6105be0a 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -17,10 +17,10 @@ use std::time::{Duration, Instant}; use bytes::Bytes; use indexmap::IndexMap; -use ql_common::{RouteId, ServiceId, StreamId, VarInt}; +use ql_common::{StreamId, VarInt}; use ql_wire::{ - RecordAck, RecordSeq, ResetTarget, SessionClose, SessionCloseCode, SessionFrame, - SessionRecordBuilder, StreamData, StreamHeader, StreamReset, StreamWindow, WireError, + LenBytes, RecordAck, RecordSeq, ResetTarget, SessionClose, SessionCloseCode, SessionFrame, + SessionRecordBuilder, StreamData, StreamReset, StreamWindow, WireError, }; use self::{ @@ -128,8 +128,7 @@ impl SessionFsm { pub fn open_stream( &mut self, - service_id: ServiceId, - route_id: RouteId, + header: Box<[u8]>, sink: E, ) -> Result, NoSessionError> where @@ -145,10 +144,7 @@ impl SessionFsm { stream_id, StreamState::new( StreamRole::Initiator, - Some(StreamHeader { - service_id, - route_id, - }), + Some(Bytes::from(header)), self.config.stream_receive_buffer_size, self.config.initial_peer_stream_receive_window, ), @@ -543,7 +539,7 @@ impl SessionFsm { stream_id, offset, header: if matches!(stream.role, StreamRole::Initiator) && candidate.offset == 0 { - stream.header + stream.header.as_ref().map(LenBytes) } else { None }, @@ -659,9 +655,9 @@ impl SessionFsm { let readable_before = stream.readable_bytes(); let was_finished = matches!(stream.inbound_state, InboundState::Finished); - let opened = match (stream.role, stream.header, header, frame_offset) { + let opened = match (stream.role, stream.header.as_ref(), header, frame_offset) { (StreamRole::Responder, None, Some(header), 0) => { - stream.header = Some(header); + stream.header = Some(header.0); true } (StreamRole::Initiator, _, Some(_), _) diff --git a/ql-fsm/src/session/state.rs b/ql-fsm/src/session/state.rs index a71ec5fc..49e21819 100644 --- a/ql-fsm/src/session/state.rs +++ b/ql-fsm/src/session/state.rs @@ -1,8 +1,9 @@ use std::time::Instant; +use bytes::Bytes; use indexmap::IndexMap; use ql_common::StreamId; -use ql_wire::{RecordSeq, ResetTarget, SessionClose, StreamHeader, StreamReset}; +use ql_wire::{RecordSeq, ResetTarget, SessionClose, StreamReset}; use super::{ ack_tracker::AckTracker, remote_stream_history::RemoteStreamHistory, stream_rx::StreamRx, @@ -46,7 +47,7 @@ pub enum TerminalFrame { #[derive(Debug)] pub struct StreamState { pub role: StreamRole, - pub header: Option, + pub header: Option, pub rx: StreamRx, pub tx: StreamTx, pub pending_reset: Option, @@ -60,14 +61,14 @@ pub struct StreamState { impl StreamState { pub fn new( role: StreamRole, - route_id: Option, + header: Option, receive_buffer_size: u32, initial_peer_stream_receive_window: u32, ) -> Self { let receive_buffer_size = receive_buffer_size as usize; Self { role, - header: route_id, + header, tx: StreamTx::new(), pending_reset: None, peer_max_offset: u64::from(initial_peer_stream_receive_window), diff --git a/ql-fsm/src/session/stream_ops.rs b/ql-fsm/src/session/stream_ops.rs index 5261edef..5ab1c1b8 100644 --- a/ql-fsm/src/session/stream_ops.rs +++ b/ql-fsm/src/session/stream_ops.rs @@ -1,5 +1,5 @@ use ql_common::{ResetCode, StreamId}; -use ql_wire::{StreamHeader, StreamReset}; +use ql_wire::StreamReset; use super::{ state::{InboundState, StreamState}, @@ -40,7 +40,7 @@ impl<'a, E: EventSink> StreamOps<'a, E> { /// returns the streams details #[inline] - pub fn header(&self) -> &StreamHeader { + pub fn header(&self) -> &[u8] { self.stream().header.as_ref().unwrap() } diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index 84386785..48e37ed5 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -1,10 +1,10 @@ use std::time::{Duration, Instant}; use bytes::Bytes; -use ql_common::{ResetCode, RouteId, ServiceId, StreamId, VarInt, QID}; +use ql_common::{ResetCode, StreamId, VarInt, QID}; use ql_wire::{ - decode_session_frames, parse_session_frames, RecordAck, RecordSeq, ResetTarget, SessionFrame, - SessionRecordBuilder, StreamData, StreamHeader, StreamReset, + decode_session_frames, parse_session_frames, LenBytes, RecordAck, RecordSeq, ResetTarget, + SessionFrame, SessionRecordBuilder, StreamData, StreamReset, }; use super::{SessionConfig, SessionEvent, SessionFsm}; @@ -22,10 +22,6 @@ fn offset(value: u64) -> VarInt { VarInt::from_u64(value).unwrap() } -fn route_id(value: u64) -> RouteId { - RouteId::from_u64(value).unwrap() -} - fn record_ack(seq: RecordSeq) -> RecordAck { RecordAck::from_ranges([seq..=seq]).unwrap() } @@ -33,11 +29,8 @@ fn record_ack(seq: RecordSeq) -> RecordAck { const REFUSED: ResetCode = ResetCode(1); const TIMEOUT: ResetCode = ResetCode(2); -fn header(value: u64) -> StreamHeader { - StreamHeader { - route_id: route_id(value), - service_id: ServiceId([0; 16]), - } +fn header(value: u64) -> LenBytes> { + LenBytes(vec![value as u8]) } // todo: remove @@ -46,7 +39,7 @@ fn opened(stream_id: StreamId) -> SessionEvent { } fn open_stream_id(fsm: &mut SessionFsm) -> StreamId { - fsm.open_stream(ServiceId([0; 16]), route_id(1), |_| {}) + fsm.open_stream(header(1).0.into_boxed_slice(), |_| {}) .unwrap() .stream_id() } @@ -451,7 +444,7 @@ fn stream_ids_follow_even_odd_xid_ordering() { }, now, ) - .open_stream(ServiceId([0; 16]), route_id(1), |_| {}) + .open_stream(header(1).0.into_boxed_slice(), |_| {}) .unwrap() .stream_id(); let odd_id = SessionFsm::new( @@ -461,7 +454,7 @@ fn stream_ids_follow_even_odd_xid_ordering() { }, now, ) - .open_stream(ServiceId([0; 16]), route_id(1), |_| {}) + .open_stream(header(1).0.into_boxed_slice(), |_| {}) .unwrap() .stream_id(); diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs index ac1fd440..76000e6a 100644 --- a/ql-fsm/src/tests/proptest.rs +++ b/ql-fsm/src/tests/proptest.rs @@ -7,13 +7,11 @@ extern crate proptest as proptest_crate; use bytes::Bytes; use proptest_crate::{collection::vec, prelude::*, test_runner::TestCaseResult}; -use ql_common::{ResetCode, RouteId, ServiceId, StreamId}; +use ql_common::{ResetCode, StreamId}; use ql_wire::WireError; use super::*; -use crate::{ - state::LinkState, Event, OpenStreamParams, PeerStatus, ReceiveError, StreamResetTarget, WriteId, -}; +use crate::{state::LinkState, Event, PeerStatus, ReceiveError, StreamResetTarget, WriteId}; const SLOT_COUNT: usize = 4; @@ -284,10 +282,7 @@ impl Runner { .harness .node_mut(*side) .fsm - .open_stream(OpenStreamParams { - service_id: ServiceId([1; 16]), - route_id: RouteId::from(1u32), - }) + .open_stream(Box::from([1])) .ok() .map(|stream| stream.stream_id()); if let Some(stream_id) = stream_id { diff --git a/ql-fsm/src/tests/session.rs b/ql-fsm/src/tests/session.rs index a70fd091..f553a9ac 100644 --- a/ql-fsm/src/tests/session.rs +++ b/ql-fsm/src/tests/session.rs @@ -1,26 +1,18 @@ use std::time::Duration; use bytes::Bytes; -use ql_common::{RouteId, ServiceId, StreamId, VarInt}; +use ql_common::{StreamId, VarInt}; use ql_wire::SessionClose; use super::*; -use crate::{ - state::LinkState, CommitReadError, Event, NoSessionError, OpenStreamParams, PeerStatus, - StreamError, -}; +use crate::{state::LinkState, CommitReadError, Event, NoSessionError, PeerStatus, StreamError}; fn stream_id(value: u32) -> StreamId { StreamId(VarInt::from_u32(value)) } fn open_stream_id(fsm: &mut QlFsm) -> StreamId { - fsm.open_stream(OpenStreamParams { - service_id: ServiceId([1; 16]), - route_id: RouteId::from(1u32), - }) - .unwrap() - .stream_id() + fsm.open_stream(Box::from([1])).unwrap().stream_id() } fn write_stream_bytes( @@ -197,10 +189,7 @@ fn disconnected_stream_operations_fail_with_no_session() { let missing = stream_id(0); assert!(matches!( - harness.a.fsm.open_stream(OpenStreamParams { - service_id: ServiceId([1; 16]), - route_id: RouteId::from(1u32), - }), + harness.a.fsm.open_stream(Box::from([1])), Err(NoSessionError) )); assert_eq!( @@ -370,10 +359,7 @@ fn close_session_disconnects_locally() { )); assert!(matches!(harness.a.fsm.state.link, LinkState::Connected(_))); assert!(matches!( - harness.a.fsm.open_stream(OpenStreamParams { - service_id: ServiceId([1; 16]), - route_id: RouteId::from(1u32), - }), + harness.a.fsm.open_stream(Box::from([1])), Err(NoSessionError) )); assert_eq!(harness.a.fsm.queue_ping(), Err(NoSessionError)); @@ -403,10 +389,7 @@ fn unpair_clears_bound_peer_and_emits_unpair_frame() { ); assert!(harness.a.fsm.peer().is_none()); assert!(matches!( - harness.a.fsm.open_stream(OpenStreamParams { - service_id: ServiceId([1; 16]), - route_id: RouteId::from(1u32), - }), + harness.a.fsm.open_stream(Box::from([1])), Err(NoSessionError) )); assert_eq!(harness.a.fsm.queue_ping(), Err(NoSessionError)); @@ -433,10 +416,7 @@ fn inbound_unpair_clears_remote_peer_binding() { ); assert!(harness.b.fsm.peer().is_none()); assert!(matches!( - harness.b.fsm.open_stream(OpenStreamParams { - service_id: ServiceId([1; 16]), - route_id: RouteId::from(1u32), - }), + harness.b.fsm.open_stream(Box::from([1])), Err(NoSessionError) )); assert!(matches!(harness.connect_ik(Side::B), Err(NoPeerError))); diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index 7dd45047..758af97a 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -6,6 +6,7 @@ mod chunk_queue; mod codec; mod error; mod framed_value; +mod metadata; mod router; mod rpc; mod stream; @@ -14,6 +15,7 @@ pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; use framed_value::*; +pub use metadata::*; pub use router::*; pub use rpc::*; pub use stream::*; diff --git a/ql-rpc/src/metadata.rs b/ql-rpc/src/metadata.rs new file mode 100644 index 00000000..77eff765 --- /dev/null +++ b/ql-rpc/src/metadata.rs @@ -0,0 +1,25 @@ +use ql_common::{RouteId, ServiceId, VarInt}; + +use crate::RouteKey; + +pub fn encode_stream_header() -> Box<[u8]> { + encode_route_key(RouteKey::new::()) +} + +pub fn encode_route_key(key: RouteKey) -> Box<[u8]> { + let mut out = Vec::with_capacity(ServiceId::SIZE + key.route_id.0.size()); + out.extend_from_slice(&key.service_id.0); + key.route_id + .0 + .write_bytes(|bytes| out.extend_from_slice(bytes)); + out.into_boxed_slice() +} + +pub fn decode_stream_header(bytes: &[u8]) -> Option { + let service_id = ServiceId(bytes.get(..ServiceId::SIZE)?.try_into().ok()?); + let (route_id, rest) = VarInt::decode_bytes(&bytes[ServiceId::SIZE..])?; + rest.is_empty().then_some(RouteKey { + service_id, + route_id: RouteId(route_id), + }) +} diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index 568e5e3d..a7030b46 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -5,7 +5,7 @@ mod config; mod mode; pub use self::{builder::*, config::*, mode::*}; -use crate::{RpcRead, RpcStream, RpcWrite}; +use crate::{decode_stream_header, RpcRead, RpcStream, RpcWrite}; pub struct Router where @@ -79,13 +79,14 @@ where let StreamInfo { qid, stream_id, - service_id, - route_id, + header, } = info; let context = Context { qid, stream_id }; - let key = RouteKey { - service_id, - route_id, + let Some(key) = decode_stream_header(&header) else { + let (reader, writer) = stream.split(); + reader.reset(ResetCode::PROTOCOL); + writer.reset(ResetCode::PROTOCOL); + return None; }; let Ok(index) = self.routes.binary_search_by_key(&key, |entry| entry.key) else { let (reader, writer) = stream.split(); diff --git a/ql-runtime/src/command.rs b/ql-runtime/src/command.rs index 72138c71..22c26897 100644 --- a/ql-runtime/src/command.rs +++ b/ql-runtime/src/command.rs @@ -1,4 +1,4 @@ -use ql_common::{ResetCode, RouteId, ServiceId, StreamId}; +use ql_common::{ResetCode, StreamId}; use ql_fsm::{NoSessionError, PairingInvite, StreamResetTarget}; use ql_wire::{PairingToken, PeerBundle, SessionCloseCode}; @@ -17,8 +17,7 @@ pub enum Command { invite: PairingInvite, }, OpenStream { - service_id: ServiceId, - route_id: RouteId, + header: Box<[u8]>, start: oneshot::Sender>, }, PollInbound { diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index e2feb2d7..9970209c 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -17,7 +17,6 @@ use async_channel::Recv; use futures_lite::future::{poll_fn, yield_now}; use ql_common::{ResetCode, ResetOrigin, StreamId, StreamInfo}; use ql_fsm::{Event, QlFsm, StreamResetEvent, StreamResetTarget, WriteId}; -use ql_wire::StreamHeader; use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; use crate::{ @@ -210,26 +209,19 @@ impl DriverState { log::info!("unpairing peer"); fsm.unpair(); } - Command::OpenStream { - service_id, - route_id, - start, - } => { - log::info!("open stream requested: route_id={route_id}"); + Command::OpenStream { header, start } => { + log::info!("open stream requested"); - let mut stream_ops = match fsm.open_stream(ql_fsm::OpenStreamParams { - service_id, - route_id, - }) { + let mut stream_ops = match fsm.open_stream(header) { Ok(stream_ops) => stream_ops, Err(error) => { - log::warn!("open stream failed: route_id={route_id}"); + log::warn!("open stream failed"); let _ = start.send(Err(error)); return; } }; let stream_id = stream_ops.stream_id(); - log::info!("open stream allocated: service_id={service_id} route_id={route_id} stream_id={stream_id}"); + log::info!("open stream allocated: stream_id={stream_id}"); let (reader, writer, reader_io, writer_io) = io::new_stream(stream_id, self.runtime_tx.clone()); self.streams.insert( @@ -362,21 +354,15 @@ impl DriverState { let qid = fsm.peer().unwrap().qid; let stream = fsm.stream(stream_id).unwrap(); - let StreamHeader { - service_id, - route_id, - } = *stream.header(); + let header = Box::<[u8]>::from(stream.header()); - log::info!( - "delivering inbound stream to platform: service_id={service_id} route_id={route_id} stream_id={stream_id}", - ); + log::info!("delivering inbound stream to platform: stream_id={stream_id}",); platform.handle_inbound( StreamInfo { qid, stream_id, - service_id, - route_id, + header, }, crate::QlStream { writer, reader }, ); diff --git a/ql-runtime/src/handle/mod.rs b/ql-runtime/src/handle/mod.rs index d9ba3f60..7256e2e8 100644 --- a/ql-runtime/src/handle/mod.rs +++ b/ql-runtime/src/handle/mod.rs @@ -1,6 +1,6 @@ use std::sync::Arc; -use ql_fsm::{NoSessionError, OpenStreamParams, PairingInvite}; +use ql_fsm::{NoSessionError, PairingInvite}; use ql_wire::{PairingToken, PeerBundle, SessionCloseCode}; use crate::command::Command; @@ -54,16 +54,10 @@ impl RuntimeHandle { } /// opens a new stream on the active encrypted session - pub async fn open_stream(&self, params: OpenStreamParams) -> Result { + pub async fn open_stream(&self, header: Box<[u8]>) -> Result { let (start_tx, start_rx) = oneshot::channel(); - let OpenStreamParams { - service_id, - route_id, - } = params; - self.send(Command::OpenStream { - route_id, - service_id, + header, start: start_tx, }); diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index d58efa17..d6247dd7 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -1,6 +1,5 @@ mod adapter; -use ql_fsm::OpenStreamParams; use ql_rpc::{download, duplex, notification, progress, request, subscription, upload, Route}; use crate::{QlStream, QlStreamError, RuntimeHandle, StreamReader, StreamWriter}; @@ -95,10 +94,7 @@ impl RpcHandle { async fn open_rpc_stream(&self) -> RpcResult { self.inner - .open_stream(OpenStreamParams { - service_id: R::SERVICE, - route_id: R::ROUTE, - }) + .open_stream(ql_rpc::encode_stream_header::()) .await .map_err(QlStreamError::from) .map_err(ql_rpc::RpcError::Transport) diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index db78ba45..c5cdf1ff 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -11,8 +11,8 @@ use std::{ use async_channel::{Receiver, Sender}; use futures_lite::Stream; -use ql_common::{RouteId, ServiceId, StreamInfo, QID}; -use ql_fsm::{OpenStreamParams, PeerStatus}; +use ql_common::{StreamInfo, QID}; +use ql_fsm::PeerStatus; use ql_wire::{ generate_identity, test_identities, MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, Nonce, PairingToken, PeerBundle, QlAead, QlHash, QlIdentity, QlKem, QlRandom, @@ -65,15 +65,8 @@ impl Side { } } -fn test_route_id() -> RouteId { - RouteId::from_u32(1) -} - -fn test_open_stream_params() -> OpenStreamParams { - OpenStreamParams { - service_id: ServiceId([1; 16]), - route_id: test_route_id(), - } +fn test_open_stream_params() -> Box<[u8]> { + Box::from([1]) } #[derive(Debug, Clone)] diff --git a/ql-wire/src/codec.rs b/ql-wire/src/codec.rs index a8a164ec..2bc60b3e 100644 --- a/ql-wire/src/codec.rs +++ b/ql-wire/src/codec.rs @@ -246,6 +246,7 @@ impl Reader { } /// bytes encoded with a VarInt len prefix +#[derive(Debug, Clone, PartialEq, Eq)] pub struct LenBytes(pub B); impl WireEncode for LenBytes diff --git a/ql-wire/src/encrypted/builder.rs b/ql-wire/src/encrypted/builder.rs index f25f6859..f1a87f3e 100644 --- a/ql-wire/src/encrypted/builder.rs +++ b/ql-wire/src/encrypted/builder.rs @@ -71,7 +71,7 @@ impl SessionRecordBuilder { self.push_frame_payload(super::SessionFrameKind::Ack, ack) } - pub fn push_stream_data(&mut self, frame: &StreamData) -> bool { + pub fn push_stream_data(&mut self, frame: &StreamData) -> bool { self.push_frame_payload(super::SessionFrameKind::StreamData, frame) } diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index ca12e9e8..2d900d22 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -1,4 +1,4 @@ -use ql_common::{RouteId, ServiceId, StreamId}; +use ql_common::StreamId; use crate::{ codec, encrypted_message::EncryptedMessage, BufView, ByteSlice, Nonce, QlCrypto, Reader, @@ -19,9 +19,7 @@ pub use stream_data::*; pub use stream_reset::*; pub use stream_window::*; -varint_wrapper_codec!(RouteId); varint_wrapper_codec!(StreamId); -array_wrapper_codec!(ServiceId); #[derive(Debug, Clone, PartialEq, Eq)] pub enum SessionFrame { diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index d160d47e..d62cfdc1 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -1,4 +1,4 @@ -use ql_common::{RouteId, ServiceId, StreamId, VarInt}; +use ql_common::{StreamId, VarInt}; use crate::{ codec::{self, LenBytes}, @@ -7,19 +7,19 @@ use crate::{ /// carries bytes for a stream and may finish that sending direction. #[derive(Debug, Clone, PartialEq, Eq)] -pub struct StreamData { +pub struct StreamData { pub stream_id: StreamId, pub offset: VarInt, - pub header: Option, + pub header: Option>, pub fin: bool, pub bytes: B, } -impl StreamData { +impl StreamData { pub const MIN_WIRE_SIZE: usize = StreamId::MAX_ENCODED_LEN + VarInt::MAX_SIZE + size_of::() - + StreamHeader::MAX_WIRE_SIZE + + VarInt::MAX_SIZE + VarInt::MAX_SIZE; } @@ -47,22 +47,23 @@ impl WireDecode for StreamData { } } -impl StreamData { +impl StreamData { pub fn into_owned(self) -> StreamData> where B: ByteSlice, + H: ByteSlice, { StreamData { stream_id: self.stream_id, offset: self.offset, - header: self.header, + header: self.header.map(|header| LenBytes(header.0.to_vec())), fin: self.fin, bytes: self.bytes.to_vec(), } } } -impl WireEncode for StreamData { +impl WireEncode for StreamData { fn encoded_len(&self) -> usize { self.stream_id.encoded_len() + self.offset.encoded_len() @@ -94,36 +95,6 @@ impl WireEncode for StreamData { } } -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct StreamHeader { - pub service_id: ServiceId, - pub route_id: RouteId, -} - -impl StreamHeader { - pub const MAX_WIRE_SIZE: usize = ServiceId::SIZE + RouteId::MAX_ENCODED_LEN; -} - -impl WireDecode for StreamHeader { - fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self { - service_id: reader.decode()?, - route_id: reader.decode()?, - }) - } -} - -impl WireEncode for StreamHeader { - fn encoded_len(&self) -> usize { - self.service_id.encoded_len() + self.route_id.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.service_id.encode(out); - self.route_id.encode(out); - } -} - mod flag { pub const FIN: u8 = 0x01; pub const HEADER: u8 = 0x02; diff --git a/ql-wire/src/varint.rs b/ql-wire/src/varint.rs index 8c94d34c..02f8d2a3 100644 --- a/ql-wire/src/varint.rs +++ b/ql-wire/src/varint.rs @@ -1,37 +1,14 @@ use bytes::BufMut; +use ql_common::VarInt; use crate::{ByteSlice, Reader, WireDecode, WireEncode, WireError}; impl WireDecode for ql_common::VarInt { fn decode(reader: &mut Reader) -> Result { let first = reader.decode::()?; - let tag = first >> 6; - let first = first & 0b0011_1111; - let value = match tag { - 0b00 => u64::from(first), - 0b01 => { - let mut buf = [0; 2]; - buf[0] = first; - buf[1] = reader.decode()?; - u64::from(u16::from_be_bytes(buf)) - } - 0b10 => { - let mut buf = [0; 4]; - buf[0] = first; - buf[1..].copy_from_slice(&reader.take_bytes(3)?); - u64::from(u32::from_be_bytes(buf)) - } - 0b11 => { - let mut buf = [0; 8]; - buf[0] = first; - buf[1..].copy_from_slice(&reader.take_bytes(7)?); - u64::from_be_bytes(buf) - } - _ => unreachable!(), - }; - - // SAFETY: the decoded value is guaranteed to fit in the 62-bit varint range. - Ok(unsafe { Self::from_u64_unchecked(value) }) + let len = VarInt::encoded_len_from_first_byte(first); + let tail = reader.take_bytes(len - 1)?; + VarInt::decode_with_first_byte(first, &tail).ok_or(WireError::InvalidPayload) } } @@ -40,15 +17,7 @@ impl WireEncode for ql_common::VarInt { self.size() } - #[allow(clippy::cast_possible_truncation)] fn encode(&self, out: &mut W) { - let x = self.into_inner(); - match self.size() { - 1 => out.put_u8(x as u8), - 2 => out.put_u16((0b01 << 14) | x as u16), - 4 => out.put_u32((0b10 << 30) | x as u32), - 8 => out.put_u64((0b11 << 62) | x), - _ => unreachable!("malformed varint"), - } + self.write_bytes(|bytes| out.put_slice(bytes)); } } From 65ae6490432970208a1befa320caa91f66e7b7cf Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 1 Jul 2026 15:36:44 -0400 Subject: [PATCH 35/59] ql-rpc: generic key --- ql-rpc/src/lib.rs | 2 - ql-rpc/src/metadata.rs | 25 ------ ql-rpc/src/router/builder.rs | 58 ++++++------- ql-rpc/src/router/mod.rs | 60 ++++++------- ql-rpc/src/rpc/mod.rs | 20 +++-- ql-runtime/src/rpc/mod.rs | 15 +++- ql-runtime/src/tests/rpc.rs | 158 +++++++++++++++++++---------------- 7 files changed, 165 insertions(+), 173 deletions(-) delete mode 100644 ql-rpc/src/metadata.rs diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index 758af97a..7dd45047 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -6,7 +6,6 @@ mod chunk_queue; mod codec; mod error; mod framed_value; -mod metadata; mod router; mod rpc; mod stream; @@ -15,7 +14,6 @@ pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; use framed_value::*; -pub use metadata::*; pub use router::*; pub use rpc::*; pub use stream::*; diff --git a/ql-rpc/src/metadata.rs b/ql-rpc/src/metadata.rs deleted file mode 100644 index 77eff765..00000000 --- a/ql-rpc/src/metadata.rs +++ /dev/null @@ -1,25 +0,0 @@ -use ql_common::{RouteId, ServiceId, VarInt}; - -use crate::RouteKey; - -pub fn encode_stream_header() -> Box<[u8]> { - encode_route_key(RouteKey::new::()) -} - -pub fn encode_route_key(key: RouteKey) -> Box<[u8]> { - let mut out = Vec::with_capacity(ServiceId::SIZE + key.route_id.0.size()); - out.extend_from_slice(&key.service_id.0); - key.route_id - .0 - .write_bytes(|bytes| out.extend_from_slice(bytes)); - out.into_boxed_slice() -} - -pub fn decode_stream_header(bytes: &[u8]) -> Option { - let service_id = ServiceId(bytes.get(..ServiceId::SIZE)?.try_into().ok()?); - let (route_id, rest) = VarInt::decode_bytes(&bytes[ServiceId::SIZE..])?; - rest.is_empty().then_some(RouteKey { - service_id, - route_id: RouteId(route_id), - }) -} diff --git a/ql-rpc/src/router/builder.rs b/ql-rpc/src/router/builder.rs index 05ff6e39..0c1f9dd7 100644 --- a/ql-rpc/src/router/builder.rs +++ b/ql-rpc/src/router/builder.rs @@ -3,24 +3,26 @@ use std::marker::PhantomData; use super::*; use crate::{ download::*, duplex::*, notification::*, progress::*, request::*, subscription::*, upload::*, - RouteKey, + Route, }; pub struct LocalRoutes; pub struct SendRoutes; -pub struct RouterBuilder +pub struct RouterBuilder where + K: RpcRouteKey, Sp: Spawner, { config: RouterConfig, spawner: Sp, - routes: Vec>, + routes: Vec>, marker: PhantomData Mode>, } -impl RouterBuilder +impl RouterBuilder where + K: RpcRouteKey, Sp: Spawner, { pub(crate) fn new(spawner: Sp) -> Self { @@ -42,8 +44,8 @@ where self } - pub fn build(mut self, state: S) -> Router { - self.routes.sort_by_key(|entry| entry.key); + pub fn build(mut self, state: S) -> Router { + self.routes.sort_by_key(|entry| entry.key.clone()); self.routes.shrink_to_fit(); Router { config: self.config, @@ -53,27 +55,24 @@ where } } - fn add_route(mut self, key: RouteKey, route: RouteFn) -> Self { + fn add_route(mut self, key: K, route: RouteFn) -> Self { if self.routes.iter().any(|entry| entry.key == key) { - panic!( - "duplicate rpc route {} for service {:?}", - key.route_id.0.into_inner(), - key.service_id.0 - ); + panic!("duplicate rpc route {key:?}"); } self.routes.push(RouteEntry::new(key, route)); self } } -impl RouterBuilder +impl RouterBuilder where + K: RpcRouteKey, Sp: LocalSpawner, St: RpcStream + 'static, { pub fn request(self) -> Self where - M: Request + 'static, + M: Request + 'static, S: RequestHandlerLocal + 'static, { add_route!(self, M, handle_request, S::handle, S::handle_error) @@ -81,7 +80,7 @@ where pub fn notification(self) -> Self where - M: Notification + 'static, + M: Notification + 'static, S: NotificationHandlerLocal + 'static, { add_route!(self, M, handle_notification, S::handle, S::handle_error) @@ -89,7 +88,7 @@ where pub fn duplex(self) -> Self where - M: Duplex + 'static, + M: Duplex + 'static, S: DuplexHandlerLocal + 'static, { add_route!(self, M, handle_duplex, S::handle) @@ -97,7 +96,7 @@ where pub fn download(self) -> Self where - M: Download + 'static, + M: Download + 'static, S: DownloadHandlerLocal + 'static, { add_route!(self, M, handle_download, S::handle, S::handle_error) @@ -105,7 +104,7 @@ where pub fn subscription(self) -> Self where - M: Subscription + 'static, + M: Subscription + 'static, S: SubscriptionHandlerLocal + 'static, { add_route!(self, M, handle_subscription, S::handle, S::handle_error) @@ -113,7 +112,7 @@ where pub fn progress(self) -> Self where - M: Progress + 'static, + M: Progress + 'static, S: ProgressHandlerLocal + 'static, { add_route!(self, M, handle_progress, S::handle, S::handle_error) @@ -121,21 +120,22 @@ where pub fn upload(self) -> Self where - M: Upload + 'static, + M: Upload + 'static, S: UploadHandlerLocal + 'static, { add_route!(self, M, handle_upload, S::handle, S::handle_error) } } -impl RouterBuilder +impl RouterBuilder where + K: RpcRouteKey, Sp: SendSpawner + Send, St: RpcStream + 'static, { pub fn request(self) -> Self where - M: Request + 'static, + M: Request + 'static, M::Request: Send + 'static, S: RequestHandler + Send + 'static, St::Reader: Send + 'static, @@ -146,7 +146,7 @@ where pub fn notification(self) -> Self where - M: Notification + 'static, + M: Notification + 'static, M::Payload: Send + 'static, S: NotificationHandler + Send + 'static, St::Reader: Send + 'static, @@ -157,7 +157,7 @@ where pub fn duplex(self) -> Self where - M: Duplex + 'static, + M: Duplex + 'static, M::InitiatorEvent: Send + 'static, M::ResponderEvent: Send + 'static, S: DuplexHandler + Send + 'static, @@ -169,7 +169,7 @@ where pub fn download(self) -> Self where - M: Download + 'static, + M: Download + 'static, M::Request: Send + 'static, S: DownloadHandler + Send + 'static, St::Reader: Send + 'static, @@ -180,7 +180,7 @@ where pub fn subscription(self) -> Self where - M: Subscription + 'static, + M: Subscription + 'static, M::Request: Send + 'static, S: SubscriptionHandler + Send + 'static, St::Reader: Send + 'static, @@ -191,7 +191,7 @@ where pub fn progress(self) -> Self where - M: Progress + 'static, + M: Progress + 'static, M::Request: Send + 'static, S: ProgressHandler + Send + 'static, St::Reader: Send + 'static, @@ -202,7 +202,7 @@ where pub fn upload(self) -> Self where - M: Upload + 'static, + M: Upload + 'static, M::Request: Send + 'static, S: UploadHandler + Send + 'static, St::Reader: Send + 'static, @@ -215,7 +215,7 @@ where macro_rules! add_route { ($builder:expr, $rpc:ty, $handler:ident, $($arg:path),+ $(,)?) => { $builder.add_route( - RouteKey::new::<$rpc>(), + <$rpc as Route>::key(), |spawner, state, context, config, stream| { spawner.spawn($handler( state, diff --git a/ql-rpc/src/router/mod.rs b/ql-rpc/src/router/mod.rs index a7030b46..0b867c6d 100644 --- a/ql-rpc/src/router/mod.rs +++ b/ql-rpc/src/router/mod.rs @@ -1,20 +1,21 @@ -use ql_common::{ResetCode, RouteId, ServiceId, StreamId, StreamInfo, QID}; +use ql_common::{ResetCode, StreamId, StreamInfo, QID}; mod builder; mod config; mod mode; pub use self::{builder::*, config::*, mode::*}; -use crate::{decode_stream_header, RpcRead, RpcStream, RpcWrite}; +use crate::{RpcRead, RpcRouteKey, RpcStream, RpcWrite}; -pub struct Router +pub struct Router where + K: RpcRouteKey, Sp: Spawner, { config: RouterConfig, state: S, spawner: Sp, - routes: Vec>, + routes: Vec>, } #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] @@ -23,56 +24,44 @@ pub struct Context { pub stream_id: StreamId, } -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -pub struct RouteKey { - pub service_id: ServiceId, - pub route_id: RouteId, -} - -impl RouteKey { - pub const fn new() -> Self { - Self { - service_id: R::SERVICE, - route_id: R::ROUTE, - } - } -} - -struct RouteEntry +struct RouteEntry where + K: RpcRouteKey, Sp: Spawner, { - key: RouteKey, + key: K, route: RouteFn, } -impl RouteEntry +impl RouteEntry where + K: RpcRouteKey, Sp: Spawner, { - fn new(key: RouteKey, route: RouteFn) -> Self { + fn new(key: K, route: RouteFn) -> Self { Self { key, route } } } -impl Router +impl Router where + K: RpcRouteKey, S: Clone + 'static, St: RpcStream, Sp: Spawner, { - pub fn builder_local(spawner: Sp) -> RouterBuilder + pub fn builder_local(spawner: Sp) -> RouterBuilder where Sp: LocalSpawner, { - RouterBuilder::::new(spawner) + RouterBuilder::::new(spawner) } - pub fn builder_send(spawner: Sp) -> RouterBuilder + pub fn builder_send(spawner: Sp) -> RouterBuilder where Sp: SendSpawner, { - RouterBuilder::::new(spawner) + RouterBuilder::::new(spawner) } pub fn handle(&self, info: StreamInfo, stream: St) -> Option { @@ -82,13 +71,16 @@ where header, } = info; let context = Context { qid, stream_id }; - let Some(key) = decode_stream_header(&header) else { + let Some(key) = K::decode(&header) else { let (reader, writer) = stream.split(); reader.reset(ResetCode::PROTOCOL); writer.reset(ResetCode::PROTOCOL); return None; }; - let Ok(index) = self.routes.binary_search_by_key(&key, |entry| entry.key) else { + let Ok(index) = self + .routes + .binary_search_by_key(&key, |entry| entry.key.clone()) + else { let (reader, writer) = stream.split(); reader.reset(ResetCode::UNKNOWN_ROUTE); writer.reset(ResetCode::UNKNOWN_ROUTE); @@ -104,11 +96,7 @@ where )) } - pub fn route_ids(&self) -> impl ExactSizeIterator + '_ { - self.routes.iter().map(|entry| entry.key.route_id) - } - - pub fn route_keys(&self) -> impl ExactSizeIterator + '_ { - self.routes.iter().map(|entry| entry.key) + pub fn route_keys(&self) -> impl ExactSizeIterator + '_ { + self.routes.iter().map(|entry| &entry.key) } } diff --git a/ql-rpc/src/rpc/mod.rs b/ql-rpc/src/rpc/mod.rs index f51803aa..1d726b92 100644 --- a/ql-rpc/src/rpc/mod.rs +++ b/ql-rpc/src/rpc/mod.rs @@ -2,10 +2,10 @@ //! //! each trait in this module names one rpc shape and the typed values that //! travel on that stream -//! route dispatch uses [`crate::RouteId`] and the submodules provide the matching -//! client and server helpers for encoding, decoding, and handler glue +//! route dispatch uses caller-defined route keys and the submodules provide the +//! matching client and server helpers for encoding, decoding, and handler glue -use ql_common::{RouteId, ServiceId}; +use bytes::BufMut; pub mod download; pub mod duplex; @@ -17,12 +17,18 @@ pub mod subscription; pub mod upload; mod utils; +pub trait RpcRouteKey: Sized + std::fmt::Debug + Ord + Clone + 'static { + fn encoded_len(&self) -> usize; + + fn encode(&self, out: &mut W); + + fn decode(bytes: &[u8]) -> Option; +} + pub trait Route { - /// service used to scope this rpc route. - const SERVICE: ServiceId; + type Key: RpcRouteKey; - /// route used to dispatch this rpc family within [`Self::SERVICE`]. - const ROUTE: RouteId; + fn key() -> Self::Key; } use utils::*; diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs index d6247dd7..e75df6a2 100644 --- a/ql-runtime/src/rpc/mod.rs +++ b/ql-runtime/src/rpc/mod.rs @@ -1,6 +1,8 @@ mod adapter; -use ql_rpc::{download, duplex, notification, progress, request, subscription, upload, Route}; +use ql_rpc::{ + download, duplex, notification, progress, request, subscription, upload, Route, RpcRouteKey, +}; use crate::{QlStream, QlStreamError, RuntimeHandle, StreamReader, StreamWriter}; @@ -92,9 +94,16 @@ impl RpcHandle { Self { inner } } - async fn open_rpc_stream(&self) -> RpcResult { + async fn open_rpc_stream(&self) -> RpcResult + where + R: Route, + R::Key: RpcRouteKey, + { + let key = R::key(); + let mut header = Vec::with_capacity(key.encoded_len()); + key.encode(&mut header); self.inner - .open_stream(ql_rpc::encode_stream_header::()) + .open_stream(header.into_boxed_slice()) .await .map_err(QlStreamError::from) .map_err(ql_rpc::RpcError::Transport) diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index 9c6b608a..b1d85108 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -7,8 +7,8 @@ use std::{ time::Duration, }; -use bytes::Bytes; -use ql_common::{ResetCode, ResetOrigin, RouteId, ServiceId}; +use bytes::{BufMut, Bytes}; +use ql_common::{ResetCode, ResetOrigin, RouteId, VarInt}; use ql_rpc::{ download::{DownloadHandlerLocal, DownloadStart}, duplex::{DuplexHandlerLocal, DuplexPeer}, @@ -23,7 +23,35 @@ use ql_rpc::{ use super::*; use crate::{QlStream, QlStreamError, StreamWriter}; -const TEST_SERVICE: ServiceId = ServiceId([7; 16]); +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +struct TestRouteKey(RouteId); + +impl ql_rpc::RpcRouteKey for TestRouteKey { + fn encoded_len(&self) -> usize { + self.0 .0.size() + } + + fn encode(&self, out: &mut W) { + self.0 .0.write_bytes(|bytes| out.put_slice(bytes)); + } + + fn decode(bytes: &[u8]) -> Option { + let (route_id, rest) = VarInt::decode_bytes(bytes)?; + rest.is_empty().then_some(Self(RouteId(route_id))) + } +} + +macro_rules! test_route { + ($route:ty, $id:expr) => { + impl ql_rpc::Route for $route { + type Key = TestRouteKey; + + fn key() -> Self::Key { + TestRouteKey(RouteId::from_u32($id)) + } + } + }; +} #[derive(Debug, Clone, Copy)] struct TokioLocalSpawner; @@ -59,10 +87,7 @@ impl SendSpawner for TokioSendSpawner { struct Echo; -impl ql_rpc::Route for Echo { - const SERVICE: ServiceId = TEST_SERVICE; - const ROUTE: RouteId = RouteId::from_u32(51); -} +test_route!(Echo, 51); impl ql_rpc::request::Request for Echo { type Error = Utf8Error; @@ -73,10 +98,7 @@ impl ql_rpc::request::Request for Echo { struct Feed; -impl ql_rpc::Route for Feed { - const SERVICE: ServiceId = TEST_SERVICE; - const ROUTE: RouteId = RouteId::from_u32(52); -} +test_route!(Feed, 52); impl ql_rpc::subscription::Subscription for Feed { type Error = core::convert::Infallible; @@ -86,10 +108,7 @@ impl ql_rpc::subscription::Subscription for Feed { struct Notice; -impl ql_rpc::Route for Notice { - const SERVICE: ServiceId = TEST_SERVICE; - const ROUTE: RouteId = RouteId::from_u32(521); -} +test_route!(Notice, 521); impl ql_rpc::notification::Notification for Notice { type Error = core::convert::Infallible; @@ -98,10 +117,7 @@ impl ql_rpc::notification::Notification for Notice { struct Download; -impl ql_rpc::Route for Download { - const SERVICE: ServiceId = TEST_SERVICE; - const ROUTE: RouteId = RouteId::from_u32(53); -} +test_route!(Download, 53); impl ql_rpc::progress::Progress for Download { type Error = core::convert::Infallible; @@ -112,10 +128,7 @@ impl ql_rpc::progress::Progress for Download { struct BlobDownload; -impl ql_rpc::Route for BlobDownload { - const SERVICE: ServiceId = TEST_SERVICE; - const ROUTE: RouteId = RouteId::from_u32(54); -} +test_route!(BlobDownload, 54); impl ql_rpc::download::Download for BlobDownload { type Error = core::convert::Infallible; @@ -126,10 +139,7 @@ impl ql_rpc::download::Download for BlobDownload { struct BlobUpload; -impl ql_rpc::Route for BlobUpload { - const SERVICE: ServiceId = TEST_SERVICE; - const ROUTE: RouteId = RouteId::from_u32(55); -} +test_route!(BlobUpload, 55); impl ql_rpc::upload::Upload for BlobUpload { type Error = core::convert::Infallible; @@ -140,10 +150,7 @@ impl ql_rpc::upload::Upload for BlobUpload { struct Chat; -impl ql_rpc::Route for Chat { - const SERVICE: ServiceId = TEST_SERVICE; - const ROUTE: RouteId = RouteId::from_u32(56); -} +test_route!(Chat, 56); impl ql_rpc::duplex::Duplex for Chat { type Error = core::convert::Infallible; @@ -177,10 +184,11 @@ async fn rpc_request() { let inbound_b = pair.take_inbound(Side::B); let seen = Arc::new(Mutex::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioSendSpawner>::builder_send(TokioSendSpawner) - .request::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::::builder_send( + TokioSendSpawner, + ) + .request::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); @@ -226,10 +234,11 @@ async fn rpc_notification() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .notification::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::::builder_local( + TokioLocalSpawner, + ) + .notification::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); @@ -280,10 +289,11 @@ async fn rpc_subscrption() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .subscription::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::::builder_local( + TokioLocalSpawner, + ) + .subscription::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); @@ -333,11 +343,12 @@ async fn rpc_router_enforces_max_request_bytes() { let mut pair = TestPair::new(default_runtime_config()); pair.connect_and_wait(Side::A).await; let inbound_b = pair.take_inbound(Side::B); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .max_request_bytes(4) - .request::() - .build(LimitedState); + let router = ql_rpc::Router::::builder_local( + TokioLocalSpawner, + ) + .max_request_bytes(4) + .request::() + .build(LimitedState); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); @@ -390,10 +401,11 @@ async fn rpc_progress() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .progress::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::::builder_local( + TokioLocalSpawner, + ) + .progress::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); @@ -453,10 +465,11 @@ async fn rpc_download() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .download::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::::builder_local( + TokioLocalSpawner, + ) + .download::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); @@ -530,10 +543,11 @@ async fn rpc_download_complete() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .download::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::::builder_local( + TokioLocalSpawner, + ) + .download::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); @@ -602,13 +616,14 @@ async fn rpc_upload() { let requests = Rc::new(RefCell::new(Vec::new())); let uploads = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .upload::() - .build(RouterState { - requests: requests.clone(), - uploads: uploads.clone(), - }); + let router = ql_rpc::Router::::builder_local( + TokioLocalSpawner, + ) + .upload::() + .build(RouterState { + requests: requests.clone(), + uploads: uploads.clone(), + }); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); @@ -678,10 +693,11 @@ async fn rpc_duplex() { let inbound_b = pair.take_inbound(Side::B); let seen = Rc::new(RefCell::new(Vec::new())); - let router = - ql_rpc::Router::<_, QlStream, TokioLocalSpawner>::builder_local(TokioLocalSpawner) - .duplex::() - .build(RouterState { seen: seen.clone() }); + let router = ql_rpc::Router::::builder_local( + TokioLocalSpawner, + ) + .duplex::() + .build(RouterState { seen: seen.clone() }); let responder = tokio::task::spawn_local(async move { let (info, stream) = inbound_b.recv().await.unwrap(); From c6a7dad6e3eaabd6c1ade9f54540af0f8d5ba5e4 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 6 Jul 2026 09:18:49 -0400 Subject: [PATCH 36/59] ql-wire: bundle metadata bytes --- ql-wire/src/identity.rs | 25 +++++++++++++++++++++++-- ql-wire/src/tests.rs | 2 ++ 2 files changed, 25 insertions(+), 2 deletions(-) diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index 90486375..f88fe2f0 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -15,6 +15,7 @@ pub struct PeerBundle { pub capabilities: u32, pub mlkem_public_key: MlKemPublicKey, pub name: QlName, + pub metadata: Box<[u8]>, } impl PeerBundle { @@ -25,7 +26,7 @@ impl PeerBundle { impl WireEncode for PeerBundle { fn encoded_len(&self) -> usize { - Self::FIXED_WIRE_SIZE + self.name.encoded_len() + Self::FIXED_WIRE_SIZE + self.name.encoded_len() + LenBytes(&*self.metadata).encoded_len() } fn encode(&self, out: &mut W) { @@ -34,6 +35,7 @@ impl WireEncode for PeerBundle { self.capabilities.encode(out); self.mlkem_public_key.encode(out); self.name.encode(out); + LenBytes(&*self.metadata).encode(out); } } @@ -45,6 +47,11 @@ impl codec::WireDecode for PeerBundle { capabilities: reader.decode()?, mlkem_public_key: reader.decode()?, name: reader.decode()?, + metadata: reader + .decode::>()? + .0 + .to_vec() + .into_boxed_slice(), }) } } @@ -56,6 +63,7 @@ pub struct QlIdentity { pub mlkem_public_key: MlKemPublicKey, pub capabilities: u32, pub name: QlName, + pub metadata: Box<[u8]>, } impl QlIdentity { @@ -76,9 +84,15 @@ impl QlIdentity { mlkem_public_key, capabilities: 0, name, + metadata: Box::default(), } } + pub fn with_metadata(mut self, metadata: impl Into>) -> Self { + self.metadata = metadata.into(); + self + } + pub fn bundle(&self) -> PeerBundle { PeerBundle { version: PeerBundle::VERSION, @@ -86,13 +100,14 @@ impl QlIdentity { capabilities: self.capabilities, mlkem_public_key: self.mlkem_public_key.clone(), name: self.name.clone(), + metadata: self.metadata.clone(), } } } impl WireEncode for QlIdentity { fn encoded_len(&self) -> usize { - Self::FIXED_WIRE_SIZE + self.name.encoded_len() + Self::FIXED_WIRE_SIZE + self.name.encoded_len() + LenBytes(&*self.metadata).encoded_len() } fn encode(&self, out: &mut W) { @@ -101,6 +116,7 @@ impl WireEncode for QlIdentity { self.mlkem_public_key.encode(out); self.capabilities.encode(out); self.name.encode(out); + LenBytes(&*self.metadata).encode(out); } } @@ -112,6 +128,11 @@ impl codec::WireDecode for QlIdentity { mlkem_public_key: reader.decode()?, capabilities: reader.decode()?, name: reader.decode()?, + metadata: reader + .decode::>()? + .0 + .to_vec() + .into_boxed_slice(), }) } } diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index 53e5bdcd..bc30f3c0 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -73,6 +73,7 @@ fn peer_bundle_round_trip() { let crypto = SoftwareCrypto; let mut identity = generate_identity(&crypto, "alice"); identity.capabilities = 1231; + identity.metadata = b"peer metadata".to_vec().into_boxed_slice(); let bundle = identity.bundle(); let encoded = bundle.encode_vec(); @@ -80,6 +81,7 @@ fn peer_bundle_round_trip() { assert_eq!(decoded, bundle); assert_eq!(&*decoded.name, "alice"); + assert_eq!(&*decoded.metadata, b"peer metadata"); } #[test] From 59f66d78394f8a21c0bc70c80a4dd2075db48604 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 6 Jul 2026 13:46:22 -0400 Subject: [PATCH 37/59] ql-common: remove service-id --- ql-common/src/lib.rs | 23 ++++------------------- ql-common/src/varint.rs | 3 ++- ql-runtime/src/tests/rpc.rs | 12 ++++++------ 3 files changed, 12 insertions(+), 26 deletions(-) diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index 373cc509..8514d6a2 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -65,25 +65,10 @@ impl QID { pub const SIZE: usize = 16; } -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -#[repr(transparent)] -pub struct ServiceId(pub [u8; 16]); - -impl ServiceId { - pub const SIZE: usize = size_of::(); -} - -impl std::fmt::Display for ServiceId { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - for byte in self.0 { - write!(f, "{byte:02x}")?; - } - Ok(()) - } -} - -varint_wrapper!(RouteId); -varint_wrapper!(StreamId); +varint_wrapper!( + /// Identifier for a stream within a QL session. + StreamId +); #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct StreamInfo { diff --git a/ql-common/src/varint.rs b/ql-common/src/varint.rs index a88d7765..09b3df01 100644 --- a/ql-common/src/varint.rs +++ b/ql-common/src/varint.rs @@ -186,7 +186,8 @@ impl std::error::Error for VarIntBoundsExceeded {} #[macro_export] macro_rules! varint_wrapper { - ($name:ident) => { + ($(#[$attr:meta])* $name:ident $(,)?) => { + $(#[$attr])* #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] #[repr(transparent)] pub struct $name(pub $crate::VarInt); diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index b1d85108..fe7a2ddc 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -8,7 +8,7 @@ use std::{ }; use bytes::{BufMut, Bytes}; -use ql_common::{ResetCode, ResetOrigin, RouteId, VarInt}; +use ql_common::{ResetCode, ResetOrigin, VarInt}; use ql_rpc::{ download::{DownloadHandlerLocal, DownloadStart}, duplex::{DuplexHandlerLocal, DuplexPeer}, @@ -24,20 +24,20 @@ use super::*; use crate::{QlStream, QlStreamError, StreamWriter}; #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] -struct TestRouteKey(RouteId); +struct TestRouteKey(VarInt); impl ql_rpc::RpcRouteKey for TestRouteKey { fn encoded_len(&self) -> usize { - self.0 .0.size() + self.0.size() } fn encode(&self, out: &mut W) { - self.0 .0.write_bytes(|bytes| out.put_slice(bytes)); + self.0.write_bytes(|bytes| out.put_slice(bytes)); } fn decode(bytes: &[u8]) -> Option { let (route_id, rest) = VarInt::decode_bytes(bytes)?; - rest.is_empty().then_some(Self(RouteId(route_id))) + rest.is_empty().then_some(Self(route_id)) } } @@ -47,7 +47,7 @@ macro_rules! test_route { type Key = TestRouteKey; fn key() -> Self::Key { - TestRouteKey(RouteId::from_u32($id)) + TestRouteKey(VarInt::from_u32($id)) } } }; From 9bc07459482f9ffee6127b5c2235394ac05c4bf5 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 6 Jul 2026 14:07:48 -0400 Subject: [PATCH 38/59] ql-common: varint conversions --- ql-common/src/varint.rs | 20 ++++++++++++++++++++ 1 file changed, 20 insertions(+) diff --git a/ql-common/src/varint.rs b/ql-common/src/varint.rs index 09b3df01..ae9800f8 100644 --- a/ql-common/src/varint.rs +++ b/ql-common/src/varint.rs @@ -222,12 +222,32 @@ macro_rules! varint_wrapper { } } + impl From<$name> for $crate::VarInt { + fn from(value: $name) -> Self { + value.0 + } + } + + impl From<$name> for u64 { + fn from(value: $name) -> Self { + value.0.into_inner() + } + } + impl From for $name { fn from(value: u32) -> Self { Self::from_u32(value) } } + impl TryFrom for $name { + type Error = $crate::VarIntBoundsExceeded; + + fn try_from(value: u64) -> Result { + Self::from_u64(value) + } + } + impl std::fmt::Display for $name { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write!(f, "{}", self.0) From 20ff66a82e833b79e5ee8e04910463f9ed9b33e2 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 8 Jul 2026 12:39:23 -0400 Subject: [PATCH 39/59] ql-rpc: simplify --- ql-rpc/src/framed_value.rs | 138 -------------------------- ql-rpc/src/lib.rs | 2 - ql-rpc/src/rpc/download/client.rs | 43 +++----- ql-rpc/src/rpc/download/server.rs | 4 +- ql-rpc/src/rpc/duplex/client.rs | 10 +- ql-rpc/src/rpc/notification/server.rs | 4 +- ql-rpc/src/rpc/parts.rs | 28 +++--- ql-rpc/src/rpc/progress/server.rs | 4 +- ql-rpc/src/rpc/request/server.rs | 4 +- ql-rpc/src/rpc/subscription/client.rs | 93 +---------------- ql-rpc/src/rpc/subscription/codec.rs | 47 --------- ql-rpc/src/rpc/subscription/mod.rs | 1 - ql-rpc/src/rpc/subscription/server.rs | 56 ++--------- ql-rpc/src/rpc/upload/server.rs | 10 +- ql-rpc/src/rpc/utils.rs | 69 ++++++------- ql-rpc/src/stream.rs | 2 +- ql-runtime/src/tests/rpc.rs | 6 +- 17 files changed, 104 insertions(+), 417 deletions(-) delete mode 100644 ql-rpc/src/framed_value.rs delete mode 100644 ql-rpc/src/rpc/subscription/codec.rs diff --git a/ql-rpc/src/framed_value.rs b/ql-rpc/src/framed_value.rs deleted file mode 100644 index bf96dc6f..00000000 --- a/ql-rpc/src/framed_value.rs +++ /dev/null @@ -1,138 +0,0 @@ -use std::marker::PhantomData; - -use bytes::Bytes; - -use crate::{chunk_queue::ChunkQueue, RpcCodec, RpcError}; - -/// reads one length-delimited rpc value from buffered byte chunks -pub struct FramedReader { - bytes: ChunkQueue, - marker: PhantomData T>, -} - -pub enum FramedReadStep { - NeedMore(FramedReader), - Value(T), -} - -pub enum FramedPrefixStep { - NeedMore(FramedReader), - Value { value: T, bytes: ChunkQueue }, -} - -impl Default for FramedReader { - fn default() -> Self { - Self { - bytes: ChunkQueue::default(), - marker: PhantomData, - } - } -} - -impl FramedReader { - pub fn push(mut self, chunk: Bytes) -> Self { - self.bytes.push(chunk); - self - } - - pub(crate) fn exceeds_total_len(&self, max_len: usize) -> Result { - let Some(len) = self.bytes.next_part_total_len()? else { - return Ok(self.bytes.remaining() > max_len); - }; - Ok(len > max_len) - } - - pub fn advance(self) -> Result, RpcError> { - match self.advance_prefix::()? { - FramedPrefixStep::NeedMore(next) => Ok(FramedReadStep::NeedMore(next)), - FramedPrefixStep::Value { value, bytes } => { - bytes.expect_empty().map_err(RpcError::Protocol)?; - Ok(FramedReadStep::Value(value)) - } - } - } - - pub fn advance_prefix(self) -> Result, RpcError> { - let mut this = self; - let Some(mut body) = this.bytes.try_take_part().map_err(RpcError::Protocol)? else { - return Ok(FramedPrefixStep::NeedMore(this)); - }; - - let value = T::decode_value(&mut body).map_err(RpcError::Codec)?; - drop(body); - Ok(FramedPrefixStep::Value { - value, - bytes: this.bytes, - }) - } -} - -#[cfg(test)] -mod tests { - use bytes::Bytes; - - use super::{FramedPrefixStep, FramedReadStep, FramedReader}; - use crate::codec::encode_value_part; - - #[test] - fn value_reader_round_trips_framed_values() { - let mut encoded = Vec::new(); - encode_value_part(&b"hello".to_vec(), &mut encoded); - - match FramedReader::>::default() - .push(Bytes::from(encoded)) - .advance::() - .unwrap() - { - FramedReadStep::Value(value) => assert_eq!(value, b"hello".to_vec()), - _ => unreachable!(), - } - } - - #[test] - fn value_reader_waits_for_complete_frame() { - let mut encoded = Vec::new(); - encode_value_part(&b"hello".to_vec(), &mut encoded); - let encoded = Bytes::from(encoded); - - let reader = match FramedReader::>::default() - .push(encoded.slice(..4)) - .advance::() - .unwrap() - { - FramedReadStep::NeedMore(next) => next, - _ => unreachable!(), - }; - - match reader - .push(encoded.slice(4..)) - .advance::() - .unwrap() - { - FramedReadStep::Value(value) => assert_eq!(value, b"hello".to_vec()), - _ => unreachable!(), - } - } - - #[test] - fn value_reader_returns_prefix_remainder() { - let mut encoded = Vec::new(); - encode_value_part(&b"hello".to_vec(), &mut encoded); - encoded.extend_from_slice(b"tail"); - - match FramedReader::>::default() - .push(Bytes::from(encoded)) - .advance_prefix::() - .unwrap() - { - FramedPrefixStep::Value { value, mut bytes } => { - assert_eq!(value, b"hello".to_vec()); - assert_eq!( - bytes.pop_front(usize::MAX), - Some(Bytes::from_static(b"tail")) - ); - } - _ => unreachable!(), - } - } -} diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs index 7dd45047..57fc4ba0 100644 --- a/ql-rpc/src/lib.rs +++ b/ql-rpc/src/lib.rs @@ -5,7 +5,6 @@ mod chunk_queue; mod codec; mod error; -mod framed_value; mod router; mod rpc; mod stream; @@ -13,7 +12,6 @@ mod stream; pub use chunk_queue::ChunkQueue; pub use codec::RpcCodec; pub use error::*; -use framed_value::*; pub use router::*; pub use rpc::*; pub use stream::*; diff --git a/ql-rpc/src/rpc/download/client.rs b/ql-rpc/src/rpc/download/client.rs index b759f10a..81bbd615 100644 --- a/ql-rpc/src/rpc/download/client.rs +++ b/ql-rpc/src/rpc/download/client.rs @@ -1,4 +1,4 @@ -use std::future::poll_fn; +use std::{future::poll_fn, marker::PhantomData}; use bytes::Bytes; use ql_common::ResetCode; @@ -6,8 +6,8 @@ use ql_common::ResetCode; use crate::{ download::Download, parts::{PartFrameReader, PartReadStep}, - rpc::{parts::FrameKind, write_eof_value}, - DropResetRead, FramedPrefixStep, FramedReader, RpcError, RpcRead, RpcStream, + rpc::{parts::FrameKind, read_framed_prefix, write_eof_value}, + DropResetRead, RpcError, RpcRead, RpcStream, }; pub async fn start( @@ -31,7 +31,7 @@ where R: RpcRead, { stream: DropResetRead, - reader: Option>, + marker: PhantomData M>, } pub struct DownloadPart<'a, M, R> @@ -60,37 +60,22 @@ where pub fn new(stream: R) -> Self { Self { stream: DropResetRead::new(stream), - reader: Some(FramedReader::default()), + marker: PhantomData, } } pub async fn start( mut self, ) -> Result<(M::ResponseHeader, DownloadReader), RpcError> { - loop { - let reader = self.reader.take().unwrap(); - let reader = match reader.advance_prefix() { - Ok(FramedPrefixStep::Value { value, bytes }) => { - return Ok(( - value, - DownloadReader { - stream: self.stream, - reader: PartFrameReader::::new(bytes), - }, - )); - } - Ok(FramedPrefixStep::NeedMore(next)) => next, - Err(error) => return Err(error), - }; - - match poll_fn(|cx| self.stream.poll_read(cx)).await { - Ok(Some(chunk)) => { - self.reader = Some(reader.push(chunk)); - } - Ok(None) => return Err(crate::Error::Truncated.into()), - Err(error) => return Err(RpcError::Transport(error)), - } - } + let (value, bytes) = + read_framed_prefix::(&mut self.stream, None).await?; + Ok(( + value, + DownloadReader { + stream: self.stream, + reader: PartFrameReader::::new(bytes), + }, + )) } pub fn reset(mut self, code: ResetCode) { diff --git a/ql-rpc/src/rpc/download/server.rs b/ql-rpc/src/rpc/download/server.rs index 84fa3c08..bfd95ae7 100644 --- a/ql-rpc/src/rpc/download/server.rs +++ b/ql-rpc/src/rpc/download/server.rs @@ -27,7 +27,9 @@ where download: DownloadStart, ); - fn handle_error(&self, _error: &RpcError) {} + fn handle_error(&self, error: &RpcError) { + let _ = error; + } } pub struct DownloadStart diff --git a/ql-rpc/src/rpc/duplex/client.rs b/ql-rpc/src/rpc/duplex/client.rs index 346c08cd..aa501e7f 100644 --- a/ql-rpc/src/rpc/duplex/client.rs +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -8,8 +8,8 @@ use bytes::Bytes; use ql_common::ResetCode; use crate::{ - codec, duplex::Duplex, write_bytes, ChunkQueue, DropResetRead, DropResetWrite, RpcCodec, - RpcError, RpcRead, RpcStream, RpcWrite, + codec, duplex::Duplex, finish_bytes, write_bytes, ChunkQueue, DropResetRead, DropResetWrite, + RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, }; pub fn start(stream: St) -> DuplexCall @@ -71,10 +71,16 @@ where write_bytes(writer, Bytes::from(encoded)).await } + /// queue a graceful write-side finish and return without waiting for transport errors pub fn finish(mut self) { self.writer.queue_finish(); } + /// queue a graceful write-side finish and wait until the transport reports it was sent + pub async fn finish_wait(mut self) -> Result<(), W::Error> { + finish_bytes(&mut self.writer).await + } + pub fn reset(mut self, code: ResetCode) { DropResetWrite::reset(&mut self.writer, code); } diff --git a/ql-rpc/src/rpc/notification/server.rs b/ql-rpc/src/rpc/notification/server.rs index ed39ca66..fc18efa5 100644 --- a/ql-rpc/src/rpc/notification/server.rs +++ b/ql-rpc/src/rpc/notification/server.rs @@ -15,7 +15,9 @@ where { async fn handle(self, context: Context, message: M::Payload); - fn handle_error(&self, _error: &RpcError) {} + fn handle_error(&self, error: &RpcError) { + let _ = error; + } } pub(crate) fn handle_notification( diff --git a/ql-rpc/src/rpc/parts.rs b/ql-rpc/src/rpc/parts.rs index f5032eb6..d206c9ca 100644 --- a/ql-rpc/src/rpc/parts.rs +++ b/ql-rpc/src/rpc/parts.rs @@ -111,19 +111,19 @@ impl PartFrameReader { } pub fn encode_part_header(part_header: &H, out: &mut (impl BufMut + AsMut<[u8]>)) { - codec::encode_tagged_value_part(FrameKind::PartHeader.tag(), part_header, out) + codec::encode_tagged_value_part(FrameKind::PartHeader.tag(), part_header, out); } pub fn encode_body_chunk(bytes: &Bytes, out: &mut (impl BufMut + AsMut<[u8]>)) { - codec::encode_tagged_value_part(FrameKind::BodyChunk.tag(), bytes, out) + codec::encode_tagged_value_part(FrameKind::BodyChunk.tag(), bytes, out); } pub fn encode_end_part(out: &mut (impl BufMut + AsMut<[u8]>)) { - encode_tagged_empty_part(FrameKind::EndPart, out) + encode_tagged_empty_part(FrameKind::EndPart, out); } pub fn encode_finish(out: &mut (impl BufMut + AsMut<[u8]>)) { - encode_tagged_empty_part(FrameKind::Finish, out) + encode_tagged_empty_part(FrameKind::Finish, out); } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -189,34 +189,34 @@ mod tests { assert_eq!(value, b"a.txt".to_vec()); } _ => unreachable!(), - }; + } match reader.advance::().unwrap() { PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"hel")), _ => unreachable!(), - }; + } match reader.advance::().unwrap() { PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"lo")), _ => unreachable!(), - }; + } match reader.advance::().unwrap() { PartReadStep::EndPart => {} _ => unreachable!(), - }; + } match reader.advance::().unwrap() { PartReadStep::PartHeader(value) => { assert_eq!(value, b"b.txt".to_vec()); } _ => unreachable!(), - }; + } match reader.advance::().unwrap() { PartReadStep::EndPart => {} _ => unreachable!(), - }; + } match reader.advance::().unwrap() { PartReadStep::Finish => {} @@ -235,7 +235,7 @@ mod tests { match reader.advance::().unwrap() { PartReadStep::NeedMore => {} _ => unreachable!(), - }; + } reader.push(encoded.slice(4..)); match reader.advance::().unwrap() { @@ -255,18 +255,18 @@ mod tests { match reader.advance::().unwrap() { PartReadStep::NeedMore => {} _ => unreachable!(), - }; + } reader.push(encoded.slice(9..11)); match reader.advance::().unwrap() { PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"he")), _ => unreachable!(), - }; + } reader.push(encoded.slice(11..)); match reader.advance::().unwrap() { PartReadStep::BodyBytes(bytes) => assert_eq!(bytes, Bytes::from_static(b"llo")), _ => unreachable!(), - }; + } } } diff --git a/ql-rpc/src/rpc/progress/server.rs b/ql-rpc/src/rpc/progress/server.rs index 42875152..476cb950 100644 --- a/ql-rpc/src/rpc/progress/server.rs +++ b/ql-rpc/src/rpc/progress/server.rs @@ -23,7 +23,9 @@ where responder: ProgressResponder, ); - fn handle_error(&self, _error: &RpcError) {} + fn handle_error(&self, error: &RpcError) { + let _ = error; + } } pub struct ProgressResponder diff --git a/ql-rpc/src/rpc/request/server.rs b/ql-rpc/src/rpc/request/server.rs index ba434c4d..c8677043 100644 --- a/ql-rpc/src/rpc/request/server.rs +++ b/ql-rpc/src/rpc/request/server.rs @@ -21,7 +21,9 @@ where responder: Response, ); - fn handle_error(&self, _error: &RpcError) {} + fn handle_error(&self, error: &RpcError) { + let _ = error; + } } pub struct Response diff --git a/ql-rpc/src/rpc/subscription/client.rs b/ql-rpc/src/rpc/subscription/client.rs index 9767211a..37d51334 100644 --- a/ql-rpc/src/rpc/subscription/client.rs +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -1,19 +1,9 @@ -use std::{ - future::poll_fn, - task::{Context, Poll}, -}; - -use ql_common::ResetCode; - use crate::{ - rpc::{ - subscription::codec::{ReadStep, ResponseReader}, - write_eof_value, - }, - subscription::Subscription, - DropResetRead, RpcError, RpcRead, RpcStream, + duplex::DuplexReceiver, rpc::write_eof_value, subscription::Subscription, RpcError, RpcStream, }; +pub type SubscriptionCall = DuplexReceiver<::Event, R>; + pub async fn start( stream: St, request: &M::Request, @@ -26,80 +16,5 @@ where write_eof_value(&mut writer, request) .await .map_err(RpcError::Transport)?; - Ok(SubscriptionCall::new(reader)) -} - -pub struct SubscriptionCall -where - M: Subscription, - R: RpcRead, -{ - stream: DropResetRead, - reader: ResponseReader, -} - -impl SubscriptionCall -where - M: Subscription, - R: RpcRead, -{ - pub fn new(stream: R) -> Self { - Self { - stream: DropResetRead::new(stream), - reader: ResponseReader::default(), - } - } - - pub async fn next_event(&mut self) -> Option>> { - poll_fn(|cx| self.poll_next_event(cx)).await - } - - pub fn poll_next_event( - &mut self, - cx: &mut Context<'_>, - ) -> Poll>>> { - if !self.stream.is_some() { - return Poll::Ready(None); - } - - loop { - match self.reader.advance() { - Ok(ReadStep::Item(value)) => return Poll::Ready(Some(Ok(value))), - Ok(ReadStep::NeedMore) => {} - Err(error) => { - self.stream.disarm(); - return Poll::Ready(Some(Err(error))); - } - } - - match self.stream.poll_read(cx) { - Poll::Ready(Ok(Some(chunk))) => { - self.reader.push(chunk); - } - Poll::Ready(Ok(None)) => { - if self.reader.is_empty() { - self.stream.disarm(); - return Poll::Ready(None); - } - self.stream.disarm(); - return Poll::Ready(Some(Err(crate::Error::Truncated.into()))); - } - Poll::Ready(Err(error)) => { - self.stream.disarm(); - return Poll::Ready(Some(Err(RpcError::Transport(error)))); - } - Poll::Pending => { - return Poll::Pending; - } - } - } - } - - pub fn reset(mut self, code: ResetCode) { - self.reset_inner(code); - } - - fn reset_inner(&mut self, code: ResetCode) { - DropResetRead::reset(&mut self.stream, code); - } + Ok(DuplexReceiver::new(reader)) } diff --git a/ql-rpc/src/rpc/subscription/codec.rs b/ql-rpc/src/rpc/subscription/codec.rs deleted file mode 100644 index 9563df0f..00000000 --- a/ql-rpc/src/rpc/subscription/codec.rs +++ /dev/null @@ -1,47 +0,0 @@ -use std::marker::PhantomData; - -use bytes::Bytes; - -use crate::{subscription::Subscription, ChunkQueue, RpcCodec, RpcError}; - -pub enum ReadStep { - NeedMore, - Item(M::Event), -} - -pub struct ResponseReader { - bytes: ChunkQueue, - marker: PhantomData M>, -} - -impl Default for ResponseReader { - fn default() -> Self { - Self { - bytes: ChunkQueue::default(), - marker: PhantomData, - } - } -} - -impl ResponseReader { - pub fn push(&mut self, chunk: Bytes) { - self.bytes.push(chunk); - } - - pub fn is_empty(&self) -> bool { - self.bytes.remaining() == 0 - } - - pub fn advance(&mut self) -> Result, RpcError> { - let Some(mut body) = self.bytes.try_take_part().map_err(RpcError::Protocol)? else { - return Ok(ReadStep::NeedMore); - }; - - let item = { - let item = M::Event::decode_value(&mut body).map_err(RpcError::Codec)?; - drop(body); - item - }; - Ok(ReadStep::Item(item)) - } -} diff --git a/ql-rpc/src/rpc/subscription/mod.rs b/ql-rpc/src/rpc/subscription/mod.rs index d1c4346f..657f96dc 100644 --- a/ql-rpc/src/rpc/subscription/mod.rs +++ b/ql-rpc/src/rpc/subscription/mod.rs @@ -2,7 +2,6 @@ use super::Route; use crate::RpcCodec; mod client; -pub(crate) mod codec; mod server; pub use self::{client::*, server::*}; diff --git a/ql-rpc/src/rpc/subscription/server.rs b/ql-rpc/src/rpc/subscription/server.rs index c0766067..97bcc8d5 100644 --- a/ql-rpc/src/rpc/subscription/server.rs +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -1,13 +1,12 @@ -use std::{future::Future, marker::PhantomData}; - -use bytes::Bytes; -use ql_common::ResetCode; +use std::future::Future; use crate::{ - codec, finish_bytes, rpc::read_eof_request, subscription::Subscription, write_bytes, Context, - DropResetWrite, RouterConfig, RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, + duplex::DuplexSender, rpc::read_eof_request, subscription::Subscription, Context, RouterConfig, + RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, }; +pub type SubscriptionResponder = DuplexSender; + #[trait_variant::make(SubscriptionHandler: Send)] pub trait SubscriptionHandlerLocal where @@ -18,46 +17,11 @@ where self, context: Context, message: M::Request, - responder: SubscriptionResponder, + responder: DuplexSender, ); - fn handle_error(&self, _error: &RpcError) {} -} - -pub struct SubscriptionResponder -where - W: RpcWrite, -{ - writer: DropResetWrite, - marker: PhantomData T>, -} - -impl SubscriptionResponder -where - T: RpcCodec, - W: RpcWrite, -{ - pub(crate) fn new(writer: W) -> Self { - Self { - writer: DropResetWrite::new(writer), - marker: PhantomData, - } - } - - pub async fn send(&mut self, event: T) -> Result<(), W::Error> { - let writer = &mut self.writer; - let mut encoded = Vec::new(); - codec::encode_value_part(&event, &mut encoded); - write_bytes(writer, Bytes::from(encoded)).await?; - Ok(()) - } - - pub async fn finish(mut self) -> Result<(), W::Error> { - finish_bytes(&mut self.writer).await - } - - pub fn reset(mut self, code: ResetCode) { - DropResetWrite::reset(&mut self.writer, code); + fn handle_error(&self, error: &RpcError) { + let _ = error; } } @@ -73,7 +37,7 @@ where Req: RpcCodec + 'static, Event: RpcCodec + 'static, St: RpcStream + 'static, - H: FnOnce(S, Context, Req, SubscriptionResponder) -> HF, + H: FnOnce(S, Context, Req, DuplexSender) -> HF, HF: Future, E: FnOnce(&S, &RpcError), { @@ -93,6 +57,6 @@ where } }; - handle(state, context, request, SubscriptionResponder::new(writer)).await; + handle(state, context, request, DuplexSender::new(writer)).await; } } diff --git a/ql-rpc/src/rpc/upload/server.rs b/ql-rpc/src/rpc/upload/server.rs index 32676154..d74145b4 100644 --- a/ql-rpc/src/rpc/upload/server.rs +++ b/ql-rpc/src/rpc/upload/server.rs @@ -7,7 +7,7 @@ use crate::{ request::Response, rpc::{ parts::{FrameKind, PartFrameReader, PartReadStep}, - read_framed_request_prefix, + read_framed_prefix, }, Context, DropResetRead, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, Upload, }; @@ -26,7 +26,9 @@ where responder: UploadResponder, ); - fn handle_error(&self, _error: &RpcError) {} + fn handle_error(&self, error: &RpcError) { + let _ = error; + } } pub struct UploadReader @@ -184,7 +186,9 @@ where async move { let (request, buffered) = - match read_framed_request_prefix::(&mut reader, config).await { + match read_framed_prefix::(&mut reader, Some(config.max_request_bytes)) + .await + { Ok(value) => value, Err(error) => { let code = error.reset_code(); diff --git a/ql-rpc/src/rpc/utils.rs b/ql-rpc/src/rpc/utils.rs index 5f210b3f..81208c00 100644 --- a/ql-rpc/src/rpc/utils.rs +++ b/ql-rpc/src/rpc/utils.rs @@ -1,8 +1,8 @@ use bytes::Bytes; use crate::{ - finish_bytes, read_bytes, write_bytes, ChunkQueue, Error, FramedPrefixStep, FramedReadStep, - FramedReader, RouterConfig, RpcCodec, RpcError, RpcRead, RpcWrite, + finish_bytes, read_bytes, write_bytes, ChunkQueue, Error, RouterConfig, RpcCodec, RpcError, + RpcRead, RpcWrite, }; pub async fn write_eof_value(writer: &mut W, value: &T) -> Result<(), W::Error> @@ -43,23 +43,8 @@ where T: RpcCodec, R: RpcRead, { - let mut value_reader = FramedReader::::default(); - let value = loop { - match value_reader.advance::() { - Ok(FramedReadStep::Value(value)) => break value, - Ok(FramedReadStep::NeedMore(next)) => value_reader = next, - Err(error) => return Err(error), - } - - match read_bytes(reader).await { - Ok(Some(chunk)) => { - value_reader = value_reader.push(chunk); - reject_oversized_frame(&value_reader, config)?; - } - Ok(None) => return Err(RpcError::Protocol(Error::Truncated)), - Err(error) => return Err(RpcError::Transport(error)), - } - }; + let (value, buffered) = read_framed_prefix(reader, Some(config.max_request_bytes)).await?; + buffered.expect_empty().map_err(RpcError::Protocol)?; match read_bytes(reader).await { Ok(None) => Ok(value), @@ -68,27 +53,27 @@ where } } -/// reads one length-delimited value and returns any bytes already buffered -pub async fn read_framed_request_prefix( +pub async fn read_framed_prefix( reader: &mut R, - config: RouterConfig, + max_len: Option, ) -> Result<(T, ChunkQueue), RpcError> where T: RpcCodec, R: RpcRead, { - let mut value_reader = FramedReader::::default(); + let mut bytes = ChunkQueue::default(); + loop { - match value_reader.advance_prefix::() { - Ok(FramedPrefixStep::Value { value, bytes }) => return Ok((value, bytes)), - Ok(FramedPrefixStep::NeedMore(next)) => value_reader = next, - Err(error) => return Err(error), + if let Some(value) = try_take_framed_value(&mut bytes)? { + return Ok((value, bytes)); } match read_bytes(reader).await { Ok(Some(chunk)) => { - value_reader = value_reader.push(chunk); - reject_oversized_frame(&value_reader, config)?; + bytes.push(chunk); + if let Some(max_len) = max_len { + reject_oversized_frame(&bytes, max_len).map_err(RpcError::Protocol)?; + } } Ok(None) => return Err(RpcError::Protocol(Error::Truncated)), Err(error) => return Err(RpcError::Transport(error)), @@ -130,18 +115,26 @@ where Ok(value) } -fn reject_oversized_frame( - value_reader: &FramedReader, - config: RouterConfig, -) -> Result<(), RpcError> +fn try_take_framed_value(bytes: &mut ChunkQueue) -> Result, RpcError> where T: RpcCodec, { - if value_reader - .exceeds_total_len(config.max_request_bytes) - .map_err(RpcError::Protocol)? - { - return Err(RpcError::Protocol(Error::LengthOverflow)); + let Some(mut body) = bytes.try_take_part().map_err(RpcError::Protocol)? else { + return Ok(None); + }; + + let value = T::decode_value(&mut body).map_err(RpcError::Codec)?; + Ok(Some(value)) +} + +fn reject_oversized_frame(bytes: &ChunkQueue, max_len: usize) -> Result<(), Error> { + let oversized = match bytes.next_part_total_len()? { + Some(len) => len > max_len, + None => bytes.remaining() > max_len, + }; + + if oversized { + return Err(Error::LengthOverflow); } Ok(()) } diff --git a/ql-rpc/src/stream.rs b/ql-rpc/src/stream.rs index 82e31994..d5da6951 100644 --- a/ql-rpc/src/stream.rs +++ b/ql-rpc/src/stream.rs @@ -167,7 +167,7 @@ mod drop { impl Drop for DropResetWrite { fn drop(&mut self) { - self.reset(ResetCode::DROPPED) + self.reset(ResetCode::DROPPED); } } } diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index fe7a2ddc..430a4395 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -277,9 +277,9 @@ async fn rpc_subscrption() { ) { let seen = self.seen.clone(); seen.borrow_mut().push(request); - let _ = response.send(b"one".to_vec()).await; - let _ = response.send(b"two".to_vec()).await; - let _ = response.finish().await; + let _ = response.send(&b"one".to_vec()).await; + let _ = response.send(&b"two".to_vec()).await; + let _ = response.finish_wait().await; } } From c3d79ba9530f7bb141cecee76a93bd12c87e2837 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 9 Jul 2026 19:55:21 -0400 Subject: [PATCH 40/59] ql-codec: impl --- Cargo.lock | 7 ++ Cargo.toml | 2 + ql-codec/Cargo.toml | 8 ++ ql-codec/src/buf_view.rs | 91 +++++++++++++++ ql-codec/src/codec.rs | 238 +++++++++++++++++++++++++++++++++++++++ ql-codec/src/error.rs | 29 +++++ ql-codec/src/lib.rs | 63 +++++++++++ ql-codec/src/reader.rs | 85 ++++++++++++++ ql-codec/src/slice.rs | 63 +++++++++++ ql-codec/src/varint.rs | 176 +++++++++++++++++++++++++++++ 10 files changed, 762 insertions(+) create mode 100644 ql-codec/Cargo.toml create mode 100644 ql-codec/src/buf_view.rs create mode 100644 ql-codec/src/codec.rs create mode 100644 ql-codec/src/error.rs create mode 100644 ql-codec/src/lib.rs create mode 100644 ql-codec/src/reader.rs create mode 100644 ql-codec/src/slice.rs create mode 100644 ql-codec/src/varint.rs diff --git a/Cargo.lock b/Cargo.lock index 00b789c1..22ceda2c 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1050,6 +1050,13 @@ dependencies = [ "syn", ] +[[package]] +name = "ql-codec" +version = "0.1.0" +dependencies = [ + "bytes", +] + [[package]] name = "ql-common" version = "0.1.0" diff --git a/Cargo.toml b/Cargo.toml index 7522bcd0..056fd430 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -3,6 +3,7 @@ resolver = "2" members = [ "backup-shard", "btp", + "ql-codec", "ql-common", "ql-fsm", "ql-rpc", @@ -23,6 +24,7 @@ rkyv = { version = "0.8" } # workspace crates backup-shard = { path = "backup-shard" } btp = { path = "btp" } +ql-codec = { path = "ql-codec" } ql-common = { path = "ql-common" } ql-fsm = { path = "ql-fsm" } ql-rpc = { path = "ql-rpc" } diff --git a/ql-codec/Cargo.toml b/ql-codec/Cargo.toml new file mode 100644 index 00000000..051af9d5 --- /dev/null +++ b/ql-codec/Cargo.toml @@ -0,0 +1,8 @@ +[package] +name = "ql-codec" +version = "0.1.0" +edition = "2021" +description = "QuantumLink binary codec primitives" + +[dependencies] +bytes = { workspace = true } diff --git a/ql-codec/src/buf_view.rs b/ql-codec/src/buf_view.rs new file mode 100644 index 00000000..5f112ec3 --- /dev/null +++ b/ql-codec/src/buf_view.rs @@ -0,0 +1,91 @@ +use bytes::{Buf, Bytes}; + +/// A byte container that can expose a replayable [`Buf`] view for encoding. +pub trait BufView { + type Buf<'a>: Buf + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_>; + + fn is_empty(&self) -> bool { + self.buf().remaining() == 0 + } +} + +impl BufView for &T { + type Buf<'a> + = T::Buf<'a> + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + (*self).buf() + } +} + +impl BufView for &mut T { + type Buf<'a> + = T::Buf<'a> + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + (**self).buf() + } +} + +impl BufView for [u8] { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self + } +} + +impl BufView for [u8; N] { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self.as_slice() + } +} + +impl BufView for Vec { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self.as_slice() + } +} + +impl BufView for Box<[u8]> { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self.as_ref() + } +} + +impl BufView for Bytes { + type Buf<'a> + = &'a [u8] + where + Self: 'a; + + fn buf(&self) -> Self::Buf<'_> { + self.as_ref() + } +} diff --git a/ql-codec/src/codec.rs b/ql-codec/src/codec.rs new file mode 100644 index 00000000..891cde89 --- /dev/null +++ b/ql-codec/src/codec.rs @@ -0,0 +1,238 @@ +use bytes::{Buf, BufMut}; + +use crate::{varint, BufView, ByteSlice, Decode, Encode, Error, Reader}; + +impl Decode for [u8; N] { + fn decode(reader: &mut Reader) -> Result { + let bytes = reader.take_n(N)?; + let mut out = [0u8; N]; + out.copy_from_slice(&bytes); + Ok(out) + } +} + +impl Encode for [u8; N] { + fn encoded_len(&self) -> usize { + N + } + + fn encode(&self, out: &mut W) { + out.put_slice(self); + } +} + +impl Decode for Box<[u8; N]> { + fn decode(reader: &mut Reader) -> Result { + let bytes = reader.take_n(N)?; + let mut out = Self::new_uninit(); + let src = bytes.as_ptr(); + let dst = out.as_mut_ptr().cast::(); + // SAFETY: `take_bytes(N)` guarantees the source has exactly `N` bytes. + unsafe { + std::ptr::copy_nonoverlapping(src, dst, N); + Ok(out.assume_init()) + } + } +} + +impl Encode for Box<[u8; N]> { + fn encoded_len(&self) -> usize { + N + } + + fn encode(&self, out: &mut W) { + out.put_slice(self.as_ref()); + } +} + +macro_rules! impl_codec { + (byte_encode: $($ty:ty),* $(,)?) => { + $( + impl Encode for $ty { + fn encoded_len(&self) -> usize { + encoded_len_bytes(self) + } + fn encode(&self, out: &mut W) { + encode_bytes(self, out); + } + } + )* + }; + (owned_byte_decode: $($ty:ty),* $(,)?) => { + $( + impl Decode for $ty { + fn decode(reader: &mut Reader) -> Result { + Ok(<$ty>::from(&*reader.take_len_prefixed()?)) + } + } + )* + }; + (fixed_integer: $($ty:ty),* $(,)?) => { + $( + impl Decode for $ty { + fn decode(reader: &mut Reader) -> Result { + Ok(Self::from_le_bytes(reader.decode()?)) + } + } + impl Encode for $ty { + fn encoded_len(&self) -> usize { + size_of::() + } + fn encode(&self, out: &mut W) { + out.put_slice(&self.to_le_bytes()); + } + } + )* + }; +} + +impl_codec!(byte_encode: [u8], Vec, Box<[u8]>, bytes::Bytes); + +impl<'a> Decode<&'a [u8]> for &'a [u8] { + fn decode(reader: &mut Reader<&'a [u8]>) -> Result { + reader.take_len_prefixed() + } +} + +impl<'a> Decode<&'a mut [u8]> for &'a mut [u8] { + fn decode(reader: &mut Reader<&'a mut [u8]>) -> Result { + reader.take_len_prefixed() + } +} + +impl Decode for bytes::Bytes { + fn decode(reader: &mut Reader) -> Result { + reader.take_len_prefixed() + } +} + +impl_codec!(owned_byte_decode: Vec, Box<[u8]>); + +impl Decode for u8 { + fn decode(reader: &mut Reader) -> Result { + reader.take_u8() + } +} + +impl Encode for u8 { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u8(*self); + } +} + +impl_codec!(fixed_integer: u16, u32, u64); + +impl Decode for bool { + fn decode(reader: &mut Reader) -> Result { + match reader.decode::()? { + 0 => Ok(false), + 1 => Ok(true), + _ => Err(Error::InvalidDiscriminant), + } + } +} + +impl Encode for bool { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u8(u8::from(*self)); + } +} + +impl Encode for Option { + fn encoded_len(&self) -> usize { + 1 + self.as_ref().map_or(0, Encode::encoded_len) + } + + fn encode(&self, out: &mut W) { + match self { + None => out.put_u8(0), + Some(inner) => { + out.put_u8(1); + inner.encode(out); + } + } + } +} + +impl> Decode for Option { + fn decode(reader: &mut Reader) -> Result { + match reader.decode::()? { + 0 => Ok(None), + 1 => Ok(Some(reader.decode::()?)), + _ => Err(Error::InvalidDiscriminant), + } + } +} + +pub fn encoded_len_bytes(bytes: &B) -> usize { + let len = bytes.buf().remaining(); + varint::encoded_len(len) + len +} + +pub fn encode_bytes(bytes: &B, out: &mut W) +where + B: BufView + ?Sized, + W: BufMut + ?Sized, +{ + let mut bytes = bytes.buf(); + varint::encode(bytes.remaining(), out); + while bytes.has_remaining() { + let chunk = bytes.chunk(); + out.put_slice(chunk); + bytes.advance(chunk.len()); + } +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use super::*; + + #[test] + fn integers_are_little_endian() { + assert_eq!(0x1234u16.encode_vec(), [0x34, 0x12]); + assert_eq!(0x1234_5678u32.encode_vec(), [0x78, 0x56, 0x34, 0x12]); + assert_eq!( + 0x0123_4567_89ab_cdefu64.encode_vec(), + [0xef, 0xcd, 0xab, 0x89, 0x67, 0x45, 0x23, 0x01] + ); + assert_eq!(u16::decode_bytes([0x34, 0x12].as_slice()), Ok(0x1234)); + assert_eq!( + u32::decode_bytes([0x78, 0x56, 0x34, 0x12].as_slice()), + Ok(0x1234_5678) + ); + } + + #[test] + fn byte_containers_use_varint_prefix() { + let bytes = [1u8, 2, 3]; + + let encoded = bytes[..].encode_vec(); + assert_eq!(encoded, [3, 1, 2, 3]); + assert_eq!(bytes.to_vec().encode_vec(), encoded); + assert_eq!(Box::<[u8]>::from(bytes).encode_vec(), encoded); + assert_eq!(Bytes::copy_from_slice(&bytes).encode_vec(), encoded); + + let decoded = <&[u8]>::decode_bytes(encoded.as_slice()).unwrap(); + assert_eq!(decoded, [1, 2, 3]); + assert_eq!( + Vec::::decode_bytes(encoded.as_slice()).unwrap(), + [1, 2, 3] + ); + assert_eq!( + Box::<[u8]>::decode_bytes(encoded.as_slice()) + .unwrap() + .as_ref(), + [1, 2, 3] + ); + } +} diff --git a/ql-codec/src/error.rs b/ql-codec/src/error.rs new file mode 100644 index 00000000..861c3c06 --- /dev/null +++ b/ql-codec/src/error.rs @@ -0,0 +1,29 @@ +use core::fmt; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Error { + InvalidData, + UnexpectedEof, + InvalidDiscriminant, + InvalidRange, + LengthOverflow, + InvalidVarint, + InvalidUtf8, +} + +impl fmt::Display for Error { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let message = match self { + Self::InvalidData => "invalid data", + Self::UnexpectedEof => "unexpected end of input", + Self::InvalidDiscriminant => "invalid discriminant", + Self::InvalidRange => "invalid range", + Self::LengthOverflow => "length overflow", + Self::InvalidVarint => "invalid varint", + Self::InvalidUtf8 => "invalid utf-8", + }; + f.write_str(message) + } +} + +impl std::error::Error for Error {} diff --git a/ql-codec/src/lib.rs b/ql-codec/src/lib.rs new file mode 100644 index 00000000..d51ca811 --- /dev/null +++ b/ql-codec/src/lib.rs @@ -0,0 +1,63 @@ +//! Small binary codec primitives shared by QuantumLink crates. + +mod buf_view; +mod codec; +mod error; +mod reader; +mod slice; +pub mod varint; + +pub use buf_view::BufView; +pub use codec::{encode_bytes, encoded_len_bytes}; +pub use error::Error; +pub use reader::Reader; +pub use slice::ByteSlice; + +pub trait Encode { + fn encoded_len(&self) -> usize; + + fn encode(&self, out: &mut W); + + fn encode_vec(&self) -> Vec { + let len = self.encoded_len(); + let mut out = Vec::with_capacity(len); + self.encode(&mut out); + assert_eq!(out.len(), len); + out + } +} + +pub trait Decode: Sized { + fn decode(reader: &mut Reader) -> Result; + + fn decode_bytes(bytes: B) -> Result { + let mut reader = Reader::new(bytes); + Self::decode(&mut reader) + } +} + +#[macro_export] +macro_rules! varint_wrapper { + ($name:ty, $inner:ty) => { + impl $name { + pub const MAX_ENCODED_LEN: usize = + <$inner as ql_codec::varint::VarInt>::MAX_ENCODED_LEN; + } + + impl ql_codec::Encode for $name { + fn encoded_len(&self) -> usize { + ql_codec::varint::encoded_len::<$inner>(self.0) + } + + fn encode(&self, out: &mut W) { + ql_codec::varint::encode::<$inner, W>(self.0, out); + } + } + + impl ql_codec::Decode for $name { + fn decode(reader: &mut ql_codec::Reader) -> Result { + Ok(Self(reader.decode_varint::<$inner>()?)) + } + } + }; +} diff --git a/ql-codec/src/reader.rs b/ql-codec/src/reader.rs new file mode 100644 index 00000000..2c15b6ce --- /dev/null +++ b/ql-codec/src/reader.rs @@ -0,0 +1,85 @@ +use core::mem::ManuallyDrop; + +use crate::{varint, ByteSlice, Decode, Error}; + +#[derive(Clone)] +pub struct Reader { + remaining: ManuallyDrop, +} + +impl Reader { + #[inline] + pub fn new(bytes: B) -> Self { + Self { + remaining: ManuallyDrop::new(bytes), + } + } + + #[inline] + pub fn is_empty(&self) -> bool { + self.remaining.is_empty() + } + + #[inline] + pub fn remaining_len(&self) -> usize { + self.remaining.len() + } + + pub fn take_n(&mut self, len: usize) -> Result { + if len > self.remaining.len() { + return Err(Error::UnexpectedEof); + } + // SAFETY: checked above + let (head, tail) = unsafe { + let remaining = ManuallyDrop::take(&mut self.remaining); + remaining.split_at_unchecked(len) + }; + self.remaining = ManuallyDrop::new(tail); + Ok(head) + } + + #[inline] + pub fn take_u8(&mut self) -> Result { + self.remaining.take_u8().ok_or(Error::UnexpectedEof) + } + + pub fn take_all(&mut self) -> B { + // SAFETY: 0 is always a valid split point + let (empty, rest) = unsafe { + let remaining = ManuallyDrop::take(&mut self.remaining); + remaining.split_at_unchecked(0) + }; + self.remaining = ManuallyDrop::new(empty); + rest + } + + pub fn take_len_prefixed(&mut self) -> Result { + let len = self.decode_varint::()?; + self.take_n(len) + } + + #[inline] + pub fn decode(&mut self) -> Result + where + T: Decode, + { + T::decode(self) + } + + #[inline] + pub fn decode_varint(&mut self) -> Result + where + T: varint::VarInt, + { + varint::decode(self) + } +} + +impl Drop for Reader { + fn drop(&mut self) { + // SAFETY: `remaining` is initialized except during `take_bytes` + unsafe { + ManuallyDrop::drop(&mut self.remaining); + } + } +} diff --git a/ql-codec/src/slice.rs b/ql-codec/src/slice.rs new file mode 100644 index 00000000..afe8a495 --- /dev/null +++ b/ql-codec/src/slice.rs @@ -0,0 +1,63 @@ +use core::{mem, ops::Deref}; + +use bytes::{Buf, Bytes}; + +/// A byte slice owner used by the codec reader +/// +/// # Safety +/// +/// `split_at_unchecked` must return byte slices matching `self[..mid]` and +/// `self[mid..]` when `mid <= self.len()` +pub unsafe trait ByteSlice: Deref + Sized { + /// splits the current byte view at `mid` without checking bounds + /// mid can be 0 + /// + /// # Safety + /// + /// `mid` must not exceed the slice length. + unsafe fn split_at_unchecked(self, mid: usize) -> (Self, Self); + + fn take_u8(&mut self) -> Option; +} + +unsafe impl ByteSlice for &[u8] { + #[inline] + unsafe fn split_at_unchecked(self, mid: usize) -> (Self, Self) { + <[u8]>::split_at_unchecked(self, mid) + } + + #[inline] + fn take_u8(&mut self) -> Option { + let (&byte, remaining) = self.split_first()?; + *self = remaining; + Some(byte) + } +} + +unsafe impl ByteSlice for &mut [u8] { + #[inline] + unsafe fn split_at_unchecked(self, mid: usize) -> (Self, Self) { + <[u8]>::split_at_mut_unchecked(self, mid) + } + + #[inline] + fn take_u8(&mut self) -> Option { + let (byte, remaining) = mem::take(self).split_first_mut()?; + let byte = *byte; + *self = remaining; + Some(byte) + } +} + +unsafe impl ByteSlice for Bytes { + #[inline] + unsafe fn split_at_unchecked(mut self, mid: usize) -> (Self, Self) { + let head = self.split_to(mid); + (head, self) + } + + #[inline] + fn take_u8(&mut self) -> Option { + Buf::try_get_u8(self).ok() + } +} diff --git a/ql-codec/src/varint.rs b/ql-codec/src/varint.rs new file mode 100644 index 00000000..1f241506 --- /dev/null +++ b/ql-codec/src/varint.rs @@ -0,0 +1,176 @@ +use bytes::BufMut; + +use crate::{ByteSlice, Error, Reader}; + +pub trait VarInt: Copy { + const MAX_ENCODED_LEN: usize; + + fn from_u8(value: u8) -> Self; + fn low_7_bits(self) -> u8; + fn shr_7(self) -> Self; + fn needs_more(self) -> bool; + fn checked_add_payload(self, payload: u8, shift: usize) -> Option; +} + +pub fn encoded_len(mut value: T) -> usize { + let mut len = 1; + while value.needs_more() { + value = value.shr_7(); + len += 1; + } + len +} + +pub fn encode(mut value: T, out: &mut W) +where + T: VarInt, + W: BufMut + ?Sized, +{ + while value.needs_more() { + out.put_u8(value.low_7_bits() | 0x80); + value = value.shr_7(); + } + out.put_u8(value.low_7_bits()); +} + +pub fn decode(reader: &mut Reader) -> Result +where + T: VarInt, + B: ByteSlice, +{ + let mut value = T::from_u8(0); + + for index in 0..T::MAX_ENCODED_LEN { + let byte = reader.decode::()?; + let payload = byte & 0x7f; + + value = value + .checked_add_payload(payload, index * 7) + .ok_or(Error::InvalidVarint)?; + + if byte & 0x80 == 0 { + if index > 0 && payload == 0 { + return Err(Error::InvalidVarint); + } + return Ok(value); + } + } + + Err(Error::InvalidVarint) +} + +macro_rules! impl_varint { + ($($ty:ty),* $(,)?) => { + $( + impl VarInt for $ty { + const MAX_ENCODED_LEN: usize = (size_of::() * 8).div_ceil(7); + + #[inline] + fn from_u8(value: u8) -> Self { + Self::from(value) + } + #[inline] + #[allow(clippy::cast_possible_truncation)] + fn low_7_bits(self) -> u8 { + (self & 0x7f) as u8 + } + #[inline] + fn shr_7(self) -> Self { + self >> 7 + } + #[inline] + fn needs_more(self) -> bool { + self >= 0x80 + } + #[inline] + fn checked_add_payload(self, payload: u8, shift: usize) -> Option { + let bit_width = size_of::() * 8; + if shift >= bit_width { + return (payload == 0).then_some(self); + } + let max_payload = Self::MAX >> shift; + let payload = Self::from(payload); + if payload > max_payload { + return None; + } + Some(self | (payload << shift)) + } + } + )* + }; +} + +impl_varint!(u8, u16, u32, u64, usize); + +#[cfg(test)] +mod tests { + use std::fmt::Debug; + + use super::*; + + fn assert_decodes(value: T) + where + T: VarInt + Debug + PartialEq, + { + let mut out = Vec::new(); + encode(value, &mut out); + assert_eq!(out.len(), encoded_len(value)); + let mut reader = Reader::new(out.as_slice()); + assert_eq!(decode(&mut reader), Ok(value)); + assert!(reader.is_empty()); + } + + fn assert_error(bytes: &[u8], error: Error) + where + T: VarInt + Debug + PartialEq, + { + let mut reader = Reader::new(bytes); + assert_eq!(decode::(&mut reader), Err(error)); + } + + fn test_type(max: T) + where + T: VarInt + Debug + PartialEq + TryFrom, + { + for value in [0, 1, 127].map(T::from_u8) { + assert_decodes(value); + } + for shift in (7..size_of::() * 8).step_by(7) { + let boundary = 1u64 << shift; + if let (Ok(before), Ok(after)) = (T::try_from(boundary - 1), T::try_from(boundary)) { + assert_decodes(before); + assert_decodes(after); + } + } + assert_decodes(max); + + let mut encoded = Vec::new(); + encode(T::from_u8(127), &mut encoded); + encoded.push(0x55); + let mut reader = Reader::new(encoded.as_slice()); + assert_eq!(decode(&mut reader), Ok(T::from_u8(127))); + assert_eq!(reader.take_all(), &[0x55]); + + assert_error::(&[0x80, 0x00], Error::InvalidVarint); + assert_error::(&[0x80], Error::UnexpectedEof); + assert_error::(&vec![0xff; T::MAX_ENCODED_LEN], Error::InvalidVarint); + + let bit_width = size_of::() * 8; + let mut overflow = vec![0x80; bit_width / 7]; + overflow.push(1 << (bit_width % 7)); + assert_error::(&overflow, Error::InvalidVarint); + } + + macro_rules! test_types { + ($($ty:ident),* $(,)?) => { + $( + #[test] + fn $ty() { + test_type::<$ty>(<$ty>::MAX); + } + )* + }; + } + + test_types!(u8, u16, u32, u64, usize); +} From 6505b918feac40a97d6aa90beb2bd623fe0e57d3 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 9 Jul 2026 14:02:19 -0400 Subject: [PATCH 41/59] ql-wire: port to codec --- Cargo.lock | 5 + ql-common/Cargo.toml | 2 + ql-common/src/lib.rs | 39 ++- ql-common/src/varint.rs | 257 ------------------- ql-wire/Cargo.toml | 1 + ql-wire/src/bytes.rs | 123 ---------- ql-wire/src/codec.rs | 285 ---------------------- ql-wire/src/encrypted/ack.rs | 188 +++++++------- ql-wire/src/encrypted/builder.rs | 8 +- ql-wire/src/encrypted/close.rs | 32 +-- ql-wire/src/encrypted/mod.rs | 32 ++- ql-wire/src/encrypted/stream_data.rs | 46 ++-- ql-wire/src/encrypted/stream_reset.rs | 40 ++- ql-wire/src/encrypted/stream_window.rs | 17 +- ql-wire/src/encrypted_message.rs | 14 +- ql-wire/src/error.rs | 44 +++- ql-wire/src/handshake/ik.rs | 93 +++---- ql-wire/src/handshake/kk.rs | 87 +++---- ql-wire/src/handshake/meta.rs | 14 +- ql-wire/src/handshake/mod.rs | 82 +++---- ql-wire/src/handshake/transport_params.rs | 8 +- ql-wire/src/handshake/xx.rs | 135 ++++------ ql-wire/src/header.rs | 29 ++- ql-wire/src/identity.rs | 68 +++--- ql-wire/src/lib.rs | 5 - ql-wire/src/macros.rs | 46 +--- ql-wire/src/pq.rs | 20 +- ql-wire/src/qid.rs | 2 - ql-wire/src/record.rs | 53 ++-- ql-wire/src/tests.rs | 115 ++++----- ql-wire/src/varint.rs | 23 -- 31 files changed, 554 insertions(+), 1359 deletions(-) delete mode 100644 ql-common/src/varint.rs delete mode 100644 ql-wire/src/bytes.rs delete mode 100644 ql-wire/src/codec.rs delete mode 100644 ql-wire/src/varint.rs diff --git a/Cargo.lock b/Cargo.lock index 22ceda2c..ec0a4a36 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1060,6 +1060,10 @@ dependencies = [ [[package]] name = "ql-common" version = "0.1.0" +dependencies = [ + "bytes", + "ql-codec", +] [[package]] name = "ql-fsm" @@ -1109,6 +1113,7 @@ dependencies = [ "getrandom 0.2.16", "libcrux-aesgcm", "libcrux-ml-kem", + "ql-codec", "ql-common", "sha2", ] diff --git a/ql-common/Cargo.toml b/ql-common/Cargo.toml index 4a44b3b4..be206ecc 100644 --- a/ql-common/Cargo.toml +++ b/ql-common/Cargo.toml @@ -6,3 +6,5 @@ description = "QuantumLink shared primitive types" license = "Proprietary" [dependencies] +bytes = { workspace = true } +ql-codec = { workspace = true } diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index 8514d6a2..53246339 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -1,11 +1,10 @@ //! Shared QuantumLink primitive types. -mod varint; -pub use varint::*; +use ql_codec::{ByteSlice, Decode, Encode, Error, Reader}; #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] -pub struct ResetCode(pub u16); +pub struct ResetCode(pub u64); impl ResetCode { /// operation was explicitly cancelled @@ -48,6 +47,8 @@ impl std::fmt::Display for ResetCode { } } +ql_codec::varint_wrapper!(ResetCode, u64); + /// origin of a stream reset: either we triggered it locally or the peer sent it. #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub enum ResetOrigin { @@ -65,10 +66,34 @@ impl QID { pub const SIZE: usize = 16; } -varint_wrapper!( - /// Identifier for a stream within a QL session. - StreamId -); +impl Encode for QID { + fn encoded_len(&self) -> usize { + Self::SIZE + } + + fn encode(&self, out: &mut W) { + self.0.encode(out); + } +} + +impl Decode for QID { + fn decode(reader: &mut Reader) -> Result { + Ok(Self(reader.decode()?)) + } +} + +/// Identifier for a stream within a QL session. +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct StreamId(pub u64); + +impl std::fmt::Display for StreamId { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } +} + +ql_codec::varint_wrapper!(StreamId, u64); #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct StreamInfo { diff --git a/ql-common/src/varint.rs b/ql-common/src/varint.rs deleted file mode 100644 index ae9800f8..00000000 --- a/ql-common/src/varint.rs +++ /dev/null @@ -1,257 +0,0 @@ -use core::fmt; - -/// An integer less than 2^62 encoded with QUIC variable-length integer rules. -#[derive(Default, Copy, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)] -pub struct VarInt(u64); - -impl VarInt { - /// The largest representable value. - pub const MAX: Self = Self((1u64 << 62) - 1); - /// The largest encoded value length. - pub const MAX_SIZE: usize = 8; - pub const MIN_SIZE: usize = 1; - - /// Construct a `VarInt` infallibly from a `u32`. - pub const fn from_u32(x: u32) -> Self { - Self(x as u64) - } - - /// Construct a `VarInt` from a `u64`. - pub const fn from_u64(x: u64) -> Result { - if x < (1u64 << 62) { - Ok(Self(x)) - } else { - Err(VarIntBoundsExceeded) - } - } - - /// Create a `VarInt` without checking the bounds. - /// - /// # Safety - /// - /// `x` must be less than 2^62. - pub const unsafe fn from_u64_unchecked(x: u64) -> Self { - Self(x) - } - - /// Extract the inner integer value. - pub const fn into_inner(self) -> u64 { - self.0 - } - - /// Return the number of bytes required to encode this value. - pub const fn size(self) -> usize { - let x = self.0; - if x < (1u64 << 6) { - 1 - } else if x < (1u64 << 14) { - 2 - } else if x < (1u64 << 30) { - 4 - } else { - 8 - } - } - - /// Return the encoded length from the first encoded byte. - pub const fn encoded_len_from_first_byte(first: u8) -> usize { - 1usize << (first >> 6) - } - - /// Encode this value by writing its encoded bytes to `write`. - #[allow(clippy::cast_possible_truncation)] - pub fn write_bytes(self, mut write: impl FnMut(&[u8])) { - let value = self.0; - match self.size() { - 1 => write(&[value as u8]), - 2 => write(&((value as u16) | 0x4000).to_be_bytes()), - 4 => write(&((value as u32) | 0x8000_0000).to_be_bytes()), - 8 => write(&(value | 0xC000_0000_0000_0000).to_be_bytes()), - _ => unreachable!(), - } - } - - /// decode a value from the start of `bytes`, returning the value and remaining bytes - pub fn decode_bytes(bytes: &[u8]) -> Option<(Self, &[u8])> { - let first = *bytes.first()?; - let len = Self::encoded_len_from_first_byte(first); - Some(( - Self::decode_with_first_byte(first, bytes.get(1..len)?)?, - &bytes[len..], - )) - } - - /// decode a value after the first encoded byte has already been consumed - pub fn decode_with_first_byte(first: u8, tail: &[u8]) -> Option { - let len = Self::encoded_len_from_first_byte(first); - if tail.len() != len - 1 { - return None; - } - let value = match len { - 1 => u64::from(first & 0x3f), - 2 => u64::from(u16::from_be_bytes([first & 0x3f, tail[0]])), - 4 => { - let bytes = [first & 0x3f, tail[0], tail[1], tail[2]]; - u64::from(u32::from_be_bytes(bytes)) - } - 8 => { - let bytes = [ - first & 0x3f, - tail[0], - tail[1], - tail[2], - tail[3], - tail[4], - tail[5], - tail[6], - ]; - u64::from_be_bytes(bytes) - } - _ => unreachable!(), - }; - Self::from_u64(value).ok() - } -} - -impl From for u64 { - fn from(value: VarInt) -> Self { - value.0 - } -} - -impl From for VarInt { - fn from(value: u8) -> Self { - Self(value.into()) - } -} - -impl From for VarInt { - fn from(value: u16) -> Self { - Self(value.into()) - } -} - -impl From for VarInt { - fn from(value: u32) -> Self { - Self(value.into()) - } -} - -impl TryFrom for VarInt { - type Error = VarIntBoundsExceeded; - - fn try_from(value: u64) -> Result { - Self::from_u64(value) - } -} - -impl TryFrom for VarInt { - type Error = VarIntBoundsExceeded; - - fn try_from(value: u128) -> Result { - Self::from_u64(value.try_into().map_err(|_| VarIntBoundsExceeded)?) - } -} - -impl TryFrom for VarInt { - type Error = VarIntBoundsExceeded; - - fn try_from(value: usize) -> Result { - Self::from_u64(value as u64) - } -} - -impl fmt::Debug for VarInt { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.fmt(f) - } -} - -impl fmt::Display for VarInt { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.fmt(f) - } -} - -#[derive(Debug, Copy, Clone, Eq, PartialEq)] -pub struct VarIntBoundsExceeded; - -impl fmt::Display for VarIntBoundsExceeded { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.write_str("value too large for varint encoding") - } -} - -impl std::error::Error for VarIntBoundsExceeded {} - -#[macro_export] -macro_rules! varint_wrapper { - ($(#[$attr:meta])* $name:ident $(,)?) => { - $(#[$attr])* - #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] - #[repr(transparent)] - pub struct $name(pub $crate::VarInt); - - impl $name { - pub const MAX_ENCODED_LEN: usize = $crate::VarInt::MAX_SIZE; - - pub const fn from_u32(value: u32) -> Self { - Self($crate::VarInt::from_u32(value)) - } - - pub const fn from_u64(value: u64) -> Result { - match $crate::VarInt::from_u64(value) { - Ok(v) => Ok(Self(v)), - Err(e) => Err(e), - } - } - - /// Create this wrapper without checking the bounds. - /// - /// # Safety - /// - /// `value` must be less than 2^62. - pub const unsafe fn from_u64_unchecked(value: u64) -> Self { - Self(unsafe { $crate::VarInt::from_u64_unchecked(value) }) - } - } - - impl From<$crate::VarInt> for $name { - fn from(value: $crate::VarInt) -> Self { - Self(value) - } - } - - impl From<$name> for $crate::VarInt { - fn from(value: $name) -> Self { - value.0 - } - } - - impl From<$name> for u64 { - fn from(value: $name) -> Self { - value.0.into_inner() - } - } - - impl From for $name { - fn from(value: u32) -> Self { - Self::from_u32(value) - } - } - - impl TryFrom for $name { - type Error = $crate::VarIntBoundsExceeded; - - fn try_from(value: u64) -> Result { - Self::from_u64(value) - } - } - - impl std::fmt::Display for $name { - fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - write!(f, "{}", self.0) - } - } - }; -} diff --git a/ql-wire/Cargo.toml b/ql-wire/Cargo.toml index 9ba6e3d1..f18f8a22 100644 --- a/ql-wire/Cargo.toml +++ b/ql-wire/Cargo.toml @@ -16,6 +16,7 @@ test-utils = [ [dependencies] bytes = { workspace = true } getrandom = { workspace = true, optional = true } +ql-codec = { workspace = true } ql-common = { workspace = true } libcrux-aesgcm = { version = "0.0.7", optional = true } libcrux-ml-kem = { version = "0.0.7", optional = true } diff --git a/ql-wire/src/bytes.rs b/ql-wire/src/bytes.rs deleted file mode 100644 index 09cd1202..00000000 --- a/ql-wire/src/bytes.rs +++ /dev/null @@ -1,123 +0,0 @@ -use core::ops::Deref; - -use bytes::{Buf, Bytes}; - -/// A mutable or immutable byte slice owner used by the wire parser. -pub trait ByteSlice: Deref + Sized { - /// Splits the current byte view at `mid`. - /// - /// Returns `Err(self)` when `mid` is out of bounds. - fn split_at(self, mid: usize) -> Result<(Self, Self), Self>; -} - -impl ByteSlice for &[u8] { - #[inline] - fn split_at(self, mid: usize) -> Result<(Self, Self), Self> { - if mid <= self.len() { - Ok(<[u8]>::split_at(self, mid)) - } else { - Err(self) - } - } -} - -impl ByteSlice for &mut [u8] { - #[inline] - fn split_at(self, mid: usize) -> Result<(Self, Self), Self> { - if mid <= self.len() { - Ok(<[u8]>::split_at_mut(self, mid)) - } else { - Err(self) - } - } -} - -impl ByteSlice for Bytes { - #[inline] - fn split_at(self, mid: usize) -> Result<(Self, Self), Self> { - if mid <= self.len() { - Ok((self.slice(..mid), self.slice(mid..))) - } else { - Err(self) - } - } -} - -/// A byte container that can expose a replayable [`Buf`] view for encoding. -pub trait BufView { - type Buf<'a>: Buf - where - Self: 'a; - - fn buf(&self) -> Self::Buf<'_>; - - fn is_empty(&self) -> bool { - self.buf().remaining() == 0 - } -} - -impl BufView for &T { - type Buf<'a> - = T::Buf<'a> - where - Self: 'a; - - fn buf(&self) -> Self::Buf<'_> { - (*self).buf() - } -} - -impl BufView for &mut T { - type Buf<'a> - = T::Buf<'a> - where - Self: 'a; - - fn buf(&self) -> Self::Buf<'_> { - (**self).buf() - } -} - -impl BufView for [u8] { - type Buf<'a> - = &'a [u8] - where - Self: 'a; - - fn buf(&self) -> Self::Buf<'_> { - self - } -} - -impl BufView for [u8; N] { - type Buf<'a> - = &'a [u8] - where - Self: 'a; - - fn buf(&self) -> Self::Buf<'_> { - self.as_slice() - } -} - -impl BufView for Vec { - type Buf<'a> - = &'a [u8] - where - Self: 'a; - - fn buf(&self) -> Self::Buf<'_> { - self.as_slice() - } -} - -impl BufView for Bytes { - type Buf<'a> - = &'a [u8] - where - Self: 'a; - - fn buf(&self) -> Self::Buf<'_> { - self.as_ref() - } -} diff --git a/ql-wire/src/codec.rs b/ql-wire/src/codec.rs deleted file mode 100644 index 2bc60b3e..00000000 --- a/ql-wire/src/codec.rs +++ /dev/null @@ -1,285 +0,0 @@ -use bytes::{Buf, BufMut}; -use ql_common::VarInt; - -use crate::{BufView, ByteSlice, WireError}; - -pub trait WireEncode { - fn encoded_len(&self) -> usize; - - fn encode(&self, out: &mut W); - - fn encode_vec(&self) -> Vec { - let mut out = Vec::with_capacity(self.encoded_len()); - self.encode(&mut out); - debug_assert_eq!(out.len(), self.encoded_len()); - out - } -} - -pub trait WireDecode: Sized { - fn decode(reader: &mut Reader) -> Result; - - fn decode_bytes(bytes: B) -> Result { - let mut reader = Reader::new(bytes); - Self::decode(&mut reader) - } - - fn decode_exact(bytes: B) -> Result { - let mut reader = Reader::new(bytes); - let value = Self::decode(&mut reader)?; - if reader.is_empty() { - Ok(value) - } else { - Err(WireError::InvalidPayload) - } - } -} - -impl WireDecode for [u8; N] { - fn decode(reader: &mut Reader) -> Result { - let bytes = reader.take_bytes(N)?; - let mut out = [0u8; N]; - out.copy_from_slice(&bytes); - Ok(out) - } -} - -impl WireEncode for [u8; N] { - fn encoded_len(&self) -> usize { - N - } - - fn encode(&self, out: &mut W) { - out.put_slice(self); - } -} - -impl WireDecode for Box<[u8; N]> { - fn decode(reader: &mut Reader) -> Result { - let bytes = reader.take_bytes(N)?; - let mut out = Self::new_uninit(); - let src = bytes.as_ptr(); - let dst = out.as_mut_ptr().cast::(); - // SAFETY: `take_bytes(N)` guarantees the source has exactly `N` bytes. - unsafe { - std::ptr::copy_nonoverlapping(src, dst, N); - Ok(out.assume_init()) - } - } -} - -impl WireEncode for Box<[u8; N]> { - fn encoded_len(&self) -> usize { - N - } - - fn encode(&self, out: &mut W) { - out.put_slice(self.as_ref()); - } -} - -impl WireEncode for [u8] { - fn encoded_len(&self) -> usize { - self.len() - } - - fn encode(&self, out: &mut W) { - out.put_slice(self); - } -} - -impl WireDecode for u8 { - fn decode(reader: &mut Reader) -> Result { - Ok(reader.take_bytes(1)?[0]) - } -} - -impl WireEncode for u8 { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u8(*self); - } -} - -impl WireDecode for u16 { - fn decode(reader: &mut Reader) -> Result { - Ok(Self::from_be_bytes(reader.decode()?)) - } -} - -impl WireEncode for u16 { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u16(*self); - } -} - -impl WireDecode for u32 { - fn decode(reader: &mut Reader) -> Result { - Ok(Self::from_be_bytes(reader.decode()?)) - } -} - -impl WireEncode for u32 { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u32(*self); - } -} - -impl WireDecode for u64 { - fn decode(reader: &mut Reader) -> Result { - Ok(Self::from_be_bytes(reader.decode()?)) - } -} - -impl WireEncode for u64 { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u64(*self); - } -} - -impl WireDecode for bool { - fn decode(reader: &mut Reader) -> Result { - match reader.decode::()? { - 0 => Ok(false), - 1 => Ok(true), - _ => Err(WireError::InvalidPayload), - } - } -} - -impl WireEncode for bool { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u8(u8::from(*self)); - } -} - -impl WireEncode for Option { - fn encoded_len(&self) -> usize { - 1 + self.as_ref().map_or(0, WireEncode::encoded_len) - } - - fn encode(&self, out: &mut W) { - match self { - None => out.put_u8(0), - Some(inner) => { - out.put_u8(1); - inner.encode(out); - } - } - } -} - -impl> WireDecode for Option { - fn decode(reader: &mut Reader) -> Result { - match reader.decode::()? { - 0 => Ok(None), - 1 => Ok(Some(reader.decode::()?)), - _ => Err(WireError::InvalidPayload), - } - } -} - -#[derive(Clone)] -pub struct Reader { - remaining: Option, -} - -impl Reader { - pub fn new(bytes: B) -> Self { - Self { - remaining: Some(bytes), - } - } - - pub fn is_empty(&self) -> bool { - self.remaining.as_ref().unwrap().is_empty() - } - - pub fn remaining_len(&self) -> usize { - self.remaining.as_ref().unwrap().len() - } - - pub fn take_bytes(&mut self, len: usize) -> Result { - let remaining = self.remaining.take().unwrap(); - match remaining.split_at(len) { - Ok((head, tail)) => { - self.remaining = Some(tail); - Ok(head) - } - Err(remaining) => { - self.remaining = Some(remaining); - Err(WireError::InvalidPayload) - } - } - } - - pub fn take_rest(&mut self) -> B { - self.take_bytes(self.remaining_len()).unwrap() - } - - #[inline] - pub fn decode(&mut self) -> Result - where - T: WireDecode, - { - T::decode(self) - } -} - -/// bytes encoded with a VarInt len prefix -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct LenBytes(pub B); - -impl WireEncode for LenBytes -where - B: BufView, -{ - fn encoded_len(&self) -> usize { - let len = self.0.buf().remaining(); - let len_prefix = VarInt::try_from(len).unwrap(); - len_prefix.encoded_len() + len - } - - fn encode(&self, out: &mut W) { - let mut bytes = self.0.buf(); - let len = bytes.remaining(); - let len_prefix = VarInt::try_from(len).unwrap(); - len_prefix.encode(out); - while bytes.has_remaining() { - let chunk = bytes.chunk(); - out.put_slice(chunk); - bytes.advance(chunk.len()); - } - } -} - -impl WireDecode for LenBytes -where - B: ByteSlice, -{ - fn decode(reader: &mut Reader) -> Result { - let len = reader.decode::()?.into_inner(); - let len = usize::try_from(len).map_err(|_| WireError::InvalidPayload)?; - let bytes = reader.take_bytes(len)?; - Ok(Self(bytes)) - } -} diff --git a/ql-wire/src/encrypted/ack.rs b/ql-wire/src/encrypted/ack.rs index 099fe42f..14855804 100644 --- a/ql-wire/src/encrypted/ack.rs +++ b/ql-wire/src/encrypted/ack.rs @@ -1,20 +1,20 @@ use std::{fmt, ops::RangeInclusive}; -use ql_common::VarInt; +use ql_codec::{ByteSlice, Encode, Error}; -use crate::{codec, ByteSlice, RecordSeq, WireEncode, WireError}; +use crate::RecordSeq; #[derive(Debug, Clone, PartialEq, Eq)] pub struct RecordAck { largest_acked: RecordSeq, - first_range_len: VarInt, + first_range_len: u64, blocks: Box<[RecordAckBlock]>, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RecordAckBlock { - pub gap: VarInt, - pub range_len: VarInt, + pub gap: u64, + pub range_len: u64, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -48,7 +48,7 @@ impl RecordAck { pub fn ranges(&self) -> RecordAckRangeIter<'_> { RecordAckRangeIter { - largest_acked: self.largest_acked.0.into_inner(), + largest_acked: self.largest_acked.0, first_range_len: Some(self.first_range_len), previous_start: None, blocks: self.blocks.iter(), @@ -56,20 +56,18 @@ impl RecordAck { } pub fn contains(&self, seq: u64) -> bool { - let Ok(seq) = RecordSeq::from_u64(seq) else { - return false; - }; + let seq = RecordSeq(seq); self.ranges().any(|range| range.contains(&seq)) } fn block_count_len(block_count: usize) -> usize { - VarInt::try_from(block_count).unwrap().encoded_len() + ql_codec::varint::encoded_len(block_count) } } impl RecordAckBlock { fn encoded_len(&self) -> usize { - self.gap.encoded_len() + self.range_len.encoded_len() + ql_codec::varint::encoded_len(self.gap) + ql_codec::varint::encoded_len(self.range_len) } } @@ -91,7 +89,7 @@ impl std::error::Error for RecordAckRangeError {} pub struct RecordAckRangeIter<'a> { largest_acked: u64, - first_range_len: Option, + first_range_len: Option, previous_start: Option, blocks: std::slice::Iter<'a, RecordAckBlock>, } @@ -102,9 +100,9 @@ impl Iterator for RecordAckRangeIter<'_> { fn next(&mut self) -> Option { if let Some(first_range_len) = self.first_range_len.take() { let end = self.largest_acked; - let start = end - first_range_len.into_inner(); + let start = end - first_range_len; self.previous_start = Some(start); - return Some(RecordSeq::from_u64(start).unwrap()..=RecordSeq::from_u64(end).unwrap()); + return Some(RecordSeq(start)..=RecordSeq(end)); } let block = self.blocks.next()?; @@ -112,18 +110,18 @@ impl Iterator for RecordAckRangeIter<'_> { .previous_start .expect("first ack range is always yielded"); // gap is encoded as missing_count - 1, so decoding steps back by gap + 2. - let end = previous_start - block.gap.into_inner() - 2; - let start = end - block.range_len.into_inner(); + let end = previous_start - block.gap - 2; + let start = end - block.range_len; self.previous_start = Some(start); - Some(RecordSeq::from_u64(start).unwrap()..=RecordSeq::from_u64(end).unwrap()) + Some(RecordSeq(start)..=RecordSeq(end)) } } -impl WireEncode for RecordAck { +impl Encode for RecordAck { fn encoded_len(&self) -> usize { self.largest_acked.encoded_len() + Self::block_count_len(self.blocks.len()) - + self.first_range_len.encoded_len() + + ql_codec::varint::encoded_len(self.first_range_len) + self .blocks .iter() @@ -133,26 +131,25 @@ impl WireEncode for RecordAck { fn encode(&self, out: &mut W) { self.largest_acked.encode(out); - VarInt::try_from(self.blocks.len()).unwrap().encode(out); - self.first_range_len.encode(out); + ql_codec::varint::encode(self.blocks.len(), out); + ql_codec::varint::encode(self.first_range_len, out); for block in &self.blocks { - block.gap.encode(out); - block.range_len.encode(out); + ql_codec::varint::encode(block.gap, out); + ql_codec::varint::encode(block.range_len, out); } } } -impl codec::WireDecode for RecordAck { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for RecordAck { + fn decode(reader: &mut ql_codec::Reader) -> Result { let largest_acked = reader.decode()?; - let block_count = usize::try_from(reader.decode::()?.into_inner()) - .map_err(|_| WireError::InvalidPayload)?; - let first_range_len = reader.decode::()?; + let block_count = reader.decode_varint::()?; + let first_range_len = reader.decode_varint()?; let mut blocks = Vec::with_capacity(block_count); for _ in 0..block_count { blocks.push(RecordAckBlock { - gap: reader.decode::()?, - range_len: reader.decode::()?, + gap: reader.decode_varint()?, + range_len: reader.decode_varint()?, }); } @@ -167,23 +164,16 @@ impl codec::WireDecode for RecordAck { let mut previous_start = ack .largest_acked .0 - .into_inner() - .checked_sub(ack.first_range_len.into_inner()) - .ok_or(WireError::InvalidPayload)?; + .checked_sub(ack.first_range_len) + .ok_or(Error::InvalidRange)?; for block in &ack.blocks { let end = previous_start - .checked_sub( - block - .gap - .into_inner() - .checked_add(2) - .ok_or(WireError::InvalidPayload)?, - ) - .ok_or(WireError::InvalidPayload)?; + .checked_sub(block.gap.checked_add(2).ok_or(Error::InvalidRange)?) + .ok_or(Error::InvalidRange)?; previous_start = end - .checked_sub(block.range_len.into_inner()) - .ok_or(WireError::InvalidPayload)?; + .checked_sub(block.range_len) + .ok_or(Error::InvalidRange)?; } } Ok(ack) @@ -193,7 +183,7 @@ impl codec::WireDecode for RecordAck { #[derive(Debug, Clone, Default, PartialEq, Eq)] pub struct RecordAckBuilder { largest_acked: Option, - first_range_len: Option, + first_range_len: Option, blocks: Vec, previous_start: Option, wire_len: usize, @@ -209,13 +199,13 @@ impl RecordAckBuilder { range: RangeInclusive, max_wire_size: usize, ) -> Result { - let start = range.start().0.into_inner(); - let end = range.end().0.into_inner(); + let start = range.start().0; + let end = range.end().0; if start > end { return Err(RecordAckRangeError::InvertedRange); } - let range_len = VarInt::from_u64(end - start).unwrap(); + let range_len = end - start; if let Some(previous_start) = self.previous_start { if end.saturating_add(1) >= previous_start { return Err(RecordAckRangeError::NotCanonical); @@ -225,10 +215,7 @@ impl RecordAckBuilder { .checked_sub(end) .and_then(|delta| delta.checked_sub(2)) .expect("canonical ack ranges stay separated by at least one sequence"); - let block = RecordAckBlock { - gap: VarInt::from_u64(gap).unwrap(), - range_len, - }; + let block = RecordAckBlock { gap, range_len }; let current_block_count_len = RecordAck::block_count_len(self.blocks.len()); let next_block_count_len = RecordAck::block_count_len(self.blocks.len() + 1); let next_wire_len = self.wire_len @@ -244,9 +231,10 @@ impl RecordAckBuilder { return Ok(true); } - let largest_acked = RecordSeq::from_u64(end).unwrap(); - let wire_len = - largest_acked.encoded_len() + RecordAck::block_count_len(0) + range_len.encoded_len(); + let largest_acked = RecordSeq(end); + let wire_len = largest_acked.encoded_len() + + RecordAck::block_count_len(0) + + ql_codec::varint::encoded_len(range_len); if wire_len > max_wire_size { return Ok(false); } @@ -272,53 +260,45 @@ impl RecordAckBuilder { } #[cfg(test)] mod tests { - use std::ops::RangeInclusive; - - use ql_common::VarInt; + use ql_codec::{Decode, Encode, Error}; use super::{RecordAck, RecordAckBlock, RecordAckBuilder, RecordAckRangeError}; - use crate::{RecordSeq, WireDecode, WireEncode, WireError}; - - fn seq(value: u64) -> RecordSeq { - RecordSeq::from_u64(value).unwrap() - } - - fn ack_range(start: u64, end: u64) -> RangeInclusive { - seq(start)..=seq(end) - } - - fn varint(value: u64) -> VarInt { - VarInt::from_u64(value).unwrap() - } + use crate::RecordSeq; #[test] fn encode_decode_round_trip() { - let ack = - RecordAck::from_ranges([ack_range(95, 100), ack_range(90, 92), ack_range(80, 80)]) - .unwrap(); + let ack = RecordAck::from_ranges([ + RecordSeq(95)..=RecordSeq(100), + RecordSeq(90)..=RecordSeq(92), + RecordSeq(80)..=RecordSeq(80), + ]) + .unwrap(); let encoded = ack.encode_vec(); - assert_eq!(RecordAck::decode_exact(encoded.as_slice()).unwrap(), ack); + assert_eq!(RecordAck::decode_bytes(encoded.as_slice()).unwrap(), ack); } #[test] fn wire_fields_match_gap_encoding() { - let ack = - RecordAck::from_ranges([ack_range(95, 100), ack_range(90, 92), ack_range(80, 80)]) - .unwrap(); + let ack = RecordAck::from_ranges([ + RecordSeq(95)..=RecordSeq(100), + RecordSeq(90)..=RecordSeq(92), + RecordSeq(80)..=RecordSeq(80), + ]) + .unwrap(); - assert_eq!(ack.largest_acked, seq(100)); - assert_eq!(ack.first_range_len, varint(5)); + assert_eq!(ack.largest_acked, RecordSeq(100)); + assert_eq!(ack.first_range_len, 5); assert_eq!( ack.blocks.as_ref(), &[ RecordAckBlock { - gap: varint(1), - range_len: varint(2), + gap: 1, + range_len: 2, }, RecordAckBlock { - gap: varint(8), - range_len: varint(0), + gap: 8, + range_len: 0, } ] ); @@ -326,14 +306,14 @@ mod tests { #[test] fn builder_stops_when_budget_is_exhausted() { - let first_only = RecordAck::from_ranges([ack_range(95, 100)]).unwrap(); + let first_only = RecordAck::from_ranges([RecordSeq(95)..=RecordSeq(100)]).unwrap(); let mut builder = RecordAckBuilder::new(); assert!(builder - .try_push_range(ack_range(95, 100), first_only.encoded_len()) + .try_push_range(RecordSeq(95)..=RecordSeq(100), first_only.encoded_len()) .unwrap()); assert!(!builder - .try_push_range(ack_range(90, 92), first_only.encoded_len()) + .try_push_range(RecordSeq(90)..=RecordSeq(92), first_only.encoded_len()) .unwrap()); assert_eq!(builder.build().unwrap(), first_only); } @@ -342,10 +322,10 @@ mod tests { fn builder_rejects_non_canonical_ranges() { let mut builder = RecordAckBuilder::new(); assert!(builder - .try_push_range(ack_range(95, 100), usize::MAX) + .try_push_range(RecordSeq(95)..=RecordSeq(100), usize::MAX) .unwrap()); assert_eq!( - builder.try_push_range(ack_range(90, 95), usize::MAX), + builder.try_push_range(RecordSeq(90)..=RecordSeq(95), usize::MAX), Err(RecordAckRangeError::NotCanonical) ); } @@ -353,7 +333,10 @@ mod tests { #[test] fn rejects_unsorted_ranges() { assert_eq!( - RecordAck::from_ranges([ack_range(90, 92), ack_range(95, 100)]), + RecordAck::from_ranges([ + RecordSeq(90)..=RecordSeq(92), + RecordSeq(95)..=RecordSeq(100) + ]), Err(RecordAckRangeError::NotCanonical) ); } @@ -361,7 +344,7 @@ mod tests { #[test] fn rejects_touching_ranges() { assert_eq!( - RecordAck::from_ranges([ack_range(10, 12), ack_range(7, 9)]), + RecordAck::from_ranges([RecordSeq(10)..=RecordSeq(12), RecordSeq(7)..=RecordSeq(9)]), Err(RecordAckRangeError::NotCanonical) ); } @@ -369,7 +352,7 @@ mod tests { #[test] fn rejects_overlapping_ranges() { assert_eq!( - RecordAck::from_ranges([ack_range(10, 12), ack_range(8, 11)]), + RecordAck::from_ranges([RecordSeq(10)..=RecordSeq(12), RecordSeq(8)..=RecordSeq(11)]), Err(RecordAckRangeError::NotCanonical) ); } @@ -377,9 +360,9 @@ mod tests { #[test] fn contains_matches_range_membership() { let ack = RecordAck::from_ranges([ - ack_range(150, 163), - ack_range(105, 110), - ack_range(100, 100), + RecordSeq(150)..=RecordSeq(163), + RecordSeq(105)..=RecordSeq(110), + RecordSeq(100)..=RecordSeq(100), ]) .unwrap(); @@ -399,7 +382,7 @@ mod tests { #[test] fn inverted_range_is_rejected() { assert_eq!( - RecordAck::from_ranges([ack_range(5, 4)]), + RecordAck::from_ranges([RecordSeq(5)..=RecordSeq(4)]), Err(RecordAckRangeError::InvertedRange) ); } @@ -415,24 +398,21 @@ mod tests { ]; assert_eq!( - RecordAck::decode_exact(encoded.as_slice()), - Err(WireError::InvalidPayload) + RecordAck::decode_bytes(encoded.as_slice()), + Err(Error::InvalidRange) ); } #[test] fn decode_rejects_truncated_payload() { - assert_eq!( - RecordAck::decode_exact(&[][..]), - Err(WireError::InvalidPayload) - ); + assert_eq!(RecordAck::decode_bytes(&[][..]), Err(Error::UnexpectedEof)); - let encoded = RecordAck::from_ranges([ack_range(42, 42)]) + let encoded = RecordAck::from_ranges([RecordSeq(42)..=RecordSeq(42)]) .unwrap() .encode_vec(); assert_eq!( - RecordAck::decode_exact(&encoded[..encoded.len() - 1]), - Err(WireError::InvalidPayload) + RecordAck::decode_bytes(&encoded[..encoded.len() - 1]), + Err(Error::UnexpectedEof) ); } } diff --git a/ql-wire/src/encrypted/builder.rs b/ql-wire/src/encrypted/builder.rs index f1a87f3e..210b113f 100644 --- a/ql-wire/src/encrypted/builder.rs +++ b/ql-wire/src/encrypted/builder.rs @@ -1,9 +1,9 @@ use bytes::BufMut; +use ql_codec::{BufView, Encode}; use super::{RecordAck, SessionClose, SessionFrame, StreamData, StreamReset, StreamWindow}; use crate::{ - BufView, Nonce, QlCrypto, RecordHeader, RecordSeq, RecordType, RouteHeader, SessionHeader, - SessionKey, WireEncode, + Nonce, QlCrypto, RecordHeader, RecordSeq, RecordType, RouteHeader, SessionHeader, SessionKey, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -109,7 +109,7 @@ impl SessionRecordBuilder { let record_header = RecordHeader::new(route, RecordType::Session); let header = SessionHeader { seq: self.seq }; let aad = header.aad(route); - let nonce = Nonce::from_counter(self.seq.0.into_inner()); + let nonce = Nonce::from_counter(self.seq.0); let auth = crypto.aes256_gcm_encrypt( session_key, &nonce, @@ -140,7 +140,7 @@ impl SessionRecordBuilder { self.push_wire_size(1, |out| out.put_u8(kind as u8)) } - fn push_frame_payload( + fn push_frame_payload( &mut self, kind: super::SessionFrameKind, payload: &T, diff --git a/ql-wire/src/encrypted/close.rs b/ql-wire/src/encrypted/close.rs index c9e0d237..77528c90 100644 --- a/ql-wire/src/encrypted/close.rs +++ b/ql-wire/src/encrypted/close.rs @@ -1,4 +1,4 @@ -use crate::{codec, codec::Reader, ByteSlice, WireEncode, WireError}; +use ql_codec::{ByteSlice, Encode}; /// closes the whole session immediately with a reset code. #[derive(Debug, Clone, PartialEq, Eq)] @@ -6,13 +6,9 @@ pub struct SessionClose { pub code: SessionCloseCode, } -impl SessionClose { - pub const WIRE_SIZE: usize = size_of::(); -} - #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] -pub struct SessionCloseCode(pub u16); +pub struct SessionCloseCode(pub u64); impl SessionCloseCode { pub const CANCELLED: Self = Self(0); @@ -20,33 +16,19 @@ impl SessionCloseCode { pub const TIMEOUT: Self = Self(2); } -impl WireEncode for SessionCloseCode { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl codec::WireDecode for SessionCloseCode { - fn decode(reader: &mut Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} +ql_codec::varint_wrapper!(SessionCloseCode, u64); -impl codec::WireDecode for SessionClose { - fn decode(reader: &mut Reader) -> Result { +impl ql_codec::Decode for SessionClose { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { code: reader.decode()?, }) } } -impl WireEncode for SessionClose { +impl Encode for SessionClose { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + self.code.encoded_len() } fn encode(&self, out: &mut W) { diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index 2d900d22..48f59b81 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -1,8 +1,8 @@ +use ql_codec::{BufView, ByteSlice, Decode, Encode, Reader}; use ql_common::StreamId; use crate::{ - codec, encrypted_message::EncryptedMessage, BufView, ByteSlice, Nonce, QlCrypto, Reader, - SessionHeader, SessionKey, WireDecode, WireEncode, WireError, + encrypted_message::EncryptedMessage, Error, Nonce, QlCrypto, SessionHeader, SessionKey, }; mod ack; @@ -19,8 +19,6 @@ pub use stream_data::*; pub use stream_reset::*; pub use stream_window::*; -varint_wrapper_codec!(StreamId); - #[derive(Debug, Clone, PartialEq, Eq)] pub enum SessionFrame { // todo: do we need ping as explicit frame? @@ -33,8 +31,8 @@ pub enum SessionFrame { Close(SessionClose), } -impl WireDecode for SessionFrame { - fn decode(reader: &mut Reader) -> Result { +impl Decode for SessionFrame { + fn decode(reader: &mut Reader) -> Result { let kind = reader.decode::()?; let frame = match kind { SessionFrameKind::Ping => Self::Ping, @@ -77,7 +75,7 @@ impl SessionFrame { } } -impl WireEncode for SessionFrame { +impl Encode for SessionFrame { fn encoded_len(&self) -> usize { 1 + match self { Self::Ping | Self::Unpair => 0, @@ -115,7 +113,7 @@ pub enum SessionFrameKind { } impl TryFrom for SessionFrameKind { - type Error = WireError; + type Error = ql_codec::Error; fn try_from(value: u8) -> Result { match value { @@ -126,13 +124,13 @@ impl TryFrom for SessionFrameKind { 5 => Ok(Self::StreamReset), 6 => Ok(Self::Close), 7 => Ok(Self::Unpair), - _ => Err(WireError::InvalidPayload), + _ => Err(ql_codec::Error::InvalidDiscriminant), } } } -impl codec::WireDecode for SessionFrameKind { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for SessionFrameKind { + fn decode(reader: &mut ql_codec::Reader) -> Result { reader.decode::()?.try_into() } } @@ -143,7 +141,7 @@ pub fn parse_session_frames(bytes: B) -> SessionFrameIter { } } -pub fn decode_session_frames(bytes: &[u8]) -> Result>>, WireError> { +pub fn decode_session_frames(bytes: &[u8]) -> Result>>, Error> { parse_session_frames(bytes) .map(|frame| frame.map(SessionFrame::into_owned)) .collect() @@ -155,13 +153,13 @@ pub struct SessionFrameIter { } impl Iterator for SessionFrameIter { - type Item = Result, WireError>; + type Item = Result, Error>; fn next(&mut self) -> Option { if self.reader.is_empty() { None } else { - Some(self.reader.decode::>()) + Some(self.reader.decode::>().map_err(Into::into)) } } } @@ -172,9 +170,9 @@ pub fn decrypt_record>( header: &SessionHeader, encrypted: EncryptedMessage, session_key: &SessionKey, -) -> Result { +) -> Result { let aad = header.aad(record_header.route); - let nonce = Nonce::from_counter(header.seq.0.into_inner()); + let nonce = Nonce::from_counter(header.seq.0); let mut ciphertext = encrypted.ciphertext; if !crypto.aes256_gcm_decrypt( session_key, @@ -183,7 +181,7 @@ pub fn decrypt_record>( ciphertext.as_mut(), &encrypted.auth, ) { - return Err(WireError::DecryptFailed); + return Err(Error::DecryptFailed); } Ok(ciphertext) } diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index d62cfdc1..d6786157 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -1,41 +1,39 @@ -use ql_common::{StreamId, VarInt}; - -use crate::{ - codec::{self, LenBytes}, - BufView, ByteSlice, WireDecode, WireEncode, WireError, +use ql_codec::{ + encode_bytes, encoded_len_bytes, varint, BufView, ByteSlice, Decode, Encode, Error, }; +use ql_common::StreamId; /// carries bytes for a stream and may finish that sending direction. #[derive(Debug, Clone, PartialEq, Eq)] pub struct StreamData { pub stream_id: StreamId, - pub offset: VarInt, - pub header: Option>, + pub offset: u64, + pub header: Option, pub fin: bool, pub bytes: B, } impl StreamData { pub const MIN_WIRE_SIZE: usize = StreamId::MAX_ENCODED_LEN - + VarInt::MAX_SIZE + + ::MAX_ENCODED_LEN + size_of::() - + VarInt::MAX_SIZE - + VarInt::MAX_SIZE; + + ::MAX_ENCODED_LEN + + ::MAX_ENCODED_LEN; } -impl WireDecode for StreamData { - fn decode(reader: &mut codec::Reader) -> Result { +impl Decode for StreamData { + fn decode(reader: &mut ql_codec::Reader) -> Result { let stream_id = reader.decode()?; - let offset: VarInt = reader.decode()?; + let offset = reader.decode_varint()?; let flags = reader.decode::()?; let fin = (flags & flag::FIN) != 0; let has_header = (flags & flag::HEADER) != 0; let header = if has_header { - Some(reader.decode()?) + Some(reader.take_len_prefixed()?) } else { None }; - let bytes = reader.decode::>()?.0; + let bytes = reader.take_len_prefixed()?; Ok(Self { stream_id, @@ -56,30 +54,30 @@ impl StreamData { StreamData { stream_id: self.stream_id, offset: self.offset, - header: self.header.map(|header| LenBytes(header.0.to_vec())), + header: self.header.map(|header| header.to_vec()), fin: self.fin, bytes: self.bytes.to_vec(), } } } -impl WireEncode for StreamData { +impl Encode for StreamData { fn encoded_len(&self) -> usize { self.stream_id.encoded_len() - + self.offset.encoded_len() + + varint::encoded_len(self.offset) + size_of::() - + self.header.as_ref().map_or(0, WireEncode::encoded_len) - + LenBytes(&self.bytes).encoded_len() + + self.header.as_ref().map_or(0, encoded_len_bytes) + + encoded_len_bytes(&self.bytes) } fn encode(&self, out: &mut W) { debug_assert!( - self.offset.into_inner() == 0 || self.header.is_none(), + self.offset == 0 || self.header.is_none(), "stream header is only valid at offset 0" ); self.stream_id.encode(out); - self.offset.encode(out); + varint::encode(self.offset, out); let mut flags = 0; if self.fin { flags |= flag::FIN; @@ -89,9 +87,9 @@ impl WireEncode for StreamData { } flags.encode(out); if let Some(header) = &self.header { - header.encode(out); + encode_bytes(header, out); } - LenBytes(&self.bytes).encode(out); + encode_bytes(&self.bytes, out); } } diff --git a/ql-wire/src/encrypted/stream_reset.rs b/ql-wire/src/encrypted/stream_reset.rs index 6ee73722..6efaec34 100644 --- a/ql-wire/src/encrypted/stream_reset.rs +++ b/ql-wire/src/encrypted/stream_reset.rs @@ -1,7 +1,8 @@ +use ql_codec::{ByteSlice, Encode}; use ql_common::ResetCode; use super::StreamId; -use crate::{codec, ByteSlice, WireEncode, WireError}; +use crate::Error; /// aborts one or both lanes of a stream with a reset code /// @@ -17,7 +18,7 @@ pub struct StreamReset { impl StreamReset {} -impl WireEncode for StreamReset { +impl Encode for StreamReset { fn encoded_len(&self) -> usize { self.stream_id.encoded_len() + self.target.encoded_len() + self.code.encoded_len() } @@ -29,8 +30,8 @@ impl WireEncode for StreamReset { } } -impl codec::WireDecode for StreamReset { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for StreamReset { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { stream_id: reader.decode()?, target: reader.decode()?, @@ -57,7 +58,7 @@ impl ResetTarget { } } -impl WireEncode for ResetTarget { +impl Encode for ResetTarget { fn encoded_len(&self) -> usize { size_of::() } @@ -68,36 +69,23 @@ impl WireEncode for ResetTarget { } impl TryFrom for ResetTarget { - type Error = WireError; + type Error = Error; fn try_from(value: u8) -> Result { match value { 1 => Ok(Self::Origin), 2 => Ok(Self::Return), 3 => Ok(Self::Both), - _ => Err(WireError::InvalidPayload), + _ => Err(Error::InvalidDiscriminant), } } } -impl codec::WireDecode for ResetTarget { - fn decode(reader: &mut codec::Reader) -> Result { - reader.decode::()?.try_into() - } -} - -impl codec::WireDecode for ResetCode { - fn decode(reader: &mut codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -impl WireEncode for ResetCode { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); +impl ql_codec::Decode for ResetTarget { + fn decode(reader: &mut ql_codec::Reader) -> Result { + reader + .decode::()? + .try_into() + .map_err(|_| ql_codec::Error::InvalidDiscriminant) } } diff --git a/ql-wire/src/encrypted/stream_window.rs b/ql-wire/src/encrypted/stream_window.rs index f932a2e6..20cbaa59 100644 --- a/ql-wire/src/encrypted/stream_window.rs +++ b/ql-wire/src/encrypted/stream_window.rs @@ -1,31 +1,30 @@ -use ql_common::VarInt; +use ql_codec::{ByteSlice, Encode, Error}; use super::StreamId; -use crate::{codec, ByteSlice, WireEncode, WireError}; /// advertises the highest byte offset the peer may send on a stream. #[derive(Debug, Clone, PartialEq, Eq)] pub struct StreamWindow { pub stream_id: StreamId, - pub maximum_offset: VarInt, + pub maximum_offset: u64, } -impl WireEncode for StreamWindow { +impl Encode for StreamWindow { fn encoded_len(&self) -> usize { - self.stream_id.encoded_len() + self.maximum_offset.encoded_len() + self.stream_id.encoded_len() + ql_codec::varint::encoded_len(self.maximum_offset) } fn encode(&self, out: &mut W) { self.stream_id.encode(out); - self.maximum_offset.encode(out); + ql_codec::varint::encode(self.maximum_offset, out); } } -impl codec::WireDecode for StreamWindow { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for StreamWindow { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { stream_id: reader.decode()?, - maximum_offset: reader.decode()?, + maximum_offset: reader.decode_varint()?, }) } } diff --git a/ql-wire/src/encrypted_message.rs b/ql-wire/src/encrypted_message.rs index 58b72591..99abda5e 100644 --- a/ql-wire/src/encrypted_message.rs +++ b/ql-wire/src/encrypted_message.rs @@ -1,4 +1,6 @@ -use crate::{codec, ByteSlice, WireDecode, WireEncode, WireError, ENCRYPTED_MESSAGE_AUTH_SIZE}; +use ql_codec::{ByteSlice, Decode, Encode}; + +use crate::ENCRYPTED_MESSAGE_AUTH_SIZE; #[derive(Debug, Clone, PartialEq, Eq)] pub struct EncryptedMessage { @@ -18,22 +20,22 @@ impl EncryptedMessage { } } -impl WireDecode for EncryptedMessage { - fn decode(reader: &mut codec::Reader) -> Result { +impl Decode for EncryptedMessage { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { auth: reader.decode()?, - ciphertext: reader.take_rest(), + ciphertext: reader.take_all(), }) } } -impl> WireEncode for EncryptedMessage { +impl> Encode for EncryptedMessage { fn encoded_len(&self) -> usize { ENCRYPTED_MESSAGE_AUTH_SIZE + self.ciphertext.as_ref().len() } fn encode(&self, out: &mut W) { self.auth.encode(out); - self.ciphertext.as_ref().encode(out); + out.put_slice(self.ciphertext.as_ref()); } } diff --git a/ql-wire/src/error.rs b/ql-wire/src/error.rs index c3c039a5..936c155d 100644 --- a/ql-wire/src/error.rs +++ b/ql-wire/src/error.rs @@ -1,33 +1,67 @@ use core::fmt; #[derive(Debug, Clone, PartialEq, Eq)] -pub enum WireError { +pub enum Error { + // codec errors + UnexpectedEof, + InvalidData, + InvalidDiscriminant, + InvalidRange, + LengthOverflow, + InvalidVarint, + InvalidUtf8, InvalidPayload, + + // protocol validation InvalidRouteHeader, InvalidHandshakeMeta, InvalidPairingId, InvalidRemoteBundle, InvalidTransportParams, - Expired, + + // cryptographic/session DecryptFailed, + Expired, + InvalidState, } -impl fmt::Display for WireError { +impl fmt::Display for Error { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { let message = match self { + Self::UnexpectedEof => "unexpected end of input", + Self::InvalidData => "invalid data", + Self::InvalidDiscriminant => "invalid discriminant", + Self::InvalidRange => "invalid range", + Self::LengthOverflow => "length overflow", + Self::InvalidVarint => "invalid varint", + Self::InvalidUtf8 => "invalid utf-8", Self::InvalidPayload => "invalid payload", Self::InvalidRouteHeader => "invalid route header", Self::InvalidHandshakeMeta => "invalid handshake meta", Self::InvalidPairingId => "invalid pairing id", Self::InvalidRemoteBundle => "invalid remote bundle", Self::InvalidTransportParams => "invalid transport params", - Self::Expired => "expired", Self::DecryptFailed => "decryption failed", + Self::Expired => "expired", Self::InvalidState => "invalid state", }; f.write_str(message) } } -impl std::error::Error for WireError {} +impl std::error::Error for Error {} + +impl From for Error { + fn from(error: ql_codec::Error) -> Self { + match error { + ql_codec::Error::InvalidData => Self::InvalidData, + ql_codec::Error::UnexpectedEof => Self::UnexpectedEof, + ql_codec::Error::InvalidDiscriminant => Self::InvalidDiscriminant, + ql_codec::Error::InvalidRange => Self::InvalidRange, + ql_codec::Error::LengthOverflow => Self::LengthOverflow, + ql_codec::Error::InvalidVarint => Self::InvalidVarint, + ql_codec::Error::InvalidUtf8 => Self::InvalidUtf8, + } + } +} diff --git a/ql-wire/src/handshake/ik.rs b/ql-wire/src/handshake/ik.rs index 113caf1d..38afd608 100644 --- a/ql-wire/src/handshake/ik.rs +++ b/ql-wire/src/handshake/ik.rs @@ -1,3 +1,5 @@ +use ql_codec::{ByteSlice, Encode}; + use super::{ decrypt_mlkem_ciphertext, decrypt_peer_bundle, encrypt_mlkem_ciphertext, encrypt_peer_bundle, finalize_handshake, generate_ephemeral_keypair, init_ik_symmetric, initialize_handshake_meta, @@ -6,8 +8,7 @@ use super::{ FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, }; use crate::{ - codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, - QlIdentity, WireEncode, WireError, + Error, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, QlIdentity, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -19,8 +20,8 @@ pub struct Ik1 { pub static_bundle: EncryptedPeerBundle, } -impl codec::WireDecode for Ik1 { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for Ik1 { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { meta: reader.decode()?, transport_params: reader.decode()?, @@ -31,7 +32,7 @@ impl codec::WireDecode for Ik1 { } } -impl WireEncode for Ik1 { +impl Encode for Ik1 { fn encoded_len(&self) -> usize { HandshakeMeta::WIRE_SIZE + TransportParams::WIRE_SIZE @@ -57,15 +58,8 @@ pub struct Ik2 { pub skem_ciphertext: EncryptedMlKemCiphertext, } -impl Ik2 { - pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE - + TransportParams::WIRE_SIZE - + MlKemCiphertext::SIZE - + EncryptedMlKemCiphertext::WIRE_SIZE; -} - -impl codec::WireDecode for Ik2 { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for Ik2 { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { meta: reader.decode()?, transport_params: reader.decode()?, @@ -75,9 +69,12 @@ impl codec::WireDecode for Ik2 { } } -impl WireEncode for Ik2 { +impl Encode for Ik2 { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + HandshakeMeta::WIRE_SIZE + + TransportParams::WIRE_SIZE + + MlKemCiphertext::SIZE + + EncryptedMlKemCiphertext::WIRE_SIZE } fn encode(&self, out: &mut W) { @@ -158,41 +155,37 @@ impl IkHandshake { self.step == IkStep::Done } - fn outbound_header(&self) -> Result { - let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + fn outbound_header(&self) -> Result { + let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; Ok(RouteHeader { sender: self.local.qid, recipient: remote_bundle.qid, }) } - fn ensure_inbound_recipient(&self, header: RouteHeader) -> Result<(), WireError> { + fn ensure_inbound_recipient(&self, header: RouteHeader) -> Result<(), Error> { if header.recipient == self.local.qid { Ok(()) } else { - Err(WireError::InvalidPayload) + Err(Error::InvalidRouteHeader) } } - fn ensure_known_remote_sender(&self, header: RouteHeader) -> Result<(), WireError> { + fn ensure_known_remote_sender(&self, header: RouteHeader) -> Result<(), Error> { if let Some(remote_bundle) = self.remote_bundle.as_ref() { if header.sender != remote_bundle.qid { - return Err(WireError::InvalidPayload); + return Err(Error::InvalidRouteHeader); } } Ok(()) } - pub fn write_1( - &mut self, - crypto: &impl QlCrypto, - meta: HandshakeMeta, - ) -> Result { + pub fn write_1(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { if self.step != IkStep::Send1 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, meta)?; - let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; let header = self.outbound_header()?; mix_hash_routed_handshake( &mut self.symmetric, @@ -225,13 +218,9 @@ impl IkHandshake { }) } - pub fn write_2( - &mut self, - crypto: &impl QlCrypto, - meta: HandshakeMeta, - ) -> Result { + pub fn write_2(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { if self.step != IkStep::Send2 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), meta)?; let header = self.outbound_header()?; @@ -243,16 +232,13 @@ impl IkHandshake { meta, self.local_transport_params, ); - let remote_ephemeral = self - .remote_ephemeral - .clone() - .ok_or(WireError::InvalidState)?; + let remote_ephemeral = self.remote_ephemeral.clone().ok_or(Error::InvalidState)?; let (ekem_ciphertext, ekem_secret) = crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); - let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; let (skem_ciphertext, skem_secret) = crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); let skem_ciphertext = @@ -274,9 +260,9 @@ impl IkHandshake { crypto: &impl QlCrypto, header: RouteHeader, message: &Ik1, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if self.step != IkStep::Recv1 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; self.ensure_inbound_recipient(header)?; @@ -302,11 +288,11 @@ impl IkHandshake { let remote_bundle = decrypt_peer_bundle(crypto, &mut self.symmetric, &message.static_bundle)?; if remote_bundle.qid != header.sender { - return Err(WireError::InvalidPayload); + return Err(Error::InvalidRemoteBundle); } match self.remote_bundle.as_ref() { Some(expected) if expected != &remote_bundle => { - return Err(WireError::InvalidPayload); + return Err(Error::InvalidRemoteBundle); } Some(_) => {} None => self.remote_bundle = Some(remote_bundle), @@ -321,9 +307,9 @@ impl IkHandshake { crypto: &impl QlCrypto, header: RouteHeader, message: &Ik2, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if self.step != IkStep::Recv2 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; self.ensure_inbound_recipient(header)?; @@ -336,10 +322,7 @@ impl IkHandshake { message.meta, message.transport_params, ); - let local_ephemeral = self - .local_ephemeral - .as_ref() - .ok_or(WireError::InvalidState)?; + let local_ephemeral = self.local_ephemeral.as_ref().ok_or(Error::InvalidState)?; self.symmetric .mix_hash(crypto, message.ekem_ciphertext.as_bytes()); let ekem_secret = @@ -357,14 +340,12 @@ impl IkHandshake { Ok(()) } - pub fn finalize(self, crypto: &impl QlCrypto) -> Result { + pub fn finalize(self, crypto: &impl QlCrypto) -> Result { if !self.is_finished() { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } - let remote_bundle = self.remote_bundle.ok_or(WireError::InvalidState)?; - let remote_transport_params = self - .remote_transport_params - .ok_or(WireError::InvalidState)?; + let remote_bundle = self.remote_bundle.ok_or(Error::InvalidState)?; + let remote_transport_params = self.remote_transport_params.ok_or(Error::InvalidState)?; Ok(finalize_handshake( crypto, &self.symmetric, diff --git a/ql-wire/src/handshake/kk.rs b/ql-wire/src/handshake/kk.rs index e71775b3..2f2db1ef 100644 --- a/ql-wire/src/handshake/kk.rs +++ b/ql-wire/src/handshake/kk.rs @@ -1,3 +1,5 @@ +use ql_codec::{ByteSlice, Encode}; + use super::{ decrypt_mlkem_ciphertext, encrypt_mlkem_ciphertext, finalize_handshake, generate_ephemeral_keypair, init_kk_symmetric, initialize_handshake_meta, mix_hash_ephemeral, @@ -5,8 +7,7 @@ use super::{ EphemeralPublicKey, FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, }; use crate::{ - codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, - QlIdentity, WireEncode, WireError, + Error, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, QlIdentity, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -17,15 +18,8 @@ pub struct Kk1 { pub ephemeral: EphemeralPublicKey, } -impl Kk1 { - pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE - + TransportParams::WIRE_SIZE - + MlKemCiphertext::SIZE - + EphemeralPublicKey::WIRE_SIZE; -} - -impl codec::WireDecode for Kk1 { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for Kk1 { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { meta: reader.decode()?, transport_params: reader.decode()?, @@ -35,9 +29,12 @@ impl codec::WireDecode for Kk1 { } } -impl WireEncode for Kk1 { +impl Encode for Kk1 { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + HandshakeMeta::WIRE_SIZE + + TransportParams::WIRE_SIZE + + MlKemCiphertext::SIZE + + EphemeralPublicKey::WIRE_SIZE } fn encode(&self, out: &mut W) { @@ -56,15 +53,8 @@ pub struct Kk2 { pub skem_ciphertext: EncryptedMlKemCiphertext, } -impl Kk2 { - pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE - + TransportParams::WIRE_SIZE - + MlKemCiphertext::SIZE - + EncryptedMlKemCiphertext::WIRE_SIZE; -} - -impl codec::WireDecode for Kk2 { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for Kk2 { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { meta: reader.decode()?, transport_params: reader.decode()?, @@ -74,9 +64,12 @@ impl codec::WireDecode for Kk2 { } } -impl WireEncode for Kk2 { +impl Encode for Kk2 { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + HandshakeMeta::WIRE_SIZE + + TransportParams::WIRE_SIZE + + MlKemCiphertext::SIZE + + EncryptedMlKemCiphertext::WIRE_SIZE } fn encode(&self, out: &mut W) { @@ -171,21 +164,17 @@ impl KkHandshake { } } - fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), WireError> { + fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), Error> { if header == self.inbound_header() { Ok(()) } else { - Err(WireError::InvalidPayload) + Err(Error::InvalidRouteHeader) } } - pub fn write_1( - &mut self, - crypto: &impl QlCrypto, - meta: HandshakeMeta, - ) -> Result { + pub fn write_1(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { if self.step != KkStep::Send1 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, meta)?; let header = self.outbound_header(); @@ -218,13 +207,9 @@ impl KkHandshake { }) } - pub fn write_2( - &mut self, - crypto: &impl QlCrypto, - meta: HandshakeMeta, - ) -> Result { + pub fn write_2(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { if self.step != KkStep::Send2 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), meta)?; let header = self.outbound_header(); @@ -236,10 +221,7 @@ impl KkHandshake { meta, self.local_transport_params, ); - let remote_ephemeral = self - .remote_ephemeral - .clone() - .ok_or(WireError::InvalidState)?; + let remote_ephemeral = self.remote_ephemeral.clone().ok_or(Error::InvalidState)?; let (ekem_ciphertext, ekem_secret) = crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); @@ -266,9 +248,9 @@ impl KkHandshake { crypto: &impl QlCrypto, header: RouteHeader, message: &Kk1, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if self.step != KkStep::Recv1 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; self.ensure_inbound_header(header)?; @@ -299,9 +281,9 @@ impl KkHandshake { crypto: &impl QlCrypto, header: RouteHeader, message: &Kk2, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if self.step != KkStep::Recv2 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; self.ensure_inbound_header(header)?; @@ -313,10 +295,7 @@ impl KkHandshake { message.meta, message.transport_params, ); - let local_ephemeral = self - .local_ephemeral - .as_ref() - .ok_or(WireError::InvalidState)?; + let local_ephemeral = self.local_ephemeral.as_ref().ok_or(Error::InvalidState)?; self.symmetric .mix_hash(crypto, message.ekem_ciphertext.as_bytes()); let ekem_secret = @@ -334,13 +313,11 @@ impl KkHandshake { Ok(()) } - pub fn finalize(self, crypto: &impl QlCrypto) -> Result { + pub fn finalize(self, crypto: &impl QlCrypto) -> Result { if !self.is_finished() { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } - let remote_transport_params = self - .remote_transport_params - .ok_or(WireError::InvalidState)?; + let remote_transport_params = self.remote_transport_params.ok_or(Error::InvalidState)?; Ok(finalize_handshake( crypto, &self.symmetric, diff --git a/ql-wire/src/handshake/meta.rs b/ql-wire/src/handshake/meta.rs index 8cb0cf97..21b70f82 100644 --- a/ql-wire/src/handshake/meta.rs +++ b/ql-wire/src/handshake/meta.rs @@ -1,4 +1,4 @@ -use crate::{codec, ByteSlice, WireEncode, WireError}; +use ql_codec::{ByteSlice, Encode}; #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] #[repr(transparent)] @@ -9,13 +9,13 @@ pub struct HandshakeMeta { pub handshake_id: HandshakeId, } -impl codec::WireDecode for HandshakeId { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for HandshakeId { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self(reader.decode()?)) } } -impl WireEncode for HandshakeId { +impl Encode for HandshakeId { fn encoded_len(&self) -> usize { size_of::() } @@ -29,7 +29,7 @@ impl HandshakeMeta { pub const WIRE_SIZE: usize = size_of::(); } -impl WireEncode for HandshakeMeta { +impl Encode for HandshakeMeta { fn encoded_len(&self) -> usize { Self::WIRE_SIZE } @@ -39,8 +39,8 @@ impl WireEncode for HandshakeMeta { } } -impl codec::WireDecode for HandshakeMeta { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for HandshakeMeta { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { handshake_id: reader.decode()?, }) diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index 0a581132..df0d93ab 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -1,7 +1,8 @@ +use ql_codec::{ByteSlice, Decode, Encode}; + use crate::{ - codec, derive_qid, ByteSlice, HandshakeKind, MlKemCiphertext, MlKemKeyPair, MlKemPublicKey, - Nonce, PeerBundle, QlCrypto, RouteHeader, SessionKey, WireDecode, WireEncode, WireError, - ENCRYPTED_MESSAGE_AUTH_SIZE, + derive_qid, Error, HandshakeKind, MlKemCiphertext, MlKemKeyPair, MlKemPublicKey, Nonce, + PeerBundle, QlCrypto, RouteHeader, SessionKey, ENCRYPTED_MESSAGE_AUTH_SIZE, }; mod ik; @@ -33,7 +34,7 @@ impl EphemeralPublicKey { pub const WIRE_SIZE: usize = MlKemPublicKey::SIZE; } -impl WireEncode for EphemeralPublicKey { +impl Encode for EphemeralPublicKey { fn encoded_len(&self) -> usize { Self::WIRE_SIZE } @@ -43,8 +44,8 @@ impl WireEncode for EphemeralPublicKey { } } -impl codec::WireDecode for EphemeralPublicKey { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for EphemeralPublicKey { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { mlkem_public_key: reader.decode()?, }) @@ -66,7 +67,7 @@ impl EncryptedMlKemCiphertext { } } -impl WireEncode for EncryptedMlKemCiphertext { +impl Encode for EncryptedMlKemCiphertext { fn encoded_len(&self) -> usize { Self::WIRE_SIZE } @@ -76,8 +77,8 @@ impl WireEncode for EncryptedMlKemCiphertext { } } -impl codec::WireDecode for EncryptedMlKemCiphertext { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for EncryptedMlKemCiphertext { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self::new(reader.decode()?)) } } @@ -91,20 +92,20 @@ impl EncryptedPeerBundle { } } -impl WireEncode for EncryptedPeerBundle { +impl Encode for EncryptedPeerBundle { fn encoded_len(&self) -> usize { self.0.len() } fn encode(&self, out: &mut W) { - self.as_bytes().encode(out); + out.put_slice(self.as_bytes()); } } -impl codec::WireDecode for EncryptedPeerBundle { - fn decode(reader: &mut codec::Reader) -> Result { - let data = reader.take_rest(); - Ok(Self(data.to_vec().into_boxed_slice())) +impl ql_codec::Decode for EncryptedPeerBundle { + fn decode(reader: &mut ql_codec::Reader) -> Result { + let data = reader.take_all(); + Ok(Self(Box::from(&*data))) } } @@ -165,8 +166,8 @@ impl CipherState { crypto: &impl QlCrypto, aad: &[u8], plaintext: &[u8], - ) -> Result, WireError> { - let key = self.key.as_ref().ok_or(WireError::InvalidState)?; + ) -> Result, Error> { + let key = self.key.as_ref().ok_or(Error::InvalidState)?; let nonce = Nonce::from_counter(self.nonce); let mut ciphertext = Vec::with_capacity(plaintext.len() + ENCRYPTED_MESSAGE_AUTH_SIZE); ciphertext.extend_from_slice(plaintext); @@ -181,19 +182,19 @@ impl CipherState { crypto: &impl QlCrypto, aad: &[u8], ciphertext: &[u8], - ) -> Result, WireError> { + ) -> Result, Error> { if ciphertext.len() < ENCRYPTED_MESSAGE_AUTH_SIZE { - return Err(WireError::InvalidPayload); + return Err(Error::InvalidPayload); } let split = ciphertext.len() - ENCRYPTED_MESSAGE_AUTH_SIZE; let (ciphertext, auth) = ciphertext.split_at(split); let mut plaintext = ciphertext.to_vec(); - let key = self.key.as_ref().ok_or(WireError::InvalidState)?; + let key = self.key.as_ref().ok_or(Error::InvalidState)?; let nonce = Nonce::from_counter(self.nonce); let mut auth_tag = [0u8; ENCRYPTED_MESSAGE_AUTH_SIZE]; auth_tag.copy_from_slice(auth); if !crypto.aes256_gcm_decrypt(key, &nonce, aad, &mut plaintext, &auth_tag) { - return Err(WireError::DecryptFailed); + return Err(Error::DecryptFailed); } self.nonce = self.nonce.wrapping_add(1); Ok(plaintext) @@ -239,7 +240,7 @@ impl SymmetricState { &mut self, crypto: &impl QlCrypto, plaintext: &[u8], - ) -> Result, WireError> { + ) -> Result, Error> { if self.cipher.has_key() { let ciphertext = self .cipher @@ -256,7 +257,7 @@ impl SymmetricState { &mut self, crypto: &impl QlCrypto, ciphertext: &[u8], - ) -> Result, WireError> { + ) -> Result, Error> { if self.cipher.has_key() { let plaintext = self .cipher @@ -373,9 +374,9 @@ fn mix_hash_handshake_preamble( fn initialize_handshake_meta( expected: &mut Option, meta: HandshakeMeta, -) -> Result<(), WireError> { +) -> Result<(), Error> { match expected { - Some(stored) if *stored != meta => Err(WireError::InvalidHandshakeMeta), + Some(stored) if *stored != meta => Err(Error::InvalidHandshakeMeta), Some(_) => Ok(()), None => { *expected = Some(meta); @@ -387,19 +388,19 @@ fn initialize_handshake_meta( fn require_handshake_meta( expected: Option<&HandshakeMeta>, meta: HandshakeMeta, -) -> Result<(), WireError> { +) -> Result<(), Error> { match expected { Some(stored) if *stored == meta => Ok(()), - _ => Err(WireError::InvalidHandshakeMeta), + _ => Err(Error::InvalidHandshakeMeta), } } fn initialize_transport_params( expected: &mut Option, transport_params: TransportParams, -) -> Result<(), WireError> { +) -> Result<(), Error> { match expected { - Some(stored) if *stored != transport_params => Err(WireError::InvalidTransportParams), + Some(stored) if *stored != transport_params => Err(Error::InvalidTransportParams), Some(_) => Ok(()), None => { *expected = Some(transport_params); @@ -411,10 +412,10 @@ fn initialize_transport_params( fn require_transport_params( expected: Option<&TransportParams>, transport_params: TransportParams, -) -> Result<(), WireError> { +) -> Result<(), Error> { match expected { Some(stored) if *stored == transport_params => Ok(()), - _ => Err(WireError::InvalidTransportParams), + _ => Err(Error::InvalidTransportParams), } } @@ -422,7 +423,7 @@ fn encrypt_peer_bundle( crypto: &impl QlCrypto, symmetric: &mut SymmetricState, bundle: &PeerBundle, -) -> Result { +) -> Result { let ciphertext = symmetric.encrypt_and_hash(crypto, &bundle.encode_vec())?; Ok(EncryptedPeerBundle(ciphertext.into_boxed_slice())) } @@ -431,12 +432,12 @@ fn decrypt_peer_bundle( crypto: &impl QlCrypto, symmetric: &mut SymmetricState, bundle: &EncryptedPeerBundle, -) -> Result { +) -> Result { let plaintext = symmetric.decrypt_and_hash(crypto, bundle.as_bytes())?; - let bundle = PeerBundle::decode_exact(plaintext.as_slice())?; + let bundle = PeerBundle::decode_bytes(plaintext.as_slice())?; let peer_qid = derive_qid(crypto, &bundle.mlkem_public_key); if peer_qid != bundle.qid { - return Err(WireError::InvalidRemoteBundle); + return Err(Error::InvalidRemoteBundle); } Ok(bundle) } @@ -445,10 +446,10 @@ fn encrypt_mlkem_ciphertext( crypto: &impl QlCrypto, symmetric: &mut SymmetricState, ciphertext: &MlKemCiphertext, -) -> Result { +) -> Result { let encrypted = symmetric.encrypt_and_hash(crypto, ciphertext.as_bytes())?; let out: Box<[u8; EncryptedMlKemCiphertext::WIRE_SIZE]> = - encrypted.try_into().map_err(|_| WireError::InvalidState)?; + encrypted.try_into().map_err(|_| Error::InvalidState)?; Ok(EncryptedMlKemCiphertext::new(out)) } @@ -456,11 +457,10 @@ fn decrypt_mlkem_ciphertext( crypto: &impl QlCrypto, symmetric: &mut SymmetricState, ciphertext: &EncryptedMlKemCiphertext, -) -> Result { +) -> Result { let plaintext = symmetric.decrypt_and_hash(crypto, ciphertext.as_bytes())?; - let out: Box<[u8; MlKemCiphertext::SIZE]> = plaintext - .try_into() - .map_err(|_| WireError::InvalidPayload)?; + let out: Box<[u8; MlKemCiphertext::SIZE]> = + plaintext.try_into().map_err(|_| Error::InvalidPayload)?; Ok(MlKemCiphertext::new(out)) } diff --git a/ql-wire/src/handshake/transport_params.rs b/ql-wire/src/handshake/transport_params.rs index bfd0d427..71766e54 100644 --- a/ql-wire/src/handshake/transport_params.rs +++ b/ql-wire/src/handshake/transport_params.rs @@ -1,4 +1,4 @@ -use crate::{codec, ByteSlice, WireEncode, WireError}; +use ql_codec::{ByteSlice, Encode}; /// Session parameters advertised in the handshake #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -11,7 +11,7 @@ impl TransportParams { pub const WIRE_SIZE: usize = size_of::(); } -impl WireEncode for TransportParams { +impl Encode for TransportParams { fn encoded_len(&self) -> usize { Self::WIRE_SIZE } @@ -29,8 +29,8 @@ impl Default for TransportParams { } } -impl codec::WireDecode for TransportParams { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for TransportParams { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { initial_stream_receive_window: reader.decode()?, }) diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs index f859f2e3..39f4b680 100644 --- a/ql-wire/src/handshake/xx.rs +++ b/ql-wire/src/handshake/xx.rs @@ -1,3 +1,4 @@ +use ql_codec::{ByteSlice, Encode}; use ql_common::QID; use super::{ @@ -9,8 +10,8 @@ use super::{ FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, }; use crate::{ - codec, ByteSlice, HandshakeKind, HandshakeMeta, MlKemCiphertext, PairingId, PairingToken, - PeerBundle, QlCrypto, QlIdentity, WireEncode, WireError, + Error, HandshakeKind, HandshakeMeta, MlKemCiphertext, PairingId, PairingToken, PeerBundle, + QlCrypto, QlIdentity, }; #[derive(Debug, Clone, PartialEq, Eq)] @@ -21,15 +22,8 @@ pub struct Xx1 { pub ephemeral: EphemeralPublicKey, } -impl Xx1 { - pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE - + PairingId::SIZE - + TransportParams::WIRE_SIZE - + EphemeralPublicKey::WIRE_SIZE; -} - -impl codec::WireDecode for Xx1 { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for Xx1 { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { meta: reader.decode()?, pairing_id: reader.decode()?, @@ -39,9 +33,12 @@ impl codec::WireDecode for Xx1 { } } -impl WireEncode for Xx1 { +impl Encode for Xx1 { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + HandshakeMeta::WIRE_SIZE + + PairingId::SIZE + + TransportParams::WIRE_SIZE + + EphemeralPublicKey::WIRE_SIZE } fn encode(&self, out: &mut W) { @@ -61,8 +58,8 @@ pub struct Xx2 { pub static_bundle: EncryptedPeerBundle, } -impl codec::WireDecode for Xx2 { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for Xx2 { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { meta: reader.decode()?, pairing_id: reader.decode()?, @@ -73,7 +70,7 @@ impl codec::WireDecode for Xx2 { } } -impl WireEncode for Xx2 { +impl Encode for Xx2 { fn encoded_len(&self) -> usize { HandshakeMeta::WIRE_SIZE + PairingId::SIZE @@ -100,8 +97,8 @@ pub struct Xx3 { pub static_bundle: EncryptedPeerBundle, } -impl codec::WireDecode for Xx3 { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for Xx3 { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { meta: reader.decode()?, pairing_id: reader.decode()?, @@ -112,7 +109,7 @@ impl codec::WireDecode for Xx3 { } } -impl WireEncode for Xx3 { +impl Encode for Xx3 { fn encoded_len(&self) -> usize { HandshakeMeta::WIRE_SIZE + PairingId::SIZE @@ -138,15 +135,8 @@ pub struct Xx4 { pub skem_ciphertext: EncryptedMlKemCiphertext, } -impl Xx4 { - pub const WIRE_SIZE: usize = HandshakeMeta::WIRE_SIZE - + PairingId::SIZE - + TransportParams::WIRE_SIZE - + EncryptedMlKemCiphertext::WIRE_SIZE; -} - -impl codec::WireDecode for Xx4 { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for Xx4 { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { meta: reader.decode()?, pairing_id: reader.decode()?, @@ -156,9 +146,12 @@ impl codec::WireDecode for Xx4 { } } -impl WireEncode for Xx4 { +impl Encode for Xx4 { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + HandshakeMeta::WIRE_SIZE + + PairingId::SIZE + + TransportParams::WIRE_SIZE + + EncryptedMlKemCiphertext::WIRE_SIZE } fn encode(&self, out: &mut W) { @@ -277,31 +270,27 @@ impl XxHandshake { crypto: &impl QlCrypto, header: RouteHeader, pairing_id: PairingId, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if header.sender != self.remote_qid || header.recipient != self.local.qid { - return Err(WireError::InvalidRouteHeader); + return Err(Error::InvalidRouteHeader); } if pairing_id != self.pairing_token.id(crypto) { - return Err(WireError::InvalidPairingId); + return Err(Error::InvalidPairingId); } Ok(()) } - fn ensure_remote_bundle(&self, bundle: &PeerBundle) -> Result<(), WireError> { + fn ensure_remote_bundle(&self, bundle: &PeerBundle) -> Result<(), Error> { if bundle.qid == self.remote_qid { Ok(()) } else { - Err(WireError::InvalidRemoteBundle) + Err(Error::InvalidRemoteBundle) } } - pub fn write_1( - &mut self, - crypto: &impl QlCrypto, - meta: HandshakeMeta, - ) -> Result { + pub fn write_1(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { if self.step != XxStep::Send1 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, meta)?; let header = self.header(); @@ -336,9 +325,9 @@ impl XxHandshake { crypto: &impl QlCrypto, header: RouteHeader, message: &Xx1, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if self.step != XxStep::Recv1 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; self.ensure_inbound_header(crypto, header, message.pairing_id)?; @@ -360,13 +349,9 @@ impl XxHandshake { Ok(()) } - pub fn write_2( - &mut self, - crypto: &impl QlCrypto, - meta: HandshakeMeta, - ) -> Result { + pub fn write_2(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { if self.step != XxStep::Send2 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), meta)?; let header = self.header(); @@ -381,10 +366,7 @@ impl XxHandshake { self.local_transport_params, ); - let remote_ephemeral = self - .remote_ephemeral - .as_ref() - .ok_or(WireError::InvalidState)?; + let remote_ephemeral = self.remote_ephemeral.as_ref().ok_or(Error::InvalidState)?; let (ekem_ciphertext, ekem_secret) = crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); @@ -407,9 +389,9 @@ impl XxHandshake { crypto: &impl QlCrypto, header: RouteHeader, message: &Xx2, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if self.step != XxStep::Recv2 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; self.ensure_inbound_header(crypto, header, message.pairing_id)?; @@ -423,10 +405,7 @@ impl XxHandshake { message.transport_params, ); - let local_ephemeral = self - .local_ephemeral - .as_ref() - .ok_or(WireError::InvalidState)?; + let local_ephemeral = self.local_ephemeral.as_ref().ok_or(Error::InvalidState)?; self.symmetric .mix_hash(crypto, message.ekem_ciphertext.as_bytes()); let ekem_secret = @@ -442,13 +421,9 @@ impl XxHandshake { Ok(()) } - pub fn write_3( - &mut self, - crypto: &impl QlCrypto, - meta: HandshakeMeta, - ) -> Result { + pub fn write_3(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { if self.step != XxStep::Send3 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), meta)?; let header = self.header(); @@ -463,7 +438,7 @@ impl XxHandshake { self.local_transport_params, ); - let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; let (skem_ciphertext, skem_secret) = crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); let skem_ciphertext = @@ -488,9 +463,9 @@ impl XxHandshake { crypto: &impl QlCrypto, header: RouteHeader, message: &Xx3, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if self.step != XxStep::Recv3 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; self.ensure_inbound_header(crypto, header, message.pairing_id)?; @@ -522,13 +497,9 @@ impl XxHandshake { Ok(()) } - pub fn write_4( - &mut self, - crypto: &impl QlCrypto, - meta: HandshakeMeta, - ) -> Result { + pub fn write_4(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { if self.step != XxStep::Send4 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), meta)?; let header = self.header(); @@ -543,7 +514,7 @@ impl XxHandshake { self.local_transport_params, ); - let remote_bundle = self.remote_bundle.as_ref().ok_or(WireError::InvalidState)?; + let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; let (skem_ciphertext, skem_secret) = crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); let skem_ciphertext = @@ -565,9 +536,9 @@ impl XxHandshake { crypto: &impl QlCrypto, header: RouteHeader, message: &Xx4, - ) -> Result<(), WireError> { + ) -> Result<(), Error> { if self.step != XxStep::Recv4 { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; self.ensure_inbound_header(crypto, header, message.pairing_id)?; @@ -595,14 +566,12 @@ impl XxHandshake { Ok(()) } - pub fn finalize(self, crypto: &impl QlCrypto) -> Result { + pub fn finalize(self, crypto: &impl QlCrypto) -> Result { if !self.is_finished() { - return Err(WireError::InvalidState); + return Err(Error::InvalidState); } - let remote_bundle = self.remote_bundle.ok_or(WireError::InvalidState)?; - let remote_transport_params = self - .remote_transport_params - .ok_or(WireError::InvalidState)?; + let remote_bundle = self.remote_bundle.ok_or(Error::InvalidState)?; + let remote_transport_params = self.remote_transport_params.ok_or(Error::InvalidState)?; Ok(finalize_handshake( crypto, &self.symmetric, diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index 79897c57..d3f45b77 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -1,7 +1,8 @@ use ::bytes::BufMut; +use ql_codec::{ByteSlice, Encode, Error}; use ql_common::QID; -use crate::{codec, ByteSlice, WireEncode, WireError, QL_WIRE_VERSION}; +use crate::QL_WIRE_VERSION; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RouteHeader { @@ -13,7 +14,7 @@ impl RouteHeader { pub const WIRE_SIZE: usize = QID::SIZE * 2; } -impl WireEncode for RouteHeader { +impl Encode for RouteHeader { fn encoded_len(&self) -> usize { Self::WIRE_SIZE } @@ -24,8 +25,8 @@ impl WireEncode for RouteHeader { } } -impl codec::WireDecode for RouteHeader { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for RouteHeader { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { sender: reader.decode()?, recipient: reader.decode()?, @@ -38,11 +39,19 @@ pub struct SessionHeader { pub seq: RecordSeq, } -ql_common::varint_wrapper!(RecordSeq); -varint_wrapper_codec!(RecordSeq); +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct RecordSeq(pub u64); + +impl std::fmt::Display for RecordSeq { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "{}", self.0) + } +} + +ql_codec::varint_wrapper!(RecordSeq, u64); impl SessionHeader { - pub const MAX_ENCODED_LEN: usize = RecordSeq::MAX_ENCODED_LEN; const AAD_DOMAIN: &[u8] = b"ql-wire:session-aad:v1"; const AAD_RECORD_KIND_SESSION: u8 = 1; @@ -63,7 +72,7 @@ impl SessionHeader { } } -impl WireEncode for SessionHeader { +impl Encode for SessionHeader { fn encoded_len(&self) -> usize { self.seq.encoded_len() } @@ -73,8 +82,8 @@ impl WireEncode for SessionHeader { } } -impl codec::WireDecode for SessionHeader { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for SessionHeader { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { seq: reader.decode()?, }) diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index f88fe2f0..e23c4493 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -1,12 +1,9 @@ use std::ops::Deref; -use codec::LenBytes; +use ql_codec::{ByteSlice, Encode}; use ql_common::QID; -use crate::{ - codec, derive_qid, ByteSlice, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash, - WireEncode, WireError, -}; +use crate::{derive_qid, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash}; #[derive(Debug, Clone, PartialEq, Eq)] pub struct PeerBundle { @@ -20,13 +17,16 @@ pub struct PeerBundle { impl PeerBundle { pub const VERSION: u16 = 1; - pub const FIXED_WIRE_SIZE: usize = - size_of::() + QID::SIZE + size_of::() + MlKemPublicKey::SIZE; } -impl WireEncode for PeerBundle { +impl Encode for PeerBundle { fn encoded_len(&self) -> usize { - Self::FIXED_WIRE_SIZE + self.name.encoded_len() + LenBytes(&*self.metadata).encoded_len() + size_of::() + + QID::SIZE + + size_of::() + + MlKemPublicKey::SIZE + + self.name.encoded_len() + + self.metadata.encoded_len() } fn encode(&self, out: &mut W) { @@ -35,23 +35,19 @@ impl WireEncode for PeerBundle { self.capabilities.encode(out); self.mlkem_public_key.encode(out); self.name.encode(out); - LenBytes(&*self.metadata).encode(out); + self.metadata.encode(out); } } -impl codec::WireDecode for PeerBundle { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for PeerBundle { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { version: reader.decode()?, qid: reader.decode()?, capabilities: reader.decode()?, mlkem_public_key: reader.decode()?, name: reader.decode()?, - metadata: reader - .decode::>()? - .0 - .to_vec() - .into_boxed_slice(), + metadata: reader.decode()?, }) } } @@ -67,9 +63,6 @@ pub struct QlIdentity { } impl QlIdentity { - pub const FIXED_WIRE_SIZE: usize = - QID::SIZE + MlKemPrivateKey::SIZE + MlKemPublicKey::SIZE + size_of::(); - pub fn new( crypto: &impl QlHash, mlkem_private_key: MlKemPrivateKey, @@ -105,9 +98,14 @@ impl QlIdentity { } } -impl WireEncode for QlIdentity { +impl Encode for QlIdentity { fn encoded_len(&self) -> usize { - Self::FIXED_WIRE_SIZE + self.name.encoded_len() + LenBytes(&*self.metadata).encoded_len() + QID::SIZE + + MlKemPrivateKey::SIZE + + MlKemPublicKey::SIZE + + size_of::() + + self.name.encoded_len() + + self.metadata.encoded_len() } fn encode(&self, out: &mut W) { @@ -116,23 +114,19 @@ impl WireEncode for QlIdentity { self.mlkem_public_key.encode(out); self.capabilities.encode(out); self.name.encode(out); - LenBytes(&*self.metadata).encode(out); + self.metadata.encode(out); } } -impl codec::WireDecode for QlIdentity { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for QlIdentity { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { qid: reader.decode()?, mlkem_private_key: MlKemPrivateKey::new(reader.decode()?), mlkem_public_key: reader.decode()?, capabilities: reader.decode()?, name: reader.decode()?, - metadata: reader - .decode::>()? - .0 - .to_vec() - .into_boxed_slice(), + metadata: reader.decode()?, }) } } @@ -153,20 +147,20 @@ impl Deref for QlName { } } -impl WireEncode for QlName { +impl Encode for QlName { fn encoded_len(&self) -> usize { - LenBytes(self.0.as_bytes()).encoded_len() + self.0.as_bytes().encoded_len() } fn encode(&self, out: &mut W) { - LenBytes(self.0.as_bytes()).encode(out) + self.0.as_bytes().encode(out); } } -impl codec::WireDecode for QlName { - fn decode(reader: &mut codec::Reader) -> Result { - let bytes = reader.decode::>()?.0; - let name = std::str::from_utf8(&bytes).map_err(|_| WireError::InvalidPayload)?; +impl ql_codec::Decode for QlName { + fn decode(reader: &mut ql_codec::Reader) -> Result { + let bytes = reader.take_len_prefixed()?; + let name = std::str::from_utf8(&bytes).map_err(|_| ql_codec::Error::InvalidUtf8)?; Ok(QlName(name.into())) } } diff --git a/ql-wire/src/lib.rs b/ql-wire/src/lib.rs index b55c3dd9..55fd3b73 100644 --- a/ql-wire/src/lib.rs +++ b/ql-wire/src/lib.rs @@ -6,8 +6,6 @@ #[macro_use] mod macros; -mod bytes; -mod codec; mod crypto; mod encrypted; mod encrypted_message; @@ -20,10 +18,7 @@ mod qid; mod record; #[cfg(any(feature = "test-utils", test))] mod testing; -mod varint; -pub use bytes::*; -pub use codec::*; pub use crypto::*; pub use encrypted::*; pub use encrypted_message::*; diff --git a/ql-wire/src/macros.rs b/ql-wire/src/macros.rs index 29fc4e7d..9124187c 100644 --- a/ql-wire/src/macros.rs +++ b/ql-wire/src/macros.rs @@ -1,43 +1,3 @@ -macro_rules! varint_wrapper_codec { - ($name:ty) => { - impl $crate::WireEncode for $name { - fn encoded_len(&self) -> usize { - self.0.size() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } - } - - impl $crate::WireDecode for $name { - fn decode(reader: &mut $crate::Reader) -> Result { - Ok(<$name>::from(reader.decode::<::ql_common::VarInt>()?)) - } - } - }; -} - -macro_rules! array_wrapper_codec { - ($name:ty) => { - impl $crate::WireEncode for $name { - fn encoded_len(&self) -> usize { - <$name>::SIZE - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } - } - - impl $crate::codec::WireDecode for $name { - fn decode(reader: &mut $crate::codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } - } - }; -} - macro_rules! array_wrapper { ($name:ident, $size:expr) => { #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] @@ -52,7 +12,7 @@ macro_rules! array_wrapper { } } - impl $crate::WireEncode for $name { + impl ql_codec::Encode for $name { fn encoded_len(&self) -> usize { Self::SIZE } @@ -62,8 +22,8 @@ macro_rules! array_wrapper { } } - impl $crate::codec::WireDecode for $name { - fn decode(reader: &mut $crate::codec::Reader) -> Result { + impl ql_codec::Decode for $name { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self(reader.decode()?)) } } diff --git a/ql-wire/src/pq.rs b/ql-wire/src/pq.rs index 8aa513bc..9cd0efcb 100644 --- a/ql-wire/src/pq.rs +++ b/ql-wire/src/pq.rs @@ -1,4 +1,4 @@ -use crate::{codec, ByteSlice, WireEncode, WireError}; +use ql_codec::{ByteSlice, Encode}; pub const ML_KEM_SUITE_TAG: &[u8] = b"ml-kem-1024"; @@ -33,7 +33,7 @@ impl Drop for SessionKey { } } -impl WireEncode for SessionKey { +impl Encode for SessionKey { fn encoded_len(&self) -> usize { Self::SIZE } @@ -43,8 +43,8 @@ impl WireEncode for SessionKey { } } -impl codec::WireDecode for SessionKey { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for SessionKey { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self(reader.decode()?)) } } @@ -70,13 +70,13 @@ impl Drop for MlKemPublicKey { } } -impl codec::WireDecode for MlKemPublicKey { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for MlKemPublicKey { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self::new(reader.decode()?)) } } -impl WireEncode for MlKemPublicKey { +impl Encode for MlKemPublicKey { fn encoded_len(&self) -> usize { Self::SIZE } @@ -128,13 +128,13 @@ impl Drop for MlKemCiphertext { } } -impl codec::WireDecode for MlKemCiphertext { - fn decode(reader: &mut codec::Reader) -> Result { +impl ql_codec::Decode for MlKemCiphertext { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self::new(reader.decode()?)) } } -impl WireEncode for MlKemCiphertext { +impl Encode for MlKemCiphertext { fn encoded_len(&self) -> usize { Self::SIZE } diff --git a/ql-wire/src/qid.rs b/ql-wire/src/qid.rs index 38672045..7d32a1f8 100644 --- a/ql-wire/src/qid.rs +++ b/ql-wire/src/qid.rs @@ -2,8 +2,6 @@ use ql_common::QID; use crate::{MlKemPublicKey, QlHash, ML_KEM_SUITE_TAG}; -array_wrapper_codec!(QID); - pub fn derive_qid(crypto: &impl QlHash, mlkem_public_key: &MlKemPublicKey) -> QID { let digest = crypto.sha256(&[ b"quantum-link qid v1", diff --git a/ql-wire/src/record.rs b/ql-wire/src/record.rs index 51a29085..b718acb7 100644 --- a/ql-wire/src/record.rs +++ b/ql-wire/src/record.rs @@ -1,31 +1,32 @@ +use ql_codec::{ByteSlice, Decode, Encode}; + use crate::{ - codec, encrypted_message::EncryptedMessage, handshake::{Ik1, Ik2, Kk1, Kk2, Xx1, Xx2, Xx3, Xx4}, - ByteSlice, RouteHeader, SessionHeader, WireDecode, WireEncode, WireError, QL_WIRE_VERSION, + Error, RouteHeader, SessionHeader, QL_WIRE_VERSION, }; pub fn encode_record(out: &mut W, header: RecordHeader, body: &T) where W: bytes::BufMut + ?Sized, - T: WireEncode + ?Sized, + T: Encode + ?Sized, { header.encode(out); body.encode(out); } -pub fn encode_record_vec(header: RecordHeader, body: &T) -> Vec { +pub fn encode_record_vec(header: RecordHeader, body: &T) -> Vec { let mut out = Vec::with_capacity(RecordHeader::WIRE_SIZE + body.encoded_len()); encode_record(&mut out, header, body); out } -pub fn decode_record(bytes: B) -> Result<(RecordHeader, T), WireError> +pub fn decode_record(bytes: B) -> Result<(RecordHeader, T), Error> where - T: WireDecode, + T: Decode, B: ByteSlice, { - let mut reader = codec::Reader::new(bytes); + let mut reader = ql_codec::Reader::new(bytes); Ok((reader.decode()?, reader.decode()?)) } @@ -48,8 +49,8 @@ impl RecordHeader { } } -impl WireDecode for RecordHeader { - fn decode(reader: &mut codec::Reader) -> Result { +impl Decode for RecordHeader { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { version: reader.decode()?, route: reader.decode()?, @@ -58,7 +59,7 @@ impl WireDecode for RecordHeader { } } -impl WireEncode for RecordHeader { +impl Encode for RecordHeader { fn encoded_len(&self) -> usize { Self::WIRE_SIZE } @@ -78,24 +79,24 @@ pub enum RecordType { } impl TryFrom for RecordType { - type Error = WireError; + type Error = ql_codec::Error; fn try_from(value: u8) -> Result { match value { 1 => Ok(Self::Handshake), 2 => Ok(Self::Session), - _ => Err(WireError::InvalidPayload), + _ => Err(ql_codec::Error::InvalidDiscriminant), } } } -impl WireDecode for RecordType { - fn decode(reader: &mut codec::Reader) -> Result { +impl Decode for RecordType { + fn decode(reader: &mut ql_codec::Reader) -> Result { reader.decode::()?.try_into() } } -impl WireEncode for RecordType { +impl Encode for RecordType { fn encoded_len(&self) -> usize { size_of::() } @@ -131,7 +132,7 @@ pub enum HandshakeKind { } impl TryFrom for HandshakeKind { - type Error = WireError; + type Error = ql_codec::Error; fn try_from(value: u8) -> Result { match value { @@ -143,18 +144,18 @@ impl TryFrom for HandshakeKind { 6 => Ok(Self::Xx2), 7 => Ok(Self::Xx3), 8 => Ok(Self::Xx4), - _ => Err(WireError::InvalidPayload), + _ => Err(ql_codec::Error::InvalidDiscriminant), } } } -impl WireDecode for HandshakeKind { - fn decode(reader: &mut codec::Reader) -> Result { +impl Decode for HandshakeKind { + fn decode(reader: &mut ql_codec::Reader) -> Result { reader.decode::()?.try_into() } } -impl WireEncode for HandshakeKind { +impl Encode for HandshakeKind { fn encoded_len(&self) -> usize { size_of::() } @@ -179,7 +180,7 @@ impl QlHandshakeRecord { } } -impl WireEncode for QlHandshakeRecord { +impl Encode for QlHandshakeRecord { fn encoded_len(&self) -> usize { self.kind().encoded_len() + match self { @@ -209,8 +210,8 @@ impl WireEncode for QlHandshakeRecord { } } -impl WireDecode for QlHandshakeRecord { - fn decode(reader: &mut codec::Reader) -> Result { +impl Decode for QlHandshakeRecord { + fn decode(reader: &mut ql_codec::Reader) -> Result { let kind = reader.decode::()?; match kind { HandshakeKind::Ik1 => Ok(Self::Ik1(reader.decode()?)), @@ -231,7 +232,7 @@ pub struct QlSessionRecord { pub payload: EncryptedMessage, } -impl> WireEncode for QlSessionRecord { +impl> Encode for QlSessionRecord { fn encoded_len(&self) -> usize { self.header.encoded_len() + self.payload.encoded_len() } @@ -251,8 +252,8 @@ impl QlSessionRecord { } } -impl WireDecode for QlSessionRecord { - fn decode(reader: &mut codec::Reader) -> Result { +impl Decode for QlSessionRecord { + fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { header: reader.decode()?, payload: reader.decode()?, diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index bc30f3c0..0873e2b7 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -1,6 +1,5 @@ -use std::ops::RangeInclusive; - -use ql_common::{ResetCode, StreamId, VarInt, QID}; +use ql_codec::{Decode, Encode}; +use ql_common::{ResetCode, StreamId, QID}; use super::*; @@ -13,14 +12,6 @@ fn decode_session_record(bytes: &[u8]) -> QlSessionRecord> { record.into_owned() } -fn varint(value: u64) -> VarInt { - VarInt::from_u64(value).unwrap() -} - -fn record_ack_range(start: u64, end: u64) -> RangeInclusive { - RecordSeq(varint(start))..=RecordSeq(varint(end)) -} - fn handshake_meta(id: u32) -> HandshakeMeta { HandshakeMeta { handshake_id: HandshakeId(id), @@ -77,7 +68,7 @@ fn peer_bundle_round_trip() { let bundle = identity.bundle(); let encoded = bundle.encode_vec(); - let decoded = PeerBundle::decode_exact(encoded.as_slice()).unwrap(); + let decoded = PeerBundle::decode_bytes(encoded.as_slice()).unwrap(); assert_eq!(decoded, bundle); assert_eq!(&*decoded.name, "alice"); @@ -177,7 +168,7 @@ fn ik_handshake_rejects_tampered_handshake_meta() { assert_eq!( initiator_state.read_2(&crypto, responder_to_initiator, &m2), - Err(WireError::InvalidHandshakeMeta) + Err(Error::InvalidHandshakeMeta) ); } @@ -214,7 +205,7 @@ fn kk_handshake_rejects_tampered_handshake_header() { assert_eq!( initiator_state.read_2(&crypto, tampered_route, &m2), - Err(WireError::InvalidPayload) + Err(Error::InvalidRouteHeader) ); } @@ -247,7 +238,7 @@ fn ik_handshake_rejects_tampered_transport_params() { assert_eq!( initiator_state.read_2(&crypto, responder_to_initiator, &m2), - Err(WireError::DecryptFailed) + Err(Error::DecryptFailed) ); } @@ -273,7 +264,7 @@ fn ik_handshake_rejects_tampered_handshake_header() { assert_eq!( responder_state.read_1(&crypto, initiator_to_responder, &m1), - Err(WireError::DecryptFailed) + Err(Error::DecryptFailed) ); } @@ -303,7 +294,7 @@ fn ik_handshake_rejects_bound_remote_bundle_mismatch() { assert_eq!( responder_state.read_1(&crypto, initiator_to_responder, &m1), - Err(WireError::InvalidPayload) + Err(Error::InvalidRouteHeader) ); } @@ -486,7 +477,7 @@ fn kk_handshake_rejects_tampered_transport_params() { assert_eq!( initiator_state.read_2(&crypto, responder_to_initiator, &m2), - Err(WireError::DecryptFailed) + Err(Error::DecryptFailed) ); } @@ -519,7 +510,7 @@ fn xx_handshake_rejects_tampered_pairing_id() { assert_eq!( responder_state.read_1(&crypto, initiator_to_responder, &m1), - Err(WireError::InvalidPairingId) + Err(Error::InvalidPairingId) ); } @@ -552,7 +543,7 @@ fn xx_handshake_rejects_tampered_sender_or_recipient() { assert_eq!( responder_state.read_1(&crypto, route, &m1), - Err(WireError::InvalidRouteHeader) + Err(Error::InvalidRouteHeader) ); let mut initiator_state = XxHandshake::new_initiator( @@ -578,7 +569,7 @@ fn xx_handshake_rejects_tampered_sender_or_recipient() { assert_eq!( responder_state.read_1(&crypto, route, &m1), - Err(WireError::InvalidRouteHeader) + Err(Error::InvalidRouteHeader) ); } @@ -625,7 +616,7 @@ fn xx_handshake_rejects_repeated_transport_param_change() { assert_eq!( responder_state.read_3(&crypto, initiator_to_responder, &m3), - Err(WireError::InvalidTransportParams) + Err(Error::InvalidTransportParams) ); } @@ -709,29 +700,28 @@ fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { #[test] fn encrypted_session_record_round_trip_authenticates_header() { let crypto = SoftwareCrypto; - let header = SessionHeader { - seq: RecordSeq(varint(11)), - }; + let header = SessionHeader { seq: RecordSeq(11) }; let session_route = route(1, 2); let body = vec![ SessionFrame::Ping, SessionFrame::Unpair, SessionFrame::Ack( - RecordAck::from_ranges([record_ack_range(20, 23), record_ack_range(12, 13)]).unwrap(), + RecordAck::from_ranges([RecordSeq(20)..=RecordSeq(23), RecordSeq(12)..=RecordSeq(13)]) + .unwrap(), ), SessionFrame::StreamWindow(StreamWindow { - stream_id: StreamId(varint(9)), - maximum_offset: varint(65_536), + stream_id: StreamId(9), + maximum_offset: 65_536, }), SessionFrame::StreamData(StreamData { - stream_id: StreamId(varint(9)), - offset: varint(1024), + stream_id: StreamId(9), + offset: 1024, header: None, bytes: b"hello".to_vec(), fin: true, }), SessionFrame::StreamReset(StreamReset { - stream_id: StreamId(varint(9)), + stream_id: StreamId(9), target: ResetTarget::Both, code: ResetCode::CANCELLED, }), @@ -775,11 +765,11 @@ fn encrypted_session_record_round_trip_authenticates_header() { encrypted.clone(), &session_key, ), - Err(WireError::DecryptFailed) + Err(Error::DecryptFailed) ); let wrong_seq_header = SessionHeader { - seq: RecordSeq(varint(header.seq.0.into_inner() + 1)), + seq: RecordSeq(header.seq.0 + 1), }; assert_eq!( encrypted::decrypt_record( @@ -789,7 +779,7 @@ fn encrypted_session_record_round_trip_authenticates_header() { encrypted, &session_key, ), - Err(WireError::DecryptFailed) + Err(Error::DecryptFailed) ); } @@ -799,6 +789,10 @@ fn protocol_record_size_breakdown() { println!("{label:<32}: {size} bytes"); } + fn record_size(record: &impl Encode) -> usize { + RecordHeader::WIRE_SIZE + record.encoded_len() + } + let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); @@ -900,42 +894,35 @@ fn protocol_record_size_breakdown() { let session_ping = encrypt_record( &crypto, session_route, - SessionHeader { - seq: RecordSeq(varint(1)), - }, + SessionHeader { seq: RecordSeq(1) }, &session.tx_key, &[SessionFrame::Ping], ); let session_ack = encrypt_record( &crypto, session_route, - SessionHeader { - seq: RecordSeq(varint(2)), - }, + SessionHeader { seq: RecordSeq(2) }, &session.tx_key, &[SessionFrame::Ack( - RecordAck::from_ranges([record_ack_range(6, 6), record_ack_range(1, 2)]).unwrap(), + RecordAck::from_ranges([RecordSeq(6)..=RecordSeq(6), RecordSeq(1)..=RecordSeq(2)]) + .unwrap(), )], ); let session_unpair = encrypt_record( &crypto, session_route, - SessionHeader { - seq: RecordSeq(varint(3)), - }, + SessionHeader { seq: RecordSeq(3) }, &session.tx_key, &[SessionFrame::Unpair], ); let session_stream_empty = encrypt_record( &crypto, session_route, - SessionHeader { - seq: RecordSeq(varint(4)), - }, + SessionHeader { seq: RecordSeq(4) }, &session.tx_key, &[SessionFrame::StreamData(StreamData { - stream_id: StreamId(varint(1)), - offset: varint(0), + stream_id: StreamId(1), + offset: 0, header: None, fin: false, bytes: Vec::new(), @@ -944,9 +931,7 @@ fn protocol_record_size_breakdown() { let session_close = encrypt_record( &crypto, session_route, - SessionHeader { - seq: RecordSeq(varint(5)), - }, + SessionHeader { seq: RecordSeq(5) }, &session.tx_key, &[SessionFrame::Close(SessionClose { code: SessionCloseCode::PROTOCOL, @@ -956,20 +941,20 @@ fn protocol_record_size_breakdown() { print_size("ql-wire peer bundle", initiator.bundle().encode_vec().len()); print_size("ql-wire mlkem public key", MlKemPublicKey::SIZE); print_size("ql-wire mlkem ciphertext", MlKemCiphertext::SIZE); - print_size("ql-wire pq ik1", ik1.encode_vec().len()); - print_size("ql-wire pq ik2", ik2.encode_vec().len()); - print_size("ql-wire pq kk1", kk1.encode_vec().len()); - print_size("ql-wire pq kk2", kk2.encode_vec().len()); - print_size("ql-wire pq xx1", xx1.encode_vec().len()); - print_size("ql-wire pq xx2", xx2.encode_vec().len()); - print_size("ql-wire pq xx3", xx3.encode_vec().len()); - print_size("ql-wire pq xx4", xx4.encode_vec().len()); - print_size("ql-wire session ping", session_ping.encode_vec().len()); - print_size("ql-wire session ack", session_ack.encode_vec().len()); - print_size("ql-wire session unpair", session_unpair.encode_vec().len()); + print_size("ql-wire pq ik1", record_size(&ik1)); + print_size("ql-wire pq ik2", record_size(&ik2)); + print_size("ql-wire pq kk1", record_size(&kk1)); + print_size("ql-wire pq kk2", record_size(&kk2)); + print_size("ql-wire pq xx1", record_size(&xx1)); + print_size("ql-wire pq xx2", record_size(&xx2)); + print_size("ql-wire pq xx3", record_size(&xx3)); + print_size("ql-wire pq xx4", record_size(&xx4)); + print_size("ql-wire session ping", record_size(&session_ping)); + print_size("ql-wire session ack", record_size(&session_ack)); + print_size("ql-wire session unpair", record_size(&session_unpair)); print_size( "ql-wire session stream empty", - session_stream_empty.encode_vec().len(), + record_size(&session_stream_empty), ); - print_size("ql-wire session close", session_close.encode_vec().len()); + print_size("ql-wire session close", record_size(&session_close)); } diff --git a/ql-wire/src/varint.rs b/ql-wire/src/varint.rs deleted file mode 100644 index 02f8d2a3..00000000 --- a/ql-wire/src/varint.rs +++ /dev/null @@ -1,23 +0,0 @@ -use bytes::BufMut; -use ql_common::VarInt; - -use crate::{ByteSlice, Reader, WireDecode, WireEncode, WireError}; - -impl WireDecode for ql_common::VarInt { - fn decode(reader: &mut Reader) -> Result { - let first = reader.decode::()?; - let len = VarInt::encoded_len_from_first_byte(first); - let tail = reader.take_bytes(len - 1)?; - VarInt::decode_with_first_byte(first, &tail).ok_or(WireError::InvalidPayload) - } -} - -impl WireEncode for ql_common::VarInt { - fn encoded_len(&self) -> usize { - self.size() - } - - fn encode(&self, out: &mut W) { - self.write_bytes(|bytes| out.put_slice(bytes)); - } -} From 9ecb1f547cdd74fcc85def89e6b2d4a0b9b4a0e9 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 9 Jul 2026 16:59:51 -0400 Subject: [PATCH 42/59] ql-fsm: update --- Cargo.lock | 1 + ql-fsm/Cargo.toml | 1 + ql-fsm/src/error.rs | 78 ++++--- ql-fsm/src/fsm.rs | 19 +- ql-fsm/src/handshake/ik.rs | 21 +- ql-fsm/src/handshake/kk.rs | 21 +- ql-fsm/src/handshake/xx.rs | 35 ++- ql-fsm/src/pairing.rs | 14 +- ql-fsm/src/session/ack_tracker.rs | 46 ++-- ql-fsm/src/session/mod.rs | 39 ++-- ql-fsm/src/session/remote_stream_history.rs | 1 - ql-fsm/src/session/stream_parity.rs | 8 +- ql-fsm/src/session/stream_tx.rs | 2 +- ql-fsm/src/session/tests.rs | 246 +++++++++++--------- ql-fsm/src/tests/handshake.rs | 10 +- ql-fsm/src/tests/proptest.rs | 28 ++- ql-fsm/src/tests/session.rs | 10 +- 17 files changed, 304 insertions(+), 276 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index ec0a4a36..3cff966f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1072,6 +1072,7 @@ dependencies = [ "bytes", "indexmap", "proptest", + "ql-codec", "ql-common", "ql-wire", ] diff --git a/ql-fsm/Cargo.toml b/ql-fsm/Cargo.toml index 93006d35..4b5a48cc 100644 --- a/ql-fsm/Cargo.toml +++ b/ql-fsm/Cargo.toml @@ -9,6 +9,7 @@ license = "Proprietary" bytes = { workspace = true } indexmap = "2" ql-common = { workspace = true } +ql-codec = { workspace = true } ql-wire = { workspace = true } [dev-dependencies] diff --git a/ql-fsm/src/error.rs b/ql-fsm/src/error.rs index 94448778..8614a3c3 100644 --- a/ql-fsm/src/error.rs +++ b/ql-fsm/src/error.rs @@ -3,58 +3,78 @@ use std::{ fmt::{Display, Formatter}, }; -use ql_wire::{PairingId, WireError}; - #[derive(Debug, Clone, PartialEq, Eq)] pub enum ReceiveError { - InvalidRecordHeader(WireError), + Wire { + stage: ReceiveStage, + source: ql_wire::Error, + }, InvalidRecordVersion, - InvalidHandshakeRecord(WireError), - InvalidSessionRecord(WireError), - InvalidSessionPayload(WireError), - InvalidIkHandshake(WireError), - InvalidKkHandshake(WireError), - InvalidXxHandshake(WireError), InvalidRemoteBundle, InvalidQid, NoPeer, NoSession, NotPairingMode, - InvalidPairingId { - expected: PairingId, - actual: PairingId, - }, + InvalidPairingId, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ReceiveStage { + RecordHeader, + HandshakeRecord, + SessionRecord, + SessionPayload, + IkHandshake, + KkHandshake, + XxHandshake, +} + +impl ReceiveError { + pub(crate) fn wire(stage: ReceiveStage, source: impl Into) -> Self { + Self::Wire { + stage, + source: source.into(), + } + } } impl Display for ReceiveError { fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { match self { - Self::InvalidRecordHeader(error) => write!(f, "invalid record header: {error}"), + Self::Wire { stage, source } => write!(f, "invalid {stage}: {source}"), Self::InvalidRecordVersion => f.write_str("invalid record version"), - Self::InvalidHandshakeRecord(error) => { - write!(f, "invalid handshake record: {error}") - } - Self::InvalidSessionRecord(error) => write!(f, "invalid session record: {error}"), - Self::InvalidSessionPayload(error) => write!(f, "invalid session payload: {error}"), - Self::InvalidIkHandshake(error) => write!(f, "invalid ik handshake: {error}"), - Self::InvalidKkHandshake(error) => write!(f, "invalid kk handshake: {error}"), - Self::InvalidXxHandshake(error) => write!(f, "invalid xx handshake: {error}"), Self::InvalidRemoteBundle => f.write_str("invalid remote bundle"), Self::InvalidQid => f.write_str("invalid qid"), Self::NoPeer => f.write_str("no bound peer"), Self::NoSession => f.write_str("no active session"), Self::NotPairingMode => f.write_str("not in pairing mode"), - Self::InvalidPairingId { expected, actual } => { - write!( - f, - "invalid pairing id: expected {expected}, actual {actual}" - ) - } + Self::InvalidPairingId => f.write_str("invalid pairing id"), } } } -impl std::error::Error for ReceiveError {} +impl Display for ReceiveStage { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + match self { + Self::RecordHeader => f.write_str("record header"), + Self::HandshakeRecord => f.write_str("handshake record"), + Self::SessionRecord => f.write_str("session record"), + Self::SessionPayload => f.write_str("session payload"), + Self::IkHandshake => f.write_str("ik handshake"), + Self::KkHandshake => f.write_str("kk handshake"), + Self::XxHandshake => f.write_str("xx handshake"), + } + } +} + +impl std::error::Error for ReceiveError { + fn source(&self) -> Option<&(dyn Error + 'static)> { + match self { + Self::Wire { source, .. } => Some(source), + _ => None, + } + } +} impl From for ReceiveError { fn from(_: NoSessionError) -> Self { diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index df74501e..0ff3ca9a 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -1,14 +1,16 @@ use std::{collections::VecDeque, time::Instant}; use bytes::Bytes; +use ql_codec::{Decode, Reader}; use ql_common::StreamId; -use ql_wire::{self as wire, QlCrypto, SessionCloseCode, WireDecode}; +use ql_wire::{self as wire, QlCrypto, SessionCloseCode}; use crate::{ handshake, session::{self, SessionEvent, TerminalFrame}, state::LinkState, - Event, NoPeerError, NoSessionError, OutboundWrite, QlFsm, ReceiveError, StreamError, WriteId, + Event, NoPeerError, NoSessionError, OutboundWrite, QlFsm, ReceiveError, ReceiveStage, + StreamError, WriteId, }; pub struct EventSink<'a> { @@ -103,9 +105,9 @@ pub fn receive( mut bytes: Vec, crypto: &impl QlCrypto, ) -> Result<(), ReceiveError> { - let mut reader = wire::Reader::new(bytes.as_mut_slice()); - let header = - wire::RecordHeader::decode(&mut reader).map_err(ReceiveError::InvalidRecordHeader)?; + let mut reader = Reader::new(bytes.as_mut_slice()); + let header = wire::RecordHeader::decode(&mut reader) + .map_err(|error| ReceiveError::wire(ReceiveStage::RecordHeader, error))?; if header.version != wire::QL_WIRE_VERSION { return Err(ReceiveError::InvalidRecordVersion); @@ -117,7 +119,7 @@ pub fn receive( match header.record_type { wire::RecordType::Handshake => { let record = wire::QlHandshakeRecord::decode(&mut reader) - .map_err(ReceiveError::InvalidHandshakeRecord)?; + .map_err(|error| ReceiveError::wire(ReceiveStage::HandshakeRecord, error))?; handshake::handle_handshake_record(fsm, crypto, header.route, &record) } wire::RecordType::Session => { @@ -129,7 +131,7 @@ pub fn receive( } let (decrypt_len, seq) = { let record = wire::QlSessionRecord::decode(&mut reader) - .map_err(ReceiveError::InvalidSessionRecord)?; + .map_err(|error| ReceiveError::wire(ReceiveStage::SessionRecord, error))?; let payload = wire::decrypt_record( crypto, &header, @@ -137,10 +139,11 @@ pub fn receive( record.payload, &conn.transport.rx_key, ) - .map_err(ReceiveError::InvalidSessionPayload)?; + .map_err(|error| ReceiveError::wire(ReceiveStage::SessionPayload, error))?; (payload.len(), record.header.seq) }; + drop(reader); let len = bytes.len(); let plaintext = Bytes::from(bytes).slice(len - decrypt_len..); let frames = wire::parse_session_frames(plaintext); diff --git a/ql-fsm/src/handshake/ik.rs b/ql-fsm/src/handshake/ik.rs index 9349718a..1b482d78 100644 --- a/ql-fsm/src/handshake/ik.rs +++ b/ql-fsm/src/handshake/ik.rs @@ -5,7 +5,7 @@ use super::{ }; use crate::{ state::{InitiatorState, LinkState}, - QlFsm, ReceiveError, + QlFsm, ReceiveError, ReceiveStage, }; pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle) { @@ -57,16 +57,14 @@ pub fn handle_ik1( ); handshake .read_1(crypto, route, message) - .map_err(ReceiveError::InvalidIkHandshake)?; + .map_err(wire_error)?; let outbound = handshake .write_2(crypto, message.meta) - .map_err(ReceiveError::InvalidIkHandshake)?; + .map_err(wire_error)?; establish_session( fsm, message.meta.handshake_id, - handshake - .finalize(crypto) - .map_err(ReceiveError::InvalidIkHandshake)?, + handshake.finalize(crypto).map_err(wire_error)?, )?; fsm.state.handshake = None; enqueue_handshake( @@ -98,7 +96,7 @@ pub fn handle_ik2( state .handshake .read_2(crypto, route, message) - .map_err(ReceiveError::InvalidIkHandshake)?; + .map_err(wire_error)?; } let LinkState::IkInitiator(state) = fsm.state.link.take() else { @@ -107,10 +105,7 @@ pub fn handle_ik2( establish_session( fsm, message.meta.handshake_id, - state - .handshake - .finalize(crypto) - .map_err(ReceiveError::InvalidIkHandshake)?, + state.handshake.finalize(crypto).map_err(wire_error)?, ) } @@ -131,3 +126,7 @@ pub fn should_ignore_inbound(fsm: &QlFsm, route: RouteHeader, message: &Ik1) -> } } } + +fn wire_error(source: ql_wire::Error) -> ReceiveError { + ReceiveError::wire(ReceiveStage::IkHandshake, source) +} diff --git a/ql-fsm/src/handshake/kk.rs b/ql-fsm/src/handshake/kk.rs index 243766ae..01468be7 100644 --- a/ql-fsm/src/handshake/kk.rs +++ b/ql-fsm/src/handshake/kk.rs @@ -5,7 +5,7 @@ use super::{ }; use crate::{ state::{InitiatorState, LinkState}, - QlFsm, ReceiveError, + QlFsm, ReceiveError, ReceiveStage, }; pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle) { @@ -59,16 +59,14 @@ pub fn handle_kk1( ); handshake .read_1(crypto, route, message) - .map_err(ReceiveError::InvalidKkHandshake)?; + .map_err(wire_error)?; let outbound = handshake .write_2(crypto, message.meta) - .map_err(ReceiveError::InvalidKkHandshake)?; + .map_err(wire_error)?; establish_session( fsm, message.meta.handshake_id, - handshake - .finalize(crypto) - .map_err(ReceiveError::InvalidKkHandshake)?, + handshake.finalize(crypto).map_err(wire_error)?, )?; fsm.state.handshake = None; enqueue_handshake( @@ -100,7 +98,7 @@ pub fn handle_kk2( state .handshake .read_2(crypto, route, message) - .map_err(ReceiveError::InvalidKkHandshake)?; + .map_err(wire_error)?; } let LinkState::KkInitiator(state) = fsm.state.link.take() else { @@ -109,10 +107,7 @@ pub fn handle_kk2( establish_session( fsm, message.meta.handshake_id, - state - .handshake - .finalize(crypto) - .map_err(ReceiveError::InvalidKkHandshake)?, + state.handshake.finalize(crypto).map_err(wire_error)?, ) } @@ -131,3 +126,7 @@ pub fn should_ignore_inbound(fsm: &QlFsm, route: RouteHeader, message: &Kk1) -> } } } + +fn wire_error(source: ql_wire::Error) -> ReceiveError { + ReceiveError::wire(ReceiveStage::KkHandshake, source) +} diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs index 7ef17f99..df6cebcc 100644 --- a/ql-fsm/src/handshake/xx.rs +++ b/ql-fsm/src/handshake/xx.rs @@ -8,7 +8,7 @@ use super::{ }; use crate::{ state::{InitiatorState, LinkState, XxResponderState}, - QlFsm, ReceiveError, + QlFsm, ReceiveError, ReceiveStage, }; pub fn start_initiator( @@ -52,10 +52,7 @@ pub fn handle_xx1( } match fsm.state.armed_pairing_token { Some(expected) if expected.id(crypto) != message.pairing_id => { - Err(ReceiveError::InvalidPairingId { - expected: expected.id(crypto), - actual: message.pairing_id, - }) + Err(ReceiveError::InvalidPairingId) } Some(token) => { reset_connected_session_if_needed(fsm); @@ -69,10 +66,10 @@ pub fn handle_xx1( ); handshake .read_1(crypto, route, message) - .map_err(ReceiveError::InvalidXxHandshake)?; + .map_err(wire_error)?; let outbound = handshake .write_2(crypto, message.meta) - .map_err(ReceiveError::InvalidXxHandshake)?; + .map_err(wire_error)?; fsm.state.link = LinkState::XxResponder(XxResponderState { handshake, handshake_meta: message.meta, @@ -111,11 +108,11 @@ pub fn handle_xx2( state .handshake .read_2(crypto, route, message) - .map_err(ReceiveError::InvalidXxHandshake)?; + .map_err(wire_error)?; let outbound = state .handshake .write_3(crypto, message.meta) - .map_err(ReceiveError::InvalidXxHandshake)?; + .map_err(wire_error)?; fsm.state.handshake = None; enqueue_handshake( fsm, @@ -147,7 +144,7 @@ pub fn handle_xx3( state .handshake .read_3(crypto, route, message) - .map_err(ReceiveError::InvalidXxHandshake)?; + .map_err(wire_error)?; let handshake_meta = state.handshake_meta; let LinkState::XxResponder(mut state) = fsm.state.link.take() else { unreachable!("active XX responder was checked above"); @@ -155,7 +152,7 @@ pub fn handle_xx3( let outbound = state .handshake .write_4(crypto, handshake_meta) - .map_err(ReceiveError::InvalidXxHandshake)?; + .map_err(wire_error)?; fsm.state.handshake = None; enqueue_handshake( fsm, @@ -168,10 +165,7 @@ pub fn handle_xx3( establish_session( fsm, message.meta.handshake_id, - state - .handshake - .finalize(crypto) - .map_err(ReceiveError::InvalidXxHandshake)?, + state.handshake.finalize(crypto).map_err(wire_error)?, ) } @@ -193,7 +187,7 @@ pub fn handle_xx4( state .handshake .read_4(crypto, route, message) - .map_err(ReceiveError::InvalidXxHandshake)?; + .map_err(wire_error)?; } let LinkState::XxInitiator(state) = fsm.state.link.take() else { @@ -202,10 +196,7 @@ pub fn handle_xx4( establish_session( fsm, message.meta.handshake_id, - state - .handshake - .finalize(crypto) - .map_err(ReceiveError::InvalidXxHandshake)?, + state.handshake.finalize(crypto).map_err(wire_error)?, ) } @@ -239,3 +230,7 @@ pub fn should_ignore_inbound( } } } + +fn wire_error(source: ql_wire::Error) -> ReceiveError { + ReceiveError::wire(ReceiveStage::XxHandshake, source) +} diff --git a/ql-fsm/src/pairing.rs b/ql-fsm/src/pairing.rs index c1f239a1..3f179d0d 100644 --- a/ql-fsm/src/pairing.rs +++ b/ql-fsm/src/pairing.rs @@ -1,5 +1,6 @@ +use ql_codec::{ByteSlice, Decode, Encode, Reader}; use ql_common::QID; -use ql_wire::{ByteSlice, PairingToken, Reader, WireDecode, WireEncode, WireError}; +use ql_wire::PairingToken; /// Out-of-band invite consumed by the initiator of an XX pairing #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -10,12 +11,11 @@ pub struct PairingInvite { impl PairingInvite { pub const VERSION: u8 = 1; - pub const WIRE_SIZE: usize = size_of::() + QID::SIZE + PairingToken::SIZE; } -impl WireEncode for PairingInvite { +impl Encode for PairingInvite { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + size_of::() + QID::SIZE + PairingToken::SIZE } fn encode(&self, out: &mut W) { @@ -25,10 +25,10 @@ impl WireEncode for PairingInvite { } } -impl WireDecode for PairingInvite { - fn decode(reader: &mut Reader) -> Result { +impl Decode for PairingInvite { + fn decode(reader: &mut Reader) -> Result { if reader.decode::()? != Self::VERSION { - return Err(WireError::InvalidPayload); + return Err(ql_codec::Error::InvalidDiscriminant); } Ok(Self { diff --git a/ql-fsm/src/session/ack_tracker.rs b/ql-fsm/src/session/ack_tracker.rs index b22dac36..7f947d2d 100644 --- a/ql-fsm/src/session/ack_tracker.rs +++ b/ql-fsm/src/session/ack_tracker.rs @@ -45,7 +45,7 @@ impl AckTracker { } pub fn insert(&mut self, seq: RecordSeq) -> ReceiveOutcome { - let seq = seq.0.into_inner(); + let seq = seq.0; let largest_accepted = self.accepted_records.max(); if largest_accepted.is_some_and(|largest| seq < self.accepted_cutoff(largest)) { return ReceiveOutcome::TooOld; @@ -163,12 +163,12 @@ fn single_range(seq: u64) -> std::ops::Range { fn to_ack_range(range: std::ops::Range) -> RangeInclusive { let end = range.end.checked_sub(1).unwrap(); - RecordSeq::from_u64(range.start).unwrap()..=RecordSeq::from_u64(end).unwrap() + RecordSeq(range.start)..=RecordSeq(end) } fn from_ack_range(range: RangeInclusive) -> std::ops::Range { - let start = range.start().0.into_inner(); - let end = range.end().0.into_inner().checked_add(1).unwrap(); + let start = range.start().0; + let end = range.end().0.checked_add(1).unwrap(); start..end } @@ -180,15 +180,11 @@ mod tests { use super::{AckTracker, PendingAck, ReceiveOutcome}; - fn seq(value: u64) -> RecordSeq { - RecordSeq::from_u64(value).unwrap() - } - fn ack_ranges(pending_ack: &PendingAck) -> Vec<(u64, u64)> { pending_ack .ack .ranges() - .map(|range| (range.start().0.into_inner(), range.end().0.into_inner())) + .map(|range| (range.start().0, range.end().0)) .collect() } @@ -197,9 +193,9 @@ mod tests { let now = Instant::now(); let mut ack_tracker = AckTracker::new(128, 8); - assert_eq!(ack_tracker.insert(seq(10)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(11)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(12)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(10)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(11)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(12)), ReceiveOutcome::New); ack_tracker.schedule_ack(now); let pending_ack = ack_tracker.pending_ack(usize::MAX).unwrap(); @@ -211,10 +207,10 @@ mod tests { let now = Instant::now(); let mut ack_tracker = AckTracker::new(128, 8); - assert_eq!(ack_tracker.insert(seq(10)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(15)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(16)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(12)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(10)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(15)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(16)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(12)), ReceiveOutcome::New); ack_tracker.schedule_ack(now + Duration::from_millis(5)); let pending_ack = ack_tracker.pending_ack(usize::MAX).unwrap(); @@ -225,10 +221,10 @@ mod tests { fn accepted_record_window_evicts_old_sequences() { let mut ack_tracker = AckTracker::new(4, 8); - assert_eq!(ack_tracker.insert(seq(10)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(15)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(10)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(15)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(10)), ReceiveOutcome::TooOld); + assert_eq!(ack_tracker.insert(RecordSeq(10)), ReceiveOutcome::TooOld); } #[test] @@ -236,9 +232,9 @@ mod tests { let now = Instant::now(); let mut ack_tracker = AckTracker::new(128, 2); - assert_eq!(ack_tracker.insert(seq(1)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(3)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(5)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(1)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(3)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(5)), ReceiveOutcome::New); ack_tracker.schedule_ack(now); let pending_ack = ack_tracker.pending_ack(usize::MAX).unwrap(); @@ -250,9 +246,9 @@ mod tests { let now = Instant::now(); let mut ack_tracker = AckTracker::new(128, 8); - assert_eq!(ack_tracker.insert(seq(1)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(3)), ReceiveOutcome::New); - assert_eq!(ack_tracker.insert(seq(5)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(1)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(3)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(5)), ReceiveOutcome::New); ack_tracker.schedule_ack(now); let first_ack = ack_tracker.pending_ack(4).unwrap(); diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index 6105be0a..e9015401 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -17,10 +17,10 @@ use std::time::{Duration, Instant}; use bytes::Bytes; use indexmap::IndexMap; -use ql_common::{StreamId, VarInt}; +use ql_common::StreamId; use ql_wire::{ - LenBytes, RecordAck, RecordSeq, ResetTarget, SessionClose, SessionCloseCode, SessionFrame, - SessionRecordBuilder, StreamData, StreamReset, StreamWindow, WireError, + RecordAck, RecordSeq, ResetTarget, SessionClose, SessionCloseCode, SessionFrame, + SessionRecordBuilder, StreamData, StreamReset, StreamWindow, }; use self::{ @@ -111,7 +111,7 @@ impl SessionFsm { last_inbound_at: now, phase: SessionPhase::Open, next_stream_ordinal: 0, - next_record_seq: RecordSeq::from_u32(0), + next_record_seq: RecordSeq(0), next_write_id: 0, tracked_records: IndexMap::default(), ack_tracker: AckTracker::new( @@ -202,7 +202,7 @@ impl SessionFsm { pub fn receive(&mut self, now: Instant, seq: RecordSeq, frames: I, sink: &mut impl EventSink) where - I: IntoIterator, WireError>>, + I: IntoIterator, ql_wire::Error>>, { if self.state.phase != SessionPhase::Open { return; @@ -490,17 +490,17 @@ impl SessionFsm { } let frame = StreamWindow { stream_id, - maximum_offset: VarInt::from_u64(stream.recv_limit()).unwrap(), + maximum_offset: stream.recv_limit(), }; if !builder.push_stream_window(&frame) { break; } stream.pending_window = false; - stream.advertised_max_offset = frame.maximum_offset.into_inner(); + stream.advertised_max_offset = frame.maximum_offset; outbound .window_updates - .push((stream_id, frame.maximum_offset.into_inner())); + .push((stream_id, frame.maximum_offset)); } } @@ -533,13 +533,11 @@ impl SessionFsm { else { continue; }; - let offset = - VarInt::from_u64(candidate.offset).expect("stream offsets must fit ql-wire varint"); let frame = StreamData { stream_id, - offset, + offset: candidate.offset, header: if matches!(stream.role, StreamRole::Initiator) && candidate.offset == 0 { - stream.header.as_ref().map(LenBytes) + stream.header.as_deref() } else { None }, @@ -580,7 +578,7 @@ impl SessionFsm { .state .tracked_records .extract_if(.., |_, record| { - record.sent_at.is_some() && ack.contains(record.seq.0.into_inner()) + record.sent_at.is_some() && ack.contains(record.seq.0) }) .map(|(_, record)| record) .collect::>(); @@ -648,7 +646,7 @@ impl SessionFsm { }, }; - let frame_offset = offset.into_inner(); + let frame_offset = offset; let Some(frame_end) = frame_offset.checked_add(bytes.len() as u64) else { return Err(()); }; @@ -657,7 +655,7 @@ impl SessionFsm { let opened = match (stream.role, stream.header.as_ref(), header, frame_offset) { (StreamRole::Responder, None, Some(header), 0) => { - stream.header = Some(header.0); + stream.header = Some(header); true } (StreamRole::Initiator, _, Some(_), _) @@ -726,7 +724,7 @@ impl SessionFsm { }; let was_full = stream.send_capacity(self.config.stream_send_buffer_size) == 0; - let maximum_offset = frame.maximum_offset.into_inner(); + let maximum_offset = frame.maximum_offset; if maximum_offset > stream.peer_max_offset { stream.peer_max_offset = maximum_offset; } @@ -948,11 +946,7 @@ fn local_stream_was_opened( stream_id: StreamId, ) -> bool { local_parity.matches(stream_id) - && stream_id.0.into_inner() - < local_parity - .make_stream_id(next_stream_ordinal) - .0 - .into_inner() + && stream_id.0 < local_parity.make_stream_id(next_stream_ordinal).0 } fn restore_tracked_record( @@ -1043,8 +1037,7 @@ fn acknowledge_tracked_frame( fn next_seq(seq: &mut RecordSeq) { *seq = seq .0 - .into_inner() .checked_add(1) - .and_then(|next| RecordSeq::from_u64(next).ok()) + .map(RecordSeq) .expect("record sequence overflow"); } diff --git a/ql-fsm/src/session/remote_stream_history.rs b/ql-fsm/src/session/remote_stream_history.rs index e9b0f5e0..1983680f 100644 --- a/ql-fsm/src/session/remote_stream_history.rs +++ b/ql-fsm/src/session/remote_stream_history.rs @@ -29,7 +29,6 @@ impl RemoteStreamHistory { fn stream_ordinal(&self, stream_id: StreamId) -> Option { let delta = stream_id .0 - .into_inner() .checked_sub(u64::from(self.parity.first_stream_id()))?; if delta % 2 != 0 { return None; diff --git a/ql-fsm/src/session/stream_parity.rs b/ql-fsm/src/session/stream_parity.rs index 8d2d74b2..7ad63957 100644 --- a/ql-fsm/src/session/stream_parity.rs +++ b/ql-fsm/src/session/stream_parity.rs @@ -1,4 +1,4 @@ -use ql_common::{StreamId, VarInt, QID}; +use ql_common::{StreamId, QID}; #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum StreamParity { @@ -23,8 +23,8 @@ impl StreamParity { pub const fn matches(self, stream_id: StreamId) -> bool { match self { - Self::Even => stream_id.0.into_inner() % 2 == 0, - Self::Odd => stream_id.0.into_inner() % 2 == 1, + Self::Even => stream_id.0 % 2 == 0, + Self::Odd => stream_id.0 % 2 == 1, } } @@ -36,7 +36,7 @@ impl StreamParity { } pub fn make_stream_id(self, ordinal: u32) -> StreamId { - StreamId(VarInt::from_u32( + StreamId(u64::from( self.first_stream_id() .saturating_add(ordinal.saturating_mul(2)), )) diff --git a/ql-fsm/src/session/stream_tx.rs b/ql-fsm/src/session/stream_tx.rs index 15533922..d929bb5f 100644 --- a/ql-fsm/src/session/stream_tx.rs +++ b/ql-fsm/src/session/stream_tx.rs @@ -1,7 +1,7 @@ use std::{collections::VecDeque, ops::Range}; use bytes::{Buf, Bytes}; -use ql_wire::BufView; +use ql_codec::BufView; use super::range_set::RangeSet; diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index 48e37ed5..bb98b82b 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -1,47 +1,20 @@ use std::time::{Duration, Instant}; use bytes::Bytes; -use ql_common::{ResetCode, StreamId, VarInt, QID}; +use ql_common::{ResetCode, StreamId, QID}; use ql_wire::{ - decode_session_frames, parse_session_frames, LenBytes, RecordAck, RecordSeq, ResetTarget, - SessionFrame, SessionRecordBuilder, StreamData, StreamReset, + decode_session_frames, parse_session_frames, RecordAck, RecordSeq, ResetTarget, SessionFrame, + SessionRecordBuilder, StreamData, StreamReset, }; use super::{SessionConfig, SessionEvent, SessionFsm}; use crate::{session::stream_parity::StreamParity, StreamResetEvent}; -fn seq(value: u64) -> RecordSeq { - RecordSeq::from_u64(value).unwrap() -} - -fn stream_id(value: u64) -> StreamId { - StreamId(VarInt::from_u64(value).unwrap()) -} - -fn offset(value: u64) -> VarInt { - VarInt::from_u64(value).unwrap() -} - -fn record_ack(seq: RecordSeq) -> RecordAck { - RecordAck::from_ranges([seq..=seq]).unwrap() -} - const REFUSED: ResetCode = ResetCode(1); const TIMEOUT: ResetCode = ResetCode(2); -fn header(value: u64) -> LenBytes> { - LenBytes(vec![value as u8]) -} - -// todo: remove -fn opened(stream_id: StreamId) -> SessionEvent { - SessionEvent::Opened(stream_id) -} - fn open_stream_id(fsm: &mut SessionFsm) -> StreamId { - fsm.open_stream(header(1).0.into_boxed_slice(), |_| {}) - .unwrap() - .stream_id() + fsm.open_stream(Box::from([1]), |_| {}).unwrap().stream_id() } fn write_stream_bytes(fsm: &mut SessionFsm, stream_id: StreamId, bytes: &[u8]) -> usize { @@ -129,8 +102,8 @@ fn outbound_record_seq_increments_monotonically() { assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"two"), 3); let (second_seq, _) = next_outbound(&mut fsm, now + Duration::from_millis(1)).unwrap(); - assert_eq!(first_seq, seq(0)); - assert_eq!(second_seq, seq(1)); + assert_eq!(first_seq, RecordSeq(0)); + assert_eq!(second_seq, RecordSeq(1)); } #[test] @@ -218,8 +191,10 @@ fn ack_reopens_write_capacity() { let mut emit = |event| events.push(event); fsm.receive( now + Duration::from_millis(1), - seq(9), - std::iter::once(Ok(SessionFrame::Ack(record_ack(record_seq)))), + RecordSeq(9), + std::iter::once(Ok(SessionFrame::Ack( + RecordAck::from_ranges([record_seq..=record_seq]).unwrap(), + ))), &mut emit, ); @@ -255,8 +230,10 @@ fn ack_of_fin_emits_outbound_finished_once() { let mut emit = |event| events.push(event); fsm.receive( now + Duration::from_millis(1), - seq(9), - std::iter::once(Ok(SessionFrame::Ack(record_ack(record_seq)))), + RecordSeq(9), + std::iter::once(Ok(SessionFrame::Ack( + RecordAck::from_ranges([record_seq..=record_seq]).unwrap(), + ))), &mut emit, ); } @@ -266,8 +243,10 @@ fn ack_of_fin_emits_outbound_finished_once() { let mut emit = |event| events.push(event); fsm.receive( now + Duration::from_millis(2), - seq(10), - std::iter::once(Ok(SessionFrame::Ack(record_ack(record_seq)))), + RecordSeq(10), + std::iter::once(Ok(SessionFrame::Ack( + RecordAck::from_ranges([record_seq..=record_seq]).unwrap(), + ))), &mut emit, ); } @@ -285,18 +264,21 @@ fn commit_stream_read_is_what_advances_stream_window() { }, now, ); - let stream_id = stream_id(1); + let stream_id = StreamId(1); let data = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: offset(0), - header: Some(header(1)), + offset: 0, + header: Some(vec![1_u8]), fin: false, bytes: b"hi".to_vec(), })]; - let events = receive_events(&mut fsm, now, seq(7), &data); + let events = receive_events(&mut fsm, now, RecordSeq(7), &data); assert_eq!( events, - vec![opened(stream_id), SessionEvent::Readable(stream_id)] + vec![ + SessionEvent::Opened(stream_id), + SessionEvent::Readable(stream_id) + ] ); let (write_id, builder) = fsm.take_next_write(now + Duration::from_millis(1)).unwrap(); @@ -334,16 +316,16 @@ fn pure_ack_only_records_are_fire_and_forget() { }; let retransmit_timeout = config.retransmit_timeout; let mut fsm = SessionFsm::new(config, now); - let stream_id = stream_id(1); + let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: offset(0), - header: Some(header(1)), + offset: 0, + header: Some(vec![1_u8]), fin: false, bytes: b"hi".to_vec(), })]; - let _ = receive_events(&mut fsm, now, seq(7), &record); + let _ = receive_events(&mut fsm, now, RecordSeq(7), &record); let (write_id, builder) = fsm.take_next_write(now + Duration::from_millis(1)).unwrap(); let ack = decode_session_frames(builder.bytes()).unwrap(); @@ -364,19 +346,22 @@ fn pure_ack_only_records_are_fire_and_forget() { fn inbound_stream_data_emits_opened_and_readable() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let stream_id = stream_id(1); + let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(ql_wire::StreamData { stream_id, - offset: offset(0), - header: Some(header(1)), + offset: 0, + header: Some(vec![1_u8]), fin: true, bytes: b"hello".to_vec(), })]; - let events = receive_events(&mut fsm, now, seq(0), &record); + let events = receive_events(&mut fsm, now, RecordSeq(0), &record); assert_eq!( events, - vec![opened(stream_id), SessionEvent::Readable(stream_id)] + vec![ + SessionEvent::Opened(stream_id), + SessionEvent::Readable(stream_id) + ] ); let mut events = Vec::new(); assert_eq!( @@ -390,19 +375,22 @@ fn inbound_stream_data_emits_opened_and_readable() { fn inbound_empty_fin_emits_finished_immediately() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let stream_id = stream_id(1); + let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: offset(0), - header: Some(header(1)), + offset: 0, + header: Some(vec![1_u8]), fin: true, bytes: Vec::new(), })]; - let events = receive_events(&mut fsm, now, seq(0), &record); + let events = receive_events(&mut fsm, now, RecordSeq(0), &record); assert_eq!( events, - vec![opened(stream_id), SessionEvent::Finished(stream_id)] + vec![ + SessionEvent::Opened(stream_id), + SessionEvent::Finished(stream_id) + ] ); } @@ -444,7 +432,7 @@ fn stream_ids_follow_even_odd_xid_ordering() { }, now, ) - .open_stream(header(1).0.into_boxed_slice(), |_| {}) + .open_stream(vec![1_u8].into_boxed_slice(), |_| {}) .unwrap() .stream_id(); let odd_id = SessionFsm::new( @@ -454,28 +442,33 @@ fn stream_ids_follow_even_odd_xid_ordering() { }, now, ) - .open_stream(header(1).0.into_boxed_slice(), |_| {}) + .open_stream(vec![1_u8].into_boxed_slice(), |_| {}) .unwrap() .stream_id(); - assert_eq!(even_id.0.into_inner() % 2, 0); - assert_eq!(odd_id.0.into_inner() % 2, 1); + assert_eq!(even_id.0 % 2, 0); + assert_eq!(odd_id.0 % 2, 1); } #[test] fn duplicate_stream_data_is_not_redelivered() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let stream_id = stream_id(1); + let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: offset(0), - header: Some(header(1)), + offset: 0, + header: Some(vec![1_u8]), fin: false, bytes: b"hi".to_vec(), })]; - let _ = receive_events(&mut fsm, now, seq(1), &record); - let _ = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); + let _ = receive_events(&mut fsm, now, RecordSeq(1), &record); + let _ = receive_events( + &mut fsm, + now + Duration::from_millis(1), + RecordSeq(2), + &record, + ); assert_eq!(read_stream_all(&mut fsm, stream_id), b"hi".to_vec()); } @@ -485,13 +478,13 @@ fn duplicate_remote_reset_after_reap_is_ignored() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); let reset = StreamReset { - stream_id: stream_id(1), + stream_id: StreamId(1), target: ResetTarget::Both, code: ResetCode(9), }; let record = vec![SessionFrame::StreamReset(reset.clone())]; - let first = receive_events(&mut fsm, now, seq(1), &record); + let first = receive_events(&mut fsm, now, RecordSeq(1), &record); assert_eq!( first, vec![SessionEvent::Reset(StreamResetEvent { @@ -501,7 +494,12 @@ fn duplicate_remote_reset_after_reap_is_ignored() { })] ); - let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); + let second = receive_events( + &mut fsm, + now + Duration::from_millis(1), + RecordSeq(2), + &record, + ); assert!(second.is_empty()); } @@ -509,7 +507,7 @@ fn duplicate_remote_reset_after_reap_is_ignored() { fn late_remote_stream_data_after_reset_is_ignored() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let stream_id = stream_id(1); + let stream_id = StreamId(1); let reset = vec![SessionFrame::StreamReset(StreamReset { stream_id, target: ResetTarget::Both, @@ -517,13 +515,13 @@ fn late_remote_stream_data_after_reset_is_ignored() { })]; let data = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: offset(0), - header: Some(header(1)), + offset: 0, + header: Some(vec![1_u8]), fin: false, bytes: b"hello".to_vec(), })]; - let first = receive_events(&mut fsm, now, seq(1), &reset); + let first = receive_events(&mut fsm, now, RecordSeq(1), &reset); assert_eq!( first, vec![SessionEvent::Reset(StreamResetEvent { @@ -533,7 +531,12 @@ fn late_remote_stream_data_after_reset_is_ignored() { })] ); - let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &data); + let second = receive_events( + &mut fsm, + now + Duration::from_millis(1), + RecordSeq(2), + &data, + ); assert!(second.is_empty()); } @@ -541,19 +544,22 @@ fn late_remote_stream_data_after_reset_is_ignored() { fn duplicate_finished_remote_data_after_reap_is_ignored() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let stream_id = stream_id(1); + let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: offset(0), - header: Some(header(1)), + offset: 0, + header: Some(vec![1_u8]), fin: true, bytes: b"hello".to_vec(), })]; - let first = receive_events(&mut fsm, now, seq(1), &record); + let first = receive_events(&mut fsm, now, RecordSeq(1), &record); assert_eq!( first, - vec![opened(stream_id), SessionEvent::Readable(stream_id)] + vec![ + SessionEvent::Opened(stream_id), + SessionEvent::Readable(stream_id) + ] ); let mut events = Vec::new(); assert_eq!( @@ -562,7 +568,12 @@ fn duplicate_finished_remote_data_after_reap_is_ignored() { ); assert_eq!(events, vec![SessionEvent::Finished(stream_id)]); - let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); + let second = receive_events( + &mut fsm, + now + Duration::from_millis(1), + RecordSeq(2), + &record, + ); assert!(second.is_empty()); } @@ -570,22 +581,30 @@ fn duplicate_finished_remote_data_after_reap_is_ignored() { fn duplicate_finished_remote_data_before_read_is_ignored() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let stream_id = stream_id(1); + let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: offset(0), - header: Some(header(1)), + offset: 0, + header: Some(vec![1_u8]), fin: true, bytes: b"hello".to_vec(), })]; - let first = receive_events(&mut fsm, now, seq(1), &record); + let first = receive_events(&mut fsm, now, RecordSeq(1), &record); assert_eq!( first, - vec![opened(stream_id), SessionEvent::Readable(stream_id)] + vec![ + SessionEvent::Opened(stream_id), + SessionEvent::Readable(stream_id) + ] ); - let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &record); + let second = receive_events( + &mut fsm, + now + Duration::from_millis(1), + RecordSeq(2), + &record, + ); assert!(second.is_empty()); let mut events = Vec::new(); assert_eq!( @@ -600,37 +619,47 @@ fn out_of_order_remote_stream_first_observations_still_open_once_each() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); let reset3 = vec![SessionFrame::StreamReset(StreamReset { - stream_id: stream_id(3), + stream_id: StreamId(3), target: ResetTarget::Both, code: REFUSED, })]; let reset1 = vec![SessionFrame::StreamReset(StreamReset { - stream_id: stream_id(1), + stream_id: StreamId(1), target: ResetTarget::Both, code: TIMEOUT, })]; - let first = receive_events(&mut fsm, now, seq(1), &reset3); + let first = receive_events(&mut fsm, now, RecordSeq(1), &reset3); assert_eq!( first, vec![SessionEvent::Reset(StreamResetEvent { - stream_id: stream_id(3), + stream_id: StreamId(3), code: REFUSED, target: crate::StreamResetTarget::Both, })] ); - let second = receive_events(&mut fsm, now + Duration::from_millis(1), seq(2), &reset1); + let second = receive_events( + &mut fsm, + now + Duration::from_millis(1), + RecordSeq(2), + &reset1, + ); assert_eq!( second, vec![SessionEvent::Reset(StreamResetEvent { - stream_id: stream_id(1), + stream_id: StreamId(1), code: TIMEOUT, target: crate::StreamResetTarget::Both, })] ); - let third = receive_events(&mut fsm, now + Duration::from_millis(2), seq(3), &reset3); + let third = receive_events( + &mut fsm, + now + Duration::from_millis(2), + RecordSeq(3), + &reset3, + ); assert!(third.is_empty()); } @@ -640,11 +669,11 @@ fn invalid_remote_stream_reset_closes_session() { let mut fsm = SessionFsm::new(SessionConfig::default(), now); let invalid = vec![SessionFrame::StreamReset(StreamReset { - stream_id: stream_id(0), + stream_id: StreamId(0), target: ResetTarget::Both, code: ResetCode(9), })]; - let events = receive_events(&mut fsm, now, seq(1), &invalid); + let events = receive_events(&mut fsm, now, RecordSeq(1), &invalid); assert_eq!( events, @@ -666,13 +695,13 @@ fn close_does_not_ack_rejected_record_seq() { ); let invalid = vec![SessionFrame::StreamData(StreamData { - stream_id: stream_id(0), - offset: offset(0), - header: Some(header(1)), + stream_id: StreamId(0), + offset: 0, + header: Some(vec![1_u8]), fin: false, bytes: b"bad".to_vec(), })]; - let events = receive_events(&mut fsm, now, seq(7), &invalid); + let events = receive_events(&mut fsm, now, RecordSeq(7), &invalid); assert_eq!( events, vec![SessionEvent::SessionClosed(ql_wire::SessionClose { @@ -684,7 +713,7 @@ fn close_does_not_ack_rejected_record_seq() { let events = receive_events( &mut fsm, now + Duration::from_millis(1), - seq(8), + RecordSeq(8), &valid_after_close, ); assert!(events.is_empty()); @@ -698,7 +727,7 @@ fn inbound_unpair_emits_final_unpair_frame() { let now = Instant::now(); let mut fsm = SessionFsm::new(SessionConfig::default(), now); - let events = receive_events(&mut fsm, now, seq(1), &[SessionFrame::Unpair]); + let events = receive_events(&mut fsm, now, RecordSeq(1), &[SessionFrame::Unpair]); assert_eq!(events, vec![SessionEvent::Unpaired]); assert!(!fsm.is_closed()); @@ -719,7 +748,7 @@ fn terminating_session_ignores_inbound_frames() { let ignored = receive_events( &mut fsm, now + Duration::from_millis(1), - seq(1), + RecordSeq(1), &[SessionFrame::Ping], ); assert!(ignored.is_empty()); @@ -751,10 +780,10 @@ fn initial_peer_stream_receive_window_limits_first_send() { let events = receive_events( &mut fsm, now + Duration::from_millis(1), - seq(9), + RecordSeq(9), &[SessionFrame::StreamWindow(ql_wire::StreamWindow { stream_id, - maximum_offset: offset(5), + maximum_offset: 5, })], ); assert!(events.is_empty()); @@ -765,7 +794,7 @@ fn initial_peer_stream_receive_window_limits_first_send() { frame, SessionFrame::StreamData(frame) if frame.stream_id == stream_id - && frame.offset == offset(3) + && frame.offset == 3 && frame.bytes.as_slice() == b"lo" ) })); @@ -808,10 +837,7 @@ fn sparse_out_of_order_ack_ranges_page_and_quiesce() { let originals = drain_outbound(&mut sender, now, 4096); assert!(originals.len() >= 64); - for (seq, record) in originals - .iter() - .filter(|(seq, _)| seq.0.into_inner() % 2 == 1) - { + for (seq, record) in originals.iter().filter(|(seq, _)| seq.0 % 2 == 1) { let _ = receive_events(&mut receiver, now, *seq, record); } diff --git a/ql-fsm/src/tests/handshake.rs b/ql-fsm/src/tests/handshake.rs index 922a046e..63a26b60 100644 --- a/ql-fsm/src/tests/handshake.rs +++ b/ql-fsm/src/tests/handshake.rs @@ -143,7 +143,7 @@ fn inbound_xx1_rejects_when_not_in_pairing_mode() { } #[test] -fn inbound_xx1_rejects_mismatched_pairing_id_with_expected_and_actual() { +fn inbound_xx1_rejects_mismatched_pairing_id() { let mut harness = Harness::paired(QlFsmConfig::default(), false, false); let expected = pairing_token(4); let actual = pairing_token(7); @@ -156,13 +156,7 @@ fn inbound_xx1_rejects_mismatched_pairing_id_with_expected_and_actual() { let Node { fsm, crypto } = &mut harness.b; let err = fsm.receive(time, xx1, crypto); - assert_eq!( - err, - Err(ReceiveError::InvalidPairingId { - expected: expected.id(&SoftwareCrypto), - actual: actual.id(&SoftwareCrypto), - }) - ); + assert_eq!(err, Err(ReceiveError::InvalidPairingId)); } #[test] diff --git a/ql-fsm/src/tests/proptest.rs b/ql-fsm/src/tests/proptest.rs index 76000e6a..3e397b2d 100644 --- a/ql-fsm/src/tests/proptest.rs +++ b/ql-fsm/src/tests/proptest.rs @@ -8,10 +8,11 @@ extern crate proptest as proptest_crate; use bytes::Bytes; use proptest_crate::{collection::vec, prelude::*, test_runner::TestCaseResult}; use ql_common::{ResetCode, StreamId}; -use ql_wire::WireError; use super::*; -use crate::{state::LinkState, Event, PeerStatus, ReceiveError, StreamResetTarget, WriteId}; +use crate::{ + state::LinkState, Event, PeerStatus, ReceiveError, ReceiveStage, StreamResetTarget, WriteId, +}; const SLOT_COUNT: usize = 4; @@ -543,15 +544,20 @@ impl Runner { ReceiveError::NoSession | ReceiveError::NoPeer | ReceiveError::InvalidRemoteBundle - | ReceiveError::InvalidSessionPayload(WireError::InvalidPayload) - | ReceiveError::InvalidSessionPayload(WireError::DecryptFailed) - | ReceiveError::InvalidIkHandshake(WireError::InvalidPayload) - | ReceiveError::InvalidIkHandshake(WireError::InvalidState) - | ReceiveError::InvalidKkHandshake(WireError::InvalidPayload) - | ReceiveError::InvalidKkHandshake(WireError::InvalidState) - | ReceiveError::InvalidXxHandshake(WireError::InvalidPayload) - | ReceiveError::InvalidXxHandshake(WireError::InvalidState) - | ReceiveError::InvalidXxHandshake(WireError::DecryptFailed) + | ReceiveError::Wire { + stage: ReceiveStage::SessionPayload, + source: ql_wire::Error::InvalidPayload | ql_wire::Error::DecryptFailed, + } + | ReceiveError::Wire { + stage: ReceiveStage::IkHandshake | ReceiveStage::KkHandshake, + source: ql_wire::Error::InvalidPayload | ql_wire::Error::InvalidState, + } + | ReceiveError::Wire { + stage: ReceiveStage::XxHandshake, + source: ql_wire::Error::InvalidPayload + | ql_wire::Error::InvalidState + | ql_wire::Error::DecryptFailed, + } ), "unexpected receive error on side {side:?}: {error:?}" ); diff --git a/ql-fsm/src/tests/session.rs b/ql-fsm/src/tests/session.rs index f553a9ac..2f7386a6 100644 --- a/ql-fsm/src/tests/session.rs +++ b/ql-fsm/src/tests/session.rs @@ -1,16 +1,12 @@ use std::time::Duration; use bytes::Bytes; -use ql_common::{StreamId, VarInt}; +use ql_common::StreamId; use ql_wire::SessionClose; use super::*; use crate::{state::LinkState, CommitReadError, Event, NoSessionError, PeerStatus, StreamError}; -fn stream_id(value: u32) -> StreamId { - StreamId(VarInt::from_u32(value)) -} - fn open_stream_id(fsm: &mut QlFsm) -> StreamId { fsm.open_stream(Box::from([1])).unwrap().stream_id() } @@ -186,7 +182,7 @@ fn simultaneous_opens_use_even_and_odd_stream_ids() { #[test] fn disconnected_stream_operations_fail_with_no_session() { let mut harness = Harness::paired_known(QlFsmConfig::default()); - let missing = stream_id(0); + let missing = StreamId(0); assert!(matches!( harness.a.fsm.open_stream(Box::from([1])), @@ -223,7 +219,7 @@ fn disconnected_stream_operations_fail_with_no_session() { #[test] fn disconnected_stream_read_accessors_return_none() { let mut harness = Harness::paired_known(QlFsmConfig::default()); - let missing = stream_id(0); + let missing = StreamId(0); assert!(matches!( harness.a.fsm.stream(missing), From d91e02a54d73c1b63b545321cf127c90796f3954 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 9 Jul 2026 19:04:09 -0400 Subject: [PATCH 43/59] ql-runtime: update --- Cargo.lock | 1 + ql-common/src/lib.rs | 9 --------- ql-runtime/Cargo.toml | 1 + ql-runtime/src/driver/mod.rs | 4 ++-- ql-runtime/src/error.rs | 11 ++++++++++- ql-runtime/src/io/inner.rs | 4 ++-- ql-runtime/src/lib.rs | 6 +++++- ql-runtime/src/tests/mod.rs | 3 ++- ql-runtime/src/tests/rpc.rs | 17 +++++++++-------- ql-runtime/src/tests/stream.rs | 4 ++-- 10 files changed, 34 insertions(+), 26 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 3cff966f..38355b76 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1099,6 +1099,7 @@ dependencies = [ "log", "loom", "oneshot", + "ql-codec", "ql-common", "ql-fsm", "ql-rpc", diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index 53246339..e01b622a 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -49,15 +49,6 @@ impl std::fmt::Display for ResetCode { ql_codec::varint_wrapper!(ResetCode, u64); -/// origin of a stream reset: either we triggered it locally or the peer sent it. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -pub enum ResetOrigin { - /// the reset code originated from the peer - Peer, - /// the reset code originated from local logic - Local, -} - #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] pub struct QID(pub [u8; Self::SIZE]); diff --git a/ql-runtime/Cargo.toml b/ql-runtime/Cargo.toml index 1fc53189..fef1396c 100644 --- a/ql-runtime/Cargo.toml +++ b/ql-runtime/Cargo.toml @@ -25,6 +25,7 @@ ql-wire = { workspace = true } [dev-dependencies] env_logger = "0.11" log = "0.4" +ql-codec = { workspace = true } ql-wire = { workspace = true, features = ["test-utils"] } tokio = { version = "1.44", features = ["macros", "rt", "time", "sync"] } diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 9970209c..42db37ae 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -15,7 +15,7 @@ use std::{ use async_channel::Recv; use futures_lite::future::{poll_fn, yield_now}; -use ql_common::{ResetCode, ResetOrigin, StreamId, StreamInfo}; +use ql_common::{ResetCode, StreamId, StreamInfo}; use ql_fsm::{Event, QlFsm, StreamResetEvent, StreamResetTarget, WriteId}; use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; @@ -23,7 +23,7 @@ use crate::{ command::Command, io, log, platform::{QlInbound, QlPlatform, QlTimer}, - QlStreamError, Runtime, + QlStreamError, ResetOrigin, Runtime, }; impl Runtime

{ diff --git a/ql-runtime/src/error.rs b/ql-runtime/src/error.rs index d947871d..329ca521 100644 --- a/ql-runtime/src/error.rs +++ b/ql-runtime/src/error.rs @@ -1,6 +1,15 @@ -use ql_common::{ResetCode, ResetOrigin}; +use ql_common::ResetCode; use ql_fsm::NoSessionError; +/// origin of a stream reset: either we triggered it locally or the peer sent it. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum ResetOrigin { + /// the reset code originated from the peer + Peer, + /// the reset code originated from local logic + Local, +} + #[derive(Debug, Clone, PartialEq, Eq)] pub enum QlStreamError { StreamReset { diff --git a/ql-runtime/src/io/inner.rs b/ql-runtime/src/io/inner.rs index 5b2e8ab5..7a03cde7 100644 --- a/ql-runtime/src/io/inner.rs +++ b/ql-runtime/src/io/inner.rs @@ -292,12 +292,12 @@ mod loom_tests { use bytes::Bytes; use loom::thread; - use ql_common::{ResetCode, ResetOrigin}; + use ql_common::ResetCode; use super::*; use crate::{ io::{sync::loom::*, Tx}, - QlStreamError, + QlStreamError, ResetOrigin, }; #[test] diff --git a/ql-runtime/src/lib.rs b/ql-runtime/src/lib.rs index 40df8dee..a24a590d 100644 --- a/ql-runtime/src/lib.rs +++ b/ql-runtime/src/lib.rs @@ -1,6 +1,10 @@ pub use ql_fsm::{NoSessionError, PairingInvite}; -pub use self::{error::QlStreamError, handle::*, platform::*}; +pub use self::{ + error::{QlStreamError, ResetOrigin}, + handle::*, + platform::*, +}; pub(crate) mod command; pub(crate) mod driver; diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index c5cdf1ff..d64cecce 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -11,12 +11,13 @@ use std::{ use async_channel::{Receiver, Sender}; use futures_lite::Stream; +use ql_codec::Decode; use ql_common::{StreamInfo, QID}; use ql_fsm::PeerStatus; use ql_wire::{ generate_identity, test_identities, MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, Nonce, PairingToken, PeerBundle, QlAead, QlHash, QlIdentity, QlKem, QlRandom, - RecordHeader, RecordType, SessionKey, SoftwareCrypto, WireDecode, + RecordHeader, RecordType, SessionKey, SoftwareCrypto, }; use tokio::{task::LocalSet, time::Sleep}; diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index 430a4395..4cef5136 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -8,7 +8,7 @@ use std::{ }; use bytes::{BufMut, Bytes}; -use ql_common::{ResetCode, ResetOrigin, VarInt}; +use ql_common::ResetCode; use ql_rpc::{ download::{DownloadHandlerLocal, DownloadStart}, duplex::{DuplexHandlerLocal, DuplexPeer}, @@ -21,23 +21,24 @@ use ql_rpc::{ }; use super::*; -use crate::{QlStream, QlStreamError, StreamWriter}; +use crate::{QlStream, QlStreamError, ResetOrigin, StreamWriter}; #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] -struct TestRouteKey(VarInt); +struct TestRouteKey(u64); impl ql_rpc::RpcRouteKey for TestRouteKey { fn encoded_len(&self) -> usize { - self.0.size() + ql_codec::varint::encoded_len(self.0) } fn encode(&self, out: &mut W) { - self.0.write_bytes(|bytes| out.put_slice(bytes)); + ql_codec::varint::encode(self.0, out); } fn decode(bytes: &[u8]) -> Option { - let (route_id, rest) = VarInt::decode_bytes(bytes)?; - rest.is_empty().then_some(Self(route_id)) + let mut reader = ql_codec::Reader::new(bytes); + let route_id = reader.decode_varint().ok()?; + Some(Self(route_id)) } } @@ -47,7 +48,7 @@ macro_rules! test_route { type Key = TestRouteKey; fn key() -> Self::Key { - TestRouteKey(VarInt::from_u32($id)) + TestRouteKey($id) } } }; diff --git a/ql-runtime/src/tests/stream.rs b/ql-runtime/src/tests/stream.rs index 151e6ed4..af94b641 100644 --- a/ql-runtime/src/tests/stream.rs +++ b/ql-runtime/src/tests/stream.rs @@ -1,10 +1,10 @@ use std::time::Duration; use bytes::Bytes; -use ql_common::{ResetCode, ResetOrigin}; +use ql_common::ResetCode; use super::*; -use crate::QlStreamError; +use crate::{QlStreamError, ResetOrigin}; #[tokio::test(flavor = "current_thread")] async fn open_stream_duplex_happy_path() { From 28e93cebcfa1b960db39025fd0bcf42c541de4be Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 9 Jul 2026 20:49:49 -0400 Subject: [PATCH 44/59] ql: merge ik and kk handshakes + remove handshake metadata wrapper --- ql-fsm/src/handshake/ik.rs | 179 ++++++---- ql-fsm/src/handshake/kk.rs | 132 ------- ql-fsm/src/handshake/mod.rs | 25 +- ql-fsm/src/handshake/xx.rs | 36 +- ql-fsm/src/state.rs | 13 +- ql-fsm/src/tests/handshake.rs | 20 +- ql-wire/src/error.rs | 6 +- ql-wire/src/handshake/{meta.rs => id.rs} | 29 +- ql-wire/src/handshake/ik.rs | 435 +++++++++++++++-------- ql-wire/src/handshake/kk.rs | 329 ----------------- ql-wire/src/handshake/mod.rs | 117 ++---- ql-wire/src/handshake/xx.rs | 222 ++++++------ ql-wire/src/record.rs | 6 +- ql-wire/src/tests.rs | 198 ++++------- 14 files changed, 659 insertions(+), 1088 deletions(-) delete mode 100644 ql-fsm/src/handshake/kk.rs rename ql-wire/src/handshake/{meta.rs => id.rs} (51%) delete mode 100644 ql-wire/src/handshake/kk.rs diff --git a/ql-fsm/src/handshake/ik.rs b/ql-fsm/src/handshake/ik.rs index 1b482d78..a588693e 100644 --- a/ql-fsm/src/handshake/ik.rs +++ b/ql-fsm/src/handshake/ik.rs @@ -1,4 +1,6 @@ -use ql_wire::{self as wire, Ik1, Ik2, PeerBundle, QlCrypto, QlHandshakeRecord, RouteHeader}; +use ql_wire::{ + self as wire, Ik1, Ik2, IkPattern, PeerBundle, QlCrypto, QlHandshakeRecord, RouteHeader, +}; use super::{ emit_peer_status, enqueue_handshake, establish_session, reset_connected_session_if_needed, @@ -8,125 +10,180 @@ use crate::{ QlFsm, ReceiveError, ReceiveStage, }; -pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle) { - let meta = super::next_handshake_meta(fsm); +pub fn start_initiator( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + peer: PeerBundle, + pattern: IkPattern, +) { + let handshake_id = super::next_handshake_id(fsm); let route = RouteHeader { sender: fsm.identity.qid, recipient: peer.qid, }; - let mut handshake = wire::IkHandshake::new_initiator( - crypto, - fsm.identity.clone(), - peer, - super::local_transport_params(fsm), - ); - let message = handshake.write_1(crypto, meta).unwrap(); - - fsm.state.link = LinkState::IkInitiator(InitiatorState { - handshake_id: meta.handshake_id, - initial_ephemeral: message.ephemeral.clone(), + let mut handshake = match pattern { + IkPattern::Ik => wire::IkHandshake::new_ik_initiator( + crypto, + fsm.identity.clone(), + peer, + super::local_transport_params(fsm), + ), + IkPattern::Kk => wire::IkHandshake::new_kk_initiator( + crypto, + fsm.identity.clone(), + peer, + super::local_transport_params(fsm), + ), + }; + let message = handshake.write_1(crypto, handshake_id).unwrap(); + let state = InitiatorState { handshake, deadline: fsm.state.now + fsm.config.handshake_timeout, - }); - enqueue_handshake(fsm, route, QlHandshakeRecord::Ik1(message)); + }; + + fsm.state.link = LinkState::IkInitiator(state); + let record = match pattern { + IkPattern::Ik => QlHandshakeRecord::Ik1(message), + IkPattern::Kk => QlHandshakeRecord::Kk1(message), + }; + enqueue_handshake(fsm, route, record); emit_peer_status(fsm, fsm.state.link.status()); } -pub fn handle_ik1( +pub fn handle_1( fsm: &mut QlFsm, crypto: &impl QlCrypto, route: RouteHeader, message: &Ik1, + pattern: IkPattern, ) -> Result<(), ReceiveError> { - if should_ignore_inbound(fsm, route, message) { + if should_ignore_inbound(fsm, route, message, pattern) { return Ok(()); } - if let Some(peer) = fsm.state.peer.as_ref() { - if route.sender != peer.qid { - return Err(ReceiveError::InvalidQid); - } + + let peer = fsm.state.peer.clone(); + if pattern == IkPattern::Kk && peer.is_none() { + return Err(ReceiveError::NoPeer); + } + if peer.as_ref().is_some_and(|peer| route.sender != peer.qid) { + return Err(ReceiveError::InvalidQid); } reset_connected_session_if_needed(fsm); - let mut handshake = wire::IkHandshake::new_responder( - crypto, - fsm.identity.clone(), - fsm.state.peer.clone(), - super::local_transport_params(fsm), - ); + let mut handshake = match (pattern, peer) { + (IkPattern::Ik, expected_remote) => wire::IkHandshake::new_ik_responder( + crypto, + fsm.identity.clone(), + expected_remote, + super::local_transport_params(fsm), + ), + (IkPattern::Kk, Some(remote_bundle)) => wire::IkHandshake::new_kk_responder( + crypto, + fsm.identity.clone(), + remote_bundle, + super::local_transport_params(fsm), + ), + (IkPattern::Kk, None) => unreachable!("KK peer was checked above"), + }; handshake .read_1(crypto, route, message) - .map_err(wire_error)?; + .map_err(|source| wire_error(pattern, source))?; let outbound = handshake - .write_2(crypto, message.meta) - .map_err(wire_error)?; + .write_2(crypto, message.handshake_id) + .map_err(|source| wire_error(pattern, source))?; establish_session( fsm, - message.meta.handshake_id, - handshake.finalize(crypto).map_err(wire_error)?, + message.handshake_id, + handshake + .finalize(crypto) + .map_err(|source| wire_error(pattern, source))?, )?; fsm.state.handshake = None; + let record = match pattern { + IkPattern::Ik => QlHandshakeRecord::Ik2(outbound), + IkPattern::Kk => QlHandshakeRecord::Kk2(outbound), + }; enqueue_handshake( fsm, RouteHeader { sender: fsm.identity.qid, recipient: route.sender, }, - QlHandshakeRecord::Ik2(outbound), + record, ); Ok(()) } -pub fn handle_ik2( +pub fn handle_2( fsm: &mut QlFsm, crypto: &impl QlCrypto, route: RouteHeader, message: &Ik2, + pattern: IkPattern, ) -> Result<(), ReceiveError> { + let LinkState::IkInitiator(state) = &mut fsm.state.link else { + return Ok(()); + }; + if state.handshake.pattern() != pattern + || state.handshake.handshake_id() != Some(message.handshake_id) { - let LinkState::IkInitiator(state) = &mut fsm.state.link else { - return Ok(()); - }; - - if message.meta.handshake_id != state.handshake_id { - return Ok(()); - } - - state - .handshake - .read_2(crypto, route, message) - .map_err(wire_error)?; + return Ok(()); } + state + .handshake + .read_2(crypto, route, message) + .map_err(|source| wire_error(pattern, source))?; let LinkState::IkInitiator(state) = fsm.state.link.take() else { - unreachable!("active IK initiator was checked above"); + unreachable!("active handshake initiator was checked above"); }; establish_session( fsm, - message.meta.handshake_id, - state.handshake.finalize(crypto).map_err(wire_error)?, + message.handshake_id, + state + .handshake + .finalize(crypto) + .map_err(|source| wire_error(pattern, source))?, ) } -pub fn should_ignore_inbound(fsm: &QlFsm, route: RouteHeader, message: &Ik1) -> bool { +fn should_ignore_inbound( + fsm: &QlFsm, + route: RouteHeader, + message: &Ik1, + pattern: IkPattern, +) -> bool { match &fsm.state.link { - LinkState::Idle - | LinkState::KkInitiator(_) - | LinkState::XxInitiator(_) - | LinkState::XxResponder(_) => false, + LinkState::Idle | LinkState::XxInitiator(_) | LinkState::XxResponder(_) => false, LinkState::Connected(_) => { - super::is_connected_replay(fsm, message.meta.handshake_id, route.sender) + super::is_connected_replay(fsm, message.handshake_id, route.sender) } LinkState::IkInitiator(state) => { - if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(route.sender) { - return false; + if state.handshake.pattern() != pattern { + return pattern == IkPattern::Kk; + } + if fsm.state.peer.as_ref().map(|peer| peer.qid) == Some(route.sender) { + super::local_start_wins( + state + .handshake + .local_ephemeral() + .expect("initiator has sent message 1"), + &message.ephemeral, + ) + } else { + false } - super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) } } } -fn wire_error(source: ql_wire::Error) -> ReceiveError { - ReceiveError::wire(ReceiveStage::IkHandshake, source) +fn wire_error(pattern: IkPattern, source: ql_wire::Error) -> ReceiveError { + ReceiveError::wire( + match pattern { + IkPattern::Ik => ReceiveStage::IkHandshake, + IkPattern::Kk => ReceiveStage::KkHandshake, + }, + source, + ) } diff --git a/ql-fsm/src/handshake/kk.rs b/ql-fsm/src/handshake/kk.rs deleted file mode 100644 index 01468be7..00000000 --- a/ql-fsm/src/handshake/kk.rs +++ /dev/null @@ -1,132 +0,0 @@ -use ql_wire::{self as wire, Kk1, Kk2, PeerBundle, QlCrypto, QlHandshakeRecord, RouteHeader}; - -use super::{ - emit_peer_status, enqueue_handshake, establish_session, reset_connected_session_if_needed, -}; -use crate::{ - state::{InitiatorState, LinkState}, - QlFsm, ReceiveError, ReceiveStage, -}; - -pub fn start_initiator(fsm: &mut QlFsm, crypto: &impl QlCrypto, peer: PeerBundle) { - let meta = super::next_handshake_meta(fsm); - let route = RouteHeader { - sender: fsm.identity.qid, - recipient: peer.qid, - }; - let mut handshake = wire::KkHandshake::new_initiator( - crypto, - fsm.identity.clone(), - peer, - super::local_transport_params(fsm), - ); - let message = handshake.write_1(crypto, meta).unwrap(); - - fsm.state.link = LinkState::KkInitiator(InitiatorState { - handshake_id: meta.handshake_id, - initial_ephemeral: message.ephemeral.clone(), - handshake, - deadline: fsm.state.now + fsm.config.handshake_timeout, - }); - enqueue_handshake(fsm, route, QlHandshakeRecord::Kk1(message)); - emit_peer_status(fsm, fsm.state.link.status()); -} - -pub fn handle_kk1( - fsm: &mut QlFsm, - crypto: &impl QlCrypto, - route: RouteHeader, - message: &Kk1, -) -> Result<(), ReceiveError> { - if should_ignore_inbound(fsm, route, message) { - return Ok(()); - } - - let Some(peer) = fsm.state.peer.clone() else { - return Err(ReceiveError::NoPeer); - }; - if route.sender != peer.qid { - return Err(ReceiveError::InvalidQid); - } - - reset_connected_session_if_needed(fsm); - - let mut handshake = wire::KkHandshake::new_responder( - crypto, - fsm.identity.clone(), - peer, - super::local_transport_params(fsm), - ); - handshake - .read_1(crypto, route, message) - .map_err(wire_error)?; - let outbound = handshake - .write_2(crypto, message.meta) - .map_err(wire_error)?; - establish_session( - fsm, - message.meta.handshake_id, - handshake.finalize(crypto).map_err(wire_error)?, - )?; - fsm.state.handshake = None; - enqueue_handshake( - fsm, - RouteHeader { - sender: fsm.identity.qid, - recipient: route.sender, - }, - QlHandshakeRecord::Kk2(outbound), - ); - Ok(()) -} - -pub fn handle_kk2( - fsm: &mut QlFsm, - crypto: &impl QlCrypto, - route: RouteHeader, - message: &Kk2, -) -> Result<(), ReceiveError> { - { - let LinkState::KkInitiator(state) = &mut fsm.state.link else { - return Ok(()); - }; - - if message.meta.handshake_id != state.handshake_id { - return Ok(()); - } - - state - .handshake - .read_2(crypto, route, message) - .map_err(wire_error)?; - } - - let LinkState::KkInitiator(state) = fsm.state.link.take() else { - unreachable!("active KK initiator was checked above"); - }; - establish_session( - fsm, - message.meta.handshake_id, - state.handshake.finalize(crypto).map_err(wire_error)?, - ) -} - -pub fn should_ignore_inbound(fsm: &QlFsm, route: RouteHeader, message: &Kk1) -> bool { - match &fsm.state.link { - LinkState::Idle | LinkState::XxInitiator(_) | LinkState::XxResponder(_) => false, - LinkState::Connected(_) => { - super::is_connected_replay(fsm, message.meta.handshake_id, route.sender) - } - LinkState::IkInitiator(_) => true, - LinkState::KkInitiator(state) => { - if fsm.state.peer.as_ref().map(|peer| peer.qid) != Some(route.sender) { - return false; - } - super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) - } - } -} - -fn wire_error(source: ql_wire::Error) -> ReceiveError { - ReceiveError::wire(ReceiveStage::KkHandshake, source) -} diff --git a/ql-fsm/src/handshake/mod.rs b/ql-fsm/src/handshake/mod.rs index 7add5539..e70c2936 100644 --- a/ql-fsm/src/handshake/mod.rs +++ b/ql-fsm/src/handshake/mod.rs @@ -1,11 +1,10 @@ mod ik; -mod kk; mod xx; use ql_common::QID; use ql_wire::{ - self as wire, EphemeralPublicKey, HandshakeId, HandshakeMeta, QlCrypto, QlHandshakeRecord, - RouteHeader, + self as wire, EphemeralPublicKey, HandshakeId, IkPattern, MlKemPublicKey, QlCrypto, + QlHandshakeRecord, RouteHeader, }; use crate::{ @@ -18,14 +17,14 @@ use crate::{ pub fn handle_connect_ik(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Result<(), NoPeerError> { let peer = fsm.state.peer.clone().ok_or(NoPeerError)?; prepare_for_outbound_connect(fsm); - ik::start_initiator(fsm, crypto, peer); + ik::start_initiator(fsm, crypto, peer, IkPattern::Ik); Ok(()) } pub fn handle_connect_kk(fsm: &mut QlFsm, crypto: &impl QlCrypto) -> Result<(), NoPeerError> { let peer = fsm.state.peer.clone().ok_or(NoPeerError)?; prepare_for_outbound_connect(fsm); - kk::start_initiator(fsm, crypto, peer); + ik::start_initiator(fsm, crypto, peer, IkPattern::Kk); Ok(()) } @@ -34,10 +33,10 @@ pub fn handle_connect_xx(fsm: &mut QlFsm, invite: crate::PairingInvite, crypto: xx::start_initiator(fsm, crypto, invite.token, invite.qid); } -pub fn next_handshake_meta(fsm: &mut QlFsm) -> HandshakeMeta { +pub fn next_handshake_id(fsm: &mut QlFsm) -> HandshakeId { let handshake_id = wire::HandshakeId(fsm.state.next_control_id); fsm.state.next_control_id = fsm.state.next_control_id.wrapping_add(1); - HandshakeMeta { handshake_id } + handshake_id } pub fn enqueue_handshake(fsm: &mut QlFsm, route: RouteHeader, record: QlHandshakeRecord) { @@ -67,10 +66,10 @@ pub fn handle_handshake_record( record: &QlHandshakeRecord, ) -> Result<(), ReceiveError> { match record { - QlHandshakeRecord::Ik1(message) => ik::handle_ik1(fsm, crypto, route, message), - QlHandshakeRecord::Ik2(message) => ik::handle_ik2(fsm, crypto, route, message), - QlHandshakeRecord::Kk1(message) => kk::handle_kk1(fsm, crypto, route, message), - QlHandshakeRecord::Kk2(message) => kk::handle_kk2(fsm, crypto, route, message), + QlHandshakeRecord::Ik1(message) => ik::handle_1(fsm, crypto, route, message, IkPattern::Ik), + QlHandshakeRecord::Ik2(message) => ik::handle_2(fsm, crypto, route, message, IkPattern::Ik), + QlHandshakeRecord::Kk1(message) => ik::handle_1(fsm, crypto, route, message, IkPattern::Kk), + QlHandshakeRecord::Kk2(message) => ik::handle_2(fsm, crypto, route, message, IkPattern::Kk), QlHandshakeRecord::Xx1(message) => xx::handle_xx1(fsm, crypto, route, message), QlHandshakeRecord::Xx2(message) => xx::handle_xx2(fsm, crypto, route, message), QlHandshakeRecord::Xx3(message) => xx::handle_xx3(fsm, crypto, route, message), @@ -159,8 +158,8 @@ pub fn reset_connected_session_if_needed(fsm: &mut QlFsm) { } } -fn local_start_wins(local: &EphemeralPublicKey, inbound: &EphemeralPublicKey) -> bool { - local.mlkem_public_key.as_bytes() <= inbound.mlkem_public_key.as_bytes() +fn local_start_wins(local: &MlKemPublicKey, inbound: &EphemeralPublicKey) -> bool { + local.as_bytes() <= inbound.mlkem_public_key.as_bytes() } fn is_connected_replay(fsm: &QlFsm, handshake_id: HandshakeId, sender: QID) -> bool { diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs index df6cebcc..3c0cbc91 100644 --- a/ql-fsm/src/handshake/xx.rs +++ b/ql-fsm/src/handshake/xx.rs @@ -17,7 +17,7 @@ pub fn start_initiator( token: PairingToken, remote_qid: QID, ) { - let meta = super::next_handshake_meta(fsm); + let handshake_id = super::next_handshake_id(fsm); let route = RouteHeader { sender: fsm.identity.qid, recipient: remote_qid, @@ -29,11 +29,9 @@ pub fn start_initiator( token, super::local_transport_params(fsm), ); - let message = handshake.write_1(crypto, meta).unwrap(); + let message = handshake.write_1(crypto, handshake_id).unwrap(); fsm.state.link = LinkState::XxInitiator(InitiatorState { - handshake_id: meta.handshake_id, - initial_ephemeral: message.ephemeral.clone(), handshake, deadline: fsm.state.now + fsm.config.handshake_timeout, }); @@ -68,11 +66,10 @@ pub fn handle_xx1( .read_1(crypto, route, message) .map_err(wire_error)?; let outbound = handshake - .write_2(crypto, message.meta) + .write_2(crypto, message.handshake_id) .map_err(wire_error)?; fsm.state.link = LinkState::XxResponder(XxResponderState { handshake, - handshake_meta: message.meta, deadline: fsm.state.now + fsm.config.handshake_timeout, }); fsm.state.handshake = None; @@ -101,7 +98,7 @@ pub fn handle_xx2( return Ok(()); }; - if message.meta.handshake_id != state.handshake_id { + if state.handshake.handshake_id() != Some(message.handshake_id) { return Ok(()); } @@ -111,7 +108,7 @@ pub fn handle_xx2( .map_err(wire_error)?; let outbound = state .handshake - .write_3(crypto, message.meta) + .write_3(crypto, message.handshake_id) .map_err(wire_error)?; fsm.state.handshake = None; enqueue_handshake( @@ -137,7 +134,7 @@ pub fn handle_xx3( return Ok(()); }; - if message.meta.handshake_id != state.handshake_meta.handshake_id { + if state.handshake.handshake_id() != Some(message.handshake_id) { return Ok(()); } @@ -145,13 +142,12 @@ pub fn handle_xx3( .handshake .read_3(crypto, route, message) .map_err(wire_error)?; - let handshake_meta = state.handshake_meta; let LinkState::XxResponder(mut state) = fsm.state.link.take() else { unreachable!("active XX responder was checked above"); }; let outbound = state .handshake - .write_4(crypto, handshake_meta) + .write_4(crypto, message.handshake_id) .map_err(wire_error)?; fsm.state.handshake = None; enqueue_handshake( @@ -164,7 +160,7 @@ pub fn handle_xx3( ); establish_session( fsm, - message.meta.handshake_id, + message.handshake_id, state.handshake.finalize(crypto).map_err(wire_error)?, ) } @@ -180,7 +176,7 @@ pub fn handle_xx4( return Ok(()); }; - if message.meta.handshake_id != state.handshake_id { + if state.handshake.handshake_id() != Some(message.handshake_id) { return Ok(()); } @@ -195,7 +191,7 @@ pub fn handle_xx4( }; establish_session( fsm, - message.meta.handshake_id, + message.handshake_id, state.handshake.finalize(crypto).map_err(wire_error)?, ) } @@ -216,9 +212,9 @@ pub fn should_ignore_inbound( match &fsm.state.link { LinkState::Idle => false, LinkState::Connected(_) => { - super::is_connected_replay(fsm, message.meta.handshake_id, route.sender) + super::is_connected_replay(fsm, message.handshake_id, route.sender) } - LinkState::IkInitiator(_) | LinkState::KkInitiator(_) | LinkState::XxResponder(_) => true, + LinkState::IkInitiator(_) | LinkState::XxResponder(_) => true, LinkState::XxInitiator(state) => { if state.handshake.pairing_id(crypto) != message.pairing_id { return false; @@ -226,7 +222,13 @@ pub fn should_ignore_inbound( if route.sender != state.handshake.remote_qid() { return false; } - super::local_start_wins(&state.initial_ephemeral, &message.ephemeral) + super::local_start_wins( + state + .handshake + .local_ephemeral() + .expect("initiator has sent message 1"), + &message.ephemeral, + ) } } } diff --git a/ql-fsm/src/state.rs b/ql-fsm/src/state.rs index 093bd4b0..8ceeb711 100644 --- a/ql-fsm/src/state.rs +++ b/ql-fsm/src/state.rs @@ -2,8 +2,8 @@ use std::time::Instant; use ql_common::QID; use ql_wire::{ - EphemeralPublicKey, HandshakeId, HandshakeMeta, IkHandshake, KkHandshake, PairingToken, - PeerBundle, QlHandshakeRecord, RouteHeader, SessionKey, TransportParams, XxHandshake, + HandshakeId, IkHandshake, PairingToken, PeerBundle, QlHandshakeRecord, RouteHeader, SessionKey, + TransportParams, XxHandshake, }; use crate::{session::SessionFsm, NoSessionError, PeerStatus}; @@ -29,7 +29,6 @@ pub struct SessionTransport { pub enum LinkState { Idle, IkInitiator(InitiatorState), - KkInitiator(InitiatorState), XxInitiator(InitiatorState), XxResponder(XxResponderState), Connected(ConnectedState), @@ -44,15 +43,12 @@ pub struct ConnectedState { #[derive(Debug, Clone)] pub struct InitiatorState { pub handshake: H, - pub handshake_id: HandshakeId, pub deadline: Instant, - pub initial_ephemeral: EphemeralPublicKey, } #[derive(Debug, Clone)] pub struct XxResponderState { pub handshake: XxHandshake, - pub handshake_meta: HandshakeMeta, pub deadline: Instant, } @@ -64,9 +60,7 @@ impl LinkState { pub fn status(&self) -> PeerStatus { match self { Self::Idle | Self::XxResponder(_) => PeerStatus::Disconnected, - Self::IkInitiator(_) | Self::KkInitiator(_) | Self::XxInitiator(_) => { - PeerStatus::Initiator - } + Self::IkInitiator(_) | Self::XxInitiator(_) => PeerStatus::Initiator, Self::Connected(_) => PeerStatus::Connected, } } @@ -96,7 +90,6 @@ impl LinkState { match self { Self::Idle | Self::Connected(_) => None, Self::IkInitiator(state) => Some(state.deadline), - Self::KkInitiator(state) => Some(state.deadline), Self::XxInitiator(state) => Some(state.deadline), Self::XxResponder(state) => Some(state.deadline), } diff --git a/ql-fsm/src/tests/handshake.rs b/ql-fsm/src/tests/handshake.rs index 63a26b60..b517dcc1 100644 --- a/ql-fsm/src/tests/handshake.rs +++ b/ql-fsm/src/tests/handshake.rs @@ -1,6 +1,6 @@ use std::time::Duration; -use ql_wire::QlHandshakeRecord; +use ql_wire::{IkPattern, QlHandshakeRecord}; use super::*; use crate::{state::LinkState, Event, NoPeerError, PeerStatus, ReceiveError}; @@ -255,7 +255,7 @@ fn connect_kk_replaces_in_flight_attempt_and_ignores_stale_reply() { harness.deliver(Side::A, stale_reply); assert!(matches!( harness.a.fsm.state.link, - LinkState::KkInitiator(_) + LinkState::IkInitiator(ref state) if state.handshake.pattern() == IkPattern::Kk )); harness.deliver(Side::B, second); @@ -370,13 +370,13 @@ fn simultaneous_ik_and_kk_connect_prefers_ik() { fn handshake_id(record: &[u8]) -> ql_wire::HandshakeId { let (_, record) = ql_wire::decode_record(record).unwrap(); match record { - ql_wire::QlHandshakeRecord::Ik1(message) => message.meta.handshake_id, - ql_wire::QlHandshakeRecord::Ik2(message) => message.meta.handshake_id, - ql_wire::QlHandshakeRecord::Kk1(message) => message.meta.handshake_id, - ql_wire::QlHandshakeRecord::Kk2(message) => message.meta.handshake_id, - ql_wire::QlHandshakeRecord::Xx1(message) => message.meta.handshake_id, - ql_wire::QlHandshakeRecord::Xx2(message) => message.meta.handshake_id, - ql_wire::QlHandshakeRecord::Xx3(message) => message.meta.handshake_id, - ql_wire::QlHandshakeRecord::Xx4(message) => message.meta.handshake_id, + ql_wire::QlHandshakeRecord::Ik1(message) => message.handshake_id, + ql_wire::QlHandshakeRecord::Ik2(message) => message.handshake_id, + ql_wire::QlHandshakeRecord::Kk1(message) => message.handshake_id, + ql_wire::QlHandshakeRecord::Kk2(message) => message.handshake_id, + ql_wire::QlHandshakeRecord::Xx1(message) => message.handshake_id, + ql_wire::QlHandshakeRecord::Xx2(message) => message.handshake_id, + ql_wire::QlHandshakeRecord::Xx3(message) => message.handshake_id, + ql_wire::QlHandshakeRecord::Xx4(message) => message.handshake_id, } } diff --git a/ql-wire/src/error.rs b/ql-wire/src/error.rs index 936c155d..71359231 100644 --- a/ql-wire/src/error.rs +++ b/ql-wire/src/error.rs @@ -14,10 +14,9 @@ pub enum Error { // protocol validation InvalidRouteHeader, - InvalidHandshakeMeta, + InvalidHandshakeId, InvalidPairingId, InvalidRemoteBundle, - InvalidTransportParams, // cryptographic/session DecryptFailed, @@ -38,10 +37,9 @@ impl fmt::Display for Error { Self::InvalidUtf8 => "invalid utf-8", Self::InvalidPayload => "invalid payload", Self::InvalidRouteHeader => "invalid route header", - Self::InvalidHandshakeMeta => "invalid handshake meta", + Self::InvalidHandshakeId => "invalid handshake id", Self::InvalidPairingId => "invalid pairing id", Self::InvalidRemoteBundle => "invalid remote bundle", - Self::InvalidTransportParams => "invalid transport params", Self::DecryptFailed => "decryption failed", Self::Expired => "expired", Self::InvalidState => "invalid state", diff --git a/ql-wire/src/handshake/meta.rs b/ql-wire/src/handshake/id.rs similarity index 51% rename from ql-wire/src/handshake/meta.rs rename to ql-wire/src/handshake/id.rs index 21b70f82..3e4915ab 100644 --- a/ql-wire/src/handshake/meta.rs +++ b/ql-wire/src/handshake/id.rs @@ -4,9 +4,8 @@ use ql_codec::{ByteSlice, Encode}; #[repr(transparent)] pub struct HandshakeId(pub u32); -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct HandshakeMeta { - pub handshake_id: HandshakeId, +impl HandshakeId { + pub const WIRE_SIZE: usize = size_of::(); } impl ql_codec::Decode for HandshakeId { @@ -16,33 +15,11 @@ impl ql_codec::Decode for HandshakeId { } impl Encode for HandshakeId { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl HandshakeMeta { - pub const WIRE_SIZE: usize = size_of::(); -} - -impl Encode for HandshakeMeta { fn encoded_len(&self) -> usize { Self::WIRE_SIZE } fn encode(&self, out: &mut W) { - self.handshake_id.encode(out); - } -} - -impl ql_codec::Decode for HandshakeMeta { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - handshake_id: reader.decode()?, - }) + self.0.encode(out); } } diff --git a/ql-wire/src/handshake/ik.rs b/ql-wire/src/handshake/ik.rs index 38afd608..c2c1d73f 100644 --- a/ql-wire/src/handshake/ik.rs +++ b/ql-wire/src/handshake/ik.rs @@ -2,91 +2,106 @@ use ql_codec::{ByteSlice, Encode}; use super::{ decrypt_mlkem_ciphertext, decrypt_peer_bundle, encrypt_mlkem_ciphertext, encrypt_peer_bundle, - finalize_handshake, generate_ephemeral_keypair, init_ik_symmetric, initialize_handshake_meta, - mix_hash_ephemeral, mix_hash_routed_handshake, require_handshake_meta, - EncryptedMlKemCiphertext, EncryptedPeerBundle, EphemeralKeyPair, EphemeralPublicKey, - FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, + finalize_handshake, generate_ephemeral_keypair, initialize_handshake_id, + mix_hash_routed_handshake, require_handshake_id, EncryptedMlKemCiphertext, EncryptedPeerBundle, + EphemeralKeyPair, EphemeralPublicKey, FinalizedHandshake, Role, RouteHeader, SymmetricState, + TransportParams, PROTOCOL_IK, PROTOCOL_KK, }; use crate::{ - Error, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, QlIdentity, + Error, HandshakeId, HandshakeKind, MlKemCiphertext, MlKemPublicKey, PeerBundle, QlCrypto, + QlIdentity, }; #[derive(Debug, Clone, PartialEq, Eq)] pub struct Ik1 { - pub meta: HandshakeMeta, + pub handshake_id: HandshakeId, pub transport_params: TransportParams, pub skem_ciphertext: MlKemCiphertext, pub ephemeral: EphemeralPublicKey, - pub static_bundle: EncryptedPeerBundle, -} - -impl ql_codec::Decode for Ik1 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - meta: reader.decode()?, - transport_params: reader.decode()?, - skem_ciphertext: reader.decode()?, - ephemeral: reader.decode()?, - static_bundle: reader.decode()?, - }) - } + pub static_bundle: Option, } impl Encode for Ik1 { fn encoded_len(&self) -> usize { - HandshakeMeta::WIRE_SIZE + HandshakeId::WIRE_SIZE + TransportParams::WIRE_SIZE + MlKemCiphertext::SIZE + EphemeralPublicKey::WIRE_SIZE - + self.static_bundle.encoded_len() + + self + .static_bundle + .as_ref() + .map_or(0, EncryptedPeerBundle::encoded_len) } fn encode(&self, out: &mut W) { - self.meta.encode(out); + self.handshake_id.encode(out); self.transport_params.encode(out); self.skem_ciphertext.encode(out); self.ephemeral.encode(out); - self.static_bundle.encode(out); + if let Some(static_bundle) = self.static_bundle.as_ref() { + static_bundle.encode(out); + } + } +} + +impl ql_codec::Decode for Ik1 { + fn decode(reader: &mut ql_codec::Reader) -> Result { + let handshake_id = reader.decode()?; + let transport_params = reader.decode()?; + let skem_ciphertext = reader.decode()?; + let ephemeral = reader.decode()?; + let static_bundle = if reader.is_empty() { + None + } else { + Some(reader.decode()?) + }; + Ok(Self { + handshake_id, + transport_params, + skem_ciphertext, + ephemeral, + static_bundle, + }) } } #[derive(Debug, Clone, PartialEq, Eq)] pub struct Ik2 { - pub meta: HandshakeMeta, + pub handshake_id: HandshakeId, pub transport_params: TransportParams, pub ekem_ciphertext: MlKemCiphertext, pub skem_ciphertext: EncryptedMlKemCiphertext, } -impl ql_codec::Decode for Ik2 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - meta: reader.decode()?, - transport_params: reader.decode()?, - ekem_ciphertext: reader.decode()?, - skem_ciphertext: reader.decode()?, - }) - } -} - impl Encode for Ik2 { fn encoded_len(&self) -> usize { - HandshakeMeta::WIRE_SIZE + HandshakeId::WIRE_SIZE + TransportParams::WIRE_SIZE + MlKemCiphertext::SIZE + EncryptedMlKemCiphertext::WIRE_SIZE } fn encode(&self, out: &mut W) { - self.meta.encode(out); + self.handshake_id.encode(out); self.transport_params.encode(out); self.ekem_ciphertext.encode(out); self.skem_ciphertext.encode(out); } } +impl ql_codec::Decode for Ik2 { + fn decode(reader: &mut ql_codec::Reader) -> Result { + Ok(Self { + handshake_id: reader.decode()?, + transport_params: reader.decode()?, + ekem_ciphertext: reader.decode()?, + skem_ciphertext: reader.decode()?, + }) + } +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] -enum IkStep { +enum Step { Send1, Recv1, Send2, @@ -94,107 +109,172 @@ enum IkStep { Done, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum IkPattern { + Ik, + Kk, +} + #[derive(Debug, Clone)] pub struct IkHandshake { + pattern: IkPattern, role: Role, - step: IkStep, + step: Step, symmetric: SymmetricState, local: QlIdentity, remote_bundle: Option, local_ephemeral: Option, remote_ephemeral: Option, - handshake_meta: Option, + handshake_id: Option, local_transport_params: TransportParams, remote_transport_params: Option, } impl IkHandshake { - pub fn new_initiator( + pub fn pattern(&self) -> IkPattern { + self.pattern + } + + pub fn handshake_id(&self) -> Option { + self.handshake_id + } + + pub fn local_ephemeral(&self) -> Option<&MlKemPublicKey> { + self.local_ephemeral + .as_ref() + .map(|keypair| &keypair.mlkem.public) + } + + pub fn new_ik_initiator( crypto: &impl QlCrypto, local: QlIdentity, remote_bundle: PeerBundle, local_transport_params: TransportParams, ) -> Self { - let symmetric = init_ik_symmetric(crypto, &remote_bundle); - Self { - role: Role::Initiator, - step: IkStep::Send1, + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_IK); + symmetric.mix_hash(crypto, &remote_bundle.encode_vec()); + Self::new( + IkPattern::Ik, + Role::Initiator, symmetric, local, - remote_bundle: Some(remote_bundle), - local_ephemeral: None, - remote_ephemeral: None, - handshake_meta: None, + Some(remote_bundle), local_transport_params, - remote_transport_params: None, - } + ) } - pub fn new_responder( + pub fn new_ik_responder( crypto: &impl QlCrypto, local: QlIdentity, expected_remote: Option, local_transport_params: TransportParams, ) -> Self { - let symmetric = init_ik_symmetric(crypto, &local.bundle()); - Self { - role: Role::Responder, - step: IkStep::Recv1, + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_IK); + symmetric.mix_hash(crypto, &local.bundle().encode_vec()); + Self::new( + IkPattern::Ik, + Role::Responder, symmetric, local, - remote_bundle: expected_remote, - local_ephemeral: None, - remote_ephemeral: None, - handshake_meta: None, + expected_remote, local_transport_params, - remote_transport_params: None, - } + ) } - pub fn is_finished(&self) -> bool { - self.step == IkStep::Done + pub fn new_kk_initiator( + crypto: &impl QlCrypto, + local: QlIdentity, + remote_bundle: PeerBundle, + local_transport_params: TransportParams, + ) -> Self { + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_KK); + symmetric.mix_hash(crypto, &local.bundle().encode_vec()); + symmetric.mix_hash(crypto, &remote_bundle.encode_vec()); + Self::new( + IkPattern::Kk, + Role::Initiator, + symmetric, + local, + Some(remote_bundle), + local_transport_params, + ) } - fn outbound_header(&self) -> Result { - let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; - Ok(RouteHeader { - sender: self.local.qid, - recipient: remote_bundle.qid, - }) + pub fn new_kk_responder( + crypto: &impl QlCrypto, + local: QlIdentity, + remote_bundle: PeerBundle, + local_transport_params: TransportParams, + ) -> Self { + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_KK); + symmetric.mix_hash(crypto, &remote_bundle.encode_vec()); + symmetric.mix_hash(crypto, &local.bundle().encode_vec()); + Self::new( + IkPattern::Kk, + Role::Responder, + symmetric, + local, + Some(remote_bundle), + local_transport_params, + ) } - fn ensure_inbound_recipient(&self, header: RouteHeader) -> Result<(), Error> { - if header.recipient == self.local.qid { - Ok(()) - } else { - Err(Error::InvalidRouteHeader) + fn new( + pattern: IkPattern, + role: Role, + symmetric: SymmetricState, + local: QlIdentity, + remote_bundle: Option, + local_transport_params: TransportParams, + ) -> Self { + Self { + pattern, + role, + step: match role { + Role::Initiator => Step::Send1, + Role::Responder => Step::Recv1, + }, + symmetric, + local, + remote_bundle, + local_ephemeral: None, + remote_ephemeral: None, + handshake_id: None, + local_transport_params, + remote_transport_params: None, } } - fn ensure_known_remote_sender(&self, header: RouteHeader) -> Result<(), Error> { - if let Some(remote_bundle) = self.remote_bundle.as_ref() { - if header.sender != remote_bundle.qid { - return Err(Error::InvalidRouteHeader); - } - } - Ok(()) + pub fn is_finished(&self) -> bool { + self.step == Step::Done } - pub fn write_1(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { - if self.step != IkStep::Send1 { + pub fn write_1( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { + if self.step != Step::Send1 { return Err(Error::InvalidState); } - initialize_handshake_meta(&mut self.handshake_meta, meta)?; + initialize_handshake_id(&mut self.handshake_id, handshake_id)?; let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; - let header = self.outbound_header()?; + let header = RouteHeader { + sender: self.local.qid, + recipient: remote_bundle.qid, + }; mix_hash_routed_handshake( &mut self.symmetric, crypto, header, - HandshakeKind::Ik1, - meta, + match self.pattern { + IkPattern::Ik => HandshakeKind::Ik1, + IkPattern::Kk => HandshakeKind::Kk1, + }, + handshake_id, self.local_transport_params, ); + let (skem_ciphertext, skem_secret) = crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); self.symmetric.mix_hash(crypto, skem_ciphertext.as_bytes()); @@ -202,77 +282,49 @@ impl IkHandshake { .mix_key_and_hash(crypto, skem_secret.as_bytes()); let local_ephemeral = generate_ephemeral_keypair(crypto); - let public = local_ephemeral.public(); - mix_hash_ephemeral(&mut self.symmetric, crypto, &public); + let ephemeral = local_ephemeral.public(); + self.symmetric.mix_hash_ephemeral(crypto, &ephemeral); - let static_bundle = encrypt_peer_bundle(crypto, &mut self.symmetric, &self.local.bundle())?; + let static_bundle = match self.pattern { + IkPattern::Ik => Some(encrypt_peer_bundle( + crypto, + &mut self.symmetric, + &self.local.bundle(), + )?), + IkPattern::Kk => None, + }; self.local_ephemeral = Some(local_ephemeral); - self.step = IkStep::Recv2; + self.step = Step::Recv2; Ok(Ik1 { - meta, + handshake_id, transport_params: self.local_transport_params, skem_ciphertext, - ephemeral: public, + ephemeral, static_bundle, }) } - pub fn write_2(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { - if self.step != IkStep::Send2 { - return Err(Error::InvalidState); - } - require_handshake_meta(self.handshake_meta.as_ref(), meta)?; - let header = self.outbound_header()?; - mix_hash_routed_handshake( - &mut self.symmetric, - crypto, - header, - HandshakeKind::Ik2, - meta, - self.local_transport_params, - ); - let remote_ephemeral = self.remote_ephemeral.clone().ok_or(Error::InvalidState)?; - let (ekem_ciphertext, ekem_secret) = - crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); - self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); - self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); - - let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; - let (skem_ciphertext, skem_secret) = - crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); - let skem_ciphertext = - encrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &skem_ciphertext)?; - self.symmetric - .mix_key_and_hash(crypto, skem_secret.as_bytes()); - - self.step = IkStep::Done; - Ok(Ik2 { - meta, - transport_params: self.local_transport_params, - ekem_ciphertext, - skem_ciphertext, - }) - } - pub fn read_1( &mut self, crypto: &impl QlCrypto, header: RouteHeader, message: &Ik1, ) -> Result<(), Error> { - if self.step != IkStep::Recv1 { + if self.step != Step::Recv1 { return Err(Error::InvalidState); } - initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; - self.ensure_inbound_recipient(header)?; - self.ensure_known_remote_sender(header)?; + initialize_handshake_id(&mut self.handshake_id, message.handshake_id)?; + self.ensure_inbound_header(header)?; mix_hash_routed_handshake( &mut self.symmetric, crypto, header, - HandshakeKind::Ik1, - message.meta, + match self.pattern { + IkPattern::Ik => HandshakeKind::Ik1, + IkPattern::Kk => HandshakeKind::Kk1, + }, + message.handshake_id, message.transport_params, ); self.symmetric @@ -282,46 +334,105 @@ impl IkHandshake { self.symmetric .mix_key_and_hash(crypto, skem_secret.as_bytes()); - mix_hash_ephemeral(&mut self.symmetric, crypto, &message.ephemeral); + self.symmetric + .mix_hash_ephemeral(crypto, &message.ephemeral); self.remote_ephemeral = Some(message.ephemeral.clone()); - let remote_bundle = - decrypt_peer_bundle(crypto, &mut self.symmetric, &message.static_bundle)?; - if remote_bundle.qid != header.sender { - return Err(Error::InvalidRemoteBundle); - } - match self.remote_bundle.as_ref() { - Some(expected) if expected != &remote_bundle => { - return Err(Error::InvalidRemoteBundle); + match (self.pattern, message.static_bundle.as_ref()) { + (IkPattern::Ik, Some(static_bundle)) => { + let remote_bundle = + decrypt_peer_bundle(crypto, &mut self.symmetric, static_bundle)?; + if remote_bundle.qid != header.sender { + return Err(Error::InvalidRemoteBundle); + } + match self.remote_bundle.as_ref() { + Some(expected) if expected != &remote_bundle => { + return Err(Error::InvalidRemoteBundle); + } + Some(_) => {} + None => self.remote_bundle = Some(remote_bundle), + } } - Some(_) => {} - None => self.remote_bundle = Some(remote_bundle), + (IkPattern::Kk, None) => {} + _ => return Err(Error::InvalidState), } + self.remote_transport_params = Some(message.transport_params); - self.step = IkStep::Send2; + self.step = Step::Send2; Ok(()) } + pub fn write_2( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { + if self.step != Step::Send2 { + return Err(Error::InvalidState); + } + require_handshake_id(self.handshake_id.as_ref(), handshake_id)?; + let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; + let header = RouteHeader { + sender: self.local.qid, + recipient: remote_bundle.qid, + }; + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + header, + match self.pattern { + IkPattern::Ik => HandshakeKind::Ik2, + IkPattern::Kk => HandshakeKind::Kk2, + }, + handshake_id, + self.local_transport_params, + ); + + let remote_ephemeral = self.remote_ephemeral.as_ref().ok_or(Error::InvalidState)?; + let (ekem_ciphertext, ekem_secret) = + crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); + self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); + self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); + + let (skem_ciphertext, skem_secret) = + crypto.mlkem_encapsulate(&remote_bundle.mlkem_public_key); + let skem_ciphertext = + encrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &skem_ciphertext)?; + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + self.step = Step::Done; + Ok(Ik2 { + handshake_id, + transport_params: self.local_transport_params, + ekem_ciphertext, + skem_ciphertext, + }) + } + pub fn read_2( &mut self, crypto: &impl QlCrypto, header: RouteHeader, message: &Ik2, ) -> Result<(), Error> { - if self.step != IkStep::Recv2 { + if self.step != Step::Recv2 { return Err(Error::InvalidState); } - require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_recipient(header)?; - self.ensure_known_remote_sender(header)?; + require_handshake_id(self.handshake_id.as_ref(), message.handshake_id)?; + self.ensure_inbound_header(header)?; mix_hash_routed_handshake( &mut self.symmetric, crypto, header, - HandshakeKind::Ik2, - message.meta, + match self.pattern { + IkPattern::Ik => HandshakeKind::Ik2, + IkPattern::Kk => HandshakeKind::Kk2, + }, + message.handshake_id, message.transport_params, ); + let local_ephemeral = self.local_ephemeral.as_ref().ok_or(Error::InvalidState)?; self.symmetric .mix_hash(crypto, message.ekem_ciphertext.as_bytes()); @@ -336,7 +447,7 @@ impl IkHandshake { .mix_key_and_hash(crypto, skem_secret.as_bytes()); self.remote_transport_params = Some(message.transport_params); - self.step = IkStep::Done; + self.step = Step::Done; Ok(()) } @@ -354,4 +465,16 @@ impl IkHandshake { remote_transport_params, )) } + + fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), Error> { + if header.recipient != self.local.qid { + return Err(Error::InvalidRouteHeader); + } + if let Some(remote_bundle) = self.remote_bundle.as_ref() { + if header.sender != remote_bundle.qid { + return Err(Error::InvalidRouteHeader); + } + } + Ok(()) + } } diff --git a/ql-wire/src/handshake/kk.rs b/ql-wire/src/handshake/kk.rs deleted file mode 100644 index 2f2db1ef..00000000 --- a/ql-wire/src/handshake/kk.rs +++ /dev/null @@ -1,329 +0,0 @@ -use ql_codec::{ByteSlice, Encode}; - -use super::{ - decrypt_mlkem_ciphertext, encrypt_mlkem_ciphertext, finalize_handshake, - generate_ephemeral_keypair, init_kk_symmetric, initialize_handshake_meta, mix_hash_ephemeral, - mix_hash_routed_handshake, require_handshake_meta, EncryptedMlKemCiphertext, EphemeralKeyPair, - EphemeralPublicKey, FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, -}; -use crate::{ - Error, HandshakeKind, HandshakeMeta, MlKemCiphertext, PeerBundle, QlCrypto, QlIdentity, -}; - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct Kk1 { - pub meta: HandshakeMeta, - pub transport_params: TransportParams, - pub skem_ciphertext: MlKemCiphertext, - pub ephemeral: EphemeralPublicKey, -} - -impl ql_codec::Decode for Kk1 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - meta: reader.decode()?, - transport_params: reader.decode()?, - skem_ciphertext: reader.decode()?, - ephemeral: reader.decode()?, - }) - } -} - -impl Encode for Kk1 { - fn encoded_len(&self) -> usize { - HandshakeMeta::WIRE_SIZE - + TransportParams::WIRE_SIZE - + MlKemCiphertext::SIZE - + EphemeralPublicKey::WIRE_SIZE - } - - fn encode(&self, out: &mut W) { - self.meta.encode(out); - self.transport_params.encode(out); - self.skem_ciphertext.encode(out); - self.ephemeral.encode(out); - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct Kk2 { - pub meta: HandshakeMeta, - pub transport_params: TransportParams, - pub ekem_ciphertext: MlKemCiphertext, - pub skem_ciphertext: EncryptedMlKemCiphertext, -} - -impl ql_codec::Decode for Kk2 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - meta: reader.decode()?, - transport_params: reader.decode()?, - ekem_ciphertext: reader.decode()?, - skem_ciphertext: reader.decode()?, - }) - } -} - -impl Encode for Kk2 { - fn encoded_len(&self) -> usize { - HandshakeMeta::WIRE_SIZE - + TransportParams::WIRE_SIZE - + MlKemCiphertext::SIZE - + EncryptedMlKemCiphertext::WIRE_SIZE - } - - fn encode(&self, out: &mut W) { - self.meta.encode(out); - self.transport_params.encode(out); - self.ekem_ciphertext.encode(out); - self.skem_ciphertext.encode(out); - } -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -enum KkStep { - Send1, - Recv1, - Send2, - Recv2, - Done, -} - -#[derive(Debug, Clone)] -pub struct KkHandshake { - role: Role, - step: KkStep, - symmetric: SymmetricState, - local: QlIdentity, - remote_bundle: PeerBundle, - local_ephemeral: Option, - remote_ephemeral: Option, - handshake_meta: Option, - local_transport_params: TransportParams, - remote_transport_params: Option, -} - -impl KkHandshake { - pub fn new_initiator( - crypto: &impl QlCrypto, - local: QlIdentity, - remote_bundle: PeerBundle, - local_transport_params: TransportParams, - ) -> Self { - let symmetric = init_kk_symmetric(crypto, &local.bundle(), &remote_bundle); - Self { - role: Role::Initiator, - step: KkStep::Send1, - symmetric, - local, - remote_bundle, - local_ephemeral: None, - remote_ephemeral: None, - handshake_meta: None, - local_transport_params, - remote_transport_params: None, - } - } - - pub fn new_responder( - crypto: &impl QlCrypto, - local: QlIdentity, - remote_bundle: PeerBundle, - local_transport_params: TransportParams, - ) -> Self { - let symmetric = init_kk_symmetric(crypto, &remote_bundle, &local.bundle()); - Self { - role: Role::Responder, - step: KkStep::Recv1, - symmetric, - local, - remote_bundle, - local_ephemeral: None, - remote_ephemeral: None, - handshake_meta: None, - local_transport_params, - remote_transport_params: None, - } - } - - pub fn is_finished(&self) -> bool { - self.step == KkStep::Done - } - - fn outbound_header(&self) -> RouteHeader { - RouteHeader { - sender: self.local.qid, - recipient: self.remote_bundle.qid, - } - } - - fn inbound_header(&self) -> RouteHeader { - RouteHeader { - sender: self.remote_bundle.qid, - recipient: self.local.qid, - } - } - - fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), Error> { - if header == self.inbound_header() { - Ok(()) - } else { - Err(Error::InvalidRouteHeader) - } - } - - pub fn write_1(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { - if self.step != KkStep::Send1 { - return Err(Error::InvalidState); - } - initialize_handshake_meta(&mut self.handshake_meta, meta)?; - let header = self.outbound_header(); - mix_hash_routed_handshake( - &mut self.symmetric, - crypto, - header, - HandshakeKind::Kk1, - meta, - self.local_transport_params, - ); - let (skem_ciphertext, skem_secret) = - crypto.mlkem_encapsulate(&self.remote_bundle.mlkem_public_key); - self.symmetric - .encrypt_and_hash(crypto, skem_ciphertext.as_bytes())?; - self.symmetric - .mix_key_and_hash(crypto, skem_secret.as_bytes()); - - let local_ephemeral = generate_ephemeral_keypair(crypto); - let public = local_ephemeral.public(); - mix_hash_ephemeral(&mut self.symmetric, crypto, &public); - - self.local_ephemeral = Some(local_ephemeral); - self.step = KkStep::Recv2; - Ok(Kk1 { - meta, - transport_params: self.local_transport_params, - skem_ciphertext, - ephemeral: public, - }) - } - - pub fn write_2(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { - if self.step != KkStep::Send2 { - return Err(Error::InvalidState); - } - require_handshake_meta(self.handshake_meta.as_ref(), meta)?; - let header = self.outbound_header(); - mix_hash_routed_handshake( - &mut self.symmetric, - crypto, - header, - HandshakeKind::Kk2, - meta, - self.local_transport_params, - ); - let remote_ephemeral = self.remote_ephemeral.clone().ok_or(Error::InvalidState)?; - let (ekem_ciphertext, ekem_secret) = - crypto.mlkem_encapsulate(&remote_ephemeral.mlkem_public_key); - self.symmetric.mix_hash(crypto, ekem_ciphertext.as_bytes()); - self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); - - let (skem_ciphertext, skem_secret) = - crypto.mlkem_encapsulate(&self.remote_bundle.mlkem_public_key); - let skem_ciphertext = - encrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &skem_ciphertext)?; - self.symmetric - .mix_key_and_hash(crypto, skem_secret.as_bytes()); - - self.step = KkStep::Done; - Ok(Kk2 { - meta, - transport_params: self.local_transport_params, - ekem_ciphertext, - skem_ciphertext, - }) - } - - pub fn read_1( - &mut self, - crypto: &impl QlCrypto, - header: RouteHeader, - message: &Kk1, - ) -> Result<(), Error> { - if self.step != KkStep::Recv1 { - return Err(Error::InvalidState); - } - initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; - self.ensure_inbound_header(header)?; - mix_hash_routed_handshake( - &mut self.symmetric, - crypto, - header, - HandshakeKind::Kk1, - message.meta, - message.transport_params, - ); - self.symmetric - .decrypt_and_hash(crypto, message.skem_ciphertext.as_bytes())?; - let skem_secret = - crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &message.skem_ciphertext); - self.symmetric - .mix_key_and_hash(crypto, skem_secret.as_bytes()); - - mix_hash_ephemeral(&mut self.symmetric, crypto, &message.ephemeral); - self.remote_ephemeral = Some(message.ephemeral.clone()); - self.remote_transport_params = Some(message.transport_params); - self.step = KkStep::Send2; - Ok(()) - } - - pub fn read_2( - &mut self, - crypto: &impl QlCrypto, - header: RouteHeader, - message: &Kk2, - ) -> Result<(), Error> { - if self.step != KkStep::Recv2 { - return Err(Error::InvalidState); - } - require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_header(header)?; - mix_hash_routed_handshake( - &mut self.symmetric, - crypto, - header, - HandshakeKind::Kk2, - message.meta, - message.transport_params, - ); - let local_ephemeral = self.local_ephemeral.as_ref().ok_or(Error::InvalidState)?; - self.symmetric - .mix_hash(crypto, message.ekem_ciphertext.as_bytes()); - let ekem_secret = - crypto.mlkem_decapsulate(&local_ephemeral.mlkem.private, &message.ekem_ciphertext); - self.symmetric.mix_key(crypto, ekem_secret.as_bytes()); - - let skem_ciphertext = - decrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &message.skem_ciphertext)?; - let skem_secret = crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &skem_ciphertext); - self.symmetric - .mix_key_and_hash(crypto, skem_secret.as_bytes()); - - self.remote_transport_params = Some(message.transport_params); - self.step = KkStep::Done; - Ok(()) - } - - pub fn finalize(self, crypto: &impl QlCrypto) -> Result { - if !self.is_finished() { - return Err(Error::InvalidState); - } - let remote_transport_params = self.remote_transport_params.ok_or(Error::InvalidState)?; - Ok(finalize_handshake( - crypto, - &self.symmetric, - self.role, - self.remote_bundle, - remote_transport_params, - )) - } -} diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index df0d93ab..b7c8a4da 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -5,16 +5,14 @@ use crate::{ PeerBundle, QlCrypto, RouteHeader, SessionKey, ENCRYPTED_MESSAGE_AUTH_SIZE, }; +mod id; mod ik; -mod kk; -mod meta; mod pairing; mod transport_params; mod xx; -pub use ik::{Ik1, Ik2, IkHandshake}; -pub use kk::{Kk1, Kk2, KkHandshake}; -pub use meta::{HandshakeId, HandshakeMeta}; +pub use id::HandshakeId; +pub use ik::{Ik1, Ik2, IkHandshake, IkPattern}; pub use pairing::{PairingId, PairingToken}; pub use transport_params::TransportParams; pub use xx::{Xx1, Xx2, Xx3, Xx4, XxHandshake}; @@ -22,7 +20,6 @@ pub use xx::{Xx1, Xx2, Xx3, Xx4, XxHandshake}; const SHA256_BLOCK_LEN: usize = 64; const PROTOCOL_IK: &[u8] = b"ql-wire:pq-ik:v1"; const PROTOCOL_KK: &[u8] = b"ql-wire:pq-kk:v1"; -const PROTOCOL_XX: &[u8] = b"ql-wire:pq-xx:v1"; const HANDSHAKE_PREAMBLE_DOMAIN: &[u8] = b"ql-wire:handshake-preamble:v1"; #[derive(Debug, Clone, PartialEq, Eq)] @@ -222,6 +219,10 @@ impl SymmetricState { self.handshake_hash = crypto.sha256(&[&self.handshake_hash, data]); } + fn mix_hash_ephemeral(&mut self, crypto: &impl QlCrypto, public: &EphemeralPublicKey) { + self.mix_hash(crypto, public.mlkem_public_key.as_bytes()); + } + fn mix_key(&mut self, crypto: &impl QlCrypto, input_key_material: &[u8]) { let (chaining_key, cipher_key) = hkdf2(crypto, &self.chaining_key, input_key_material); self.chaining_key = chaining_key; @@ -236,6 +237,10 @@ impl SymmetricState { self.cipher.initialize_key(cipher_key); } + fn mix_psk_pairing_token(&mut self, crypto: &impl QlCrypto, pairing_token: PairingToken) { + self.mix_key_and_hash(crypto, &pairing_token.psk(crypto)); + } + fn encrypt_and_hash( &mut self, crypto: &impl QlCrypto, @@ -281,55 +286,18 @@ impl SymmetricState { } } -fn init_kk_symmetric( - crypto: &impl QlCrypto, - initiator_bundle: &PeerBundle, - responder_bundle: &PeerBundle, -) -> SymmetricState { - let mut symmetric = SymmetricState::new(crypto, PROTOCOL_KK); - symmetric.mix_hash(crypto, &initiator_bundle.encode_vec()); - symmetric.mix_hash(crypto, &responder_bundle.encode_vec()); - symmetric -} - -fn init_ik_symmetric(crypto: &impl QlCrypto, responder_bundle: &PeerBundle) -> SymmetricState { - let mut symmetric = SymmetricState::new(crypto, PROTOCOL_IK); - symmetric.mix_hash(crypto, &responder_bundle.encode_vec()); - symmetric -} - -fn init_xx_symmetric(crypto: &impl QlCrypto) -> SymmetricState { - SymmetricState::new(crypto, PROTOCOL_XX) -} - -fn mix_psk_pairing_token( - symmetric: &mut SymmetricState, - crypto: &impl QlCrypto, - pairing_token: PairingToken, -) { - symmetric.mix_key_and_hash(crypto, &pairing_token.psk(crypto)); -} - fn generate_ephemeral_keypair(crypto: &impl QlCrypto) -> EphemeralKeyPair { EphemeralKeyPair { mlkem: crypto.mlkem_generate_keypair(), } } -fn mix_hash_ephemeral( - symmetric: &mut SymmetricState, - crypto: &impl QlCrypto, - public: &EphemeralPublicKey, -) { - symmetric.mix_hash(crypto, public.mlkem_public_key.as_bytes()); -} - fn mix_hash_routed_handshake( symmetric: &mut SymmetricState, crypto: &impl QlCrypto, header: RouteHeader, kind: HandshakeKind, - meta: HandshakeMeta, + handshake_id: HandshakeId, transport_params: TransportParams, ) { mix_hash_handshake_preamble( @@ -337,7 +305,7 @@ fn mix_hash_routed_handshake( crypto, &header.encode_vec(), kind, - meta, + handshake_id, transport_params, ); } @@ -347,13 +315,20 @@ fn mix_hash_pairing_handshake( crypto: &impl QlCrypto, header: RouteHeader, kind: HandshakeKind, - meta: HandshakeMeta, + handshake_id: HandshakeId, pairing_id: PairingId, transport_params: TransportParams, ) { let mut preamble = header.encode_vec(); pairing_id.encode(&mut preamble); - mix_hash_handshake_preamble(symmetric, crypto, &preamble, kind, meta, transport_params); + mix_hash_handshake_preamble( + symmetric, + crypto, + &preamble, + kind, + handshake_id, + transport_params, + ); } fn mix_hash_handshake_preamble( @@ -361,61 +336,37 @@ fn mix_hash_handshake_preamble( crypto: &impl QlCrypto, header: &[u8], kind: HandshakeKind, - meta: HandshakeMeta, + handshake_id: HandshakeId, transport_params: TransportParams, ) { symmetric.mix_hash(crypto, HANDSHAKE_PREAMBLE_DOMAIN); symmetric.mix_hash(crypto, header); symmetric.mix_hash(crypto, &[kind as u8]); - symmetric.mix_hash(crypto, &meta.encode_vec()); + symmetric.mix_hash(crypto, &handshake_id.encode_vec()); symmetric.mix_hash(crypto, &transport_params.encode_vec()); } -fn initialize_handshake_meta( - expected: &mut Option, - meta: HandshakeMeta, +fn initialize_handshake_id( + expected: &mut Option, + handshake_id: HandshakeId, ) -> Result<(), Error> { match expected { - Some(stored) if *stored != meta => Err(Error::InvalidHandshakeMeta), + Some(stored) if *stored != handshake_id => Err(Error::InvalidHandshakeId), Some(_) => Ok(()), None => { - *expected = Some(meta); + *expected = Some(handshake_id); Ok(()) } } } -fn require_handshake_meta( - expected: Option<&HandshakeMeta>, - meta: HandshakeMeta, -) -> Result<(), Error> { - match expected { - Some(stored) if *stored == meta => Ok(()), - _ => Err(Error::InvalidHandshakeMeta), - } -} - -fn initialize_transport_params( - expected: &mut Option, - transport_params: TransportParams, -) -> Result<(), Error> { - match expected { - Some(stored) if *stored != transport_params => Err(Error::InvalidTransportParams), - Some(_) => Ok(()), - None => { - *expected = Some(transport_params); - Ok(()) - } - } -} - -fn require_transport_params( - expected: Option<&TransportParams>, - transport_params: TransportParams, +fn require_handshake_id( + expected: Option<&HandshakeId>, + handshake_id: HandshakeId, ) -> Result<(), Error> { match expected { - Some(stored) if *stored == transport_params => Ok(()), - _ => Err(Error::InvalidTransportParams), + Some(stored) if *stored == handshake_id => Ok(()), + _ => Err(Error::InvalidHandshakeId), } } diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs index 39f4b680..22e607ee 100644 --- a/ql-wire/src/handshake/xx.rs +++ b/ql-wire/src/handshake/xx.rs @@ -3,20 +3,21 @@ use ql_common::QID; use super::{ decrypt_mlkem_ciphertext, decrypt_peer_bundle, encrypt_mlkem_ciphertext, encrypt_peer_bundle, - finalize_handshake, generate_ephemeral_keypair, init_xx_symmetric, initialize_handshake_meta, - initialize_transport_params, mix_hash_ephemeral, mix_hash_pairing_handshake, - mix_psk_pairing_token, require_handshake_meta, require_transport_params, - EncryptedMlKemCiphertext, EncryptedPeerBundle, EphemeralKeyPair, EphemeralPublicKey, - FinalizedHandshake, Role, RouteHeader, SymmetricState, TransportParams, + finalize_handshake, generate_ephemeral_keypair, initialize_handshake_id, + mix_hash_pairing_handshake, require_handshake_id, EncryptedMlKemCiphertext, + EncryptedPeerBundle, EphemeralKeyPair, EphemeralPublicKey, FinalizedHandshake, Role, + RouteHeader, SymmetricState, TransportParams, }; use crate::{ - Error, HandshakeKind, HandshakeMeta, MlKemCiphertext, PairingId, PairingToken, PeerBundle, - QlCrypto, QlIdentity, + Error, HandshakeId, HandshakeKind, MlKemCiphertext, MlKemPublicKey, PairingId, PairingToken, + PeerBundle, QlCrypto, QlIdentity, }; +const PROTOCOL_XX: &[u8] = b"ql-wire:pq-xx:v1"; + #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx1 { - pub meta: HandshakeMeta, + pub handshake_id: HandshakeId, pub pairing_id: PairingId, pub transport_params: TransportParams, pub ephemeral: EphemeralPublicKey, @@ -25,7 +26,7 @@ pub struct Xx1 { impl ql_codec::Decode for Xx1 { fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { - meta: reader.decode()?, + handshake_id: reader.decode()?, pairing_id: reader.decode()?, transport_params: reader.decode()?, ephemeral: reader.decode()?, @@ -35,14 +36,14 @@ impl ql_codec::Decode for Xx1 { impl Encode for Xx1 { fn encoded_len(&self) -> usize { - HandshakeMeta::WIRE_SIZE + HandshakeId::WIRE_SIZE + PairingId::SIZE + TransportParams::WIRE_SIZE + EphemeralPublicKey::WIRE_SIZE } fn encode(&self, out: &mut W) { - self.meta.encode(out); + self.handshake_id.encode(out); self.pairing_id.encode(out); self.transport_params.encode(out); self.ephemeral.encode(out); @@ -51,8 +52,7 @@ impl Encode for Xx1 { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx2 { - pub meta: HandshakeMeta, - pub pairing_id: PairingId, + pub handshake_id: HandshakeId, pub transport_params: TransportParams, pub ekem_ciphertext: MlKemCiphertext, pub static_bundle: EncryptedPeerBundle, @@ -61,8 +61,7 @@ pub struct Xx2 { impl ql_codec::Decode for Xx2 { fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { - meta: reader.decode()?, - pairing_id: reader.decode()?, + handshake_id: reader.decode()?, transport_params: reader.decode()?, ekem_ciphertext: reader.decode()?, static_bundle: reader.decode()?, @@ -72,16 +71,14 @@ impl ql_codec::Decode for Xx2 { impl Encode for Xx2 { fn encoded_len(&self) -> usize { - HandshakeMeta::WIRE_SIZE - + PairingId::SIZE + HandshakeId::WIRE_SIZE + TransportParams::WIRE_SIZE + MlKemCiphertext::SIZE + self.static_bundle.encoded_len() } fn encode(&self, out: &mut W) { - self.meta.encode(out); - self.pairing_id.encode(out); + self.handshake_id.encode(out); self.transport_params.encode(out); self.ekem_ciphertext.encode(out); self.static_bundle.encode(out); @@ -90,9 +87,7 @@ impl Encode for Xx2 { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx3 { - pub meta: HandshakeMeta, - pub pairing_id: PairingId, - pub transport_params: TransportParams, + pub handshake_id: HandshakeId, pub skem_ciphertext: EncryptedMlKemCiphertext, pub static_bundle: EncryptedPeerBundle, } @@ -100,9 +95,7 @@ pub struct Xx3 { impl ql_codec::Decode for Xx3 { fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { - meta: reader.decode()?, - pairing_id: reader.decode()?, - transport_params: reader.decode()?, + handshake_id: reader.decode()?, skem_ciphertext: reader.decode()?, static_bundle: reader.decode()?, }) @@ -111,17 +104,13 @@ impl ql_codec::Decode for Xx3 { impl Encode for Xx3 { fn encoded_len(&self) -> usize { - HandshakeMeta::WIRE_SIZE - + PairingId::SIZE - + TransportParams::WIRE_SIZE + HandshakeId::WIRE_SIZE + EncryptedMlKemCiphertext::WIRE_SIZE + self.static_bundle.encoded_len() } fn encode(&self, out: &mut W) { - self.meta.encode(out); - self.pairing_id.encode(out); - self.transport_params.encode(out); + self.handshake_id.encode(out); self.skem_ciphertext.encode(out); self.static_bundle.encode(out); } @@ -129,18 +118,14 @@ impl Encode for Xx3 { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx4 { - pub meta: HandshakeMeta, - pub pairing_id: PairingId, - pub transport_params: TransportParams, + pub handshake_id: HandshakeId, pub skem_ciphertext: EncryptedMlKemCiphertext, } impl ql_codec::Decode for Xx4 { fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { - meta: reader.decode()?, - pairing_id: reader.decode()?, - transport_params: reader.decode()?, + handshake_id: reader.decode()?, skem_ciphertext: reader.decode()?, }) } @@ -148,16 +133,11 @@ impl ql_codec::Decode for Xx4 { impl Encode for Xx4 { fn encoded_len(&self) -> usize { - HandshakeMeta::WIRE_SIZE - + PairingId::SIZE - + TransportParams::WIRE_SIZE - + EncryptedMlKemCiphertext::WIRE_SIZE + HandshakeId::WIRE_SIZE + EncryptedMlKemCiphertext::WIRE_SIZE } fn encode(&self, out: &mut W) { - self.meta.encode(out); - self.pairing_id.encode(out); - self.transport_params.encode(out); + self.handshake_id.encode(out); self.skem_ciphertext.encode(out); } } @@ -186,7 +166,7 @@ pub struct XxHandshake { remote_bundle: Option, local_ephemeral: Option, remote_ephemeral: Option, - handshake_meta: Option, + handshake_id: Option, local_transport_params: TransportParams, remote_transport_params: Option, } @@ -202,14 +182,14 @@ impl XxHandshake { Self { role: Role::Initiator, step: XxStep::Send1, - symmetric: init_xx_symmetric(crypto), + symmetric: SymmetricState::new(crypto, PROTOCOL_XX), local, remote_qid, pairing_token, remote_bundle: None, local_ephemeral: None, remote_ephemeral: None, - handshake_meta: None, + handshake_id: None, local_transport_params, remote_transport_params: None, } @@ -225,14 +205,14 @@ impl XxHandshake { Self { role: Role::Responder, step: XxStep::Recv1, - symmetric: init_xx_symmetric(crypto), + symmetric: SymmetricState::new(crypto, PROTOCOL_XX), local, remote_qid, pairing_token, remote_bundle: None, local_ephemeral: None, remote_ephemeral: None, - handshake_meta: None, + handshake_id: None, local_transport_params, remote_transport_params: None, } @@ -250,6 +230,16 @@ impl XxHandshake { self.pairing_token.id(crypto) } + pub fn handshake_id(&self) -> Option { + self.handshake_id + } + + pub fn local_ephemeral(&self) -> Option<&MlKemPublicKey> { + self.local_ephemeral + .as_ref() + .map(|keypair| &keypair.mlkem.public) + } + pub fn remote_qid(&self) -> QID { self.remote_qid } @@ -265,18 +255,10 @@ impl XxHandshake { } } - fn ensure_inbound_header( - &self, - crypto: &impl QlCrypto, - header: RouteHeader, - pairing_id: PairingId, - ) -> Result<(), Error> { + fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), Error> { if header.sender != self.remote_qid || header.recipient != self.local.qid { return Err(Error::InvalidRouteHeader); } - if pairing_id != self.pairing_token.id(crypto) { - return Err(Error::InvalidPairingId); - } Ok(()) } @@ -288,11 +270,15 @@ impl XxHandshake { } } - pub fn write_1(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { + pub fn write_1( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { if self.step != XxStep::Send1 { return Err(Error::InvalidState); } - initialize_handshake_meta(&mut self.handshake_meta, meta)?; + initialize_handshake_id(&mut self.handshake_id, handshake_id)?; let header = self.header(); let pairing_id = self.pairing_token.id(crypto); mix_hash_pairing_handshake( @@ -300,20 +286,21 @@ impl XxHandshake { crypto, header, HandshakeKind::Xx1, - meta, + handshake_id, pairing_id, self.local_transport_params, ); - mix_psk_pairing_token(&mut self.symmetric, crypto, self.pairing_token); + self.symmetric + .mix_psk_pairing_token(crypto, self.pairing_token); let local_ephemeral = generate_ephemeral_keypair(crypto); let ephemeral = local_ephemeral.public(); - mix_hash_ephemeral(&mut self.symmetric, crypto, &ephemeral); + self.symmetric.mix_hash_ephemeral(crypto, &ephemeral); self.local_ephemeral = Some(local_ephemeral); self.step = XxStep::Recv2; Ok(Xx1 { - meta, + handshake_id, pairing_id, transport_params: self.local_transport_params, ephemeral, @@ -329,40 +316,48 @@ impl XxHandshake { if self.step != XxStep::Recv1 { return Err(Error::InvalidState); } - initialize_handshake_meta(&mut self.handshake_meta, message.meta)?; - self.ensure_inbound_header(crypto, header, message.pairing_id)?; + initialize_handshake_id(&mut self.handshake_id, message.handshake_id)?; + self.ensure_inbound_header(header)?; + if message.pairing_id != self.pairing_token.id(crypto) { + return Err(Error::InvalidPairingId); + } mix_hash_pairing_handshake( &mut self.symmetric, crypto, header, HandshakeKind::Xx1, - message.meta, + message.handshake_id, message.pairing_id, message.transport_params, ); - mix_psk_pairing_token(&mut self.symmetric, crypto, self.pairing_token); - mix_hash_ephemeral(&mut self.symmetric, crypto, &message.ephemeral); + self.symmetric + .mix_psk_pairing_token(crypto, self.pairing_token); + self.symmetric + .mix_hash_ephemeral(crypto, &message.ephemeral); self.remote_ephemeral = Some(message.ephemeral.clone()); - initialize_transport_params(&mut self.remote_transport_params, message.transport_params)?; + self.remote_transport_params = Some(message.transport_params); self.step = XxStep::Send2; Ok(()) } - pub fn write_2(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { + pub fn write_2( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { if self.step != XxStep::Send2 { return Err(Error::InvalidState); } - require_handshake_meta(self.handshake_meta.as_ref(), meta)?; + require_handshake_id(self.handshake_id.as_ref(), handshake_id)?; let header = self.header(); - let pairing_id = self.pairing_token.id(crypto); mix_hash_pairing_handshake( &mut self.symmetric, crypto, header, HandshakeKind::Xx2, - meta, - pairing_id, + handshake_id, + self.pairing_token.id(crypto), self.local_transport_params, ); @@ -376,8 +371,7 @@ impl XxHandshake { self.step = XxStep::Recv3; Ok(Xx2 { - meta, - pairing_id, + handshake_id, transport_params: self.local_transport_params, ekem_ciphertext, static_bundle, @@ -393,15 +387,15 @@ impl XxHandshake { if self.step != XxStep::Recv2 { return Err(Error::InvalidState); } - require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_header(crypto, header, message.pairing_id)?; + require_handshake_id(self.handshake_id.as_ref(), message.handshake_id)?; + self.ensure_inbound_header(header)?; mix_hash_pairing_handshake( &mut self.symmetric, crypto, header, HandshakeKind::Xx2, - message.meta, - message.pairing_id, + message.handshake_id, + self.pairing_token.id(crypto), message.transport_params, ); @@ -416,25 +410,28 @@ impl XxHandshake { decrypt_peer_bundle(crypto, &mut self.symmetric, &message.static_bundle)?; self.ensure_remote_bundle(&remote_bundle)?; self.remote_bundle = Some(remote_bundle); - initialize_transport_params(&mut self.remote_transport_params, message.transport_params)?; + self.remote_transport_params = Some(message.transport_params); self.step = XxStep::Send3; Ok(()) } - pub fn write_3(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { + pub fn write_3( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { if self.step != XxStep::Send3 { return Err(Error::InvalidState); } - require_handshake_meta(self.handshake_meta.as_ref(), meta)?; + require_handshake_id(self.handshake_id.as_ref(), handshake_id)?; let header = self.header(); - let pairing_id = self.pairing_token.id(crypto); mix_hash_pairing_handshake( &mut self.symmetric, crypto, header, HandshakeKind::Xx3, - meta, - pairing_id, + handshake_id, + self.pairing_token.id(crypto), self.local_transport_params, ); @@ -450,9 +447,7 @@ impl XxHandshake { self.step = XxStep::Recv4; Ok(Xx3 { - meta, - pairing_id, - transport_params: self.local_transport_params, + handshake_id, skem_ciphertext, static_bundle, }) @@ -467,20 +462,17 @@ impl XxHandshake { if self.step != XxStep::Recv3 { return Err(Error::InvalidState); } - require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_header(crypto, header, message.pairing_id)?; - require_transport_params( - self.remote_transport_params.as_ref(), - message.transport_params, - )?; + require_handshake_id(self.handshake_id.as_ref(), message.handshake_id)?; + self.ensure_inbound_header(header)?; + let remote_transport_params = self.remote_transport_params.ok_or(Error::InvalidState)?; mix_hash_pairing_handshake( &mut self.symmetric, crypto, header, HandshakeKind::Xx3, - message.meta, - message.pairing_id, - message.transport_params, + message.handshake_id, + self.pairing_token.id(crypto), + remote_transport_params, ); let skem_ciphertext = @@ -497,20 +489,23 @@ impl XxHandshake { Ok(()) } - pub fn write_4(&mut self, crypto: &impl QlCrypto, meta: HandshakeMeta) -> Result { + pub fn write_4( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { if self.step != XxStep::Send4 { return Err(Error::InvalidState); } - require_handshake_meta(self.handshake_meta.as_ref(), meta)?; + require_handshake_id(self.handshake_id.as_ref(), handshake_id)?; let header = self.header(); - let pairing_id = self.pairing_token.id(crypto); mix_hash_pairing_handshake( &mut self.symmetric, crypto, header, HandshakeKind::Xx4, - meta, - pairing_id, + handshake_id, + self.pairing_token.id(crypto), self.local_transport_params, ); @@ -524,9 +519,7 @@ impl XxHandshake { self.step = XxStep::Done; Ok(Xx4 { - meta, - pairing_id, - transport_params: self.local_transport_params, + handshake_id, skem_ciphertext, }) } @@ -540,20 +533,17 @@ impl XxHandshake { if self.step != XxStep::Recv4 { return Err(Error::InvalidState); } - require_handshake_meta(self.handshake_meta.as_ref(), message.meta)?; - self.ensure_inbound_header(crypto, header, message.pairing_id)?; - require_transport_params( - self.remote_transport_params.as_ref(), - message.transport_params, - )?; + require_handshake_id(self.handshake_id.as_ref(), message.handshake_id)?; + self.ensure_inbound_header(header)?; + let remote_transport_params = self.remote_transport_params.ok_or(Error::InvalidState)?; mix_hash_pairing_handshake( &mut self.symmetric, crypto, header, HandshakeKind::Xx4, - message.meta, - message.pairing_id, - message.transport_params, + message.handshake_id, + self.pairing_token.id(crypto), + remote_transport_params, ); let skem_ciphertext = diff --git a/ql-wire/src/record.rs b/ql-wire/src/record.rs index b718acb7..3eba0fb8 100644 --- a/ql-wire/src/record.rs +++ b/ql-wire/src/record.rs @@ -2,7 +2,7 @@ use ql_codec::{ByteSlice, Decode, Encode}; use crate::{ encrypted_message::EncryptedMessage, - handshake::{Ik1, Ik2, Kk1, Kk2, Xx1, Xx2, Xx3, Xx4}, + handshake::{Ik1, Ik2, Xx1, Xx2, Xx3, Xx4}, Error, RouteHeader, SessionHeader, QL_WIRE_VERSION, }; @@ -110,8 +110,8 @@ impl Encode for RecordType { pub enum QlHandshakeRecord { Ik1(Ik1), Ik2(Ik2), - Kk1(Kk1), - Kk2(Kk2), + Kk1(Ik1), + Kk2(Ik2), Xx1(Xx1), Xx2(Xx2), Xx3(Xx3), diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index 0873e2b7..3ef370b3 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -12,10 +12,8 @@ fn decode_session_record(bytes: &[u8]) -> QlSessionRecord> { record.into_owned() } -fn handshake_meta(id: u32) -> HandshakeMeta { - HandshakeMeta { - handshake_id: HandshakeId(id), - } +fn handshake_id(id: u32) -> HandshakeId { + HandshakeId(id) } fn handshake_transport_params(window: u32) -> TransportParams { @@ -78,13 +76,13 @@ fn peer_bundle_round_trip() { #[test] fn handshake_record_round_trip_supports_ik_kk_and_xx() { let ik = QlHandshakeRecord::Ik1(Ik1 { - meta: handshake_meta(1), + handshake_id: handshake_id(1), transport_params: handshake_transport_params(65_536), skem_ciphertext: MlKemCiphertext::new(Box::new([7; MlKemCiphertext::SIZE])), ephemeral: EphemeralPublicKey { mlkem_public_key: MlKemPublicKey::new(Box::new([9; MlKemPublicKey::SIZE])), }, - static_bundle: EncryptedPeerBundle(vec![13; 64].into_boxed_slice()), + static_bundle: Some(EncryptedPeerBundle(vec![13; 64].into_boxed_slice())), }); let ik_route = route(1, 2); let ik_encoded = encode_record_vec(RecordHeader::new(ik_route, RecordType::Handshake), &ik); @@ -98,13 +96,14 @@ fn handshake_record_round_trip_supports_ik_kk_and_xx() { ); assert_eq!(decode_handshake_record(ik_encoded.as_slice()), ik); - let kk = QlHandshakeRecord::Kk1(Kk1 { - meta: handshake_meta(2), + let kk = QlHandshakeRecord::Kk1(Ik1 { + handshake_id: handshake_id(2), transport_params: handshake_transport_params(131_072), skem_ciphertext: MlKemCiphertext::new(Box::new([11; MlKemCiphertext::SIZE])), ephemeral: EphemeralPublicKey { mlkem_public_key: MlKemPublicKey::new(Box::new([15; MlKemPublicKey::SIZE])), }, + static_bundle: None, }); let kk_route = route(1, 2); let kk_encoded = encode_record_vec(RecordHeader::new(kk_route, RecordType::Handshake), &kk); @@ -119,7 +118,7 @@ fn handshake_record_round_trip_supports_ik_kk_and_xx() { assert_eq!(decode_handshake_record(kk_encoded.as_slice()), kk); let xx = QlHandshakeRecord::Xx1(Xx1 { - meta: handshake_meta(3), + handshake_id: handshake_id(3), pairing_id: PairingId([3; PairingId::SIZE]), transport_params: handshake_transport_params(196_608), ephemeral: EphemeralPublicKey { @@ -140,35 +139,31 @@ fn handshake_record_round_trip_supports_ik_kk_and_xx() { } #[test] -fn ik_handshake_rejects_tampered_handshake_meta() { +fn ik_handshake_rejects_tampered_handshake_id() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); - let mut initiator_state = IkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_ik_initiator( &crypto, initiator, responder.bundle(), TransportParams::default(), ); let mut responder_state = - IkHandshake::new_responder(&crypto, responder, None, TransportParams::default()); + IkHandshake::new_ik_responder(&crypto, responder, None, TransportParams::default()); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(77)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(77)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let mut m2 = responder_state - .write_2(&crypto, handshake_meta(77)) - .unwrap(); - m2.meta.handshake_id = HandshakeId(78); + let mut m2 = responder_state.write_2(&crypto, handshake_id(77)).unwrap(); + m2.handshake_id = HandshakeId(78); assert_eq!( initiator_state.read_2(&crypto, responder_to_initiator, &m2), - Err(Error::InvalidHandshakeMeta) + Err(Error::InvalidHandshakeId) ); } @@ -178,29 +173,25 @@ fn kk_handshake_rejects_tampered_handshake_header() { let (initiator, responder) = test_identities(&crypto); let (initiator_to_responder, _) = identity_routes(&initiator, &responder); - let mut initiator_state = KkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_kk_initiator( &crypto, initiator.clone(), responder.bundle(), TransportParams::default(), ); - let mut responder_state = KkHandshake::new_responder( + let mut responder_state = IkHandshake::new_kk_responder( &crypto, responder, initiator.bundle(), TransportParams::default(), ); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(88)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(88)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let m2 = responder_state - .write_2(&crypto, handshake_meta(88)) - .unwrap(); + let m2 = responder_state.write_2(&crypto, handshake_id(88)).unwrap(); let tampered_route = route(9, 1); assert_eq!( @@ -215,25 +206,21 @@ fn ik_handshake_rejects_tampered_transport_params() { let (initiator, responder) = test_identities(&crypto); let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); - let mut initiator_state = IkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_ik_initiator( &crypto, initiator, responder.bundle(), handshake_transport_params(4096), ); let mut responder_state = - IkHandshake::new_responder(&crypto, responder, None, handshake_transport_params(8192)); + IkHandshake::new_ik_responder(&crypto, responder, None, handshake_transport_params(8192)); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(89)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(89)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let mut m2 = responder_state - .write_2(&crypto, handshake_meta(89)) - .unwrap(); + let mut m2 = responder_state.write_2(&crypto, handshake_id(89)).unwrap(); m2.transport_params.initial_stream_receive_window += 1; assert_eq!( @@ -248,18 +235,16 @@ fn ik_handshake_rejects_tampered_handshake_header() { let (initiator, responder) = test_identities(&crypto); let (mut initiator_to_responder, _) = identity_routes(&initiator, &responder); - let mut initiator_state = IkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_ik_initiator( &crypto, initiator, responder.bundle(), TransportParams::default(), ); let mut responder_state = - IkHandshake::new_responder(&crypto, responder, None, TransportParams::default()); + IkHandshake::new_ik_responder(&crypto, responder, None, TransportParams::default()); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(90)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(90)).unwrap(); initiator_to_responder.sender = QID([9; QID::SIZE]); assert_eq!( @@ -275,22 +260,20 @@ fn ik_handshake_rejects_bound_remote_bundle_mismatch() { let bogus = generate_identity(&crypto, "bogus"); let (initiator_to_responder, _) = identity_routes(&initiator, &responder); - let mut initiator_state = IkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_ik_initiator( &crypto, initiator, responder.bundle(), TransportParams::default(), ); - let mut responder_state = IkHandshake::new_responder( + let mut responder_state = IkHandshake::new_ik_responder( &crypto, responder, Some(bogus.bundle()), TransportParams::default(), ); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(91)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(91)).unwrap(); assert_eq!( responder_state.read_1(&crypto, initiator_to_responder, &m1), @@ -306,25 +289,21 @@ fn ik_handshake_round_trip_derives_matching_transport_and_learns_remote() { let initiator_params = handshake_transport_params(4096); let responder_params = handshake_transport_params(8192); - let mut initiator_state = IkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_ik_initiator( &crypto, initiator.clone(), responder.bundle(), initiator_params, ); let mut responder_state = - IkHandshake::new_responder(&crypto, responder.clone(), None, responder_params); + IkHandshake::new_ik_responder(&crypto, responder.clone(), None, responder_params); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(11)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(11)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let m2 = responder_state - .write_2(&crypto, handshake_meta(11)) - .unwrap(); + let m2 = responder_state.write_2(&crypto, handshake_id(11)).unwrap(); initiator_state .read_2(&crypto, responder_to_initiator, &m2) .unwrap(); @@ -352,29 +331,25 @@ fn ik_handshake_round_trip_derives_matching_transport_with_bound_responder() { let initiator_params = handshake_transport_params(16_384); let responder_params = handshake_transport_params(32_768); - let mut initiator_state = IkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_ik_initiator( &crypto, initiator.clone(), responder.bundle(), initiator_params, ); - let mut responder_state = IkHandshake::new_responder( + let mut responder_state = IkHandshake::new_ik_responder( &crypto, responder.clone(), Some(initiator.bundle()), responder_params, ); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(12)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(12)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let m2 = responder_state - .write_2(&crypto, handshake_meta(12)) - .unwrap(); + let m2 = responder_state.write_2(&crypto, handshake_id(12)).unwrap(); initiator_state .read_2(&crypto, responder_to_initiator, &m2) .unwrap(); @@ -402,29 +377,25 @@ fn kk_handshake_round_trip_derives_matching_transport() { let initiator_params = handshake_transport_params(24_576); let responder_params = handshake_transport_params(49_152); - let mut initiator_state = KkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_kk_initiator( &crypto, initiator.clone(), responder.bundle(), initiator_params, ); - let mut responder_state = KkHandshake::new_responder( + let mut responder_state = IkHandshake::new_kk_responder( &crypto, responder.clone(), initiator.bundle(), responder_params, ); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(21)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(21)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let m2 = responder_state - .write_2(&crypto, handshake_meta(21)) - .unwrap(); + let m2 = responder_state.write_2(&crypto, handshake_id(21)).unwrap(); initiator_state .read_2(&crypto, responder_to_initiator, &m2) .unwrap(); @@ -450,29 +421,25 @@ fn kk_handshake_rejects_tampered_transport_params() { let (initiator, responder) = test_identities(&crypto); let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); - let mut initiator_state = KkHandshake::new_initiator( + let mut initiator_state = IkHandshake::new_kk_initiator( &crypto, initiator.clone(), responder.bundle(), handshake_transport_params(12288), ); - let mut responder_state = KkHandshake::new_responder( + let mut responder_state = IkHandshake::new_kk_responder( &crypto, responder, initiator.bundle(), handshake_transport_params(24576), ); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(22)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(22)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let mut m2 = responder_state - .write_2(&crypto, handshake_meta(22)) - .unwrap(); + let mut m2 = responder_state.write_2(&crypto, handshake_id(22)).unwrap(); m2.transport_params.initial_stream_receive_window += 1; assert_eq!( @@ -503,9 +470,7 @@ fn xx_handshake_rejects_tampered_pairing_id() { TransportParams::default(), ); - let mut m1 = initiator_state - .write_1(&crypto, handshake_meta(31)) - .unwrap(); + let mut m1 = initiator_state.write_1(&crypto, handshake_id(31)).unwrap(); m1.pairing_id = PairingId([8; PairingId::SIZE]); assert_eq!( @@ -535,9 +500,7 @@ fn xx_handshake_rejects_tampered_sender_or_recipient() { TransportParams::default(), ); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(31)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(31)).unwrap(); let (mut route, _) = identity_routes(&initiator, &responder); route.sender = responder.qid; @@ -561,9 +524,7 @@ fn xx_handshake_rejects_tampered_sender_or_recipient() { TransportParams::default(), ); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(31)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(31)).unwrap(); let (mut route, _) = identity_routes(&initiator, &responder); route.recipient = initiator.qid; @@ -574,7 +535,7 @@ fn xx_handshake_rejects_tampered_sender_or_recipient() { } #[test] -fn xx_handshake_rejects_repeated_transport_param_change() { +fn xx_handshake_rejects_tampered_transport_params() { let crypto = SoftwareCrypto; let (initiator, responder) = test_identities(&crypto); let token = PairingToken([9; PairingToken::SIZE]); @@ -595,28 +556,17 @@ fn xx_handshake_rejects_repeated_transport_param_change() { handshake_transport_params(24_576), ); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(32)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(32)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let m2 = responder_state - .write_2(&crypto, handshake_meta(32)) - .unwrap(); - initiator_state - .read_2(&crypto, responder_to_initiator, &m2) - .unwrap(); - - let mut m3 = initiator_state - .write_3(&crypto, handshake_meta(32)) - .unwrap(); - m3.transport_params.initial_stream_receive_window += 1; + let mut m2 = responder_state.write_2(&crypto, handshake_id(32)).unwrap(); + m2.transport_params.initial_stream_receive_window += 1; assert_eq!( - responder_state.read_3(&crypto, initiator_to_responder, &m3), - Err(Error::InvalidTransportParams) + initiator_state.read_2(&crypto, responder_to_initiator, &m2), + Err(Error::DecryptFailed) ); } @@ -651,33 +601,25 @@ fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { assert!(initiator_state.remote_bundle().is_none()); assert!(responder_state.remote_bundle().is_none()); - let m1 = initiator_state - .write_1(&crypto, handshake_meta(33)) - .unwrap(); + let m1 = initiator_state.write_1(&crypto, handshake_id(33)).unwrap(); responder_state .read_1(&crypto, initiator_to_responder, &m1) .unwrap(); - let m2 = responder_state - .write_2(&crypto, handshake_meta(33)) - .unwrap(); + let m2 = responder_state.write_2(&crypto, handshake_id(33)).unwrap(); initiator_state .read_2(&crypto, responder_to_initiator, &m2) .unwrap(); assert_eq!(initiator_state.remote_bundle(), Some(&responder.bundle())); assert!(responder_state.remote_bundle().is_none()); - let m3 = initiator_state - .write_3(&crypto, handshake_meta(33)) - .unwrap(); + let m3 = initiator_state.write_3(&crypto, handshake_id(33)).unwrap(); responder_state .read_3(&crypto, initiator_to_responder, &m3) .unwrap(); assert_eq!(responder_state.remote_bundle(), Some(&initiator.bundle())); - let m4 = responder_state - .write_4(&crypto, handshake_meta(33)) - .unwrap(); + let m4 = responder_state.write_4(&crypto, handshake_id(33)).unwrap(); initiator_state .read_4(&crypto, responder_to_initiator, &m4) .unwrap(); @@ -797,21 +739,21 @@ fn protocol_record_size_breakdown() { let (initiator, responder) = test_identities(&crypto); let (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); - let mut ik_initiator = IkHandshake::new_initiator( + let mut ik_initiator = IkHandshake::new_ik_initiator( &crypto, initiator.clone(), responder.bundle(), TransportParams::default(), ); let mut ik_responder = - IkHandshake::new_responder(&crypto, responder.clone(), None, TransportParams::default()); + IkHandshake::new_ik_responder(&crypto, responder.clone(), None, TransportParams::default()); - let ik1 = ik_initiator.write_1(&crypto, handshake_meta(101)).unwrap(); + let ik1 = ik_initiator.write_1(&crypto, handshake_id(101)).unwrap(); ik_responder .read_1(&crypto, initiator_to_responder, &ik1) .unwrap(); - let ik2 = ik_responder.write_2(&crypto, handshake_meta(101)).unwrap(); + let ik2 = ik_responder.write_2(&crypto, handshake_id(101)).unwrap(); ik_initiator .read_2(&crypto, responder_to_initiator, &ik2) .unwrap(); @@ -819,25 +761,25 @@ fn protocol_record_size_breakdown() { let ik1 = QlHandshakeRecord::Ik1(ik1); let ik2 = QlHandshakeRecord::Ik2(ik2); - let mut kk_initiator = KkHandshake::new_initiator( + let mut kk_initiator = IkHandshake::new_kk_initiator( &crypto, initiator.clone(), responder.bundle(), TransportParams::default(), ); - let mut kk_responder = KkHandshake::new_responder( + let mut kk_responder = IkHandshake::new_kk_responder( &crypto, responder.clone(), initiator.bundle(), TransportParams::default(), ); - let kk1 = kk_initiator.write_1(&crypto, handshake_meta(201)).unwrap(); + let kk1 = kk_initiator.write_1(&crypto, handshake_id(201)).unwrap(); kk_responder .read_1(&crypto, initiator_to_responder, &kk1) .unwrap(); - let kk2 = kk_responder.write_2(&crypto, handshake_meta(201)).unwrap(); + let kk2 = kk_responder.write_2(&crypto, handshake_id(201)).unwrap(); kk_initiator .read_2(&crypto, responder_to_initiator, &kk2) .unwrap(); @@ -861,22 +803,22 @@ fn protocol_record_size_breakdown() { TransportParams::default(), ); - let xx1 = xx_initiator.write_1(&crypto, handshake_meta(301)).unwrap(); + let xx1 = xx_initiator.write_1(&crypto, handshake_id(301)).unwrap(); xx_responder .read_1(&crypto, initiator_to_responder, &xx1) .unwrap(); - let xx2 = xx_responder.write_2(&crypto, handshake_meta(301)).unwrap(); + let xx2 = xx_responder.write_2(&crypto, handshake_id(301)).unwrap(); xx_initiator .read_2(&crypto, responder_to_initiator, &xx2) .unwrap(); - let xx3 = xx_initiator.write_3(&crypto, handshake_meta(301)).unwrap(); + let xx3 = xx_initiator.write_3(&crypto, handshake_id(301)).unwrap(); xx_responder .read_3(&crypto, initiator_to_responder, &xx3) .unwrap(); - let xx4 = xx_responder.write_4(&crypto, handshake_meta(301)).unwrap(); + let xx4 = xx_responder.write_4(&crypto, handshake_id(301)).unwrap(); xx_initiator .read_4(&crypto, responder_to_initiator, &xx4) .unwrap(); From 4ee0787c9f8ea42d50bfeb72379389038b05a1bc Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Fri, 10 Jul 2026 10:04:27 -0400 Subject: [PATCH 45/59] ql: remove bundle metadata and name wrapper --- ql-codec/src/codec.rs | 44 ++++++++++++++++++++++++++++++++++ ql-wire/src/identity.rs | 53 +++-------------------------------------- ql-wire/src/tests.rs | 4 +--- 3 files changed, 48 insertions(+), 53 deletions(-) diff --git a/ql-codec/src/codec.rs b/ql-codec/src/codec.rs index 891cde89..e54d3926 100644 --- a/ql-codec/src/codec.rs +++ b/ql-codec/src/codec.rs @@ -108,6 +108,41 @@ impl Decode for bytes::Bytes { impl_codec!(owned_byte_decode: Vec, Box<[u8]>); +impl Encode for str { + fn encoded_len(&self) -> usize { + self.as_bytes().encoded_len() + } + + fn encode(&self, out: &mut W) { + self.as_bytes().encode(out); + } +} + +impl<'a> Decode<&'a [u8]> for &'a str { + fn decode(reader: &mut Reader<&'a [u8]>) -> Result { + std::str::from_utf8(reader.take_len_prefixed()?).map_err(|_| Error::InvalidUtf8) + } +} + +impl Encode for String { + fn encoded_len(&self) -> usize { + self.as_str().encoded_len() + } + + fn encode(&self, out: &mut W) { + self.as_str().encode(out); + } +} + +impl Decode for String { + fn decode(reader: &mut Reader) -> Result { + let bytes = reader.take_len_prefixed()?; + std::str::from_utf8(&bytes) + .map(str::to_owned) + .map_err(|_| Error::InvalidUtf8) + } +} + impl Decode for u8 { fn decode(reader: &mut Reader) -> Result { reader.take_u8() @@ -234,5 +269,14 @@ mod tests { .as_ref(), [1, 2, 3] ); + + let encoded = String::from("hello").encode_vec(); + assert_eq!(encoded, [5, b'h', b'e', b'l', b'l', b'o']); + assert_eq!(<&str>::decode_bytes(encoded.as_slice()).unwrap(), "hello"); + assert_eq!(String::decode_bytes(encoded.as_slice()).unwrap(), "hello"); + assert_eq!( + String::decode_bytes([1, 0xff].as_slice()), + Err(Error::InvalidUtf8) + ); } } diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index e23c4493..85d96ea4 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -1,5 +1,3 @@ -use std::ops::Deref; - use ql_codec::{ByteSlice, Encode}; use ql_common::QID; @@ -11,8 +9,7 @@ pub struct PeerBundle { pub qid: QID, pub capabilities: u32, pub mlkem_public_key: MlKemPublicKey, - pub name: QlName, - pub metadata: Box<[u8]>, + pub name: String, } impl PeerBundle { @@ -26,7 +23,6 @@ impl Encode for PeerBundle { + size_of::() + MlKemPublicKey::SIZE + self.name.encoded_len() - + self.metadata.encoded_len() } fn encode(&self, out: &mut W) { @@ -35,7 +31,6 @@ impl Encode for PeerBundle { self.capabilities.encode(out); self.mlkem_public_key.encode(out); self.name.encode(out); - self.metadata.encode(out); } } @@ -47,7 +42,6 @@ impl ql_codec::Decode for PeerBundle { capabilities: reader.decode()?, mlkem_public_key: reader.decode()?, name: reader.decode()?, - metadata: reader.decode()?, }) } } @@ -58,8 +52,7 @@ pub struct QlIdentity { pub mlkem_private_key: MlKemPrivateKey, pub mlkem_public_key: MlKemPublicKey, pub capabilities: u32, - pub name: QlName, - pub metadata: Box<[u8]>, + pub name: String, } impl QlIdentity { @@ -69,23 +62,16 @@ impl QlIdentity { mlkem_public_key: MlKemPublicKey, name: impl Into, ) -> Self { - let name = QlName(name.into()); let qid = derive_qid(crypto, &mlkem_public_key); Self { qid, mlkem_private_key, mlkem_public_key, capabilities: 0, - name, - metadata: Box::default(), + name: name.into(), } } - pub fn with_metadata(mut self, metadata: impl Into>) -> Self { - self.metadata = metadata.into(); - self - } - pub fn bundle(&self) -> PeerBundle { PeerBundle { version: PeerBundle::VERSION, @@ -93,7 +79,6 @@ impl QlIdentity { capabilities: self.capabilities, mlkem_public_key: self.mlkem_public_key.clone(), name: self.name.clone(), - metadata: self.metadata.clone(), } } } @@ -105,7 +90,6 @@ impl Encode for QlIdentity { + MlKemPublicKey::SIZE + size_of::() + self.name.encoded_len() - + self.metadata.encoded_len() } fn encode(&self, out: &mut W) { @@ -114,7 +98,6 @@ impl Encode for QlIdentity { self.mlkem_public_key.encode(out); self.capabilities.encode(out); self.name.encode(out); - self.metadata.encode(out); } } @@ -126,7 +109,6 @@ impl ql_codec::Decode for QlIdentity { mlkem_public_key: reader.decode()?, capabilities: reader.decode()?, name: reader.decode()?, - metadata: reader.decode()?, }) } } @@ -135,32 +117,3 @@ pub fn generate_identity(crypto: &impl QlCrypto, name: impl Into) -> QlI let MlKemKeyPair { private, public } = crypto.mlkem_generate_keypair(); QlIdentity::new(crypto, private, public, name) } - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct QlName(pub String); - -impl Deref for QlName { - type Target = str; - - fn deref(&self) -> &Self::Target { - &self.0 - } -} - -impl Encode for QlName { - fn encoded_len(&self) -> usize { - self.0.as_bytes().encoded_len() - } - - fn encode(&self, out: &mut W) { - self.0.as_bytes().encode(out); - } -} - -impl ql_codec::Decode for QlName { - fn decode(reader: &mut ql_codec::Reader) -> Result { - let bytes = reader.take_len_prefixed()?; - let name = std::str::from_utf8(&bytes).map_err(|_| ql_codec::Error::InvalidUtf8)?; - Ok(QlName(name.into())) - } -} diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index 3ef370b3..d5f0ae63 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -62,15 +62,13 @@ fn peer_bundle_round_trip() { let crypto = SoftwareCrypto; let mut identity = generate_identity(&crypto, "alice"); identity.capabilities = 1231; - identity.metadata = b"peer metadata".to_vec().into_boxed_slice(); let bundle = identity.bundle(); let encoded = bundle.encode_vec(); let decoded = PeerBundle::decode_bytes(encoded.as_slice()).unwrap(); assert_eq!(decoded, bundle); - assert_eq!(&*decoded.name, "alice"); - assert_eq!(&*decoded.metadata, b"peer metadata"); + assert_eq!(decoded.name, "alice"); } #[test] From e09d7e7fdc0f2dd007a71a93af8998bed2b80720 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Mon, 20 Jul 2026 20:48:11 -0400 Subject: [PATCH 46/59] ql-wire: allow borrowed identity --- ql-fsm/src/handshake/ik.rs | 4 ++-- ql-wire/src/handshake/ik.rs | 40 ++++++++++++++++++++----------------- 2 files changed, 24 insertions(+), 20 deletions(-) diff --git a/ql-fsm/src/handshake/ik.rs b/ql-fsm/src/handshake/ik.rs index a588693e..1b1c141e 100644 --- a/ql-fsm/src/handshake/ik.rs +++ b/ql-fsm/src/handshake/ik.rs @@ -74,13 +74,13 @@ pub fn handle_1( let mut handshake = match (pattern, peer) { (IkPattern::Ik, expected_remote) => wire::IkHandshake::new_ik_responder( crypto, - fsm.identity.clone(), + &fsm.identity, expected_remote, super::local_transport_params(fsm), ), (IkPattern::Kk, Some(remote_bundle)) => wire::IkHandshake::new_kk_responder( crypto, - fsm.identity.clone(), + &fsm.identity, remote_bundle, super::local_transport_params(fsm), ), diff --git a/ql-wire/src/handshake/ik.rs b/ql-wire/src/handshake/ik.rs index c2c1d73f..9b30b7b0 100644 --- a/ql-wire/src/handshake/ik.rs +++ b/ql-wire/src/handshake/ik.rs @@ -1,3 +1,5 @@ +use std::borrow::Borrow; + use ql_codec::{ByteSlice, Encode}; use super::{ @@ -116,12 +118,12 @@ pub enum IkPattern { } #[derive(Debug, Clone)] -pub struct IkHandshake { +pub struct IkHandshake { pattern: IkPattern, role: Role, step: Step, symmetric: SymmetricState, - local: QlIdentity, + local: I, remote_bundle: Option, local_ephemeral: Option, remote_ephemeral: Option, @@ -130,7 +132,7 @@ pub struct IkHandshake { remote_transport_params: Option, } -impl IkHandshake { +impl> IkHandshake { pub fn pattern(&self) -> IkPattern { self.pattern } @@ -147,7 +149,7 @@ impl IkHandshake { pub fn new_ik_initiator( crypto: &impl QlCrypto, - local: QlIdentity, + local: I, remote_bundle: PeerBundle, local_transport_params: TransportParams, ) -> Self { @@ -165,12 +167,12 @@ impl IkHandshake { pub fn new_ik_responder( crypto: &impl QlCrypto, - local: QlIdentity, + local: I, expected_remote: Option, local_transport_params: TransportParams, ) -> Self { let mut symmetric = SymmetricState::new(crypto, PROTOCOL_IK); - symmetric.mix_hash(crypto, &local.bundle().encode_vec()); + symmetric.mix_hash(crypto, &local.borrow().bundle().encode_vec()); Self::new( IkPattern::Ik, Role::Responder, @@ -183,12 +185,12 @@ impl IkHandshake { pub fn new_kk_initiator( crypto: &impl QlCrypto, - local: QlIdentity, + local: I, remote_bundle: PeerBundle, local_transport_params: TransportParams, ) -> Self { let mut symmetric = SymmetricState::new(crypto, PROTOCOL_KK); - symmetric.mix_hash(crypto, &local.bundle().encode_vec()); + symmetric.mix_hash(crypto, &local.borrow().bundle().encode_vec()); symmetric.mix_hash(crypto, &remote_bundle.encode_vec()); Self::new( IkPattern::Kk, @@ -202,13 +204,13 @@ impl IkHandshake { pub fn new_kk_responder( crypto: &impl QlCrypto, - local: QlIdentity, + local: I, remote_bundle: PeerBundle, local_transport_params: TransportParams, ) -> Self { let mut symmetric = SymmetricState::new(crypto, PROTOCOL_KK); symmetric.mix_hash(crypto, &remote_bundle.encode_vec()); - symmetric.mix_hash(crypto, &local.bundle().encode_vec()); + symmetric.mix_hash(crypto, &local.borrow().bundle().encode_vec()); Self::new( IkPattern::Kk, Role::Responder, @@ -223,7 +225,7 @@ impl IkHandshake { pattern: IkPattern, role: Role, symmetric: SymmetricState, - local: QlIdentity, + local: I, remote_bundle: Option, local_transport_params: TransportParams, ) -> Self { @@ -244,7 +246,6 @@ impl IkHandshake { remote_transport_params: None, } } - pub fn is_finished(&self) -> bool { self.step == Step::Done } @@ -259,8 +260,9 @@ impl IkHandshake { } initialize_handshake_id(&mut self.handshake_id, handshake_id)?; let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; + let local = self.local.borrow(); let header = RouteHeader { - sender: self.local.qid, + sender: local.qid, recipient: remote_bundle.qid, }; mix_hash_routed_handshake( @@ -289,7 +291,7 @@ impl IkHandshake { IkPattern::Ik => Some(encrypt_peer_bundle( crypto, &mut self.symmetric, - &self.local.bundle(), + &local.bundle(), )?), IkPattern::Kk => None, }; @@ -316,6 +318,7 @@ impl IkHandshake { } initialize_handshake_id(&mut self.handshake_id, message.handshake_id)?; self.ensure_inbound_header(header)?; + let local = self.local.borrow(); mix_hash_routed_handshake( &mut self.symmetric, crypto, @@ -330,7 +333,7 @@ impl IkHandshake { self.symmetric .mix_hash(crypto, message.skem_ciphertext.as_bytes()); let skem_secret = - crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &message.skem_ciphertext); + crypto.mlkem_decapsulate(&local.mlkem_private_key, &message.skem_ciphertext); self.symmetric .mix_key_and_hash(crypto, skem_secret.as_bytes()); @@ -373,7 +376,7 @@ impl IkHandshake { require_handshake_id(self.handshake_id.as_ref(), handshake_id)?; let remote_bundle = self.remote_bundle.as_ref().ok_or(Error::InvalidState)?; let header = RouteHeader { - sender: self.local.qid, + sender: self.local.borrow().qid, recipient: remote_bundle.qid, }; mix_hash_routed_handshake( @@ -442,7 +445,8 @@ impl IkHandshake { let skem_ciphertext = decrypt_mlkem_ciphertext(crypto, &mut self.symmetric, &message.skem_ciphertext)?; - let skem_secret = crypto.mlkem_decapsulate(&self.local.mlkem_private_key, &skem_ciphertext); + let skem_secret = + crypto.mlkem_decapsulate(&self.local.borrow().mlkem_private_key, &skem_ciphertext); self.symmetric .mix_key_and_hash(crypto, skem_secret.as_bytes()); @@ -467,7 +471,7 @@ impl IkHandshake { } fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), Error> { - if header.recipient != self.local.qid { + if header.recipient != self.local.borrow().qid { return Err(Error::InvalidRouteHeader); } if let Some(remote_bundle) = self.remote_bundle.as_ref() { From d74b4313bbfb1d0722153eee2420dcf4e3db69ab Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 21 Jul 2026 09:46:28 -0400 Subject: [PATCH 47/59] ql-wire: peer challenge --- ql-wire/src/handshake/challenge.rs | 142 +++++++++++++++++++++++++++++ ql-wire/src/handshake/mod.rs | 11 +-- ql-wire/src/identity.rs | 9 +- 3 files changed, 155 insertions(+), 7 deletions(-) create mode 100644 ql-wire/src/handshake/challenge.rs diff --git a/ql-wire/src/handshake/challenge.rs b/ql-wire/src/handshake/challenge.rs new file mode 100644 index 00000000..54f4c848 --- /dev/null +++ b/ql-wire/src/handshake/challenge.rs @@ -0,0 +1,142 @@ +use std::borrow::Borrow; + +use ql_codec::Reader; +use ql_common::QID; + +use super::{HandshakeId, IkHandshake, TransportParams}; +use crate::{ + decrypt_record, encode_record_vec, parse_session_frames, Error, PeerBundle, QlCrypto, + QlHandshakeRecord, QlIdentity, QlSessionRecord, RecordHeader, RecordSeq, RecordType, + RouteHeader, SessionFrame, SessionKey, SessionRecordBuilder, QL_WIRE_VERSION, +}; + +pub struct PeerChallenge { + handshake: IkHandshake, +} + +impl> PeerChallenge { + pub fn new( + crypto: &impl QlCrypto, + local: I, + remote: PeerBundle, + handshake_id: HandshakeId, + ) -> Result<(Self, Vec), Error> { + let route = RouteHeader { + sender: local.borrow().qid, + recipient: remote.qid, + }; + let mut handshake = + IkHandshake::new_kk_initiator(crypto, local, remote, TransportParams::default()); + let request = handshake.write_1(crypto, handshake_id)?; + let request = encode_record_vec( + RecordHeader::new(route, RecordType::Handshake), + &QlHandshakeRecord::Kk1(request), + ); + Ok((Self { handshake }, request)) + } + + pub fn verify( + mut self, + crypto: &impl QlCrypto, + response: &[u8], + ) -> Result<(QID, Vec), Error> { + let (header, response) = decode_handshake(response)?; + let QlHandshakeRecord::Kk2(response) = response else { + return Err(Error::InvalidPayload); + }; + self.handshake.read_2(crypto, header.route, &response)?; + let session = self.handshake.finalize(crypto)?; + let qid = session.remote_bundle.qid; + let route = RouteHeader { + sender: header.route.recipient, + recipient: qid, + }; + let mut confirmation = + SessionRecordBuilder::new(RecordSeq(0), SessionRecordBuilder::MIN_CAPACITY + 1); + confirmation.push_ping(); + Ok((qid, confirmation.encrypt(crypto, route, &session.tx_key))) + } +} + +pub struct PendingChallengeConfirmation { + route: RouteHeader, + rx_key: SessionKey, +} + +impl PendingChallengeConfirmation { + pub fn verify(self, crypto: &impl QlCrypto, confirmation: &mut [u8]) -> Result<(), Error> { + let mut reader = Reader::new(confirmation); + let header = reader.decode::()?; + if header.version != QL_WIRE_VERSION + || header.route != self.route + || header.record_type != RecordType::Session + { + return Err(Error::InvalidPayload); + } + let record = reader.decode::>()?; + if record.header.seq != RecordSeq(0) { + return Err(Error::InvalidPayload); + } + let payload = decrypt_record( + crypto, + &header, + &record.header, + record.payload, + &self.rx_key, + )?; + let mut frames = parse_session_frames(payload); + if !matches!(frames.next().transpose()?, Some(SessionFrame::Ping)) + || frames.next().is_some() + { + return Err(Error::InvalidPayload); + } + Ok(()) + } +} + +pub fn answer_peer_challenge>( + crypto: &impl QlCrypto, + local: I, + challenger: PeerBundle, + request: &[u8], +) -> Result<(Vec, PendingChallengeConfirmation), Error> { + let (header, request) = decode_handshake(request)?; + let QlHandshakeRecord::Kk1(request) = request else { + return Err(Error::InvalidPayload); + }; + let mut handshake = + IkHandshake::new_kk_responder(crypto, local, challenger, TransportParams::default()); + handshake.read_1(crypto, header.route, &request)?; + let response = handshake.write_2(crypto, request.handshake_id)?; + let session = handshake.finalize(crypto)?; + let response = encode_record_vec( + RecordHeader::new( + RouteHeader { + sender: header.route.recipient, + recipient: header.route.sender, + }, + RecordType::Handshake, + ), + &QlHandshakeRecord::Kk2(response), + ); + Ok(( + response, + PendingChallengeConfirmation { + route: header.route, + rx_key: session.rx_key, + }, + )) +} + +fn decode_handshake(bytes: &[u8]) -> Result<(RecordHeader, QlHandshakeRecord), Error> { + let mut reader = Reader::new(bytes); + let header = reader.decode::()?; + if header.version != QL_WIRE_VERSION || header.record_type != RecordType::Handshake { + return Err(Error::InvalidPayload); + } + let record = reader.decode()?; + if !reader.is_empty() { + return Err(Error::InvalidPayload); + } + Ok((header, record)) +} diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index b7c8a4da..db585fc5 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -1,16 +1,18 @@ use ql_codec::{ByteSlice, Decode, Encode}; use crate::{ - derive_qid, Error, HandshakeKind, MlKemCiphertext, MlKemKeyPair, MlKemPublicKey, Nonce, - PeerBundle, QlCrypto, RouteHeader, SessionKey, ENCRYPTED_MESSAGE_AUTH_SIZE, + Error, HandshakeKind, MlKemCiphertext, MlKemKeyPair, MlKemPublicKey, Nonce, PeerBundle, + QlCrypto, RouteHeader, SessionKey, ENCRYPTED_MESSAGE_AUTH_SIZE, }; +mod challenge; mod id; mod ik; mod pairing; mod transport_params; mod xx; +pub use challenge::{answer_peer_challenge, PeerChallenge, PendingChallengeConfirmation}; pub use id::HandshakeId; pub use ik::{Ik1, Ik2, IkHandshake, IkPattern}; pub use pairing::{PairingId, PairingToken}; @@ -386,10 +388,7 @@ fn decrypt_peer_bundle( ) -> Result { let plaintext = symmetric.decrypt_and_hash(crypto, bundle.as_bytes())?; let bundle = PeerBundle::decode_bytes(plaintext.as_slice())?; - let peer_qid = derive_qid(crypto, &bundle.mlkem_public_key); - if peer_qid != bundle.qid { - return Err(Error::InvalidRemoteBundle); - } + bundle.validate(crypto)?; Ok(bundle) } diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index 85d96ea4..731a91f8 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -1,7 +1,7 @@ use ql_codec::{ByteSlice, Encode}; use ql_common::QID; -use crate::{derive_qid, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash}; +use crate::{derive_qid, Error, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash}; #[derive(Debug, Clone, PartialEq, Eq)] pub struct PeerBundle { @@ -14,6 +14,13 @@ pub struct PeerBundle { impl PeerBundle { pub const VERSION: u16 = 1; + + pub fn validate(&self, crypto: &impl QlHash) -> Result<(), Error> { + if self.version != Self::VERSION || self.qid != derive_qid(crypto, &self.mlkem_public_key) { + return Err(Error::InvalidRemoteBundle); + } + Ok(()) + } } impl Encode for PeerBundle { From c35c115013b0280139b2838047392e3737db4c79 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 21 Jul 2026 10:54:35 -0400 Subject: [PATCH 48/59] fsm: retransmit time inference --- ql-fsm/src/lib.rs | 2 +- ql-fsm/src/session/mod.rs | 98 ++++++++++++--------- ql-fsm/src/session/state.rs | 8 +- ql-fsm/src/session/tests.rs | 156 ++++++++++++++++++++++++++++++++-- ql-fsm/src/session/tracked.rs | 53 +++++++++++- 5 files changed, 263 insertions(+), 54 deletions(-) diff --git a/ql-fsm/src/lib.rs b/ql-fsm/src/lib.rs index a10b1f11..02964bdb 100644 --- a/ql-fsm/src/lib.rs +++ b/ql-fsm/src/lib.rs @@ -172,7 +172,7 @@ pub struct QlFsmConfig { pub handshake_timeout: Duration, /// delay before sending a pure record ack pub session_record_ack_delay: Duration, - /// how long to wait before resending unacked session records + /// initial wait before resending unacked session records pub session_record_retransmit_timeout: Duration, /// idle delay before sending a keepalive ping pub session_keepalive_interval: Duration, diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index e9015401..aea6f6d7 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -28,7 +28,7 @@ use self::{ remote_stream_history::RemoteStreamHistory, state::{InboundState, OutboundState, SessionPhase, SessionState, StreamRole, StreamState}, stream_tx::StreamTxRange, - tracked::{TrackedFrame, TrackedRecord, TrackedStreamData}, + tracked::{LossRecovery, TrackedFrame, TrackedRecord, TrackedStreamData}, }; use crate::{NoSessionError, StreamError, StreamResetEvent, StreamResetTarget}; @@ -53,7 +53,7 @@ impl Default for SessionConfig { local_parity: StreamParity::Even, record_max_size: 8 * 1024, ack_delay: Duration::from_millis(5), - retransmit_timeout: Duration::from_millis(150), + retransmit_timeout: Duration::from_secs(1), keepalive_interval: Duration::from_secs(10), peer_timeout: Duration::from_secs(30), stream_send_buffer_size: 16 * 1024, @@ -114,6 +114,7 @@ impl SessionFsm { next_record_seq: RecordSeq(0), next_write_id: 0, tracked_records: IndexMap::default(), + loss_recovery: LossRecovery::new(config.retransmit_timeout), ack_tracker: AckTracker::new( config.accepted_record_window, config.pending_ack_range_limit, @@ -235,7 +236,7 @@ impl SessionFsm { self.unpair(sink); return; } - SessionFrame::Ack(ack) => self.process_record_ack(&ack, sink), + SessionFrame::Ack(ack) => self.process_record_ack(now, &ack, sink), SessionFrame::StreamData(frame) => { if self.handle_stream_data(frame, sink).is_err() { self.close(SessionCloseCode::PROTOCOL, sink); @@ -291,7 +292,7 @@ impl SessionFsm { &mut self.state.ack_tracker, &mut self.state.pending_ping, &mut self.state.streams, - record, + &record, ); } } @@ -320,15 +321,12 @@ impl SessionFsm { return None; } let ack_deadline = self.state.ack_tracker.ack_deadline(); + let rto = self.state.loss_recovery.rto(); let retransmit_deadline = self .state .tracked_records .values() - .filter_map(|record| { - record - .sent_at - .map(|sent_at| sent_at + self.config.retransmit_timeout) - }) + .filter_map(|record| record.sent_at.map(|sent_at| sent_at + rto)) .min(); let is_open = self.state.phase.is_open(); let keepalive_deadline = @@ -354,6 +352,8 @@ impl SessionFsm { } pub fn take_next_write(&mut self, now: Instant) -> Option<(Option, SessionRecordBuilder)> { + const TRACKED_RECORD_LIMIT: usize = 64; + match &self.state.phase { SessionPhase::Terminating(frame) => { let seq = self.state.next_record_seq; @@ -377,12 +377,26 @@ impl SessionFsm { } self.collect_timeouts(now); + // ack-only records need no tracking slot and prevent two full peers from stalling + if self.state.tracked_records.len() >= TRACKED_RECORD_LIMIT { + let seq = self.state.next_record_seq; + let mut builder = SessionRecordBuilder::new(seq, self.config.record_max_size); + let pending_ack = self.pending_ack(builder.remaining_capacity())?; + if pending_ack.due_at > now || !builder.push_ack(&pending_ack.ack) { + return None; + } + self.state.ack_tracker.on_ack_emitted(&pending_ack); + next_seq(&mut self.state.next_record_seq); + return Some((None, builder)); + } + let (builder, outbound) = self.build_next_record(now)?; let should_track = outbound.ping_included || !outbound.window_updates.is_empty() || !outbound.frames.is_empty(); let write_id = should_track.then(|| { + debug_assert!(self.state.tracked_records.len() < TRACKED_RECORD_LIMIT); let write_id = self.state.next_write_id; self.state.next_write_id = self.state.next_write_id.wrapping_add(1); self.state.tracked_records.insert(write_id, outbound); @@ -412,7 +426,6 @@ impl SessionFsm { } self.push_next_pending_stream_window(&mut builder, &mut outbound); - self.push_next_stream_data(&mut builder, &mut outbound); if let Some(pending_ack) = self.pending_ack(builder.remaining_capacity()) { @@ -572,27 +585,23 @@ impl SessionFsm { } } - fn process_record_ack(&mut self, ack: &RecordAck, sink: &mut impl EventSink) { + fn process_record_ack(&mut self, now: Instant, ack: &RecordAck, sink: &mut impl EventSink) { let stream_send_buffer_size = self.config.stream_send_buffer_size; - let acked_records = self - .state - .tracked_records - .extract_if(.., |_, record| { - record.sent_at.is_some() && ack.contains(record.seq.0) - }) - .map(|(_, record)| record) - .collect::>(); - - for record in acked_records { + let mut latest_sent_at = None; + let state = &mut self.state; + for (_, record) in state.tracked_records.extract_if(.., |_, record| { + record.sent_at.is_some() && ack.contains(record.seq.0) + }) { + latest_sent_at = latest_sent_at.max(record.sent_at); for frame in &record.frames { - acknowledge_tracked_frame( - &mut self.state.streams, - stream_send_buffer_size, - frame, - sink, - ); + acknowledge_tracked_frame(&mut state.streams, stream_send_buffer_size, frame, sink); } } + if let Some(sent_at) = latest_sent_at { + state + .loss_recovery + .on_ack(now.saturating_duration_since(sent_at)); + } self.reap_reapable_streams(); } @@ -610,20 +619,25 @@ impl SessionFsm { } fn collect_timeouts(&mut self, now: Instant) { - let retransmit_timeout = self.config.retransmit_timeout; - for (_, record) in self.state.tracked_records.extract_if(.., |_, record| { - record - .sent_at - .is_some_and(|sent_at| sent_at + retransmit_timeout <= now) + let rto = self.state.loss_recovery.rto(); + let mut timed_out = false; + let state = &mut self.state; + for (_, record) in state.tracked_records.extract_if(.., |_, record| { + record.sent_at.is_some_and(|sent_at| sent_at + rto <= now) }) { restore_tracked_record( now, - &mut self.state.ack_tracker, - &mut self.state.pending_ping, - &mut self.state.streams, - record, + &mut state.ack_tracker, + &mut state.pending_ping, + &mut state.streams, + &record, ); + timed_out = true; } + if timed_out { + state.loss_recovery.on_timeout(); + } + self.reap_reapable_streams(); } fn handle_stream_data( @@ -954,7 +968,7 @@ fn restore_tracked_record( ack_tracker: &mut AckTracker, pending_ping: &mut bool, streams: &mut IndexMap, - record: TrackedRecord, + record: &TrackedRecord, ) { if let Some(ack) = &record.ack { ack_tracker.restore_acked_ranges(ack, now); @@ -962,22 +976,22 @@ fn restore_tracked_record( if record.ping_included { *pending_ping = true; } - for (stream_id, maximum_offset) in record.window_updates { + for &(stream_id, maximum_offset) in &record.window_updates { if let Some(stream) = streams.get_mut(&stream_id) { if stream.recv_limit() >= maximum_offset { stream.pending_window = true; } } } - for frame in record.frames { + for frame in &record.frames { requeue_tracked_frame(streams, frame); } } -fn requeue_tracked_frame(streams: &mut IndexMap, frame: TrackedFrame) { +fn requeue_tracked_frame(streams: &mut IndexMap, frame: &TrackedFrame) { match frame { - TrackedFrame::StreamReset(reset) => restore_stream_reset(streams, reset), - TrackedFrame::StreamData(frame) => restore_stream_data(streams, frame), + TrackedFrame::StreamReset(reset) => restore_stream_reset(streams, reset.clone()), + TrackedFrame::StreamData(frame) => restore_stream_data(streams, *frame), } } diff --git a/ql-fsm/src/session/state.rs b/ql-fsm/src/session/state.rs index 49e21819..7e2e8a07 100644 --- a/ql-fsm/src/session/state.rs +++ b/ql-fsm/src/session/state.rs @@ -6,8 +6,11 @@ use ql_common::StreamId; use ql_wire::{RecordSeq, ResetTarget, SessionClose, StreamReset}; use super::{ - ack_tracker::AckTracker, remote_stream_history::RemoteStreamHistory, stream_rx::StreamRx, - stream_tx::StreamTx, tracked::TrackedRecord, + ack_tracker::AckTracker, + remote_stream_history::RemoteStreamHistory, + stream_rx::StreamRx, + stream_tx::StreamTx, + tracked::{LossRecovery, TrackedRecord}, }; pub struct SessionState { @@ -18,6 +21,7 @@ pub struct SessionState { pub next_record_seq: RecordSeq, pub next_write_id: u64, pub tracked_records: IndexMap, + pub loss_recovery: LossRecovery, pub ack_tracker: AckTracker, pub pending_ping: bool, pub streams: IndexMap, diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index bb98b82b..6732844b 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -109,20 +109,154 @@ fn outbound_record_seq_increments_monotonically() { #[test] fn retransmit_uses_new_record_seq() { let now = Instant::now(); - let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let mut fsm = SessionFsm::new( + SessionConfig { + retransmit_timeout: Duration::from_millis(100), + ..SessionConfig::default() + }, + now, + ); let stream_id = open_stream_id(&mut fsm); assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"retry"), 5); let (first_seq, first) = next_outbound(&mut fsm, now).unwrap(); let mut emit = |_| {}; - fsm.on_timer(now + Duration::from_millis(200), &mut emit); - let (retried_seq, retried) = next_outbound(&mut fsm, now + Duration::from_millis(200)).unwrap(); + fsm.on_timer(now + Duration::from_millis(101), &mut emit); + let (retried_seq, retried) = next_outbound(&mut fsm, now + Duration::from_millis(101)).unwrap(); assert_ne!(first_seq, retried_seq); assert_eq!(first, retried); } +#[test] +fn retransmitted_record_ack_releases_stream_data() { + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + retransmit_timeout: Duration::from_millis(20), + stream_send_buffer_size: 4, + ..SessionConfig::default() + }, + now, + ); + let stream_id = open_stream_id(&mut fsm); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"data"), 4); + let (first_seq, _) = next_outbound(&mut fsm, now).unwrap(); + + let mut emit = |_| {}; + fsm.on_timer(now + Duration::from_millis(21), &mut emit); + let (retried_seq, _) = next_outbound(&mut fsm, now + Duration::from_millis(21)).unwrap(); + assert_ne!(first_seq, retried_seq); + + let mut events = Vec::new(); + fsm.receive( + now + Duration::from_millis(22), + RecordSeq(9), + std::iter::once(Ok(SessionFrame::Ack( + RecordAck::from_ranges([retried_seq..=retried_seq]).unwrap(), + ))), + &mut |event| events.push(event), + ); + + assert!(events.contains(&SessionEvent::Writable(stream_id))); + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"z"), 1); +} + +#[test] +fn acknowledged_rtt_updates_retransmit_timeout() { + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + retransmit_timeout: Duration::from_millis(100), + ..SessionConfig::default() + }, + now, + ); + let stream_id = open_stream_id(&mut fsm); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"first"), 5); + let (first_seq, _) = next_outbound(&mut fsm, now).unwrap(); + fsm.receive( + now + Duration::from_millis(80), + RecordSeq(9), + std::iter::once(Ok(SessionFrame::Ack( + RecordAck::from_ranges([first_seq..=first_seq]).unwrap(), + ))), + &mut |_| {}, + ); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"second"), 6); + next_outbound(&mut fsm, now + Duration::from_millis(80)).unwrap(); + + let mut emit = |_| {}; + fsm.on_timer(now + Duration::from_millis(181), &mut emit); + assert!(next_outbound(&mut fsm, now + Duration::from_millis(181)).is_none()); + + fsm.on_timer(now + Duration::from_millis(321), &mut emit); + assert!(next_outbound(&mut fsm, now + Duration::from_millis(321)).is_some()); +} + +#[test] +fn retransmit_timeout_backs_off() { + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + retransmit_timeout: Duration::from_millis(20), + ..SessionConfig::default() + }, + now, + ); + let stream_id = open_stream_id(&mut fsm); + + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"retry"), 5); + next_outbound(&mut fsm, now).unwrap(); + + let mut emit = |_| {}; + fsm.on_timer(now + Duration::from_millis(21), &mut emit); + next_outbound(&mut fsm, now + Duration::from_millis(21)).unwrap(); + + fsm.on_timer(now + Duration::from_millis(42), &mut emit); + assert!(next_outbound(&mut fsm, now + Duration::from_millis(42)).is_none()); + + fsm.on_timer(now + Duration::from_millis(62), &mut emit); + assert!(next_outbound(&mut fsm, now + Duration::from_millis(62)).is_some()); +} + +#[test] +fn tracked_record_count_is_bounded() { + const PAYLOAD_LEN: usize = 1024; + + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + record_max_size: SessionRecordBuilder::MIN_CAPACITY + + 1 + + StreamData::>::MIN_WIRE_SIZE + + 1, + stream_send_buffer_size: PAYLOAD_LEN, + initial_peer_stream_receive_window: PAYLOAD_LEN as u32, + ..SessionConfig::default() + }, + now, + ); + let stream_id = open_stream_id(&mut fsm); + assert_eq!( + write_stream_bytes(&mut fsm, stream_id, &[b'x'; PAYLOAD_LEN]), + PAYLOAD_LEN + ); + + let mut count = 0; + while next_outbound(&mut fsm, now).is_some() { + count += 1; + assert!(count <= 64); + } + + assert_eq!(count, 64); + assert_eq!(fsm.state.tracked_records.len(), 64); +} + #[test] fn lost_record_on_one_stream_does_not_block_another_stream() { const PAYLOAD_LEN: usize = 40; @@ -397,7 +531,13 @@ fn inbound_empty_fin_emits_finished_immediately() { #[test] fn remote_stream_reset_is_reliable_and_retried() { let now = Instant::now(); - let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let mut fsm = SessionFsm::new( + SessionConfig { + retransmit_timeout: Duration::from_millis(100), + ..SessionConfig::default() + }, + now, + ); let stream_id = open_stream_id(&mut fsm); fsm.stream(stream_id, |_| {}) @@ -413,9 +553,9 @@ fn remote_stream_reset_is_reliable_and_retried() { )); let mut emit = |_| {}; - fsm.on_timer(now + Duration::from_millis(200), &mut emit); + fsm.on_timer(now + Duration::from_millis(101), &mut emit); let (_retried_seq, retried) = - next_outbound(&mut fsm, now + Duration::from_millis(200)).unwrap(); + next_outbound(&mut fsm, now + Duration::from_millis(101)).unwrap(); assert_eq!(first, retried); } @@ -828,14 +968,14 @@ fn sparse_out_of_order_ack_ranges_page_and_quiesce() { let mut receiver = SessionFsm::new(receiver_config, now); let stream_id = open_stream_id(&mut sender); - let payload = vec![b'x'; 2048]; + let payload = vec![b'x'; 1200]; assert_eq!( write_stream_bytes(&mut sender, stream_id, &payload), payload.len() ); let originals = drain_outbound(&mut sender, now, 4096); - assert!(originals.len() >= 64); + assert_eq!(originals.len(), 64); for (seq, record) in originals.iter().filter(|(seq, _)| seq.0 % 2 == 1) { let _ = receive_events(&mut receiver, now, *seq, record); diff --git a/ql-fsm/src/session/tracked.rs b/ql-fsm/src/session/tracked.rs index e74b2928..f8a234ca 100644 --- a/ql-fsm/src/session/tracked.rs +++ b/ql-fsm/src/session/tracked.rs @@ -1,6 +1,6 @@ //! outbound record tracking state for ack and retransmit handling -use std::time::Instant; +use std::time::{Duration, Instant}; use ql_common::StreamId; use ql_wire::{RecordAck, RecordSeq, StreamReset}; @@ -28,3 +28,54 @@ pub struct TrackedStreamData { pub len: usize, pub fin: bool, } + +/// estimates RTO from smoothed RTT and variance +/// backs off on timeout, and resets on ACK +pub struct LossRecovery { + smoothed_rtt: Option, + rtt_variance: Duration, + base_rto: Duration, + rto: Duration, +} + +impl LossRecovery { + const MIN_RTO: Duration = Duration::from_millis(10); + const MAX_RTO: Duration = Duration::from_secs(30); + + pub fn new(initial_rto: Duration) -> Self { + let rto = initial_rto.clamp(Self::MIN_RTO, Self::MAX_RTO); + Self { + smoothed_rtt: None, + rtt_variance: Duration::ZERO, + base_rto: rto, + rto, + } + } + + pub fn rto(&self) -> Duration { + self.rto + } + + pub fn on_ack(&mut self, sample: Duration) { + let (smoothed_rtt, rtt_variance) = match self.smoothed_rtt { + Some(smoothed_rtt) => ( + smoothed_rtt.saturating_mul(7).saturating_add(sample) / 8, + self.rtt_variance + .saturating_mul(3) + .saturating_add(smoothed_rtt.abs_diff(sample)) + / 4, + ), + None => (sample, sample / 2), + }; + self.smoothed_rtt = Some(smoothed_rtt); + self.rtt_variance = rtt_variance; + self.base_rto = smoothed_rtt + .saturating_add(rtt_variance.saturating_mul(4)) + .clamp(Self::MIN_RTO, Self::MAX_RTO); + self.rto = self.base_rto; + } + + pub fn on_timeout(&mut self) { + self.rto = self.rto.saturating_mul(2).min(Self::MAX_RTO); + } +} From 14123dc0fc38167b6331272d6b6b7cfe08466554 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Tue, 21 Jul 2026 11:27:34 -0400 Subject: [PATCH 49/59] ql-runtime: separate crypto getter --- ql-runtime/src/driver/mod.rs | 8 ++--- ql-runtime/src/driver/test.rs | 5 +++ ql-runtime/src/platform.rs | 4 ++- ql-runtime/src/tests/mod.rs | 60 ++++------------------------------- 4 files changed, 19 insertions(+), 58 deletions(-) diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs index 42db37ae..a00289e1 100644 --- a/ql-runtime/src/driver/mod.rs +++ b/ql-runtime/src/driver/mod.rs @@ -81,7 +81,7 @@ impl Runtime

{ } DriverStep::Inbound(bytes) => { log::trace!("received transport frame: len={}", bytes.len()); - if let Err(e) = fsm.receive(Instant::now(), bytes, &platform) { + if let Err(e) = fsm.receive(Instant::now(), bytes, platform.crypto()) { log::info!("receive rejected frame: error={e:?}"); platform.handle_recv_error(e); } @@ -185,7 +185,7 @@ impl DriverState { } Command::Connect => { log::info!("starting IK connect"); - if fsm.connect_ik(Instant::now(), platform).is_err() { + if fsm.connect_ik(Instant::now(), platform.crypto()).is_err() { log::warn!("IK connect ignored: no bound peer"); } } @@ -199,7 +199,7 @@ impl DriverState { } Command::StartPairing { invite } => { log::info!(" starting XX pairing"); - fsm.connect_xx(Instant::now(), invite, platform); + fsm.connect_xx(Instant::now(), invite, platform.crypto()); } Command::CloseSession { code } => { log::info!("closing session: code={code:?}"); @@ -476,7 +476,7 @@ impl DriverState { ) -> bool { let mut filled = false; while in_flight.len() < self.max_concurrent_message_writes { - let Some(write) = fsm.take_next_write(Instant::now(), platform) else { + let Some(write) = fsm.take_next_write(Instant::now(), platform.crypto()) else { break; }; filled = true; diff --git a/ql-runtime/src/driver/test.rs b/ql-runtime/src/driver/test.rs index dab61c0c..c30c0813 100644 --- a/ql-runtime/src/driver/test.rs +++ b/ql-runtime/src/driver/test.rs @@ -20,10 +20,15 @@ impl crate::platform::QlTimer for NoopTimer { } impl QlPlatform for NoopCrypto { + type Crypto = Self; type Timer = NoopTimer; type WriteMessageFut<'a> = std::future::Ready; type Inbound = NoopInbound; + fn crypto(&self) -> &Self::Crypto { + self + } + fn write_message(&self, _message: Vec) -> Self::WriteMessageFut<'_> { std::future::ready(true) } diff --git a/ql-runtime/src/platform.rs b/ql-runtime/src/platform.rs index 18ad6672..2e50bc65 100644 --- a/ql-runtime/src/platform.rs +++ b/ql-runtime/src/platform.rs @@ -18,13 +18,15 @@ pub trait QlInbound { fn poll_recv(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll>; } -pub trait QlPlatform: QlCrypto { +pub trait QlPlatform { + type Crypto: QlCrypto; type Timer: QlTimer; type WriteMessageFut<'a>: Future + Unpin + 'a where Self: 'a; type Inbound: QlInbound; + fn crypto(&self) -> &Self::Crypto; fn write_message(&self, message: Vec) -> Self::WriteMessageFut<'_>; /// Returns the platform's inbound transport poller. /// diff --git a/ql-runtime/src/tests/mod.rs b/ql-runtime/src/tests/mod.rs index d64cecce..e356d7f9 100644 --- a/ql-runtime/src/tests/mod.rs +++ b/ql-runtime/src/tests/mod.rs @@ -15,9 +15,8 @@ use ql_codec::Decode; use ql_common::{StreamInfo, QID}; use ql_fsm::PeerStatus; use ql_wire::{ - generate_identity, test_identities, MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, - MlKemPublicKey, Nonce, PairingToken, PeerBundle, QlAead, QlHash, QlIdentity, QlKem, QlRandom, - RecordHeader, RecordType, SessionKey, SoftwareCrypto, + generate_identity, test_identities, PairingToken, PeerBundle, QlIdentity, RecordHeader, + RecordType, SoftwareCrypto, }; use tokio::{task::LocalSet, time::Sleep}; @@ -336,61 +335,16 @@ impl QlTimer for TokioTimer { } } -impl QlRandom for TestPlatform { - fn fill_random_bytes(&self, data: &mut [u8]) { - self.crypto.fill_random_bytes(data); - } -} - -impl QlHash for TestPlatform { - fn sha256(&self, parts: &[&[u8]]) -> [u8; 32] { - self.crypto.sha256(parts) - } -} - -impl QlAead for TestPlatform { - fn aes256_gcm_encrypt( - &self, - key: &SessionKey, - nonce: &Nonce, - aad: &[u8], - buffer: &mut [u8], - ) -> [u8; ql_wire::ENCRYPTED_MESSAGE_AUTH_SIZE] { - self.crypto.aes256_gcm_encrypt(key, nonce, aad, buffer) - } - - fn aes256_gcm_decrypt( - &self, - key: &SessionKey, - nonce: &Nonce, - aad: &[u8], - buffer: &mut [u8], - auth_tag: &[u8; ql_wire::ENCRYPTED_MESSAGE_AUTH_SIZE], - ) -> bool { - self.crypto - .aes256_gcm_decrypt(key, nonce, aad, buffer, auth_tag) - } -} - -impl QlKem for TestPlatform { - fn mlkem_generate_keypair(&self) -> MlKemKeyPair { - self.crypto.mlkem_generate_keypair() - } - - fn mlkem_encapsulate(&self, public_key: &MlKemPublicKey) -> (MlKemCiphertext, SessionKey) { - self.crypto.mlkem_encapsulate(public_key) - } - - fn mlkem_decapsulate(&self, pk: &MlKemPrivateKey, cipher: &MlKemCiphertext) -> SessionKey { - self.crypto.mlkem_decapsulate(pk, cipher) - } -} - impl crate::platform::QlPlatform for TestPlatform { + type Crypto = SoftwareCrypto; type Timer = TokioTimer; type WriteMessageFut<'a> = Pin + Send + 'a>>; type Inbound = TestInbound; + fn crypto(&self) -> &Self::Crypto { + &self.crypto + } + fn write_message(&self, message: Vec) -> Self::WriteMessageFut<'_> { let outbound = self.outbound.clone(); let write_delay = self.write_delay; From a0a0408110619222299201c7471255021890506a Mon Sep 17 00:00:00 2001 From: Alex's Bot Date: Wed, 22 Jul 2026 14:48:18 +0000 Subject: [PATCH 50/59] ql-codec: make ByteSlice safe Splitting through `&mut self` means the reader never has to move `B` out of itself, so `ManuallyDrop` and the `Drop` impl go with it. Losing the `Drop` impl ends the reader's borrow of `bytes` in `fsm::receive` at its last use, making the `drop(reader)` there dead code. Co-Authored-By: Claude Opus 4.8 (1M context) --- ql-codec/src/reader.rs | 33 ++++----------------------------- ql-codec/src/slice.rs | 39 ++++++++++++++++++--------------------- ql-fsm/src/fsm.rs | 1 - 3 files changed, 22 insertions(+), 51 deletions(-) diff --git a/ql-codec/src/reader.rs b/ql-codec/src/reader.rs index 2c15b6ce..68beb6a1 100644 --- a/ql-codec/src/reader.rs +++ b/ql-codec/src/reader.rs @@ -1,18 +1,14 @@ -use core::mem::ManuallyDrop; - use crate::{varint, ByteSlice, Decode, Error}; #[derive(Clone)] pub struct Reader { - remaining: ManuallyDrop, + remaining: B, } impl Reader { #[inline] pub fn new(bytes: B) -> Self { - Self { - remaining: ManuallyDrop::new(bytes), - } + Self { remaining: bytes } } #[inline] @@ -29,13 +25,7 @@ impl Reader { if len > self.remaining.len() { return Err(Error::UnexpectedEof); } - // SAFETY: checked above - let (head, tail) = unsafe { - let remaining = ManuallyDrop::take(&mut self.remaining); - remaining.split_at_unchecked(len) - }; - self.remaining = ManuallyDrop::new(tail); - Ok(head) + Ok(self.remaining.split_off_front(len)) } #[inline] @@ -44,13 +34,7 @@ impl Reader { } pub fn take_all(&mut self) -> B { - // SAFETY: 0 is always a valid split point - let (empty, rest) = unsafe { - let remaining = ManuallyDrop::take(&mut self.remaining); - remaining.split_at_unchecked(0) - }; - self.remaining = ManuallyDrop::new(empty); - rest + self.remaining.split_off_front(self.remaining.len()) } pub fn take_len_prefixed(&mut self) -> Result { @@ -74,12 +58,3 @@ impl Reader { varint::decode(self) } } - -impl Drop for Reader { - fn drop(&mut self) { - // SAFETY: `remaining` is initialized except during `take_bytes` - unsafe { - ManuallyDrop::drop(&mut self.remaining); - } - } -} diff --git a/ql-codec/src/slice.rs b/ql-codec/src/slice.rs index afe8a495..9fafbdbf 100644 --- a/ql-codec/src/slice.rs +++ b/ql-codec/src/slice.rs @@ -3,27 +3,23 @@ use core::{mem, ops::Deref}; use bytes::{Buf, Bytes}; /// A byte slice owner used by the codec reader -/// -/// # Safety -/// -/// `split_at_unchecked` must return byte slices matching `self[..mid]` and -/// `self[mid..]` when `mid <= self.len()` -pub unsafe trait ByteSlice: Deref + Sized { - /// splits the current byte view at `mid` without checking bounds - /// mid can be 0 +pub trait ByteSlice: Deref + Sized { + /// splits `self[..mid]` off the front, leaving `self[mid..]` behind /// - /// # Safety + /// # Panics /// - /// `mid` must not exceed the slice length. - unsafe fn split_at_unchecked(self, mid: usize) -> (Self, Self); + /// Panics if `mid` exceeds the slice length. + fn split_off_front(&mut self, mid: usize) -> Self; fn take_u8(&mut self) -> Option; } -unsafe impl ByteSlice for &[u8] { +impl ByteSlice for &[u8] { #[inline] - unsafe fn split_at_unchecked(self, mid: usize) -> (Self, Self) { - <[u8]>::split_at_unchecked(self, mid) + fn split_off_front(&mut self, mid: usize) -> Self { + let (head, tail) = self.split_at(mid); + *self = tail; + head } #[inline] @@ -34,10 +30,12 @@ unsafe impl ByteSlice for &[u8] { } } -unsafe impl ByteSlice for &mut [u8] { +impl ByteSlice for &mut [u8] { #[inline] - unsafe fn split_at_unchecked(self, mid: usize) -> (Self, Self) { - <[u8]>::split_at_mut_unchecked(self, mid) + fn split_off_front(&mut self, mid: usize) -> Self { + let (head, tail) = mem::take(self).split_at_mut(mid); + *self = tail; + head } #[inline] @@ -49,11 +47,10 @@ unsafe impl ByteSlice for &mut [u8] { } } -unsafe impl ByteSlice for Bytes { +impl ByteSlice for Bytes { #[inline] - unsafe fn split_at_unchecked(mut self, mid: usize) -> (Self, Self) { - let head = self.split_to(mid); - (head, self) + fn split_off_front(&mut self, mid: usize) -> Self { + self.split_to(mid) } #[inline] diff --git a/ql-fsm/src/fsm.rs b/ql-fsm/src/fsm.rs index 0ff3ca9a..a532c9bd 100644 --- a/ql-fsm/src/fsm.rs +++ b/ql-fsm/src/fsm.rs @@ -143,7 +143,6 @@ pub fn receive( (payload.len(), record.header.seq) }; - drop(reader); let len = bytes.len(); let plaintext = Bytes::from(bytes).slice(len - decrypt_len..); let frames = wire::parse_session_frames(plaintext); From c4c01fad961cb65ed35b9b636930a39752f03a5c Mon Sep 17 00:00:00 2001 From: Alex's Bot Date: Wed, 22 Jul 2026 16:06:17 +0000 Subject: [PATCH 51/59] ql: wrap varint fields in a Varint type The plain integer codecs are fixed width, so `decode::()` and `decode_varint::()` are one word apart and produce different wire formats, with nothing in a bare `u64` field to say which one applies. The `VarInt` trait is renamed `Primitive` so it doesn't read as a typo for the `Varint` struct. Co-Authored-By: Claude Opus 4.8 (1M context) --- ql-codec/src/lib.rs | 3 +- ql-codec/src/reader.rs | 2 +- ql-codec/src/varint.rs | 63 ++++++++++++++++++++++---- ql-fsm/src/session/mod.rs | 13 +++--- ql-fsm/src/session/tests.rs | 23 +++++----- ql-runtime/src/tests/rpc.rs | 9 ++-- ql-wire/src/encrypted/ack.rs | 63 ++++++++++++++------------ ql-wire/src/encrypted/stream_data.rs | 18 ++++---- ql-wire/src/encrypted/stream_window.rs | 10 ++-- ql-wire/src/tests.rs | 8 ++-- 10 files changed, 132 insertions(+), 80 deletions(-) diff --git a/ql-codec/src/lib.rs b/ql-codec/src/lib.rs index d51ca811..cfb3c649 100644 --- a/ql-codec/src/lib.rs +++ b/ql-codec/src/lib.rs @@ -12,6 +12,7 @@ pub use codec::{encode_bytes, encoded_len_bytes}; pub use error::Error; pub use reader::Reader; pub use slice::ByteSlice; +pub use varint::Varint; pub trait Encode { fn encoded_len(&self) -> usize; @@ -41,7 +42,7 @@ macro_rules! varint_wrapper { ($name:ty, $inner:ty) => { impl $name { pub const MAX_ENCODED_LEN: usize = - <$inner as ql_codec::varint::VarInt>::MAX_ENCODED_LEN; + <$inner as ql_codec::varint::Primitive>::MAX_ENCODED_LEN; } impl ql_codec::Encode for $name { diff --git a/ql-codec/src/reader.rs b/ql-codec/src/reader.rs index 68beb6a1..d534cc3a 100644 --- a/ql-codec/src/reader.rs +++ b/ql-codec/src/reader.rs @@ -53,7 +53,7 @@ impl Reader { #[inline] pub fn decode_varint(&mut self) -> Result where - T: varint::VarInt, + T: varint::Primitive, { varint::decode(self) } diff --git a/ql-codec/src/varint.rs b/ql-codec/src/varint.rs index 1f241506..0e5e1cb9 100644 --- a/ql-codec/src/varint.rs +++ b/ql-codec/src/varint.rs @@ -1,8 +1,18 @@ +use core::{fmt, ops::Deref}; + use bytes::BufMut; -use crate::{ByteSlice, Error, Reader}; +use crate::{ByteSlice, Decode, Encode, Error, Reader}; + +/// An integer field carried as a varint +/// +/// The plain integer codecs are fixed width, so a field only encodes as a +/// varint when it is wrapped here. +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] +#[repr(transparent)] +pub struct Varint(pub T); -pub trait VarInt: Copy { +pub trait Primitive: Copy { const MAX_ENCODED_LEN: usize; fn from_u8(value: u8) -> Self; @@ -12,7 +22,42 @@ pub trait VarInt: Copy { fn checked_add_payload(self, payload: u8, shift: usize) -> Option; } -pub fn encoded_len(mut value: T) -> usize { +impl Varint { + pub const MAX_ENCODED_LEN: usize = T::MAX_ENCODED_LEN; +} + +impl Deref for Varint { + type Target = T; + + #[inline] + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl fmt::Display for Varint { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(f) + } +} + +impl Encode for Varint { + fn encoded_len(&self) -> usize { + encoded_len(self.0) + } + + fn encode(&self, out: &mut W) { + encode(self.0, out); + } +} + +impl Decode for Varint { + fn decode(reader: &mut Reader) -> Result { + self::decode(reader).map(Self) + } +} + +pub fn encoded_len(mut value: T) -> usize { let mut len = 1; while value.needs_more() { value = value.shr_7(); @@ -23,7 +68,7 @@ pub fn encoded_len(mut value: T) -> usize { pub fn encode(mut value: T, out: &mut W) where - T: VarInt, + T: Primitive, W: BufMut + ?Sized, { while value.needs_more() { @@ -35,7 +80,7 @@ where pub fn decode(reader: &mut Reader) -> Result where - T: VarInt, + T: Primitive, B: ByteSlice, { let mut value = T::from_u8(0); @@ -62,7 +107,7 @@ where macro_rules! impl_varint { ($($ty:ty),* $(,)?) => { $( - impl VarInt for $ty { + impl Primitive for $ty { const MAX_ENCODED_LEN: usize = (size_of::() * 8).div_ceil(7); #[inline] @@ -110,7 +155,7 @@ mod tests { fn assert_decodes(value: T) where - T: VarInt + Debug + PartialEq, + T: Primitive + Debug + PartialEq, { let mut out = Vec::new(); encode(value, &mut out); @@ -122,7 +167,7 @@ mod tests { fn assert_error(bytes: &[u8], error: Error) where - T: VarInt + Debug + PartialEq, + T: Primitive + Debug + PartialEq, { let mut reader = Reader::new(bytes); assert_eq!(decode::(&mut reader), Err(error)); @@ -130,7 +175,7 @@ mod tests { fn test_type(max: T) where - T: VarInt + Debug + PartialEq + TryFrom, + T: Primitive + Debug + PartialEq + TryFrom, { for value in [0, 1, 127].map(T::from_u8) { assert_decodes(value); diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index aea6f6d7..637edecd 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -17,6 +17,7 @@ use std::time::{Duration, Instant}; use bytes::Bytes; use indexmap::IndexMap; +use ql_codec::Varint; use ql_common::StreamId; use ql_wire::{ RecordAck, RecordSeq, ResetTarget, SessionClose, SessionCloseCode, SessionFrame, @@ -503,17 +504,17 @@ impl SessionFsm { } let frame = StreamWindow { stream_id, - maximum_offset: stream.recv_limit(), + maximum_offset: Varint(stream.recv_limit()), }; if !builder.push_stream_window(&frame) { break; } stream.pending_window = false; - stream.advertised_max_offset = frame.maximum_offset; + stream.advertised_max_offset = *frame.maximum_offset; outbound .window_updates - .push((stream_id, frame.maximum_offset)); + .push((stream_id, *frame.maximum_offset)); } } @@ -548,7 +549,7 @@ impl SessionFsm { }; let frame = StreamData { stream_id, - offset: candidate.offset, + offset: Varint(candidate.offset), header: if matches!(stream.role, StreamRole::Initiator) && candidate.offset == 0 { stream.header.as_deref() } else { @@ -660,7 +661,7 @@ impl SessionFsm { }, }; - let frame_offset = offset; + let frame_offset = *offset; let Some(frame_end) = frame_offset.checked_add(bytes.len() as u64) else { return Err(()); }; @@ -738,7 +739,7 @@ impl SessionFsm { }; let was_full = stream.send_capacity(self.config.stream_send_buffer_size) == 0; - let maximum_offset = frame.maximum_offset; + let maximum_offset = *frame.maximum_offset; if maximum_offset > stream.peer_max_offset { stream.peer_max_offset = maximum_offset; } diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index 6732844b..d1649bf4 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -1,6 +1,7 @@ use std::time::{Duration, Instant}; use bytes::Bytes; +use ql_codec::Varint; use ql_common::{ResetCode, StreamId, QID}; use ql_wire::{ decode_session_frames, parse_session_frames, RecordAck, RecordSeq, ResetTarget, SessionFrame, @@ -401,7 +402,7 @@ fn commit_stream_read_is_what_advances_stream_window() { let stream_id = StreamId(1); let data = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: false, bytes: b"hi".to_vec(), @@ -453,7 +454,7 @@ fn pure_ack_only_records_are_fire_and_forget() { let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: false, bytes: b"hi".to_vec(), @@ -483,7 +484,7 @@ fn inbound_stream_data_emits_opened_and_readable() { let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(ql_wire::StreamData { stream_id, - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: true, bytes: b"hello".to_vec(), @@ -512,7 +513,7 @@ fn inbound_empty_fin_emits_finished_immediately() { let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: true, bytes: Vec::new(), @@ -597,7 +598,7 @@ fn duplicate_stream_data_is_not_redelivered() { let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: false, bytes: b"hi".to_vec(), @@ -655,7 +656,7 @@ fn late_remote_stream_data_after_reset_is_ignored() { })]; let data = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: false, bytes: b"hello".to_vec(), @@ -687,7 +688,7 @@ fn duplicate_finished_remote_data_after_reap_is_ignored() { let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: true, bytes: b"hello".to_vec(), @@ -724,7 +725,7 @@ fn duplicate_finished_remote_data_before_read_is_ignored() { let stream_id = StreamId(1); let record = vec![SessionFrame::StreamData(StreamData { stream_id, - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: true, bytes: b"hello".to_vec(), @@ -836,7 +837,7 @@ fn close_does_not_ack_rejected_record_seq() { let invalid = vec![SessionFrame::StreamData(StreamData { stream_id: StreamId(0), - offset: 0, + offset: Varint(0), header: Some(vec![1_u8]), fin: false, bytes: b"bad".to_vec(), @@ -923,7 +924,7 @@ fn initial_peer_stream_receive_window_limits_first_send() { RecordSeq(9), &[SessionFrame::StreamWindow(ql_wire::StreamWindow { stream_id, - maximum_offset: 5, + maximum_offset: Varint(5), })], ); assert!(events.is_empty()); @@ -934,7 +935,7 @@ fn initial_peer_stream_receive_window_limits_first_send() { frame, SessionFrame::StreamData(frame) if frame.stream_id == stream_id - && frame.offset == 3 + && *frame.offset == 3 && frame.bytes.as_slice() == b"lo" ) })); diff --git a/ql-runtime/src/tests/rpc.rs b/ql-runtime/src/tests/rpc.rs index 4cef5136..32599389 100644 --- a/ql-runtime/src/tests/rpc.rs +++ b/ql-runtime/src/tests/rpc.rs @@ -8,6 +8,7 @@ use std::{ }; use bytes::{BufMut, Bytes}; +use ql_codec::Encode; use ql_common::ResetCode; use ql_rpc::{ download::{DownloadHandlerLocal, DownloadStart}, @@ -28,17 +29,17 @@ struct TestRouteKey(u64); impl ql_rpc::RpcRouteKey for TestRouteKey { fn encoded_len(&self) -> usize { - ql_codec::varint::encoded_len(self.0) + ql_codec::Varint(self.0).encoded_len() } fn encode(&self, out: &mut W) { - ql_codec::varint::encode(self.0, out); + ql_codec::Varint(self.0).encode(out); } fn decode(bytes: &[u8]) -> Option { let mut reader = ql_codec::Reader::new(bytes); - let route_id = reader.decode_varint().ok()?; - Some(Self(route_id)) + let route_id = reader.decode::>().ok()?; + Some(Self(*route_id)) } } diff --git a/ql-wire/src/encrypted/ack.rs b/ql-wire/src/encrypted/ack.rs index 14855804..02eda89e 100644 --- a/ql-wire/src/encrypted/ack.rs +++ b/ql-wire/src/encrypted/ack.rs @@ -1,20 +1,20 @@ use std::{fmt, ops::RangeInclusive}; -use ql_codec::{ByteSlice, Encode, Error}; +use ql_codec::{ByteSlice, Encode, Error, Varint}; use crate::RecordSeq; #[derive(Debug, Clone, PartialEq, Eq)] pub struct RecordAck { largest_acked: RecordSeq, - first_range_len: u64, + first_range_len: Varint, blocks: Box<[RecordAckBlock]>, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RecordAckBlock { - pub gap: u64, - pub range_len: u64, + pub gap: Varint, + pub range_len: Varint, } #[derive(Debug, Clone, Copy, PartialEq, Eq)] @@ -49,7 +49,7 @@ impl RecordAck { pub fn ranges(&self) -> RecordAckRangeIter<'_> { RecordAckRangeIter { largest_acked: self.largest_acked.0, - first_range_len: Some(self.first_range_len), + first_range_len: Some(*self.first_range_len), previous_start: None, blocks: self.blocks.iter(), } @@ -61,13 +61,13 @@ impl RecordAck { } fn block_count_len(block_count: usize) -> usize { - ql_codec::varint::encoded_len(block_count) + Varint(block_count).encoded_len() } } impl RecordAckBlock { fn encoded_len(&self) -> usize { - ql_codec::varint::encoded_len(self.gap) + ql_codec::varint::encoded_len(self.range_len) + self.gap.encoded_len() + self.range_len.encoded_len() } } @@ -110,8 +110,8 @@ impl Iterator for RecordAckRangeIter<'_> { .previous_start .expect("first ack range is always yielded"); // gap is encoded as missing_count - 1, so decoding steps back by gap + 2. - let end = previous_start - block.gap - 2; - let start = end - block.range_len; + let end = previous_start - *block.gap - 2; + let start = end - *block.range_len; self.previous_start = Some(start); Some(RecordSeq(start)..=RecordSeq(end)) } @@ -121,7 +121,7 @@ impl Encode for RecordAck { fn encoded_len(&self) -> usize { self.largest_acked.encoded_len() + Self::block_count_len(self.blocks.len()) - + ql_codec::varint::encoded_len(self.first_range_len) + + self.first_range_len.encoded_len() + self .blocks .iter() @@ -131,11 +131,11 @@ impl Encode for RecordAck { fn encode(&self, out: &mut W) { self.largest_acked.encode(out); - ql_codec::varint::encode(self.blocks.len(), out); - ql_codec::varint::encode(self.first_range_len, out); + Varint(self.blocks.len()).encode(out); + self.first_range_len.encode(out); for block in &self.blocks { - ql_codec::varint::encode(block.gap, out); - ql_codec::varint::encode(block.range_len, out); + block.gap.encode(out); + block.range_len.encode(out); } } } @@ -143,13 +143,13 @@ impl Encode for RecordAck { impl ql_codec::Decode for RecordAck { fn decode(reader: &mut ql_codec::Reader) -> Result { let largest_acked = reader.decode()?; - let block_count = reader.decode_varint::()?; - let first_range_len = reader.decode_varint()?; + let block_count = *reader.decode::>()?; + let first_range_len = reader.decode()?; let mut blocks = Vec::with_capacity(block_count); for _ in 0..block_count { blocks.push(RecordAckBlock { - gap: reader.decode_varint()?, - range_len: reader.decode_varint()?, + gap: reader.decode()?, + range_len: reader.decode()?, }); } @@ -164,7 +164,7 @@ impl ql_codec::Decode for RecordAck { let mut previous_start = ack .largest_acked .0 - .checked_sub(ack.first_range_len) + .checked_sub(*ack.first_range_len) .ok_or(Error::InvalidRange)?; for block in &ack.blocks { @@ -172,7 +172,7 @@ impl ql_codec::Decode for RecordAck { .checked_sub(block.gap.checked_add(2).ok_or(Error::InvalidRange)?) .ok_or(Error::InvalidRange)?; previous_start = end - .checked_sub(block.range_len) + .checked_sub(*block.range_len) .ok_or(Error::InvalidRange)?; } } @@ -183,7 +183,7 @@ impl ql_codec::Decode for RecordAck { #[derive(Debug, Clone, Default, PartialEq, Eq)] pub struct RecordAckBuilder { largest_acked: Option, - first_range_len: Option, + first_range_len: Option>, blocks: Vec, previous_start: Option, wire_len: usize, @@ -215,7 +215,10 @@ impl RecordAckBuilder { .checked_sub(end) .and_then(|delta| delta.checked_sub(2)) .expect("canonical ack ranges stay separated by at least one sequence"); - let block = RecordAckBlock { gap, range_len }; + let block = RecordAckBlock { + gap: Varint(gap), + range_len: Varint(range_len), + }; let current_block_count_len = RecordAck::block_count_len(self.blocks.len()); let next_block_count_len = RecordAck::block_count_len(self.blocks.len() + 1); let next_wire_len = self.wire_len @@ -234,13 +237,13 @@ impl RecordAckBuilder { let largest_acked = RecordSeq(end); let wire_len = largest_acked.encoded_len() + RecordAck::block_count_len(0) - + ql_codec::varint::encoded_len(range_len); + + Varint(range_len).encoded_len(); if wire_len > max_wire_size { return Ok(false); } self.largest_acked = Some(largest_acked); - self.first_range_len = Some(range_len); + self.first_range_len = Some(Varint(range_len)); self.previous_start = Some(start); self.wire_len = wire_len; Ok(true) @@ -260,7 +263,7 @@ impl RecordAckBuilder { } #[cfg(test)] mod tests { - use ql_codec::{Decode, Encode, Error}; + use ql_codec::{Decode, Encode, Error, Varint}; use super::{RecordAck, RecordAckBlock, RecordAckBuilder, RecordAckRangeError}; use crate::RecordSeq; @@ -288,17 +291,17 @@ mod tests { .unwrap(); assert_eq!(ack.largest_acked, RecordSeq(100)); - assert_eq!(ack.first_range_len, 5); + assert_eq!(ack.first_range_len, Varint(5)); assert_eq!( ack.blocks.as_ref(), &[ RecordAckBlock { - gap: 1, - range_len: 2, + gap: Varint(1), + range_len: Varint(2), }, RecordAckBlock { - gap: 8, - range_len: 0, + gap: Varint(8), + range_len: Varint(0), } ] ); diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index d6786157..9bb90f7c 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -1,5 +1,5 @@ use ql_codec::{ - encode_bytes, encoded_len_bytes, varint, BufView, ByteSlice, Decode, Encode, Error, + encode_bytes, encoded_len_bytes, BufView, ByteSlice, Decode, Encode, Error, Varint, }; use ql_common::StreamId; @@ -7,7 +7,7 @@ use ql_common::StreamId; #[derive(Debug, Clone, PartialEq, Eq)] pub struct StreamData { pub stream_id: StreamId, - pub offset: u64, + pub offset: Varint, pub header: Option, pub fin: bool, pub bytes: B, @@ -15,16 +15,16 @@ pub struct StreamData { impl StreamData { pub const MIN_WIRE_SIZE: usize = StreamId::MAX_ENCODED_LEN - + ::MAX_ENCODED_LEN + + Varint::::MAX_ENCODED_LEN + size_of::() - + ::MAX_ENCODED_LEN - + ::MAX_ENCODED_LEN; + + Varint::::MAX_ENCODED_LEN + + Varint::::MAX_ENCODED_LEN; } impl Decode for StreamData { fn decode(reader: &mut ql_codec::Reader) -> Result { let stream_id = reader.decode()?; - let offset = reader.decode_varint()?; + let offset = reader.decode()?; let flags = reader.decode::()?; let fin = (flags & flag::FIN) != 0; let has_header = (flags & flag::HEADER) != 0; @@ -64,7 +64,7 @@ impl StreamData { impl Encode for StreamData { fn encoded_len(&self) -> usize { self.stream_id.encoded_len() - + varint::encoded_len(self.offset) + + self.offset.encoded_len() + size_of::() + self.header.as_ref().map_or(0, encoded_len_bytes) + encoded_len_bytes(&self.bytes) @@ -72,12 +72,12 @@ impl Encode for StreamData { fn encode(&self, out: &mut W) { debug_assert!( - self.offset == 0 || self.header.is_none(), + *self.offset == 0 || self.header.is_none(), "stream header is only valid at offset 0" ); self.stream_id.encode(out); - varint::encode(self.offset, out); + self.offset.encode(out); let mut flags = 0; if self.fin { flags |= flag::FIN; diff --git a/ql-wire/src/encrypted/stream_window.rs b/ql-wire/src/encrypted/stream_window.rs index 20cbaa59..2910f022 100644 --- a/ql-wire/src/encrypted/stream_window.rs +++ b/ql-wire/src/encrypted/stream_window.rs @@ -1,4 +1,4 @@ -use ql_codec::{ByteSlice, Encode, Error}; +use ql_codec::{ByteSlice, Encode, Error, Varint}; use super::StreamId; @@ -6,17 +6,17 @@ use super::StreamId; #[derive(Debug, Clone, PartialEq, Eq)] pub struct StreamWindow { pub stream_id: StreamId, - pub maximum_offset: u64, + pub maximum_offset: Varint, } impl Encode for StreamWindow { fn encoded_len(&self) -> usize { - self.stream_id.encoded_len() + ql_codec::varint::encoded_len(self.maximum_offset) + self.stream_id.encoded_len() + self.maximum_offset.encoded_len() } fn encode(&self, out: &mut W) { self.stream_id.encode(out); - ql_codec::varint::encode(self.maximum_offset, out); + self.maximum_offset.encode(out); } } @@ -24,7 +24,7 @@ impl ql_codec::Decode for StreamWindow { fn decode(reader: &mut ql_codec::Reader) -> Result { Ok(Self { stream_id: reader.decode()?, - maximum_offset: reader.decode_varint()?, + maximum_offset: reader.decode()?, }) } } diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index d5f0ae63..c201c988 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -1,4 +1,4 @@ -use ql_codec::{Decode, Encode}; +use ql_codec::{Decode, Encode, Varint}; use ql_common::{ResetCode, StreamId, QID}; use super::*; @@ -651,11 +651,11 @@ fn encrypted_session_record_round_trip_authenticates_header() { ), SessionFrame::StreamWindow(StreamWindow { stream_id: StreamId(9), - maximum_offset: 65_536, + maximum_offset: Varint(65_536), }), SessionFrame::StreamData(StreamData { stream_id: StreamId(9), - offset: 1024, + offset: Varint(1024), header: None, bytes: b"hello".to_vec(), fin: true, @@ -862,7 +862,7 @@ fn protocol_record_size_breakdown() { &session.tx_key, &[SessionFrame::StreamData(StreamData { stream_id: StreamId(1), - offset: 0, + offset: Varint(0), header: None, fin: false, bytes: Vec::new(), From 7301b7e83fae62e9e135af2df2b1b6440d11bd7f Mon Sep 17 00:00:00 2001 From: Alex's Bot Date: Wed, 22 Jul 2026 16:21:21 +0000 Subject: [PATCH 52/59] ql: standardize codec impls Groundwork for generating these with a macro: `encoded_len` sums the fields instead of restating a constant, discriminant enums all fail with ql_codec::Error, and both byte-container encode bounds are now BufView. Nothing checked `encoded_len` against `encode` before, since encode_record_vec only uses it to size the buffer. The new test covers every type and passes against the pre-change impls, so this is byte-identical. Co-Authored-By: Claude Opus 4.8 (1M context) --- ql-codec/src/codec.rs | 12 +- ql-codec/src/lib.rs | 2 +- ql-common/src/lib.rs | 2 +- ql-wire/src/encrypted/mod.rs | 29 ++-- ql-wire/src/encrypted/stream_reset.rs | 20 +-- ql-wire/src/encrypted_message.rs | 9 +- ql-wire/src/handshake/id.rs | 2 +- ql-wire/src/handshake/ik.rs | 16 +-- ql-wire/src/handshake/mod.rs | 8 +- ql-wire/src/handshake/transport_params.rs | 2 +- ql-wire/src/handshake/xx.rs | 20 +-- ql-wire/src/header.rs | 2 +- ql-wire/src/identity.rs | 8 +- ql-wire/src/macros.rs | 2 +- ql-wire/src/pq.rs | 14 +- ql-wire/src/record.rs | 8 +- ql-wire/src/tests.rs | 165 ++++++++++++++++++++++ 17 files changed, 248 insertions(+), 73 deletions(-) diff --git a/ql-codec/src/codec.rs b/ql-codec/src/codec.rs index e54d3926..7e037f1f 100644 --- a/ql-codec/src/codec.rs +++ b/ql-codec/src/codec.rs @@ -217,8 +217,18 @@ where B: BufView + ?Sized, W: BufMut + ?Sized, { + varint::encode(bytes.buf().remaining(), out); + encode_bytes_raw(bytes, out); +} + +/// Writes the bytes without a length prefix, for a field that runs to the end of the message. +pub fn encode_bytes_raw(bytes: &B, out: &mut W) +where + B: BufView + ?Sized, + W: BufMut + ?Sized, +{ + // `BufMut::put` needs a sized writer, so feed it a chunk at a time. let mut bytes = bytes.buf(); - varint::encode(bytes.remaining(), out); while bytes.has_remaining() { let chunk = bytes.chunk(); out.put_slice(chunk); diff --git a/ql-codec/src/lib.rs b/ql-codec/src/lib.rs index cfb3c649..c79127ab 100644 --- a/ql-codec/src/lib.rs +++ b/ql-codec/src/lib.rs @@ -8,7 +8,7 @@ mod slice; pub mod varint; pub use buf_view::BufView; -pub use codec::{encode_bytes, encoded_len_bytes}; +pub use codec::{encode_bytes, encode_bytes_raw, encoded_len_bytes}; pub use error::Error; pub use reader::Reader; pub use slice::ByteSlice; diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index e01b622a..75a2603b 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -59,7 +59,7 @@ impl QID { impl Encode for QID { fn encoded_len(&self) -> usize { - Self::SIZE + self.0.encoded_len() } fn encode(&self, out: &mut W) { diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index 48f59b81..b870a111 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -77,18 +77,19 @@ impl SessionFrame { impl Encode for SessionFrame { fn encoded_len(&self) -> usize { - 1 + match self { - Self::Ping | Self::Unpair => 0, - Self::Ack(frame) => frame.encoded_len(), - Self::StreamData(frame) => frame.encoded_len(), - Self::StreamWindow(frame) => frame.encoded_len(), - Self::StreamReset(frame) => frame.encoded_len(), - Self::Close(frame) => frame.encoded_len(), - } + self.kind().encoded_len() + + match self { + Self::Ping | Self::Unpair => 0, + Self::Ack(frame) => frame.encoded_len(), + Self::StreamData(frame) => frame.encoded_len(), + Self::StreamWindow(frame) => frame.encoded_len(), + Self::StreamReset(frame) => frame.encoded_len(), + Self::Close(frame) => frame.encoded_len(), + } } fn encode(&self, out: &mut W) { - out.put_u8(self.kind() as u8); + self.kind().encode(out); match self { Self::Ping | Self::Unpair => {} Self::Ack(frame) => frame.encode(out), @@ -135,6 +136,16 @@ impl ql_codec::Decode for SessionFrameKind { } } +impl Encode for SessionFrameKind { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + out.put_u8(*self as u8); + } +} + pub fn parse_session_frames(bytes: B) -> SessionFrameIter { SessionFrameIter { reader: Reader::new(bytes), diff --git a/ql-wire/src/encrypted/stream_reset.rs b/ql-wire/src/encrypted/stream_reset.rs index 6efaec34..2f6c4bbb 100644 --- a/ql-wire/src/encrypted/stream_reset.rs +++ b/ql-wire/src/encrypted/stream_reset.rs @@ -2,7 +2,6 @@ use ql_codec::{ByteSlice, Encode}; use ql_common::ResetCode; use super::StreamId; -use crate::Error; /// aborts one or both lanes of a stream with a reset code /// @@ -16,8 +15,6 @@ pub struct StreamReset { pub code: ResetCode, } -impl StreamReset {} - impl Encode for StreamReset { fn encoded_len(&self) -> usize { self.stream_id.encoded_len() + self.target.encoded_len() + self.code.encoded_len() @@ -52,40 +49,31 @@ pub enum ResetTarget { Both = 3, } -impl ResetTarget { - pub const fn to_wire(self) -> u8 { - self as u8 - } -} - impl Encode for ResetTarget { fn encoded_len(&self) -> usize { size_of::() } fn encode(&self, out: &mut W) { - self.to_wire().encode(out); + out.put_u8(*self as u8); } } impl TryFrom for ResetTarget { - type Error = Error; + type Error = ql_codec::Error; fn try_from(value: u8) -> Result { match value { 1 => Ok(Self::Origin), 2 => Ok(Self::Return), 3 => Ok(Self::Both), - _ => Err(Error::InvalidDiscriminant), + _ => Err(ql_codec::Error::InvalidDiscriminant), } } } impl ql_codec::Decode for ResetTarget { fn decode(reader: &mut ql_codec::Reader) -> Result { - reader - .decode::()? - .try_into() - .map_err(|_| ql_codec::Error::InvalidDiscriminant) + reader.decode::()?.try_into() } } diff --git a/ql-wire/src/encrypted_message.rs b/ql-wire/src/encrypted_message.rs index 99abda5e..0b2ddfe0 100644 --- a/ql-wire/src/encrypted_message.rs +++ b/ql-wire/src/encrypted_message.rs @@ -1,4 +1,5 @@ -use ql_codec::{ByteSlice, Decode, Encode}; +use bytes::Buf; +use ql_codec::{encode_bytes_raw, BufView, ByteSlice, Decode, Encode}; use crate::ENCRYPTED_MESSAGE_AUTH_SIZE; @@ -29,13 +30,13 @@ impl Decode for EncryptedMessage { } } -impl> Encode for EncryptedMessage { +impl Encode for EncryptedMessage { fn encoded_len(&self) -> usize { - ENCRYPTED_MESSAGE_AUTH_SIZE + self.ciphertext.as_ref().len() + self.auth.encoded_len() + self.ciphertext.buf().remaining() } fn encode(&self, out: &mut W) { self.auth.encode(out); - out.put_slice(self.ciphertext.as_ref()); + encode_bytes_raw(&self.ciphertext, out); } } diff --git a/ql-wire/src/handshake/id.rs b/ql-wire/src/handshake/id.rs index 3e4915ab..f0a9dfdb 100644 --- a/ql-wire/src/handshake/id.rs +++ b/ql-wire/src/handshake/id.rs @@ -16,7 +16,7 @@ impl ql_codec::Decode for HandshakeId { impl Encode for HandshakeId { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + self.0.encoded_len() } fn encode(&self, out: &mut W) { diff --git a/ql-wire/src/handshake/ik.rs b/ql-wire/src/handshake/ik.rs index 9b30b7b0..6adeb8bf 100644 --- a/ql-wire/src/handshake/ik.rs +++ b/ql-wire/src/handshake/ik.rs @@ -25,10 +25,10 @@ pub struct Ik1 { impl Encode for Ik1 { fn encoded_len(&self) -> usize { - HandshakeId::WIRE_SIZE - + TransportParams::WIRE_SIZE - + MlKemCiphertext::SIZE - + EphemeralPublicKey::WIRE_SIZE + self.handshake_id.encoded_len() + + self.transport_params.encoded_len() + + self.skem_ciphertext.encoded_len() + + self.ephemeral.encoded_len() + self .static_bundle .as_ref() @@ -77,10 +77,10 @@ pub struct Ik2 { impl Encode for Ik2 { fn encoded_len(&self) -> usize { - HandshakeId::WIRE_SIZE - + TransportParams::WIRE_SIZE - + MlKemCiphertext::SIZE - + EncryptedMlKemCiphertext::WIRE_SIZE + self.handshake_id.encoded_len() + + self.transport_params.encoded_len() + + self.ekem_ciphertext.encoded_len() + + self.skem_ciphertext.encoded_len() } fn encode(&self, out: &mut W) { diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index db585fc5..e1803723 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -35,7 +35,7 @@ impl EphemeralPublicKey { impl Encode for EphemeralPublicKey { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + self.mlkem_public_key.encoded_len() } fn encode(&self, out: &mut W) { @@ -68,17 +68,17 @@ impl EncryptedMlKemCiphertext { impl Encode for EncryptedMlKemCiphertext { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + self.0.encoded_len() } fn encode(&self, out: &mut W) { - self.0.as_ref().encode(out); + self.0.encode(out); } } impl ql_codec::Decode for EncryptedMlKemCiphertext { fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self::new(reader.decode()?)) + Ok(Self(reader.decode()?)) } } diff --git a/ql-wire/src/handshake/transport_params.rs b/ql-wire/src/handshake/transport_params.rs index 71766e54..9d485b6b 100644 --- a/ql-wire/src/handshake/transport_params.rs +++ b/ql-wire/src/handshake/transport_params.rs @@ -13,7 +13,7 @@ impl TransportParams { impl Encode for TransportParams { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + self.initial_stream_receive_window.encoded_len() } fn encode(&self, out: &mut W) { diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs index 22e607ee..564f9f96 100644 --- a/ql-wire/src/handshake/xx.rs +++ b/ql-wire/src/handshake/xx.rs @@ -36,10 +36,10 @@ impl ql_codec::Decode for Xx1 { impl Encode for Xx1 { fn encoded_len(&self) -> usize { - HandshakeId::WIRE_SIZE - + PairingId::SIZE - + TransportParams::WIRE_SIZE - + EphemeralPublicKey::WIRE_SIZE + self.handshake_id.encoded_len() + + self.pairing_id.encoded_len() + + self.transport_params.encoded_len() + + self.ephemeral.encoded_len() } fn encode(&self, out: &mut W) { @@ -71,9 +71,9 @@ impl ql_codec::Decode for Xx2 { impl Encode for Xx2 { fn encoded_len(&self) -> usize { - HandshakeId::WIRE_SIZE - + TransportParams::WIRE_SIZE - + MlKemCiphertext::SIZE + self.handshake_id.encoded_len() + + self.transport_params.encoded_len() + + self.ekem_ciphertext.encoded_len() + self.static_bundle.encoded_len() } @@ -104,8 +104,8 @@ impl ql_codec::Decode for Xx3 { impl Encode for Xx3 { fn encoded_len(&self) -> usize { - HandshakeId::WIRE_SIZE - + EncryptedMlKemCiphertext::WIRE_SIZE + self.handshake_id.encoded_len() + + self.skem_ciphertext.encoded_len() + self.static_bundle.encoded_len() } @@ -133,7 +133,7 @@ impl ql_codec::Decode for Xx4 { impl Encode for Xx4 { fn encoded_len(&self) -> usize { - HandshakeId::WIRE_SIZE + EncryptedMlKemCiphertext::WIRE_SIZE + self.handshake_id.encoded_len() + self.skem_ciphertext.encoded_len() } fn encode(&self, out: &mut W) { diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index d3f45b77..a89e357b 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -16,7 +16,7 @@ impl RouteHeader { impl Encode for RouteHeader { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + self.sender.encoded_len() + self.recipient.encoded_len() } fn encode(&self, out: &mut W) { diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index 731a91f8..d7e5ae50 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -25,10 +25,10 @@ impl PeerBundle { impl Encode for PeerBundle { fn encoded_len(&self) -> usize { - size_of::() - + QID::SIZE - + size_of::() - + MlKemPublicKey::SIZE + self.version.encoded_len() + + self.qid.encoded_len() + + self.capabilities.encoded_len() + + self.mlkem_public_key.encoded_len() + self.name.encoded_len() } diff --git a/ql-wire/src/macros.rs b/ql-wire/src/macros.rs index 9124187c..7fcb2d29 100644 --- a/ql-wire/src/macros.rs +++ b/ql-wire/src/macros.rs @@ -14,7 +14,7 @@ macro_rules! array_wrapper { impl ql_codec::Encode for $name { fn encoded_len(&self) -> usize { - Self::SIZE + self.0.encoded_len() } fn encode(&self, out: &mut W) { diff --git a/ql-wire/src/pq.rs b/ql-wire/src/pq.rs index 9cd0efcb..8f5cc87b 100644 --- a/ql-wire/src/pq.rs +++ b/ql-wire/src/pq.rs @@ -35,7 +35,7 @@ impl Drop for SessionKey { impl Encode for SessionKey { fn encoded_len(&self) -> usize { - Self::SIZE + self.0.encoded_len() } fn encode(&self, out: &mut W) { @@ -72,17 +72,17 @@ impl Drop for MlKemPublicKey { impl ql_codec::Decode for MlKemPublicKey { fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self::new(reader.decode()?)) + Ok(Self(reader.decode()?)) } } impl Encode for MlKemPublicKey { fn encoded_len(&self) -> usize { - Self::SIZE + self.0.encoded_len() } fn encode(&self, out: &mut W) { - self.0.as_ref().encode(out); + self.0.encode(out); } } @@ -130,17 +130,17 @@ impl Drop for MlKemCiphertext { impl ql_codec::Decode for MlKemCiphertext { fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self::new(reader.decode()?)) + Ok(Self(reader.decode()?)) } } impl Encode for MlKemCiphertext { fn encoded_len(&self) -> usize { - Self::SIZE + self.0.encoded_len() } fn encode(&self, out: &mut W) { - self.0.as_ref().encode(out); + self.0.encode(out); } } diff --git a/ql-wire/src/record.rs b/ql-wire/src/record.rs index 3eba0fb8..6d5278c8 100644 --- a/ql-wire/src/record.rs +++ b/ql-wire/src/record.rs @@ -1,4 +1,4 @@ -use ql_codec::{ByteSlice, Decode, Encode}; +use ql_codec::{BufView, ByteSlice, Decode, Encode}; use crate::{ encrypted_message::EncryptedMessage, @@ -61,11 +61,11 @@ impl Decode for RecordHeader { impl Encode for RecordHeader { fn encoded_len(&self) -> usize { - Self::WIRE_SIZE + self.version.encoded_len() + self.route.encoded_len() + self.record_type.encoded_len() } fn encode(&self, out: &mut W) { - out.put_u8(self.version); + self.version.encode(out); self.route.encode(out); self.record_type.encode(out); } @@ -232,7 +232,7 @@ pub struct QlSessionRecord { pub payload: EncryptedMessage, } -impl> Encode for QlSessionRecord { +impl Encode for QlSessionRecord { fn encoded_len(&self) -> usize { self.header.encoded_len() + self.payload.encoded_len() } diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index c201c988..88444f4a 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -898,3 +898,168 @@ fn protocol_record_size_breakdown() { ); print_size("ql-wire session close", record_size(&session_close)); } + +/// `encode_vec` asserts the buffer ends up `encoded_len()` long, so encoding is the check. +fn assert_encoded_len(label: &str, value: &T) { + let encoded = value.encode_vec(); + assert_eq!(encoded.len(), value.encoded_len(), "{label}"); +} + +fn mlkem_public_key() -> MlKemPublicKey { + MlKemPublicKey::new(Box::new([7u8; MlKemPublicKey::SIZE])) +} + +fn mlkem_ciphertext() -> MlKemCiphertext { + MlKemCiphertext::new(Box::new([9u8; MlKemCiphertext::SIZE])) +} + +fn encrypted_mlkem_ciphertext() -> EncryptedMlKemCiphertext { + EncryptedMlKemCiphertext::new(Box::new([3u8; EncryptedMlKemCiphertext::WIRE_SIZE])) +} + +fn encrypted_peer_bundle() -> EncryptedPeerBundle { + EncryptedPeerBundle(Box::from(&[1u8, 2, 3, 4, 5][..])) +} + +fn ephemeral_public_key() -> EphemeralPublicKey { + EphemeralPublicKey { + mlkem_public_key: mlkem_public_key(), + } +} + +#[test] +fn encoded_len_matches_encoding() { + let params = handshake_transport_params(1024); + + assert_encoded_len("QID", &QID([1u8; QID::SIZE])); + assert_encoded_len("HandshakeId", &handshake_id(0xdead_beef)); + assert_encoded_len("PairingId", &PairingId([4u8; PairingId::SIZE])); + assert_encoded_len("SessionKey", &SessionKey([5u8; SessionKey::SIZE])); + assert_encoded_len("MlKemPublicKey", &mlkem_public_key()); + assert_encoded_len("MlKemCiphertext", &mlkem_ciphertext()); + assert_encoded_len("EncryptedMlKemCiphertext", &encrypted_mlkem_ciphertext()); + assert_encoded_len("EncryptedPeerBundle", &encrypted_peer_bundle()); + assert_encoded_len("EphemeralPublicKey", &ephemeral_public_key()); + assert_encoded_len("TransportParams", ¶ms); + assert_encoded_len("RouteHeader", &route(1, 2)); + assert_encoded_len( + "RecordHeader", + &RecordHeader::new(route(1, 2), RecordType::Handshake), + ); + assert_encoded_len("RecordType", &RecordType::Session); + assert_encoded_len("HandshakeKind", &HandshakeKind::Xx4); + assert_encoded_len("SessionHeader", &SessionHeader { seq: RecordSeq(7) }); + assert_encoded_len("ResetTarget", &ResetTarget::Both); + assert_encoded_len( + "PeerBundle", + &PeerBundle { + version: 1, + qid: QID([3u8; QID::SIZE]), + capabilities: 0xffff_ffff, + mlkem_public_key: mlkem_public_key(), + name: "device".to_owned(), + }, + ); + assert_encoded_len( + "StreamReset", + &StreamReset { + stream_id: StreamId(9), + target: ResetTarget::Origin, + code: ResetCode::TIMEOUT, + }, + ); + assert_encoded_len( + "StreamWindow", + &StreamWindow { + stream_id: StreamId(9), + maximum_offset: Varint(1 << 40), + }, + ); + assert_encoded_len( + "SessionClose", + &SessionClose { + code: SessionCloseCode::PROTOCOL, + }, + ); + let payload = EncryptedMessage { + auth: [6u8; ENCRYPTED_MESSAGE_AUTH_SIZE], + ciphertext: vec![21u8; 37], + }; + assert_encoded_len("EncryptedMessage", &payload); + assert_encoded_len( + "QlSessionRecord", + &QlSessionRecord { + header: SessionHeader { seq: RecordSeq(7) }, + payload, + }, + ); + assert_encoded_len("SessionFrame::Ping", &SessionFrame::>::Ping); + assert_encoded_len( + "SessionFrame::Close", + &SessionFrame::>::Close(SessionClose { + code: SessionCloseCode::PROTOCOL, + }), + ); + assert_encoded_len( + "Xx1", + &Xx1 { + handshake_id: handshake_id(1), + pairing_id: PairingId([4u8; PairingId::SIZE]), + transport_params: params, + ephemeral: ephemeral_public_key(), + }, + ); + assert_encoded_len( + "Xx2", + &Xx2 { + handshake_id: handshake_id(2), + transport_params: params, + ekem_ciphertext: mlkem_ciphertext(), + static_bundle: encrypted_peer_bundle(), + }, + ); + assert_encoded_len( + "Xx3", + &Xx3 { + handshake_id: handshake_id(3), + skem_ciphertext: encrypted_mlkem_ciphertext(), + static_bundle: encrypted_peer_bundle(), + }, + ); + assert_encoded_len( + "Xx4", + &Xx4 { + handshake_id: handshake_id(4), + skem_ciphertext: encrypted_mlkem_ciphertext(), + }, + ); + assert_encoded_len( + "Ik1", + &Ik1 { + handshake_id: handshake_id(5), + transport_params: params, + skem_ciphertext: mlkem_ciphertext(), + ephemeral: ephemeral_public_key(), + static_bundle: Some(encrypted_peer_bundle()), + }, + ); + assert_encoded_len( + "Ik1 without bundle", + &Ik1 { + handshake_id: handshake_id(5), + transport_params: params, + skem_ciphertext: mlkem_ciphertext(), + ephemeral: ephemeral_public_key(), + static_bundle: None, + }, + ); + assert_encoded_len( + "Ik2", + &Ik2 { + handshake_id: handshake_id(6), + transport_params: params, + ekem_ciphertext: mlkem_ciphertext(), + skem_ciphertext: encrypted_mlkem_ciphertext(), + }, + ); +} From 27bdf5447a92637a25a4d02fc9e83fc2f16d1da1 Mon Sep 17 00:00:00 2001 From: Alex's Bot Date: Wed, 22 Jul 2026 17:50:20 +0000 Subject: [PATCH 53/59] ql: generate codec impls with macros Most Encode/Decode impls just walked the fields in order. What stays hand-written reads or writes something that isn't a field: a flags byte (StreamData), a trailing optional (Ik1), a block count (RecordAck), or the rest of the buffer (EncryptedMessage, EncryptedPeerBundle). PairingInvite loses its codec entirely, since nothing ever called it, and array_wrapper is gone because codec_newtype covers it without #[macro_use]. Co-Authored-By: Claude Opus 4.8 (1M context) --- ql-codec/src/lib.rs | 1 + ql-codec/src/macros.rs | 381 ++++++++++++++++++++++ ql-common/src/lib.rs | 26 +- ql-fsm/src/pairing.rs | 30 -- ql-wire/src/crypto.rs | 8 +- ql-wire/src/encrypted/close.rs | 30 +- ql-wire/src/encrypted/mod.rs | 123 +------ ql-wire/src/encrypted/stream_reset.rs | 91 ++---- ql-wire/src/encrypted/stream_window.rs | 32 +- ql-wire/src/handshake/id.rs | 26 +- ql-wire/src/handshake/ik.rs | 39 +-- ql-wire/src/handshake/mod.rs | 46 +-- ql-wire/src/handshake/pairing.rs | 18 +- ql-wire/src/handshake/transport_params.rs | 32 +- ql-wire/src/handshake/xx.rs | 141 ++------ ql-wire/src/header.rs | 56 +--- ql-wire/src/identity.rs | 91 +----- ql-wire/src/lib.rs | 2 - ql-wire/src/macros.rs | 31 -- ql-wire/src/pq.rs | 70 +--- ql-wire/src/record.rs | 229 ++----------- 21 files changed, 579 insertions(+), 924 deletions(-) create mode 100644 ql-codec/src/macros.rs delete mode 100644 ql-wire/src/macros.rs diff --git a/ql-codec/src/lib.rs b/ql-codec/src/lib.rs index c79127ab..e522c7fc 100644 --- a/ql-codec/src/lib.rs +++ b/ql-codec/src/lib.rs @@ -3,6 +3,7 @@ mod buf_view; mod codec; mod error; +mod macros; mod reader; mod slice; pub mod varint; diff --git a/ql-codec/src/macros.rs b/ql-codec/src/macros.rs new file mode 100644 index 00000000..ad9a61f8 --- /dev/null +++ b/ql-codec/src/macros.rs @@ -0,0 +1,381 @@ +/// Defines a newtype and encodes it exactly as the value it wraps. +#[macro_export] +macro_rules! codec_newtype { + ( + $(#[$meta:meta])* + $vis:vis struct $name:ident($field_vis:vis $inner:ty); + ) => { + $(#[$meta])* + $vis struct $name($field_vis $inner); + + impl $crate::Encode for $name { + fn encoded_len(&self) -> usize { + $crate::Encode::encoded_len(&self.0) + } + + fn encode(&self, out: &mut W) { + $crate::Encode::encode(&self.0, out); + } + } + + impl $crate::Decode for $name { + fn decode(reader: &mut $crate::Reader) -> Result { + Ok(Self(reader.decode()?)) + } + } + }; +} + +/// Defines a struct and encodes its fields back to back in declaration order. +/// +/// A single generic parameter is taken to be the byte container, bound as +/// [`BufView`](crate::BufView) for encoding and [`ByteSlice`](crate::ByteSlice) for decoding. It +/// may only appear nested inside another codec type, and decoding ties it to the reader's own +/// container, because a borrowed field can only be taken from the reader it is read out of. +/// Anything whose fields do not map one to one onto the wire needs a hand-written impl. +#[macro_export] +macro_rules! codec_struct { + ( + $(#[$meta:meta])* + $vis:vis struct $name:ident $(<$bytes:ident>)? { + $($(#[$field_meta:meta])* $field_vis:vis $field:ident: $ty:ty),* $(,)? + } + ) => { + $(#[$meta])* + $vis struct $name$(<$bytes>)? { + $($(#[$field_meta])* $field_vis $field: $ty,)* + } + + impl$(<$bytes: $crate::BufView>)? $crate::Encode for $name$(<$bytes>)? { + fn encoded_len(&self) -> usize { + $($crate::Encode::encoded_len(&self.$field) +)* 0 + } + + fn encode(&self, out: &mut W) { + $($crate::Encode::encode(&self.$field, out);)* + } + } + + $crate::__codec_struct_decode!($name$(<$bytes>)?, $($field),*); + }; +} + +/// The one impl that cannot be written once: `Decode` always needs a byte container to name, so +/// a generic struct reuses its own parameter while a plain one introduces a fresh name. +#[macro_export] +#[doc(hidden)] +macro_rules! __codec_struct_decode { + ($name:ident<$bytes:ident>, $($field:ident),* $(,)?) => { + impl<$bytes: $crate::ByteSlice> $crate::Decode<$bytes> for $name<$bytes> { + fn decode(reader: &mut $crate::Reader<$bytes>) -> Result { + Ok(Self { $($field: reader.decode()?,)* }) + } + } + }; + + ($name:ident, $($field:ident),* $(,)?) => { + impl $crate::Decode for $name { + fn decode(reader: &mut $crate::Reader) -> Result { + Ok(Self { $($field: reader.decode()?,)* }) + } + } + }; +} + +/// Defines an enum carried on the wire as a `u8` discriminant. +/// +/// `enum Frame as FrameKind` additionally defines the discriminant enum and a `kind` accessor, +/// and encodes each variant's payload after the discriminant. +/// +/// ``` +/// use ql_codec::{Decode, Encode}; +/// +/// ql_codec::codec_enum! { +/// #[derive(Debug, Clone, Copy, PartialEq, Eq)] +/// pub enum CloseReason { +/// Done = 1, +/// Refused = 2, +/// } +/// } +/// +/// ql_codec::codec_enum! { +/// #[derive(Debug, PartialEq)] +/// pub enum Frame as FrameKind { +/// Ping = 1, +/// Close(CloseReason) = 3, +/// } +/// } +/// +/// // A unit variant is just its discriminant. +/// assert_eq!(Frame::Ping.encode_vec(), [1]); +/// +/// // A payload follows the discriminant, and `kind()` reads it back without the payload. +/// let frame = Frame::Close(CloseReason::Refused); +/// assert_eq!(frame.kind(), FrameKind::Close); +/// assert_eq!(frame.encode_vec(), [3, 2]); +/// assert_eq!(Frame::decode_bytes(&[3, 2][..]).unwrap(), frame); +/// +/// assert_eq!( +/// Frame::decode_bytes(&[9][..]), +/// Err(ql_codec::Error::InvalidDiscriminant), +/// ); +/// ``` +#[macro_export] +macro_rules! codec_enum { + ( + $(#[$meta:meta])* + $vis:vis enum $name:ident $(<$bytes:ident>)? as $kind:ident { + $($(#[$variant_meta:meta])* $variant:ident $(($payload:ty))? = $value:literal),* $(,)? + } + ) => { + $crate::codec_enum! { + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + $vis enum $kind { + $($variant = $value,)* + } + } + + $(#[$meta])* + $vis enum $name$(<$bytes>)? { + $($(#[$variant_meta])* $variant $(($payload))?,)* + } + + impl$(<$bytes>)? $name$(<$bytes>)? { + $vis fn kind(&self) -> $kind { + match self { + $(Self::$variant { .. } => $kind::$variant,)* + } + } + } + + impl$(<$bytes: $crate::BufView>)? $crate::Encode for $name$(<$bytes>)? { + #[allow(unreachable_patterns)] + fn encoded_len(&self) -> usize { + $crate::Encode::encoded_len(&self.kind()) + + match self { + $($(Self::$variant(payload) => + <$payload as $crate::Encode>::encoded_len(payload),)?)* + _ => 0, + } + } + + #[allow(unreachable_patterns)] + fn encode(&self, out: &mut W) { + $crate::Encode::encode(&self.kind(), out); + match self { + $($(Self::$variant(payload) => + <$payload as $crate::Encode>::encode(payload, out),)?)* + _ => {} + } + } + } + + $crate::__codec_enum_decode! { + $name$(<$bytes>)?, $kind, $($variant $(($payload))?),* + } + }; + + ( + $(#[$meta:meta])* + $vis:vis enum $name:ident { + $($(#[$variant_meta:meta])* $variant:ident = $value:literal),* $(,)? + } + ) => { + $(#[$meta])* + #[repr(u8)] + $vis enum $name { + $($(#[$variant_meta])* $variant = $value,)* + } + + impl TryFrom for $name { + type Error = $crate::Error; + + fn try_from(value: u8) -> Result { + match value { + $($value => Ok(Self::$variant),)* + _ => Err($crate::Error::InvalidDiscriminant), + } + } + } + + impl $crate::Encode for $name { + fn encoded_len(&self) -> usize { + size_of::() + } + + fn encode(&self, out: &mut W) { + ::bytes::BufMut::put_u8(out, *self as u8); + } + } + + impl $crate::Decode for $name { + fn decode(reader: &mut $crate::Reader) -> Result { + reader.decode::()?.try_into() + } + } + }; +} + +/// The one impl that cannot be written once: `Decode` always needs a byte container to name, so +/// a generic enum reuses its own parameter while a plain one introduces a fresh name. +#[macro_export] +#[doc(hidden)] +macro_rules! __codec_enum_decode { + ($name:ident<$bytes:ident>, $kind:ident, $($variant:ident $(($payload:ty))?),* $(,)?) => { + impl<$bytes: $crate::ByteSlice> $crate::Decode<$bytes> for $name<$bytes> { + fn decode(reader: &mut $crate::Reader<$bytes>) -> Result { + Ok(match reader.decode::<$kind>()? { + $($kind::$variant => Self::$variant $((reader.decode::<$payload>()?))?,)* + }) + } + } + }; + + ($name:ident, $kind:ident, $($variant:ident $(($payload:ty))?),* $(,)?) => { + impl $crate::Decode for $name { + fn decode(reader: &mut $crate::Reader) -> Result { + Ok(match reader.decode::<$kind>()? { + $($kind::$variant => Self::$variant $((reader.decode::<$payload>()?))?,)* + }) + } + } + }; +} + +#[cfg(test)] +mod tests { + use bytes::Bytes; + + use crate::{Decode, Encode, Error}; + + codec_newtype! { + /// A newtype carrying a fixed-size array. + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub struct Tag(pub [u8; 4]); + } + + codec_enum! { + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub enum Flavour { + Sweet = 1, + Salty = 3, + } + } + + codec_struct! { + /// Fields encode in declaration order. + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Plain { + pub tag: Tag, + pub flavour: Flavour, + pub name: String, + } + } + + codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Wrapped { + pub tag: Tag, + pub body: Nested, + } + } + + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Nested(pub B); + + impl Encode for Nested { + fn encoded_len(&self) -> usize { + crate::encoded_len_bytes(&self.0) + } + + fn encode(&self, out: &mut W) { + crate::encode_bytes(&self.0, out); + } + } + + impl Decode for Nested { + fn decode(reader: &mut crate::Reader) -> Result { + Ok(Self(reader.take_len_prefixed()?)) + } + } + + codec_enum! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub enum Frame as FrameKind { + Ping = 1, + Plain(Plain) = 2, + Body(Wrapped) = 3, + Pong = 4, + } + } + + codec_enum! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub enum Message as MessageKind { + Empty = 1, + Plain(Plain) = 2, + } + } + + fn plain() -> Plain { + Plain { + tag: Tag([1, 2, 3, 4]), + flavour: Flavour::Salty, + name: "hello".to_owned(), + } + } + + #[test] + fn struct_fields_encode_in_declaration_order() { + let encoded = plain().encode_vec(); + assert_eq!(&encoded[..4], &[1, 2, 3, 4]); + assert_eq!(encoded[4], 3); + assert_eq!(Plain::decode_bytes(encoded.as_slice()).unwrap(), plain()); + + let tag = Tag([9; 4]); + assert_eq!(Tag::decode_bytes(tag.encode_vec().as_slice()).unwrap(), tag); + } + + #[test] + fn generic_struct_decodes_from_its_own_container() { + let value = Wrapped { + tag: Tag([5; 4]), + body: Nested(Bytes::from_static(b"body")), + }; + let encoded = Bytes::from(value.encode_vec()); + assert_eq!(Wrapped::::decode_bytes(encoded).unwrap(), value); + } + + #[test] + fn unknown_discriminant_is_rejected() { + assert_eq!(Flavour::Salty.encode_vec(), [3]); + assert_eq!( + Flavour::decode_bytes(&[2u8][..]), + Err(Error::InvalidDiscriminant) + ); + assert_eq!( + Message::decode_bytes(&[9u8][..]), + Err(Error::InvalidDiscriminant) + ); + } + + #[test] + fn payload_enum_writes_the_kind_first() { + let frame = Frame::::Plain(plain()); + assert_eq!(frame.kind(), FrameKind::Plain); + assert_eq!(frame.encode_vec()[0], 2); + assert_eq!(Frame::::Pong.encode_vec(), [4]); + + let encoded = Bytes::from(frame.encode_vec()); + assert_eq!(Frame::::decode_bytes(encoded).unwrap(), frame); + } + + #[test] + fn payload_enum_without_generics_round_trips() { + for message in [Message::Empty, Message::Plain(plain())] { + let encoded = message.encode_vec(); + assert_eq!(Message::decode_bytes(encoded.as_slice()).unwrap(), message); + } + assert_eq!(Message::Plain(plain()).kind(), MessageKind::Plain); + } +} diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index 75a2603b..ef159103 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -1,7 +1,5 @@ //! Shared QuantumLink primitive types. -use ql_codec::{ByteSlice, Decode, Encode, Error, Reader}; - #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] pub struct ResetCode(pub u64); @@ -49,30 +47,16 @@ impl std::fmt::Display for ResetCode { ql_codec::varint_wrapper!(ResetCode, u64); -#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] -#[repr(transparent)] -pub struct QID(pub [u8; Self::SIZE]); +ql_codec::codec_newtype! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + #[repr(transparent)] + pub struct QID(pub [u8; Self::SIZE]); +} impl QID { pub const SIZE: usize = 16; } -impl Encode for QID { - fn encoded_len(&self) -> usize { - self.0.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl Decode for QID { - fn decode(reader: &mut Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - /// Identifier for a stream within a QL session. #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] #[repr(transparent)] diff --git a/ql-fsm/src/pairing.rs b/ql-fsm/src/pairing.rs index 3f179d0d..fb9b0532 100644 --- a/ql-fsm/src/pairing.rs +++ b/ql-fsm/src/pairing.rs @@ -1,4 +1,3 @@ -use ql_codec::{ByteSlice, Decode, Encode, Reader}; use ql_common::QID; use ql_wire::PairingToken; @@ -8,32 +7,3 @@ pub struct PairingInvite { pub qid: QID, pub token: PairingToken, } - -impl PairingInvite { - pub const VERSION: u8 = 1; -} - -impl Encode for PairingInvite { - fn encoded_len(&self) -> usize { - size_of::() + QID::SIZE + PairingToken::SIZE - } - - fn encode(&self, out: &mut W) { - Self::VERSION.encode(out); - self.qid.encode(out); - self.token.encode(out); - } -} - -impl Decode for PairingInvite { - fn decode(reader: &mut Reader) -> Result { - if reader.decode::()? != Self::VERSION { - return Err(ql_codec::Error::InvalidDiscriminant); - } - - Ok(Self { - qid: reader.decode()?, - token: reader.decode()?, - }) - } -} diff --git a/ql-wire/src/crypto.rs b/ql-wire/src/crypto.rs index 888a756a..0bc11c3a 100644 --- a/ql-wire/src/crypto.rs +++ b/ql-wire/src/crypto.rs @@ -3,9 +3,15 @@ use crate::{ ENCRYPTED_MESSAGE_AUTH_SIZE, }; -array_wrapper!(Nonce, 12); +ql_codec::codec_newtype! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + #[repr(transparent)] + pub struct Nonce(pub [u8; Self::SIZE]); +} impl Nonce { + pub const SIZE: usize = 12; + pub fn from_counter(counter: u64) -> Self { let mut nonce = [0u8; Self::SIZE]; nonce[4..].copy_from_slice(&counter.to_le_bytes()); diff --git a/ql-wire/src/encrypted/close.rs b/ql-wire/src/encrypted/close.rs index 77528c90..0b9f8d6d 100644 --- a/ql-wire/src/encrypted/close.rs +++ b/ql-wire/src/encrypted/close.rs @@ -1,9 +1,9 @@ -use ql_codec::{ByteSlice, Encode}; - -/// closes the whole session immediately with a reset code. -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct SessionClose { - pub code: SessionCloseCode, +ql_codec::codec_struct! { + /// closes the whole session immediately with a reset code. + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct SessionClose { + pub code: SessionCloseCode, + } } #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] @@ -17,21 +17,3 @@ impl SessionCloseCode { } ql_codec::varint_wrapper!(SessionCloseCode, u64); - -impl ql_codec::Decode for SessionClose { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - code: reader.decode()?, - }) - } -} - -impl Encode for SessionClose { - fn encoded_len(&self) -> usize { - self.code.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.code.encode(out); - } -} diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index b870a111..4f349033 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -1,4 +1,4 @@ -use ql_codec::{BufView, ByteSlice, Decode, Encode, Reader}; +use ql_codec::{ByteSlice, Reader}; use ql_common::StreamId; use crate::{ @@ -19,45 +19,17 @@ pub use stream_data::*; pub use stream_reset::*; pub use stream_window::*; -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum SessionFrame { - // todo: do we need ping as explicit frame? - Ping, - Unpair, - Ack(RecordAck), - StreamData(StreamData), - StreamWindow(StreamWindow), - StreamReset(StreamReset), - Close(SessionClose), -} - -impl Decode for SessionFrame { - fn decode(reader: &mut Reader) -> Result { - let kind = reader.decode::()?; - let frame = match kind { - SessionFrameKind::Ping => Self::Ping, - SessionFrameKind::Unpair => Self::Unpair, - SessionFrameKind::Ack => Self::Ack(reader.decode::()?), - SessionFrameKind::StreamData => Self::StreamData(reader.decode::>()?), - SessionFrameKind::StreamWindow => Self::StreamWindow(reader.decode::()?), - SessionFrameKind::StreamReset => Self::StreamReset(reader.decode::()?), - SessionFrameKind::Close => Self::Close(reader.decode::()?), - }; - Ok(frame) - } -} - -impl SessionFrame { - fn kind(&self) -> SessionFrameKind { - match self { - Self::Ping => SessionFrameKind::Ping, - Self::Unpair => SessionFrameKind::Unpair, - Self::Ack(_) => SessionFrameKind::Ack, - Self::StreamData(_) => SessionFrameKind::StreamData, - Self::StreamWindow(_) => SessionFrameKind::StreamWindow, - Self::StreamReset(_) => SessionFrameKind::StreamReset, - Self::Close(_) => SessionFrameKind::Close, - } +ql_codec::codec_enum! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub enum SessionFrame as SessionFrameKind { + // todo: do we need ping as explicit frame? + Ping = 1, + Ack(RecordAck) = 2, + StreamData(StreamData) = 3, + StreamWindow(StreamWindow) = 4, + StreamReset(StreamReset) = 5, + Close(SessionClose) = 6, + Unpair = 7, } } @@ -75,77 +47,6 @@ impl SessionFrame { } } -impl Encode for SessionFrame { - fn encoded_len(&self) -> usize { - self.kind().encoded_len() - + match self { - Self::Ping | Self::Unpair => 0, - Self::Ack(frame) => frame.encoded_len(), - Self::StreamData(frame) => frame.encoded_len(), - Self::StreamWindow(frame) => frame.encoded_len(), - Self::StreamReset(frame) => frame.encoded_len(), - Self::Close(frame) => frame.encoded_len(), - } - } - - fn encode(&self, out: &mut W) { - self.kind().encode(out); - match self { - Self::Ping | Self::Unpair => {} - Self::Ack(frame) => frame.encode(out), - Self::StreamData(frame) => frame.encode(out), - Self::StreamWindow(frame) => frame.encode(out), - Self::StreamReset(frame) => frame.encode(out), - Self::Close(frame) => frame.encode(out), - } - } -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -#[repr(u8)] -pub enum SessionFrameKind { - Ping = 1, - Ack = 2, - StreamData = 3, - StreamWindow = 4, - StreamReset = 5, - Close = 6, - Unpair = 7, -} - -impl TryFrom for SessionFrameKind { - type Error = ql_codec::Error; - - fn try_from(value: u8) -> Result { - match value { - 1 => Ok(Self::Ping), - 2 => Ok(Self::Ack), - 3 => Ok(Self::StreamData), - 4 => Ok(Self::StreamWindow), - 5 => Ok(Self::StreamReset), - 6 => Ok(Self::Close), - 7 => Ok(Self::Unpair), - _ => Err(ql_codec::Error::InvalidDiscriminant), - } - } -} - -impl ql_codec::Decode for SessionFrameKind { - fn decode(reader: &mut ql_codec::Reader) -> Result { - reader.decode::()?.try_into() - } -} - -impl Encode for SessionFrameKind { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u8(*self as u8); - } -} - pub fn parse_session_frames(bytes: B) -> SessionFrameIter { SessionFrameIter { reader: Reader::new(bytes), diff --git a/ql-wire/src/encrypted/stream_reset.rs b/ql-wire/src/encrypted/stream_reset.rs index 2f6c4bbb..30c342e7 100644 --- a/ql-wire/src/encrypted/stream_reset.rs +++ b/ql-wire/src/encrypted/stream_reset.rs @@ -1,79 +1,30 @@ -use ql_codec::{ByteSlice, Encode}; use ql_common::ResetCode; use super::StreamId; -/// aborts one or both lanes of a stream with a reset code -/// -/// stream origin is the peer that opened the stream -/// origin lane carries bytes sent by the stream origin -/// return lane carries bytes sent back toward the stream origin -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct StreamReset { - pub stream_id: StreamId, - pub target: ResetTarget, - pub code: ResetCode, -} - -impl Encode for StreamReset { - fn encoded_len(&self) -> usize { - self.stream_id.encoded_len() + self.target.encoded_len() + self.code.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.stream_id.encode(out); - self.target.encode(out); - self.code.encode(out); - } -} - -impl ql_codec::Decode for StreamReset { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - stream_id: reader.decode()?, - target: reader.decode()?, - code: reader.decode()?, - }) - } -} - -/// selects which stream lane a [`StreamReset`] applies to -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -#[repr(u8)] -pub enum ResetTarget { - /// reset the lane sent by the stream origin - Origin = 1, - /// reset the lane sent back toward the stream origin - Return = 2, - /// reset both stream lanes - Both = 3, -} - -impl Encode for ResetTarget { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u8(*self as u8); - } -} - -impl TryFrom for ResetTarget { - type Error = ql_codec::Error; - - fn try_from(value: u8) -> Result { - match value { - 1 => Ok(Self::Origin), - 2 => Ok(Self::Return), - 3 => Ok(Self::Both), - _ => Err(ql_codec::Error::InvalidDiscriminant), - } +ql_codec::codec_struct! { + /// aborts one or both lanes of a stream with a reset code + /// + /// stream origin is the peer that opened the stream + /// origin lane carries bytes sent by the stream origin + /// return lane carries bytes sent back toward the stream origin + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct StreamReset { + pub stream_id: StreamId, + pub target: ResetTarget, + pub code: ResetCode, } } -impl ql_codec::Decode for ResetTarget { - fn decode(reader: &mut ql_codec::Reader) -> Result { - reader.decode::()?.try_into() +ql_codec::codec_enum! { + /// selects which stream lane a [`StreamReset`] applies to + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub enum ResetTarget { + /// reset the lane sent by the stream origin + Origin = 1, + /// reset the lane sent back toward the stream origin + Return = 2, + /// reset both stream lanes + Both = 3, } } diff --git a/ql-wire/src/encrypted/stream_window.rs b/ql-wire/src/encrypted/stream_window.rs index 2910f022..67ecba25 100644 --- a/ql-wire/src/encrypted/stream_window.rs +++ b/ql-wire/src/encrypted/stream_window.rs @@ -1,30 +1,12 @@ -use ql_codec::{ByteSlice, Encode, Error, Varint}; +use ql_codec::Varint; use super::StreamId; -/// advertises the highest byte offset the peer may send on a stream. -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct StreamWindow { - pub stream_id: StreamId, - pub maximum_offset: Varint, -} - -impl Encode for StreamWindow { - fn encoded_len(&self) -> usize { - self.stream_id.encoded_len() + self.maximum_offset.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.stream_id.encode(out); - self.maximum_offset.encode(out); - } -} - -impl ql_codec::Decode for StreamWindow { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - stream_id: reader.decode()?, - maximum_offset: reader.decode()?, - }) +ql_codec::codec_struct! { + /// advertises the highest byte offset the peer may send on a stream. + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct StreamWindow { + pub stream_id: StreamId, + pub maximum_offset: Varint, } } diff --git a/ql-wire/src/handshake/id.rs b/ql-wire/src/handshake/id.rs index f0a9dfdb..82c5a644 100644 --- a/ql-wire/src/handshake/id.rs +++ b/ql-wire/src/handshake/id.rs @@ -1,25 +1,9 @@ -use ql_codec::{ByteSlice, Encode}; - -#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] -#[repr(transparent)] -pub struct HandshakeId(pub u32); +ql_codec::codec_newtype! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] + #[repr(transparent)] + pub struct HandshakeId(pub u32); +} impl HandshakeId { pub const WIRE_SIZE: usize = size_of::(); } - -impl ql_codec::Decode for HandshakeId { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -impl Encode for HandshakeId { - fn encoded_len(&self) -> usize { - self.0.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} diff --git a/ql-wire/src/handshake/ik.rs b/ql-wire/src/handshake/ik.rs index 6adeb8bf..6de2999a 100644 --- a/ql-wire/src/handshake/ik.rs +++ b/ql-wire/src/handshake/ik.rs @@ -67,38 +67,13 @@ impl ql_codec::Decode for Ik1 { } } -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct Ik2 { - pub handshake_id: HandshakeId, - pub transport_params: TransportParams, - pub ekem_ciphertext: MlKemCiphertext, - pub skem_ciphertext: EncryptedMlKemCiphertext, -} - -impl Encode for Ik2 { - fn encoded_len(&self) -> usize { - self.handshake_id.encoded_len() - + self.transport_params.encoded_len() - + self.ekem_ciphertext.encoded_len() - + self.skem_ciphertext.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.handshake_id.encode(out); - self.transport_params.encode(out); - self.ekem_ciphertext.encode(out); - self.skem_ciphertext.encode(out); - } -} - -impl ql_codec::Decode for Ik2 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - handshake_id: reader.decode()?, - transport_params: reader.decode()?, - ekem_ciphertext: reader.decode()?, - skem_ciphertext: reader.decode()?, - }) +ql_codec::codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Ik2 { + pub handshake_id: HandshakeId, + pub transport_params: TransportParams, + pub ekem_ciphertext: MlKemCiphertext, + pub skem_ciphertext: EncryptedMlKemCiphertext, } } diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index e1803723..f8191dec 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -24,36 +24,22 @@ const PROTOCOL_IK: &[u8] = b"ql-wire:pq-ik:v1"; const PROTOCOL_KK: &[u8] = b"ql-wire:pq-kk:v1"; const HANDSHAKE_PREAMBLE_DOMAIN: &[u8] = b"ql-wire:handshake-preamble:v1"; -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct EphemeralPublicKey { - pub mlkem_public_key: MlKemPublicKey, +ql_codec::codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct EphemeralPublicKey { + pub mlkem_public_key: MlKemPublicKey, + } } impl EphemeralPublicKey { pub const WIRE_SIZE: usize = MlKemPublicKey::SIZE; } -impl Encode for EphemeralPublicKey { - fn encoded_len(&self) -> usize { - self.mlkem_public_key.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.mlkem_public_key.encode(out); - } -} - -impl ql_codec::Decode for EphemeralPublicKey { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - mlkem_public_key: reader.decode()?, - }) - } +ql_codec::codec_newtype! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct EncryptedMlKemCiphertext(pub Box<[u8; Self::WIRE_SIZE]>); } -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct EncryptedMlKemCiphertext(pub Box<[u8; Self::WIRE_SIZE]>); - impl EncryptedMlKemCiphertext { pub const WIRE_SIZE: usize = MlKemCiphertext::SIZE + ENCRYPTED_MESSAGE_AUTH_SIZE; @@ -66,22 +52,6 @@ impl EncryptedMlKemCiphertext { } } -impl Encode for EncryptedMlKemCiphertext { - fn encoded_len(&self) -> usize { - self.0.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -impl ql_codec::Decode for EncryptedMlKemCiphertext { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - #[derive(Debug, Clone, PartialEq, Eq)] pub struct EncryptedPeerBundle(pub Box<[u8]>); diff --git a/ql-wire/src/handshake/pairing.rs b/ql-wire/src/handshake/pairing.rs index ea305a0b..ce966908 100644 --- a/ql-wire/src/handshake/pairing.rs +++ b/ql-wire/src/handshake/pairing.rs @@ -5,9 +5,15 @@ use crate::QlCrypto; const PAIRING_ID_DOMAIN: &[u8] = b"ql-wire:pairing-id:v1"; const PAIRING_PSK_DOMAIN: &[u8] = b"ql-wire:pairing-psk:v1"; -array_wrapper!(PairingToken, 16); +ql_codec::codec_newtype! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + #[repr(transparent)] + pub struct PairingToken(pub [u8; Self::SIZE]); +} impl PairingToken { + pub const SIZE: usize = 16; + pub fn id(&self, crypto: &impl QlCrypto) -> PairingId { let hash = crypto.sha256(&[PAIRING_ID_DOMAIN, &self.0]); let mut id = [0u8; PairingId::SIZE]; @@ -29,7 +35,15 @@ impl Display for PairingToken { } } -array_wrapper!(PairingId, 16); +ql_codec::codec_newtype! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + #[repr(transparent)] + pub struct PairingId(pub [u8; Self::SIZE]); +} + +impl PairingId { + pub const SIZE: usize = 16; +} impl Display for PairingId { fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result { diff --git a/ql-wire/src/handshake/transport_params.rs b/ql-wire/src/handshake/transport_params.rs index 9d485b6b..c1d9deb9 100644 --- a/ql-wire/src/handshake/transport_params.rs +++ b/ql-wire/src/handshake/transport_params.rs @@ -1,26 +1,16 @@ -use ql_codec::{ByteSlice, Encode}; - -/// Session parameters advertised in the handshake -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct TransportParams { - /// Initial per-stream receive credit granted to the remote peer - pub initial_stream_receive_window: u32, +ql_codec::codec_struct! { + /// Session parameters advertised in the handshake + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub struct TransportParams { + /// Initial per-stream receive credit granted to the remote peer + pub initial_stream_receive_window: u32, + } } impl TransportParams { pub const WIRE_SIZE: usize = size_of::(); } -impl Encode for TransportParams { - fn encoded_len(&self) -> usize { - self.initial_stream_receive_window.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.initial_stream_receive_window.encode(out); - } -} - impl Default for TransportParams { fn default() -> Self { Self { @@ -28,11 +18,3 @@ impl Default for TransportParams { } } } - -impl ql_codec::Decode for TransportParams { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - initial_stream_receive_window: reader.decode()?, - }) - } -} diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs index 564f9f96..00f8c7bf 100644 --- a/ql-wire/src/handshake/xx.rs +++ b/ql-wire/src/handshake/xx.rs @@ -1,4 +1,3 @@ -use ql_codec::{ByteSlice, Encode}; use ql_common::QID; use super::{ @@ -15,130 +14,40 @@ use crate::{ const PROTOCOL_XX: &[u8] = b"ql-wire:pq-xx:v1"; -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct Xx1 { - pub handshake_id: HandshakeId, - pub pairing_id: PairingId, - pub transport_params: TransportParams, - pub ephemeral: EphemeralPublicKey, -} - -impl ql_codec::Decode for Xx1 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - handshake_id: reader.decode()?, - pairing_id: reader.decode()?, - transport_params: reader.decode()?, - ephemeral: reader.decode()?, - }) - } -} - -impl Encode for Xx1 { - fn encoded_len(&self) -> usize { - self.handshake_id.encoded_len() - + self.pairing_id.encoded_len() - + self.transport_params.encoded_len() - + self.ephemeral.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.handshake_id.encode(out); - self.pairing_id.encode(out); - self.transport_params.encode(out); - self.ephemeral.encode(out); - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct Xx2 { - pub handshake_id: HandshakeId, - pub transport_params: TransportParams, - pub ekem_ciphertext: MlKemCiphertext, - pub static_bundle: EncryptedPeerBundle, -} - -impl ql_codec::Decode for Xx2 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - handshake_id: reader.decode()?, - transport_params: reader.decode()?, - ekem_ciphertext: reader.decode()?, - static_bundle: reader.decode()?, - }) - } -} - -impl Encode for Xx2 { - fn encoded_len(&self) -> usize { - self.handshake_id.encoded_len() - + self.transport_params.encoded_len() - + self.ekem_ciphertext.encoded_len() - + self.static_bundle.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.handshake_id.encode(out); - self.transport_params.encode(out); - self.ekem_ciphertext.encode(out); - self.static_bundle.encode(out); - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct Xx3 { - pub handshake_id: HandshakeId, - pub skem_ciphertext: EncryptedMlKemCiphertext, - pub static_bundle: EncryptedPeerBundle, -} - -impl ql_codec::Decode for Xx3 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - handshake_id: reader.decode()?, - skem_ciphertext: reader.decode()?, - static_bundle: reader.decode()?, - }) +ql_codec::codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Xx1 { + pub handshake_id: HandshakeId, + pub pairing_id: PairingId, + pub transport_params: TransportParams, + pub ephemeral: EphemeralPublicKey, } } -impl Encode for Xx3 { - fn encoded_len(&self) -> usize { - self.handshake_id.encoded_len() - + self.skem_ciphertext.encoded_len() - + self.static_bundle.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.handshake_id.encode(out); - self.skem_ciphertext.encode(out); - self.static_bundle.encode(out); +ql_codec::codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Xx2 { + pub handshake_id: HandshakeId, + pub transport_params: TransportParams, + pub ekem_ciphertext: MlKemCiphertext, + pub static_bundle: EncryptedPeerBundle, } } -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct Xx4 { - pub handshake_id: HandshakeId, - pub skem_ciphertext: EncryptedMlKemCiphertext, -} - -impl ql_codec::Decode for Xx4 { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - handshake_id: reader.decode()?, - skem_ciphertext: reader.decode()?, - }) +ql_codec::codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Xx3 { + pub handshake_id: HandshakeId, + pub skem_ciphertext: EncryptedMlKemCiphertext, + pub static_bundle: EncryptedPeerBundle, } } -impl Encode for Xx4 { - fn encoded_len(&self) -> usize { - self.handshake_id.encoded_len() + self.skem_ciphertext.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.handshake_id.encode(out); - self.skem_ciphertext.encode(out); +ql_codec::codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Xx4 { + pub handshake_id: HandshakeId, + pub skem_ciphertext: EncryptedMlKemCiphertext, } } diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index a89e357b..6399193d 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -1,44 +1,28 @@ use ::bytes::BufMut; -use ql_codec::{ByteSlice, Encode, Error}; +use ql_codec::Encode; use ql_common::QID; use crate::QL_WIRE_VERSION; -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct RouteHeader { - pub sender: QID, - pub recipient: QID, +ql_codec::codec_struct! { + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub struct RouteHeader { + pub sender: QID, + pub recipient: QID, + } } impl RouteHeader { pub const WIRE_SIZE: usize = QID::SIZE * 2; } -impl Encode for RouteHeader { - fn encoded_len(&self) -> usize { - self.sender.encoded_len() + self.recipient.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.sender.encode(out); - self.recipient.encode(out); - } -} - -impl ql_codec::Decode for RouteHeader { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - sender: reader.decode()?, - recipient: reader.decode()?, - }) +ql_codec::codec_struct! { + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub struct SessionHeader { + pub seq: RecordSeq, } } -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct SessionHeader { - pub seq: RecordSeq, -} - #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] #[repr(transparent)] pub struct RecordSeq(pub u64); @@ -71,21 +55,3 @@ impl SessionHeader { aad } } - -impl Encode for SessionHeader { - fn encoded_len(&self) -> usize { - self.seq.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.seq.encode(out); - } -} - -impl ql_codec::Decode for SessionHeader { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - seq: reader.decode()?, - }) - } -} diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index d7e5ae50..5db2e524 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -1,15 +1,16 @@ -use ql_codec::{ByteSlice, Encode}; use ql_common::QID; use crate::{derive_qid, Error, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash}; -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct PeerBundle { - pub version: u16, - pub qid: QID, - pub capabilities: u32, - pub mlkem_public_key: MlKemPublicKey, - pub name: String, +ql_codec::codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct PeerBundle { + pub version: u16, + pub qid: QID, + pub capabilities: u32, + pub mlkem_public_key: MlKemPublicKey, + pub name: String, + } } impl PeerBundle { @@ -23,45 +24,17 @@ impl PeerBundle { } } -impl Encode for PeerBundle { - fn encoded_len(&self) -> usize { - self.version.encoded_len() - + self.qid.encoded_len() - + self.capabilities.encoded_len() - + self.mlkem_public_key.encoded_len() - + self.name.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.version.encode(out); - self.qid.encode(out); - self.capabilities.encode(out); - self.mlkem_public_key.encode(out); - self.name.encode(out); - } -} - -impl ql_codec::Decode for PeerBundle { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - version: reader.decode()?, - qid: reader.decode()?, - capabilities: reader.decode()?, - mlkem_public_key: reader.decode()?, - name: reader.decode()?, - }) +ql_codec::codec_struct! { + #[derive(Debug, Clone)] + pub struct QlIdentity { + pub qid: QID, + pub mlkem_private_key: MlKemPrivateKey, + pub mlkem_public_key: MlKemPublicKey, + pub capabilities: u32, + pub name: String, } } -#[derive(Debug, Clone)] -pub struct QlIdentity { - pub qid: QID, - pub mlkem_private_key: MlKemPrivateKey, - pub mlkem_public_key: MlKemPublicKey, - pub capabilities: u32, - pub name: String, -} - impl QlIdentity { pub fn new( crypto: &impl QlHash, @@ -90,36 +63,6 @@ impl QlIdentity { } } -impl Encode for QlIdentity { - fn encoded_len(&self) -> usize { - QID::SIZE - + MlKemPrivateKey::SIZE - + MlKemPublicKey::SIZE - + size_of::() - + self.name.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.qid.encode(out); - self.mlkem_private_key.as_bytes().encode(out); - self.mlkem_public_key.encode(out); - self.capabilities.encode(out); - self.name.encode(out); - } -} - -impl ql_codec::Decode for QlIdentity { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - qid: reader.decode()?, - mlkem_private_key: MlKemPrivateKey::new(reader.decode()?), - mlkem_public_key: reader.decode()?, - capabilities: reader.decode()?, - name: reader.decode()?, - }) - } -} - pub fn generate_identity(crypto: &impl QlCrypto, name: impl Into) -> QlIdentity { let MlKemKeyPair { private, public } = crypto.mlkem_generate_keypair(); QlIdentity::new(crypto, private, public, name) diff --git a/ql-wire/src/lib.rs b/ql-wire/src/lib.rs index 55fd3b73..4b415562 100644 --- a/ql-wire/src/lib.rs +++ b/ql-wire/src/lib.rs @@ -4,8 +4,6 @@ #![allow(clippy::too_many_arguments)] -#[macro_use] -mod macros; mod crypto; mod encrypted; mod encrypted_message; diff --git a/ql-wire/src/macros.rs b/ql-wire/src/macros.rs deleted file mode 100644 index 7fcb2d29..00000000 --- a/ql-wire/src/macros.rs +++ /dev/null @@ -1,31 +0,0 @@ -macro_rules! array_wrapper { - ($name:ident, $size:expr) => { - #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] - #[repr(transparent)] - pub struct $name(pub [u8; Self::SIZE]); - - impl $name { - pub const SIZE: usize = $size; - - pub const fn as_bytes(&self) -> &[u8; Self::SIZE] { - &self.0 - } - } - - impl ql_codec::Encode for $name { - fn encoded_len(&self) -> usize { - self.0.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } - } - - impl ql_codec::Decode for $name { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } - } - }; -} diff --git a/ql-wire/src/pq.rs b/ql-wire/src/pq.rs index 8f5cc87b..c5c0c434 100644 --- a/ql-wire/src/pq.rs +++ b/ql-wire/src/pq.rs @@ -1,5 +1,3 @@ -use ql_codec::{ByteSlice, Encode}; - pub const ML_KEM_SUITE_TAG: &[u8] = b"ml-kem-1024"; // ql-wire fixes the protocol to ML-KEM-1024 on the wire, but the host @@ -10,8 +8,10 @@ const ML_KEM_1024_PUBLIC_KEY_SIZE: usize = 1568; const ML_KEM_1024_PRIVATE_KEY_SIZE: usize = 3168; const ML_KEM_1024_CIPHERTEXT_SIZE: usize = 1568; -#[derive(Debug, Clone, PartialEq, Eq, Hash)] -pub struct SessionKey(pub [u8; Self::SIZE]); +ql_codec::codec_newtype! { + #[derive(Debug, Clone, PartialEq, Eq, Hash)] + pub struct SessionKey(pub [u8; Self::SIZE]); +} impl SessionKey { pub const SIZE: usize = ML_KEM_1024_SHARED_SECRET_SIZE; @@ -33,25 +33,11 @@ impl Drop for SessionKey { } } -impl Encode for SessionKey { - fn encoded_len(&self) -> usize { - self.0.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } +ql_codec::codec_newtype! { + #[derive(Debug, Clone, PartialEq, Eq, Hash)] + pub struct MlKemPublicKey(Box<[u8; MlKemPublicKey::SIZE]>); } -impl ql_codec::Decode for SessionKey { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -#[derive(Debug, Clone, PartialEq, Eq, Hash)] -pub struct MlKemPublicKey(Box<[u8; MlKemPublicKey::SIZE]>); - impl MlKemPublicKey { pub const SIZE: usize = ML_KEM_1024_PUBLIC_KEY_SIZE; @@ -70,25 +56,11 @@ impl Drop for MlKemPublicKey { } } -impl ql_codec::Decode for MlKemPublicKey { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } +ql_codec::codec_newtype! { + #[derive(Debug, Clone, PartialEq, Eq, Hash)] + pub struct MlKemPrivateKey(Box<[u8; MlKemPrivateKey::SIZE]>); } -impl Encode for MlKemPublicKey { - fn encoded_len(&self) -> usize { - self.0.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - -#[derive(Debug, Clone, PartialEq, Eq, Hash)] -pub struct MlKemPrivateKey(Box<[u8; MlKemPrivateKey::SIZE]>); - impl MlKemPrivateKey { pub const SIZE: usize = ML_KEM_1024_PRIVATE_KEY_SIZE; @@ -107,8 +79,10 @@ impl Drop for MlKemPrivateKey { } } -#[derive(Debug, Clone, PartialEq, Eq, Hash)] -pub struct MlKemCiphertext(Box<[u8; MlKemCiphertext::SIZE]>); +ql_codec::codec_newtype! { + #[derive(Debug, Clone, PartialEq, Eq, Hash)] + pub struct MlKemCiphertext(Box<[u8; MlKemCiphertext::SIZE]>); +} impl MlKemCiphertext { pub const SIZE: usize = ML_KEM_1024_CIPHERTEXT_SIZE; @@ -128,22 +102,6 @@ impl Drop for MlKemCiphertext { } } -impl ql_codec::Decode for MlKemCiphertext { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self(reader.decode()?)) - } -} - -impl Encode for MlKemCiphertext { - fn encoded_len(&self) -> usize { - self.0.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.0.encode(out); - } -} - #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct MlKemKeyPair { pub private: MlKemPrivateKey, diff --git a/ql-wire/src/record.rs b/ql-wire/src/record.rs index 6d5278c8..3e2945b6 100644 --- a/ql-wire/src/record.rs +++ b/ql-wire/src/record.rs @@ -1,4 +1,4 @@ -use ql_codec::{BufView, ByteSlice, Decode, Encode}; +use ql_codec::{ByteSlice, Decode, Encode}; use crate::{ encrypted_message::EncryptedMessage, @@ -30,11 +30,13 @@ where Ok((reader.decode()?, reader.decode()?)) } -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct RecordHeader { - pub version: u8, - pub route: RouteHeader, - pub record_type: RecordType, +ql_codec::codec_struct! { + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub struct RecordHeader { + pub version: u8, + pub route: RouteHeader, + pub record_type: RecordType, + } } impl RecordHeader { @@ -49,197 +51,33 @@ impl RecordHeader { } } -impl Decode for RecordHeader { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - version: reader.decode()?, - route: reader.decode()?, - record_type: reader.decode()?, - }) - } -} - -impl Encode for RecordHeader { - fn encoded_len(&self) -> usize { - self.version.encoded_len() + self.route.encoded_len() + self.record_type.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.version.encode(out); - self.route.encode(out); - self.record_type.encode(out); - } -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -#[repr(u8)] -pub enum RecordType { - Handshake = 1, - Session = 2, -} - -impl TryFrom for RecordType { - type Error = ql_codec::Error; - - fn try_from(value: u8) -> Result { - match value { - 1 => Ok(Self::Handshake), - 2 => Ok(Self::Session), - _ => Err(ql_codec::Error::InvalidDiscriminant), - } - } -} - -impl Decode for RecordType { - fn decode(reader: &mut ql_codec::Reader) -> Result { - reader.decode::()?.try_into() - } -} - -impl Encode for RecordType { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u8(*self as u8); - } -} - -#[derive(Debug, Clone, PartialEq, Eq)] -pub enum QlHandshakeRecord { - Ik1(Ik1), - Ik2(Ik2), - Kk1(Ik1), - Kk2(Ik2), - Xx1(Xx1), - Xx2(Xx2), - Xx3(Xx3), - Xx4(Xx4), -} - -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -#[repr(u8)] -pub enum HandshakeKind { - Ik1 = 1, - Ik2 = 2, - Kk1 = 3, - Kk2 = 4, - Xx1 = 5, - Xx2 = 6, - Xx3 = 7, - Xx4 = 8, -} - -impl TryFrom for HandshakeKind { - type Error = ql_codec::Error; - - fn try_from(value: u8) -> Result { - match value { - 1 => Ok(Self::Ik1), - 2 => Ok(Self::Ik2), - 3 => Ok(Self::Kk1), - 4 => Ok(Self::Kk2), - 5 => Ok(Self::Xx1), - 6 => Ok(Self::Xx2), - 7 => Ok(Self::Xx3), - 8 => Ok(Self::Xx4), - _ => Err(ql_codec::Error::InvalidDiscriminant), - } - } -} - -impl Decode for HandshakeKind { - fn decode(reader: &mut ql_codec::Reader) -> Result { - reader.decode::()?.try_into() - } -} - -impl Encode for HandshakeKind { - fn encoded_len(&self) -> usize { - size_of::() - } - - fn encode(&self, out: &mut W) { - out.put_u8(*self as u8); - } -} - -impl QlHandshakeRecord { - pub fn kind(&self) -> HandshakeKind { - match self { - Self::Ik1(_) => HandshakeKind::Ik1, - Self::Ik2(_) => HandshakeKind::Ik2, - Self::Kk1(_) => HandshakeKind::Kk1, - Self::Kk2(_) => HandshakeKind::Kk2, - Self::Xx1(_) => HandshakeKind::Xx1, - Self::Xx2(_) => HandshakeKind::Xx2, - Self::Xx3(_) => HandshakeKind::Xx3, - Self::Xx4(_) => HandshakeKind::Xx4, - } - } -} - -impl Encode for QlHandshakeRecord { - fn encoded_len(&self) -> usize { - self.kind().encoded_len() - + match self { - Self::Ik1(message) => message.encoded_len(), - Self::Ik2(message) => message.encoded_len(), - Self::Kk1(message) => message.encoded_len(), - Self::Kk2(message) => message.encoded_len(), - Self::Xx1(message) => message.encoded_len(), - Self::Xx2(message) => message.encoded_len(), - Self::Xx3(message) => message.encoded_len(), - Self::Xx4(message) => message.encoded_len(), - } - } - - fn encode(&self, out: &mut W) { - self.kind().encode(out); - match self { - Self::Ik1(message) => message.encode(out), - Self::Ik2(message) => message.encode(out), - Self::Kk1(message) => message.encode(out), - Self::Kk2(message) => message.encode(out), - Self::Xx1(message) => message.encode(out), - Self::Xx2(message) => message.encode(out), - Self::Xx3(message) => message.encode(out), - Self::Xx4(message) => message.encode(out), - } +ql_codec::codec_enum! { + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub enum RecordType { + Handshake = 1, + Session = 2, } } -impl Decode for QlHandshakeRecord { - fn decode(reader: &mut ql_codec::Reader) -> Result { - let kind = reader.decode::()?; - match kind { - HandshakeKind::Ik1 => Ok(Self::Ik1(reader.decode()?)), - HandshakeKind::Ik2 => Ok(Self::Ik2(reader.decode()?)), - HandshakeKind::Kk1 => Ok(Self::Kk1(reader.decode()?)), - HandshakeKind::Kk2 => Ok(Self::Kk2(reader.decode()?)), - HandshakeKind::Xx1 => Ok(Self::Xx1(reader.decode()?)), - HandshakeKind::Xx2 => Ok(Self::Xx2(reader.decode()?)), - HandshakeKind::Xx3 => Ok(Self::Xx3(reader.decode()?)), - HandshakeKind::Xx4 => Ok(Self::Xx4(reader.decode()?)), - } +ql_codec::codec_enum! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub enum QlHandshakeRecord as HandshakeKind { + Ik1(Ik1) = 1, + Ik2(Ik2) = 2, + Kk1(Ik1) = 3, + Kk2(Ik2) = 4, + Xx1(Xx1) = 5, + Xx2(Xx2) = 6, + Xx3(Xx3) = 7, + Xx4(Xx4) = 8, } } -#[derive(Debug, Clone, PartialEq, Eq)] -pub struct QlSessionRecord { - pub header: SessionHeader, - pub payload: EncryptedMessage, -} - -impl Encode for QlSessionRecord { - fn encoded_len(&self) -> usize { - self.header.encoded_len() + self.payload.encoded_len() - } - - fn encode(&self, out: &mut W) { - self.header.encode(out); - self.payload.encode(out); +ql_codec::codec_struct! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct QlSessionRecord { + pub header: SessionHeader, + pub payload: EncryptedMessage, } } @@ -251,12 +89,3 @@ impl QlSessionRecord { } } } - -impl Decode for QlSessionRecord { - fn decode(reader: &mut ql_codec::Reader) -> Result { - Ok(Self { - header: reader.decode()?, - payload: reader.decode()?, - }) - } -} From edaf6160e657489d8e01bccd30d08d5e50869ef3 Mon Sep 17 00:00:00 2001 From: Alex's Bot Date: Wed, 22 Jul 2026 18:07:00 +0000 Subject: [PATCH 54/59] ql: consolidate the wire size constants MIN_WIRE_SIZE was a sum of maximums subtracted from the remaining record capacity, so the name had it backwards, and bounding the length prefixes as usize made it 41 on the host but 31 on the 32-bit target, giving the simulator and the device different framing. Three other WIRE_SIZE constants lost their last caller when `encoded_len` started summing fields; the survivors are the ones no value is in hand for. Co-Authored-By: Claude Opus 4.8 (1M context) --- ql-codec/src/lib.rs | 5 ----- ql-fsm/src/session/mod.rs | 2 +- ql-fsm/src/session/tests.rs | 6 +++--- ql-wire/src/encrypted/builder.rs | 7 ++++--- ql-wire/src/encrypted/stream_data.rs | 16 +++++++++++----- ql-wire/src/handshake/id.rs | 4 ---- ql-wire/src/handshake/mod.rs | 14 +++++--------- ql-wire/src/handshake/transport_params.rs | 4 ---- ql-wire/src/header.rs | 6 +----- ql-wire/src/record.rs | 5 +++-- ql-wire/src/tests.rs | 2 +- 11 files changed, 29 insertions(+), 42 deletions(-) diff --git a/ql-codec/src/lib.rs b/ql-codec/src/lib.rs index e522c7fc..53e9966a 100644 --- a/ql-codec/src/lib.rs +++ b/ql-codec/src/lib.rs @@ -41,11 +41,6 @@ pub trait Decode: Sized { #[macro_export] macro_rules! varint_wrapper { ($name:ty, $inner:ty) => { - impl $name { - pub const MAX_ENCODED_LEN: usize = - <$inner as ql_codec::varint::Primitive>::MAX_ENCODED_LEN; - } - impl ql_codec::Encode for $name { fn encoded_len(&self) -> usize { ql_codec::varint::encoded_len::<$inner>(self.0) diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index 637edecd..e7f6bf6d 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -523,7 +523,7 @@ impl SessionFsm { builder: &mut SessionRecordBuilder, outbound: &mut TrackedRecord, ) { - const OVERHEAD: usize = 1 + StreamData::>::MIN_WIRE_SIZE; + const OVERHEAD: usize = 1 + StreamData::>::MAX_WIRE_OVERHEAD; let len = self.state.streams.len(); if len == 0 { diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index d1649bf4..d8595144 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -234,7 +234,7 @@ fn tracked_record_count_is_bounded() { SessionConfig { record_max_size: SessionRecordBuilder::MIN_CAPACITY + 1 - + StreamData::>::MIN_WIRE_SIZE + + StreamData::>::MAX_WIRE_OVERHEAD + 1, stream_send_buffer_size: PAYLOAD_LEN, initial_peer_stream_receive_window: PAYLOAD_LEN as u32, @@ -267,7 +267,7 @@ fn lost_record_on_one_stream_does_not_block_another_stream() { SessionConfig { record_max_size: SessionRecordBuilder::MIN_CAPACITY + 1 // discriminator byte - + StreamData::>::MIN_WIRE_SIZE + + StreamData::>::MAX_WIRE_OVERHEAD + PAYLOAD_LEN, ..SessionConfig::default() }, @@ -948,7 +948,7 @@ fn sparse_out_of_order_ack_ranges_page_and_quiesce() { local_parity: StreamParity::Even, record_max_size: SessionRecordBuilder::MIN_CAPACITY + 1 // discriminator byte - + StreamData::>::MIN_WIRE_SIZE + + StreamData::>::MAX_WIRE_OVERHEAD + 10, // keeps stream-data records tiny enough to force ACK paging ack_delay: Duration::from_millis(5), retransmit_timeout: Duration::from_millis(25), diff --git a/ql-wire/src/encrypted/builder.rs b/ql-wire/src/encrypted/builder.rs index 210b113f..82ccef0b 100644 --- a/ql-wire/src/encrypted/builder.rs +++ b/ql-wire/src/encrypted/builder.rs @@ -1,5 +1,5 @@ use bytes::BufMut; -use ql_codec::{BufView, Encode}; +use ql_codec::{BufView, Encode, Varint}; use super::{RecordAck, SessionClose, SessionFrame, StreamData, StreamReset, StreamWindow}; use crate::{ @@ -15,8 +15,9 @@ pub struct SessionRecordBuilder { } impl SessionRecordBuilder { - pub const MIN_CAPACITY: usize = - RecordHeader::WIRE_SIZE + RecordSeq::MAX_ENCODED_LEN + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; + pub const MIN_CAPACITY: usize = RecordHeader::WIRE_SIZE + + Varint::::MAX_ENCODED_LEN + + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; pub fn new(seq: RecordSeq, max_capacity: usize) -> Self { let prefix_len = diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index 9bb90f7c..101cb32d 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -14,11 +14,17 @@ pub struct StreamData { } impl StreamData { - pub const MIN_WIRE_SIZE: usize = StreamId::MAX_ENCODED_LEN - + Varint::::MAX_ENCODED_LEN - + size_of::() - + Varint::::MAX_ENCODED_LEN - + Varint::::MAX_ENCODED_LEN; + /// Largest framing overhead of a stream data frame, excluding the header and payload bytes + /// that its two length prefixes measure. + /// + /// The terms follow the field order in `encode`. Lengths are bounded as `u64` rather than + /// `usize` so the figure does not shrink on a 32-bit target, where framing would otherwise + /// differ from the host. + pub const MAX_WIRE_OVERHEAD: usize = Varint::::MAX_ENCODED_LEN // stream id + + Varint::::MAX_ENCODED_LEN // offset + + size_of::() // flags + + Varint::::MAX_ENCODED_LEN // header length + + Varint::::MAX_ENCODED_LEN; // payload length } impl Decode for StreamData { diff --git a/ql-wire/src/handshake/id.rs b/ql-wire/src/handshake/id.rs index 82c5a644..200d3420 100644 --- a/ql-wire/src/handshake/id.rs +++ b/ql-wire/src/handshake/id.rs @@ -3,7 +3,3 @@ ql_codec::codec_newtype! { #[repr(transparent)] pub struct HandshakeId(pub u32); } - -impl HandshakeId { - pub const WIRE_SIZE: usize = size_of::(); -} diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index f8191dec..0385f275 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -31,23 +31,19 @@ ql_codec::codec_struct! { } } -impl EphemeralPublicKey { - pub const WIRE_SIZE: usize = MlKemPublicKey::SIZE; -} - ql_codec::codec_newtype! { #[derive(Debug, Clone, PartialEq, Eq)] - pub struct EncryptedMlKemCiphertext(pub Box<[u8; Self::WIRE_SIZE]>); + pub struct EncryptedMlKemCiphertext(pub Box<[u8; Self::SIZE]>); } impl EncryptedMlKemCiphertext { - pub const WIRE_SIZE: usize = MlKemCiphertext::SIZE + ENCRYPTED_MESSAGE_AUTH_SIZE; + pub const SIZE: usize = MlKemCiphertext::SIZE + ENCRYPTED_MESSAGE_AUTH_SIZE; - pub fn new(data: Box<[u8; Self::WIRE_SIZE]>) -> Self { + pub fn new(data: Box<[u8; Self::SIZE]>) -> Self { Self(data) } - pub fn as_bytes(&self) -> &[u8; Self::WIRE_SIZE] { + pub fn as_bytes(&self) -> &[u8; Self::SIZE] { self.0.as_ref() } } @@ -368,7 +364,7 @@ fn encrypt_mlkem_ciphertext( ciphertext: &MlKemCiphertext, ) -> Result { let encrypted = symmetric.encrypt_and_hash(crypto, ciphertext.as_bytes())?; - let out: Box<[u8; EncryptedMlKemCiphertext::WIRE_SIZE]> = + let out: Box<[u8; EncryptedMlKemCiphertext::SIZE]> = encrypted.try_into().map_err(|_| Error::InvalidState)?; Ok(EncryptedMlKemCiphertext::new(out)) } diff --git a/ql-wire/src/handshake/transport_params.rs b/ql-wire/src/handshake/transport_params.rs index c1d9deb9..eb2e0a6a 100644 --- a/ql-wire/src/handshake/transport_params.rs +++ b/ql-wire/src/handshake/transport_params.rs @@ -7,10 +7,6 @@ ql_codec::codec_struct! { } } -impl TransportParams { - pub const WIRE_SIZE: usize = size_of::(); -} - impl Default for TransportParams { fn default() -> Self { Self { diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index 6399193d..3e8c8fff 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -12,10 +12,6 @@ ql_codec::codec_struct! { } } -impl RouteHeader { - pub const WIRE_SIZE: usize = QID::SIZE * 2; -} - ql_codec::codec_struct! { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct SessionHeader { @@ -43,7 +39,7 @@ impl SessionHeader { let aad_len = Self::AAD_DOMAIN.len() + size_of::() + size_of::() - + RouteHeader::WIRE_SIZE + + route.encoded_len() + self.seq.encoded_len(); let mut aad = Vec::with_capacity(aad_len); aad.put_slice(Self::AAD_DOMAIN); diff --git a/ql-wire/src/record.rs b/ql-wire/src/record.rs index 3e2945b6..ceb5f1e6 100644 --- a/ql-wire/src/record.rs +++ b/ql-wire/src/record.rs @@ -1,4 +1,5 @@ use ql_codec::{ByteSlice, Decode, Encode}; +use ql_common::QID; use crate::{ encrypted_message::EncryptedMessage, @@ -16,7 +17,7 @@ where } pub fn encode_record_vec(header: RecordHeader, body: &T) -> Vec { - let mut out = Vec::with_capacity(RecordHeader::WIRE_SIZE + body.encoded_len()); + let mut out = Vec::with_capacity(header.encoded_len() + body.encoded_len()); encode_record(&mut out, header, body); out } @@ -40,7 +41,7 @@ ql_codec::codec_struct! { } impl RecordHeader { - pub const WIRE_SIZE: usize = size_of::() + RouteHeader::WIRE_SIZE + size_of::(); + pub const WIRE_SIZE: usize = size_of::() + QID::SIZE * 2 + size_of::(); pub fn new(route: RouteHeader, record_type: RecordType) -> Self { Self { diff --git a/ql-wire/src/tests.rs b/ql-wire/src/tests.rs index 88444f4a..680dc141 100644 --- a/ql-wire/src/tests.rs +++ b/ql-wire/src/tests.rs @@ -914,7 +914,7 @@ fn mlkem_ciphertext() -> MlKemCiphertext { } fn encrypted_mlkem_ciphertext() -> EncryptedMlKemCiphertext { - EncryptedMlKemCiphertext::new(Box::new([3u8; EncryptedMlKemCiphertext::WIRE_SIZE])) + EncryptedMlKemCiphertext::new(Box::new([3u8; EncryptedMlKemCiphertext::SIZE])) } fn encrypted_peer_bundle() -> EncryptedPeerBundle { From 0219878634d39137e951c3b0fbfcce17a0a75d95 Mon Sep 17 00:00:00 2001 From: Alex's Bot Date: Wed, 22 Jul 2026 18:13:46 +0000 Subject: [PATCH 55/59] ql: length-prefix byte fields as u32 A usize varint is capped at ten bytes on a 64-bit host and five on the 32-bit target, so a length a host emits can be rejected outright by the device's decoder. Only per-record lengths and counts move; offsets, sequence numbers and ack spans stay u64, since those accumulate. Encoded bytes are unchanged for anything under 4 GiB, as a varint does not depend on the width it was declared with. Co-Authored-By: Claude Opus 4.8 (1M context) --- ql-codec/src/codec.rs | 10 ++++++++-- ql-codec/src/reader.rs | 4 ++-- ql-wire/src/encrypted/ack.rs | 14 ++++++++++---- ql-wire/src/encrypted/stream_data.rs | 4 ++-- 4 files changed, 22 insertions(+), 10 deletions(-) diff --git a/ql-codec/src/codec.rs b/ql-codec/src/codec.rs index 7e037f1f..af4a0d77 100644 --- a/ql-codec/src/codec.rs +++ b/ql-codec/src/codec.rs @@ -207,9 +207,15 @@ impl> Decode for Option { } } +/// Length prefixes are `u32`, so a 64-bit host cannot emit one a 32-bit peer would reject: a +/// `usize` varint is capped at ten bytes on one target and five on the other. +fn len_prefix(len: usize) -> u32 { + u32::try_from(len).expect("byte field longer than u32::MAX") +} + pub fn encoded_len_bytes(bytes: &B) -> usize { let len = bytes.buf().remaining(); - varint::encoded_len(len) + len + varint::encoded_len(len_prefix(len)) + len } pub fn encode_bytes(bytes: &B, out: &mut W) @@ -217,7 +223,7 @@ where B: BufView + ?Sized, W: BufMut + ?Sized, { - varint::encode(bytes.buf().remaining(), out); + varint::encode(len_prefix(bytes.buf().remaining()), out); encode_bytes_raw(bytes, out); } diff --git a/ql-codec/src/reader.rs b/ql-codec/src/reader.rs index d534cc3a..35ed7033 100644 --- a/ql-codec/src/reader.rs +++ b/ql-codec/src/reader.rs @@ -38,8 +38,8 @@ impl Reader { } pub fn take_len_prefixed(&mut self) -> Result { - let len = self.decode_varint::()?; - self.take_n(len) + let len = self.decode_varint::()?; + self.take_n(len as usize) } #[inline] diff --git a/ql-wire/src/encrypted/ack.rs b/ql-wire/src/encrypted/ack.rs index 02eda89e..8be99b93 100644 --- a/ql-wire/src/encrypted/ack.rs +++ b/ql-wire/src/encrypted/ack.rs @@ -60,8 +60,14 @@ impl RecordAck { self.ranges().any(|range| range.contains(&seq)) } - fn block_count_len(block_count: usize) -> usize { - Varint(block_count).encoded_len() + /// The count is carried as a `u32`, so it encodes the same on a 32-bit target as on the host. + /// Blocks are two bytes each at minimum, so a record can never hold `u32::MAX` of them. + fn block_count(blocks: usize) -> u32 { + u32::try_from(blocks).expect("record ack blocks are bounded by the record size") + } + + fn block_count_len(blocks: usize) -> usize { + Varint(Self::block_count(blocks)).encoded_len() } } @@ -131,7 +137,7 @@ impl Encode for RecordAck { fn encode(&self, out: &mut W) { self.largest_acked.encode(out); - Varint(self.blocks.len()).encode(out); + Varint(Self::block_count(self.blocks.len())).encode(out); self.first_range_len.encode(out); for block in &self.blocks { block.gap.encode(out); @@ -143,7 +149,7 @@ impl Encode for RecordAck { impl ql_codec::Decode for RecordAck { fn decode(reader: &mut ql_codec::Reader) -> Result { let largest_acked = reader.decode()?; - let block_count = *reader.decode::>()?; + let block_count = *reader.decode::>()? as usize; let first_range_len = reader.decode()?; let mut blocks = Vec::with_capacity(block_count); for _ in 0..block_count { diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs index 101cb32d..1f211795 100644 --- a/ql-wire/src/encrypted/stream_data.rs +++ b/ql-wire/src/encrypted/stream_data.rs @@ -23,8 +23,8 @@ impl StreamData { pub const MAX_WIRE_OVERHEAD: usize = Varint::::MAX_ENCODED_LEN // stream id + Varint::::MAX_ENCODED_LEN // offset + size_of::() // flags - + Varint::::MAX_ENCODED_LEN // header length - + Varint::::MAX_ENCODED_LEN; // payload length + + Varint::::MAX_ENCODED_LEN // header length + + Varint::::MAX_ENCODED_LEN; // payload length } impl Decode for StreamData { From d3b35f4b4e916a8803b25b568b6cfab22ca014f8 Mon Sep 17 00:00:00 2001 From: Alex's Bot Date: Wed, 22 Jul 2026 18:33:40 +0000 Subject: [PATCH 56/59] ql-fsm: budget the stream header before polling The payload budget reserved room for the header's length prefix but not the header itself, so a stream opened with a header larger than the remaining record space built a frame that would not fit, and pushing it tripped the "builder has capacity" assert. Reserving stops once offset 0 is acked rather than once the header is sent, because until then a loss there is retransmitted and carries the header again. Co-Authored-By: Claude Opus 4.8 (1M context) --- ql-fsm/src/session/mod.rs | 15 ++++++++++----- ql-fsm/src/session/stream_tx.rs | 7 +++++++ ql-fsm/src/session/tests.rs | 22 ++++++++++++++++++++++ 3 files changed, 39 insertions(+), 5 deletions(-) diff --git a/ql-fsm/src/session/mod.rs b/ql-fsm/src/session/mod.rs index e7f6bf6d..a24da78f 100644 --- a/ql-fsm/src/session/mod.rs +++ b/ql-fsm/src/session/mod.rs @@ -543,6 +543,15 @@ impl SessionFsm { if matches!(stream.outbound_state, OutboundState::Closed) { continue; } + // The header shares the frame with the payload, so it has to come out of the same + // budget, and that budget is set before poll_transmit picks the range. + let header = match stream.role { + StreamRole::Initiator if stream.tx.can_send_header() => stream.header.as_deref(), + _ => None, + }; + let Some(max_payload) = max_payload.checked_sub(header.map_or(0, <[u8]>::len)) else { + continue; + }; let Some(candidate) = stream.tx.poll_transmit(max_payload, stream.peer_max_offset) else { continue; @@ -550,11 +559,7 @@ impl SessionFsm { let frame = StreamData { stream_id, offset: Varint(candidate.offset), - header: if matches!(stream.role, StreamRole::Initiator) && candidate.offset == 0 { - stream.header.as_deref() - } else { - None - }, + header: if candidate.offset == 0 { header } else { None }, fin: candidate.fin, bytes: stream.tx.ranged_bytes(candidate), }; diff --git a/ql-fsm/src/session/stream_tx.rs b/ql-fsm/src/session/stream_tx.rs index d929bb5f..ca29906e 100644 --- a/ql-fsm/src/session/stream_tx.rs +++ b/ql-fsm/src/session/stream_tx.rs @@ -152,6 +152,13 @@ impl StreamTx { self.base_offset + self.buffered_len as u64 } + /// Whether a frame can still start at offset 0, the only one that carries the stream header. + /// + /// Stays true until offset 0 is acked, because until then it can still be retransmitted. + pub fn can_send_header(&self) -> bool { + self.base_offset == 0 + } + pub fn is_empty(&self) -> bool { self.buffered_len == 0 && self.final_offset.is_none() } diff --git a/ql-fsm/src/session/tests.rs b/ql-fsm/src/session/tests.rs index d8595144..6599ce8d 100644 --- a/ql-fsm/src/session/tests.rs +++ b/ql-fsm/src/session/tests.rs @@ -1022,3 +1022,25 @@ fn sparse_out_of_order_ack_ranges_page_and_quiesce() { assert!(next_outbound(&mut sender, final_now).is_none()); assert!(next_outbound(&mut receiver, final_now).is_none()); } + +#[test] +fn stream_header_larger_than_the_record_budget_does_not_panic() { + let now = Instant::now(); + let record_max_size = SessionRecordBuilder::MIN_CAPACITY + 256; + let mut fsm = SessionFsm::new( + SessionConfig { + record_max_size, + ..SessionConfig::default() + }, + now, + ); + + // The header rides in the same frame as the payload, so one this large leaves no room. + let stream_id = fsm + .open_stream(Box::from(vec![7u8; 256]), |_| {}) + .unwrap() + .stream_id(); + assert_eq!(write_stream_bytes(&mut fsm, stream_id, b"payload"), 7); + + assert!(next_outbound(&mut fsm, now).is_none()); +} From c079ee8cf005a85a9f279fb28c968ddb89dc76d7 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 22 Jul 2026 15:12:15 -0400 Subject: [PATCH 57/59] ql-codec: single codec! macro --- ql-codec/src/codec.rs | 1 + ql-codec/src/macros.rs | 284 +++++----------------- ql-common/src/lib.rs | 2 +- ql-wire/src/crypto.rs | 2 +- ql-wire/src/encrypted/close.rs | 2 +- ql-wire/src/encrypted/mod.rs | 2 +- ql-wire/src/encrypted/stream_reset.rs | 4 +- ql-wire/src/encrypted/stream_window.rs | 2 +- ql-wire/src/handshake/id.rs | 2 +- ql-wire/src/handshake/ik.rs | 2 +- ql-wire/src/handshake/mod.rs | 4 +- ql-wire/src/handshake/pairing.rs | 4 +- ql-wire/src/handshake/transport_params.rs | 2 +- ql-wire/src/handshake/xx.rs | 8 +- ql-wire/src/header.rs | 4 +- ql-wire/src/identity.rs | 4 +- ql-wire/src/pq.rs | 8 +- ql-wire/src/record.rs | 8 +- 18 files changed, 90 insertions(+), 255 deletions(-) diff --git a/ql-codec/src/codec.rs b/ql-codec/src/codec.rs index af4a0d77..e89af01e 100644 --- a/ql-codec/src/codec.rs +++ b/ql-codec/src/codec.rs @@ -209,6 +209,7 @@ impl> Decode for Option { /// Length prefixes are `u32`, so a 64-bit host cannot emit one a 32-bit peer would reject: a /// `usize` varint is capped at ten bytes on one target and five on the other. +#[track_caller] fn len_prefix(len: usize) -> u32 { u32::try_from(len).expect("byte field longer than u32::MAX") } diff --git a/ql-codec/src/macros.rs b/ql-codec/src/macros.rs index ad9a61f8..50cee548 100644 --- a/ql-codec/src/macros.rs +++ b/ql-codec/src/macros.rs @@ -1,6 +1,28 @@ -/// Defines a newtype and encodes it exactly as the value it wraps. +/// generates `Encode` and `Decode` for newtypes, structs, and enums. +/// +/// newtypes encode as their wrapped value, structs in field order, and enums as a `u8` discriminant +/// followed by an optional payload; the lone generic parameter is treated as the reader's byte +/// container. +/// +/// `enum Frame as FrameKind` also generates `FrameKind` and `Frame::kind()`: +/// ``` +/// use ql_codec::Encode; +/// +/// ql_codec::codec! { +/// #[derive(Debug, PartialEq)] +/// pub enum Frame as FrameKind { +/// Ping = 1, +/// Close(u16) = 2, +/// } +/// } +/// +/// let frame = Frame::Close(7); +/// assert_eq!(frame.kind(), FrameKind::Close); +/// assert_eq!(frame.encode_vec(), [2, 7, 0]); +/// ``` #[macro_export] -macro_rules! codec_newtype { +macro_rules! codec { + // newtype ( $(#[$meta:meta])* $vis:vis struct $name:ident($field_vis:vis $inner:ty); @@ -24,17 +46,8 @@ macro_rules! codec_newtype { } } }; -} -/// Defines a struct and encodes its fields back to back in declaration order. -/// -/// A single generic parameter is taken to be the byte container, bound as -/// [`BufView`](crate::BufView) for encoding and [`ByteSlice`](crate::ByteSlice) for decoding. It -/// may only appear nested inside another codec type, and decoding ties it to the reader's own -/// container, because a borrowed field can only be taken from the reader it is read out of. -/// Anything whose fields do not map one to one onto the wire needs a hand-written impl. -#[macro_export] -macro_rules! codec_struct { + // struct with an optional byte container ( $(#[$meta:meta])* $vis:vis struct $name:ident $(<$bytes:ident>)? { @@ -56,79 +69,17 @@ macro_rules! codec_struct { } } - $crate::__codec_struct_decode!($name$(<$bytes>)?, $($field),*); + $crate::codec!(@struct_decode $name$(<$bytes>)?, $($field),*); }; -} -/// The one impl that cannot be written once: `Decode` always needs a byte container to name, so -/// a generic struct reuses its own parameter while a plain one introduces a fresh name. -#[macro_export] -#[doc(hidden)] -macro_rules! __codec_struct_decode { - ($name:ident<$bytes:ident>, $($field:ident),* $(,)?) => { - impl<$bytes: $crate::ByteSlice> $crate::Decode<$bytes> for $name<$bytes> { - fn decode(reader: &mut $crate::Reader<$bytes>) -> Result { - Ok(Self { $($field: reader.decode()?,)* }) - } - } - }; - - ($name:ident, $($field:ident),* $(,)?) => { - impl $crate::Decode for $name { - fn decode(reader: &mut $crate::Reader) -> Result { - Ok(Self { $($field: reader.decode()?,)* }) - } - } - }; -} - -/// Defines an enum carried on the wire as a `u8` discriminant. -/// -/// `enum Frame as FrameKind` additionally defines the discriminant enum and a `kind` accessor, -/// and encodes each variant's payload after the discriminant. -/// -/// ``` -/// use ql_codec::{Decode, Encode}; -/// -/// ql_codec::codec_enum! { -/// #[derive(Debug, Clone, Copy, PartialEq, Eq)] -/// pub enum CloseReason { -/// Done = 1, -/// Refused = 2, -/// } -/// } -/// -/// ql_codec::codec_enum! { -/// #[derive(Debug, PartialEq)] -/// pub enum Frame as FrameKind { -/// Ping = 1, -/// Close(CloseReason) = 3, -/// } -/// } -/// -/// // A unit variant is just its discriminant. -/// assert_eq!(Frame::Ping.encode_vec(), [1]); -/// -/// // A payload follows the discriminant, and `kind()` reads it back without the payload. -/// let frame = Frame::Close(CloseReason::Refused); -/// assert_eq!(frame.kind(), FrameKind::Close); -/// assert_eq!(frame.encode_vec(), [3, 2]); -/// assert_eq!(Frame::decode_bytes(&[3, 2][..]).unwrap(), frame); -/// -/// assert_eq!( -/// Frame::decode_bytes(&[9][..]), -/// Err(ql_codec::Error::InvalidDiscriminant), -/// ); -/// ``` -#[macro_export] -macro_rules! codec_enum { + // payload enum with a separate discriminant enum ( $(#[$meta:meta])* $vis:vis enum $name:ident $(<$bytes:ident>)? as $kind:ident { $($(#[$variant_meta:meta])* $variant:ident $(($payload:ty))? = $value:literal),* $(,)? } ) => { - $crate::codec_enum! { + $crate::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq)] $vis enum $kind { $($variant = $value,)* @@ -170,11 +121,12 @@ macro_rules! codec_enum { } } - $crate::__codec_enum_decode! { - $name$(<$bytes>)?, $kind, $($variant $(($payload))?),* + $crate::codec! { + @enum_decode $name$(<$bytes>)?, $kind, $($variant $(($payload))? = $value),* } }; + // u8 discriminant enum ( $(#[$meta:meta])* $vis:vis enum $name:ident { @@ -214,14 +166,30 @@ macro_rules! codec_enum { } } }; -} -/// The one impl that cannot be written once: `Decode` always needs a byte container to name, so -/// a generic enum reuses its own parameter while a plain one introduces a fresh name. -#[macro_export] -#[doc(hidden)] -macro_rules! __codec_enum_decode { - ($name:ident<$bytes:ident>, $kind:ident, $($variant:ident $(($payload:ty))?),* $(,)?) => { + // struct decoding with its byte container + (@struct_decode $name:ident<$bytes:ident>, $($field:ident),* $(,)?) => { + impl<$bytes: $crate::ByteSlice> $crate::Decode<$bytes> for $name<$bytes> { + fn decode(reader: &mut $crate::Reader<$bytes>) -> Result { + Ok(Self { $($field: reader.decode()?,)* }) + } + } + }; + + // struct decoding with a fresh byte container + (@struct_decode $name:ident, $($field:ident),* $(,)?) => { + impl $crate::Decode for $name { + fn decode(reader: &mut $crate::Reader) -> Result { + Ok(Self { $($field: reader.decode()?,)* }) + } + } + }; + + // payload enum decoding with its byte container + ( + @enum_decode $name:ident<$bytes:ident>, $kind:ident, + $($variant:ident $(($payload:ty))? = $value:literal),* $(,)? + ) => { impl<$bytes: $crate::ByteSlice> $crate::Decode<$bytes> for $name<$bytes> { fn decode(reader: &mut $crate::Reader<$bytes>) -> Result { Ok(match reader.decode::<$kind>()? { @@ -231,7 +199,11 @@ macro_rules! __codec_enum_decode { } }; - ($name:ident, $kind:ident, $($variant:ident $(($payload:ty))?),* $(,)?) => { + // payload enum decoding with a fresh byte container + ( + @enum_decode $name:ident, $kind:ident, + $($variant:ident $(($payload:ty))? = $value:literal),* $(,)? + ) => { impl $crate::Decode for $name { fn decode(reader: &mut $crate::Reader) -> Result { Ok(match reader.decode::<$kind>()? { @@ -241,141 +213,3 @@ macro_rules! __codec_enum_decode { } }; } - -#[cfg(test)] -mod tests { - use bytes::Bytes; - - use crate::{Decode, Encode, Error}; - - codec_newtype! { - /// A newtype carrying a fixed-size array. - #[derive(Debug, Clone, Copy, PartialEq, Eq)] - pub struct Tag(pub [u8; 4]); - } - - codec_enum! { - #[derive(Debug, Clone, Copy, PartialEq, Eq)] - pub enum Flavour { - Sweet = 1, - Salty = 3, - } - } - - codec_struct! { - /// Fields encode in declaration order. - #[derive(Debug, Clone, PartialEq, Eq)] - pub struct Plain { - pub tag: Tag, - pub flavour: Flavour, - pub name: String, - } - } - - codec_struct! { - #[derive(Debug, Clone, PartialEq, Eq)] - pub struct Wrapped { - pub tag: Tag, - pub body: Nested, - } - } - - #[derive(Debug, Clone, PartialEq, Eq)] - pub struct Nested(pub B); - - impl Encode for Nested { - fn encoded_len(&self) -> usize { - crate::encoded_len_bytes(&self.0) - } - - fn encode(&self, out: &mut W) { - crate::encode_bytes(&self.0, out); - } - } - - impl Decode for Nested { - fn decode(reader: &mut crate::Reader) -> Result { - Ok(Self(reader.take_len_prefixed()?)) - } - } - - codec_enum! { - #[derive(Debug, Clone, PartialEq, Eq)] - pub enum Frame as FrameKind { - Ping = 1, - Plain(Plain) = 2, - Body(Wrapped) = 3, - Pong = 4, - } - } - - codec_enum! { - #[derive(Debug, Clone, PartialEq, Eq)] - pub enum Message as MessageKind { - Empty = 1, - Plain(Plain) = 2, - } - } - - fn plain() -> Plain { - Plain { - tag: Tag([1, 2, 3, 4]), - flavour: Flavour::Salty, - name: "hello".to_owned(), - } - } - - #[test] - fn struct_fields_encode_in_declaration_order() { - let encoded = plain().encode_vec(); - assert_eq!(&encoded[..4], &[1, 2, 3, 4]); - assert_eq!(encoded[4], 3); - assert_eq!(Plain::decode_bytes(encoded.as_slice()).unwrap(), plain()); - - let tag = Tag([9; 4]); - assert_eq!(Tag::decode_bytes(tag.encode_vec().as_slice()).unwrap(), tag); - } - - #[test] - fn generic_struct_decodes_from_its_own_container() { - let value = Wrapped { - tag: Tag([5; 4]), - body: Nested(Bytes::from_static(b"body")), - }; - let encoded = Bytes::from(value.encode_vec()); - assert_eq!(Wrapped::::decode_bytes(encoded).unwrap(), value); - } - - #[test] - fn unknown_discriminant_is_rejected() { - assert_eq!(Flavour::Salty.encode_vec(), [3]); - assert_eq!( - Flavour::decode_bytes(&[2u8][..]), - Err(Error::InvalidDiscriminant) - ); - assert_eq!( - Message::decode_bytes(&[9u8][..]), - Err(Error::InvalidDiscriminant) - ); - } - - #[test] - fn payload_enum_writes_the_kind_first() { - let frame = Frame::::Plain(plain()); - assert_eq!(frame.kind(), FrameKind::Plain); - assert_eq!(frame.encode_vec()[0], 2); - assert_eq!(Frame::::Pong.encode_vec(), [4]); - - let encoded = Bytes::from(frame.encode_vec()); - assert_eq!(Frame::::decode_bytes(encoded).unwrap(), frame); - } - - #[test] - fn payload_enum_without_generics_round_trips() { - for message in [Message::Empty, Message::Plain(plain())] { - let encoded = message.encode_vec(); - assert_eq!(Message::decode_bytes(encoded.as_slice()).unwrap(), message); - } - assert_eq!(Message::Plain(plain()).kind(), MessageKind::Plain); - } -} diff --git a/ql-common/src/lib.rs b/ql-common/src/lib.rs index ef159103..1f661aa0 100644 --- a/ql-common/src/lib.rs +++ b/ql-common/src/lib.rs @@ -47,7 +47,7 @@ impl std::fmt::Display for ResetCode { ql_codec::varint_wrapper!(ResetCode, u64); -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] pub struct QID(pub [u8; Self::SIZE]); diff --git a/ql-wire/src/crypto.rs b/ql-wire/src/crypto.rs index 0bc11c3a..380e98de 100644 --- a/ql-wire/src/crypto.rs +++ b/ql-wire/src/crypto.rs @@ -3,7 +3,7 @@ use crate::{ ENCRYPTED_MESSAGE_AUTH_SIZE, }; -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] pub struct Nonce(pub [u8; Self::SIZE]); diff --git a/ql-wire/src/encrypted/close.rs b/ql-wire/src/encrypted/close.rs index 0b9f8d6d..1b76f3a9 100644 --- a/ql-wire/src/encrypted/close.rs +++ b/ql-wire/src/encrypted/close.rs @@ -1,4 +1,4 @@ -ql_codec::codec_struct! { +ql_codec::codec! { /// closes the whole session immediately with a reset code. #[derive(Debug, Clone, PartialEq, Eq)] pub struct SessionClose { diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs index 4f349033..efd1c918 100644 --- a/ql-wire/src/encrypted/mod.rs +++ b/ql-wire/src/encrypted/mod.rs @@ -19,7 +19,7 @@ pub use stream_data::*; pub use stream_reset::*; pub use stream_window::*; -ql_codec::codec_enum! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub enum SessionFrame as SessionFrameKind { // todo: do we need ping as explicit frame? diff --git a/ql-wire/src/encrypted/stream_reset.rs b/ql-wire/src/encrypted/stream_reset.rs index 30c342e7..711d786b 100644 --- a/ql-wire/src/encrypted/stream_reset.rs +++ b/ql-wire/src/encrypted/stream_reset.rs @@ -2,7 +2,7 @@ use ql_common::ResetCode; use super::StreamId; -ql_codec::codec_struct! { +ql_codec::codec! { /// aborts one or both lanes of a stream with a reset code /// /// stream origin is the peer that opened the stream @@ -16,7 +16,7 @@ ql_codec::codec_struct! { } } -ql_codec::codec_enum! { +ql_codec::codec! { /// selects which stream lane a [`StreamReset`] applies to #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum ResetTarget { diff --git a/ql-wire/src/encrypted/stream_window.rs b/ql-wire/src/encrypted/stream_window.rs index 67ecba25..49b61342 100644 --- a/ql-wire/src/encrypted/stream_window.rs +++ b/ql-wire/src/encrypted/stream_window.rs @@ -2,7 +2,7 @@ use ql_codec::Varint; use super::StreamId; -ql_codec::codec_struct! { +ql_codec::codec! { /// advertises the highest byte offset the peer may send on a stream. #[derive(Debug, Clone, PartialEq, Eq)] pub struct StreamWindow { diff --git a/ql-wire/src/handshake/id.rs b/ql-wire/src/handshake/id.rs index 200d3420..4f34602f 100644 --- a/ql-wire/src/handshake/id.rs +++ b/ql-wire/src/handshake/id.rs @@ -1,4 +1,4 @@ -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)] #[repr(transparent)] pub struct HandshakeId(pub u32); diff --git a/ql-wire/src/handshake/ik.rs b/ql-wire/src/handshake/ik.rs index 6de2999a..ec8da42b 100644 --- a/ql-wire/src/handshake/ik.rs +++ b/ql-wire/src/handshake/ik.rs @@ -67,7 +67,7 @@ impl ql_codec::Decode for Ik1 { } } -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Ik2 { pub handshake_id: HandshakeId, diff --git a/ql-wire/src/handshake/mod.rs b/ql-wire/src/handshake/mod.rs index 0385f275..8e3c7c9e 100644 --- a/ql-wire/src/handshake/mod.rs +++ b/ql-wire/src/handshake/mod.rs @@ -24,14 +24,14 @@ const PROTOCOL_IK: &[u8] = b"ql-wire:pq-ik:v1"; const PROTOCOL_KK: &[u8] = b"ql-wire:pq-kk:v1"; const HANDSHAKE_PREAMBLE_DOMAIN: &[u8] = b"ql-wire:handshake-preamble:v1"; -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct EphemeralPublicKey { pub mlkem_public_key: MlKemPublicKey, } } -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct EncryptedMlKemCiphertext(pub Box<[u8; Self::SIZE]>); } diff --git a/ql-wire/src/handshake/pairing.rs b/ql-wire/src/handshake/pairing.rs index ce966908..54e22efa 100644 --- a/ql-wire/src/handshake/pairing.rs +++ b/ql-wire/src/handshake/pairing.rs @@ -5,7 +5,7 @@ use crate::QlCrypto; const PAIRING_ID_DOMAIN: &[u8] = b"ql-wire:pairing-id:v1"; const PAIRING_PSK_DOMAIN: &[u8] = b"ql-wire:pairing-psk:v1"; -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] pub struct PairingToken(pub [u8; Self::SIZE]); @@ -35,7 +35,7 @@ impl Display for PairingToken { } } -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] #[repr(transparent)] pub struct PairingId(pub [u8; Self::SIZE]); diff --git a/ql-wire/src/handshake/transport_params.rs b/ql-wire/src/handshake/transport_params.rs index eb2e0a6a..b0c6cb5e 100644 --- a/ql-wire/src/handshake/transport_params.rs +++ b/ql-wire/src/handshake/transport_params.rs @@ -1,4 +1,4 @@ -ql_codec::codec_struct! { +ql_codec::codec! { /// Session parameters advertised in the handshake #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct TransportParams { diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs index 00f8c7bf..34c256e1 100644 --- a/ql-wire/src/handshake/xx.rs +++ b/ql-wire/src/handshake/xx.rs @@ -14,7 +14,7 @@ use crate::{ const PROTOCOL_XX: &[u8] = b"ql-wire:pq-xx:v1"; -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx1 { pub handshake_id: HandshakeId, @@ -24,7 +24,7 @@ ql_codec::codec_struct! { } } -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx2 { pub handshake_id: HandshakeId, @@ -34,7 +34,7 @@ ql_codec::codec_struct! { } } -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx3 { pub handshake_id: HandshakeId, @@ -43,7 +43,7 @@ ql_codec::codec_struct! { } } -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct Xx4 { pub handshake_id: HandshakeId, diff --git a/ql-wire/src/header.rs b/ql-wire/src/header.rs index 3e8c8fff..a40d96f5 100644 --- a/ql-wire/src/header.rs +++ b/ql-wire/src/header.rs @@ -4,7 +4,7 @@ use ql_common::QID; use crate::QL_WIRE_VERSION; -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RouteHeader { pub sender: QID, @@ -12,7 +12,7 @@ ql_codec::codec_struct! { } } -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct SessionHeader { pub seq: RecordSeq, diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs index 5db2e524..ee39e371 100644 --- a/ql-wire/src/identity.rs +++ b/ql-wire/src/identity.rs @@ -2,7 +2,7 @@ use ql_common::QID; use crate::{derive_qid, Error, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash}; -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct PeerBundle { pub version: u16, @@ -24,7 +24,7 @@ impl PeerBundle { } } -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone)] pub struct QlIdentity { pub qid: QID, diff --git a/ql-wire/src/pq.rs b/ql-wire/src/pq.rs index c5c0c434..afd43f91 100644 --- a/ql-wire/src/pq.rs +++ b/ql-wire/src/pq.rs @@ -8,7 +8,7 @@ const ML_KEM_1024_PUBLIC_KEY_SIZE: usize = 1568; const ML_KEM_1024_PRIVATE_KEY_SIZE: usize = 3168; const ML_KEM_1024_CIPHERTEXT_SIZE: usize = 1568; -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct SessionKey(pub [u8; Self::SIZE]); } @@ -33,7 +33,7 @@ impl Drop for SessionKey { } } -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct MlKemPublicKey(Box<[u8; MlKemPublicKey::SIZE]>); } @@ -56,7 +56,7 @@ impl Drop for MlKemPublicKey { } } -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct MlKemPrivateKey(Box<[u8; MlKemPrivateKey::SIZE]>); } @@ -79,7 +79,7 @@ impl Drop for MlKemPrivateKey { } } -ql_codec::codec_newtype! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub struct MlKemCiphertext(Box<[u8; MlKemCiphertext::SIZE]>); } diff --git a/ql-wire/src/record.rs b/ql-wire/src/record.rs index ceb5f1e6..6484b7d5 100644 --- a/ql-wire/src/record.rs +++ b/ql-wire/src/record.rs @@ -31,7 +31,7 @@ where Ok((reader.decode()?, reader.decode()?)) } -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub struct RecordHeader { pub version: u8, @@ -52,7 +52,7 @@ impl RecordHeader { } } -ql_codec::codec_enum! { +ql_codec::codec! { #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum RecordType { Handshake = 1, @@ -60,7 +60,7 @@ ql_codec::codec_enum! { } } -ql_codec::codec_enum! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub enum QlHandshakeRecord as HandshakeKind { Ik1(Ik1) = 1, @@ -74,7 +74,7 @@ ql_codec::codec_enum! { } } -ql_codec::codec_struct! { +ql_codec::codec! { #[derive(Debug, Clone, PartialEq, Eq)] pub struct QlSessionRecord { pub header: SessionHeader, From 85853eaa2a25035a1676ed840ac075891489fcf3 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Wed, 22 Jul 2026 15:20:10 -0400 Subject: [PATCH 58/59] ql-codec: useless tests --- ql-codec/src/codec.rs | 55 ------------------------------------------- 1 file changed, 55 deletions(-) diff --git a/ql-codec/src/codec.rs b/ql-codec/src/codec.rs index e89af01e..2a116083 100644 --- a/ql-codec/src/codec.rs +++ b/ql-codec/src/codec.rs @@ -242,58 +242,3 @@ where bytes.advance(chunk.len()); } } - -#[cfg(test)] -mod tests { - use bytes::Bytes; - - use super::*; - - #[test] - fn integers_are_little_endian() { - assert_eq!(0x1234u16.encode_vec(), [0x34, 0x12]); - assert_eq!(0x1234_5678u32.encode_vec(), [0x78, 0x56, 0x34, 0x12]); - assert_eq!( - 0x0123_4567_89ab_cdefu64.encode_vec(), - [0xef, 0xcd, 0xab, 0x89, 0x67, 0x45, 0x23, 0x01] - ); - assert_eq!(u16::decode_bytes([0x34, 0x12].as_slice()), Ok(0x1234)); - assert_eq!( - u32::decode_bytes([0x78, 0x56, 0x34, 0x12].as_slice()), - Ok(0x1234_5678) - ); - } - - #[test] - fn byte_containers_use_varint_prefix() { - let bytes = [1u8, 2, 3]; - - let encoded = bytes[..].encode_vec(); - assert_eq!(encoded, [3, 1, 2, 3]); - assert_eq!(bytes.to_vec().encode_vec(), encoded); - assert_eq!(Box::<[u8]>::from(bytes).encode_vec(), encoded); - assert_eq!(Bytes::copy_from_slice(&bytes).encode_vec(), encoded); - - let decoded = <&[u8]>::decode_bytes(encoded.as_slice()).unwrap(); - assert_eq!(decoded, [1, 2, 3]); - assert_eq!( - Vec::::decode_bytes(encoded.as_slice()).unwrap(), - [1, 2, 3] - ); - assert_eq!( - Box::<[u8]>::decode_bytes(encoded.as_slice()) - .unwrap() - .as_ref(), - [1, 2, 3] - ); - - let encoded = String::from("hello").encode_vec(); - assert_eq!(encoded, [5, b'h', b'e', b'l', b'l', b'o']); - assert_eq!(<&str>::decode_bytes(encoded.as_slice()).unwrap(), "hello"); - assert_eq!(String::decode_bytes(encoded.as_slice()).unwrap(), "hello"); - assert_eq!( - String::decode_bytes([1, 0xff].as_slice()), - Err(Error::InvalidUtf8) - ); - } -} From 89d09eecdaecb8cd4ebbc0a40452e6080578a9b8 Mon Sep 17 00:00:00 2001 From: Nico Burniske Date: Thu, 23 Jul 2026 09:16:27 -0400 Subject: [PATCH 59/59] ql-fsm: restore pairing invite codec --- ql-fsm/src/pairing.rs | 17 ++++++++++++----- ql-fsm/src/tests/handshake.rs | 1 + ql-fsm/src/tests/mod.rs | 1 + ql-runtime/src/tests/handshake.rs | 2 ++ 4 files changed, 16 insertions(+), 5 deletions(-) diff --git a/ql-fsm/src/pairing.rs b/ql-fsm/src/pairing.rs index fb9b0532..12eec3e8 100644 --- a/ql-fsm/src/pairing.rs +++ b/ql-fsm/src/pairing.rs @@ -1,9 +1,16 @@ use ql_common::QID; use ql_wire::PairingToken; -/// Out-of-band invite consumed by the initiator of an XX pairing -#[derive(Debug, Clone, Copy, PartialEq, Eq)] -pub struct PairingInvite { - pub qid: QID, - pub token: PairingToken, +ql_codec::codec! { + /// Out-of-band invite consumed by the initiator of an XX pairing + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub struct PairingInvite { + pub version: u8, + pub qid: QID, + pub token: PairingToken, + } +} + +impl PairingInvite { + pub const VERSION: u8 = 1; } diff --git a/ql-fsm/src/tests/handshake.rs b/ql-fsm/src/tests/handshake.rs index b517dcc1..3ffe083a 100644 --- a/ql-fsm/src/tests/handshake.rs +++ b/ql-fsm/src/tests/handshake.rs @@ -106,6 +106,7 @@ fn connect_methods_require_bound_peer() { fsm.connect_xx( time, PairingInvite { + version: PairingInvite::VERSION, qid: ql_common::QID([2; ql_common::QID::SIZE]), token: pairing_token(2), }, diff --git a/ql-fsm/src/tests/mod.rs b/ql-fsm/src/tests/mod.rs index d7003faa..60986ae3 100644 --- a/ql-fsm/src/tests/mod.rs +++ b/ql-fsm/src/tests/mod.rs @@ -203,6 +203,7 @@ impl Harness { fsm.connect_xx( time, PairingInvite { + version: PairingInvite::VERSION, qid: remote_qid, token, }, diff --git a/ql-runtime/src/tests/handshake.rs b/ql-runtime/src/tests/handshake.rs index fb915d4d..805371e0 100644 --- a/ql-runtime/src/tests/handshake.rs +++ b/ql-runtime/src/tests/handshake.rs @@ -139,6 +139,7 @@ async fn start_pairing_round_trip_connects_when_armed() { handle_b.arm_pairing(token); handle_a.start_pairing(PairingInvite { + version: PairingInvite::VERSION, qid: identity_b.qid, token, }); @@ -168,6 +169,7 @@ async fn start_pairing_does_not_connect_when_unarmed() { spawn_forwarder(outbound_b, inbound_a_tx); handle_a.start_pairing(PairingInvite { + version: PairingInvite::VERSION, qid: identity_b.qid, token, });