diff --git a/Cargo.lock b/Cargo.lock index f144305f..38355b76 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,17 +26,6 @@ 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" @@ -60,60 +33,75 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e999941b234f3131b00bc13c22d06e8c5ff726d1b6318ac7eb276997bbb4fef0" [[package]] -name = "android_log-sys" -version = "0.3.2" +name = "android_system_properties" +version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "84521a3cf562bc62942e294181d9eef17eb38ceb8c68677bc49f144e4c3d4f8d" +checksum = "819e7219dbd41043ac279b19830f2efc897156490d7fd6ea916720117ee66311" +dependencies = [ + "libc", +] [[package]] -name = "android_logger" -version = "0.15.1" +name = "anstream" +version = "1.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dbb4e440d04be07da1f1bf44fb4495ebd58669372fe0cffa6e48595ac5bd88a3" +checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" dependencies = [ - "android_log-sys", - "env_filter", - "log", + "anstyle", + "anstyle-parse", + "anstyle-query", + "anstyle-wincon", + "colorchoice", + "is_terminal_polyfill", + "utf8parse", ] [[package]] -name = "android_system_properties" -version = "0.1.5" +name = "anstyle" +version = "1.0.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "819e7219dbd41043ac279b19830f2efc897156490d7fd6ea916720117ee66311" -dependencies = [ - "libc", -] +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" [[package]] -name = "anyhow" -version = "1.0.99" +name = "anstyle-parse" +version = "1.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b0674a1ddeecb70197781e945de4b3b8ffb61fa939a5597bcf48503737663100" +checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" +dependencies = [ + "utf8parse", +] [[package]] -name = "argon2" -version = "0.5.3" +name = "anstyle-query" +version = "1.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c3610892ee6e0cbce8ae2700349fcf8f98adb0dbfbee85aec3c9179d29cc072" +checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" dependencies = [ - "base64ct", - "blake2", - "cpufeatures", - "password-hash", + "windows-sys 0.61.2", ] [[package]] -name = "arrayvec" -version = "0.7.6" +name = "anstyle-wincon" +version = "3.0.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7c02d123df017efcdfbd739ef81735b36c5ba83ec3c59c80a9d7ecc718f92e50" +checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" +dependencies = [ + "anstyle", + "once_cell_polyfill", + "windows-sys 0.61.2", +] [[package]] -name = "atomic" -version = "0.5.3" +name = "async-channel" +version = "2.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c59bdb34bc650a32731b31bd8f0829cc15d24a708ee31559e0bb34f2bc320cba" +checksum = "924ed96dd52d1b75e9c1a3e6275715fd320f5f9439fb5a4a11fa51f4221158d2" +dependencies = [ + "concurrent-queue", + "event-listener-strategy", + "futures-core", + "pin-project-lite", +] [[package]] name = "autocfg" @@ -130,7 +118,7 @@ dependencies = [ "addr2line", "cfg-if", "libc", - "miniz_oxide 0.8.9", + "miniz_oxide", "object", "rustc-demangle", "windows-targets", @@ -147,197 +135,25 @@ dependencies = [ ] [[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 = "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" +name = "bit-set" +version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5d7066118b13d4b20b23645932dfb3a81ce7e29f95726c2036fa33cd7b092501" +checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3" dependencies = [ - "bitcoin-private", + "bit-vec", ] [[package]] -name = "bitcoin_hashes" -version = "0.14.0" +name = "bit-vec" +version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bb18c03d0db0247e147a21a6faafd5a7eb851c743db062de72018b6b7e8e4d16" -dependencies = [ - "bitcoin-io", - "hex-conservative", -] +checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" [[package]] name = "bitflags" -version = "2.9.3" +version = "2.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "34efbcccd345379ca2868b2b2c9d3782e9cc58ba87bc7d79d5b53d9c9ae6f25d" - -[[package]] -name = "blake2" -version = "0.10.6" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "46502ad458c9a52b69d4d4d32775c788b7a1b85e8bc9d482d92250fc0e3f8efe" -dependencies = [ - "digest", -] +checksum = "843867be96c8daad0d758b57df9392b6d8d271134fce549de6ce169ff98a92af" [[package]] name = "block-buffer" @@ -355,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" @@ -391,7 +201,7 @@ checksum = "89385e82b5d1821d2219e0b095efa2cc1f246cbf99080f3be46a1a85c0d392d9" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -411,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" @@ -432,8 +236,6 @@ version = "1.2.34" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42bc4aea80032b7bf409b0bc7ccad88853858911b7713a8062fdc0623867bedc" dependencies = [ - "jobserver", - "libc", "shlex", ] @@ -443,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" @@ -477,20 +255,24 @@ dependencies = [ "iana-time-zone", "js-sys", "num-traits", - "serde", "wasm-bindgen", - "windows-link", + "windows-link 0.1.3", ] [[package]] -name = "cipher" -version = "0.4.4" +name = "colorchoice" +version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "773f3b9af64447d2ce9850330c473515014aa235e6a783b02db81ff39e4a3dad" +checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" + +[[package]] +name = "concurrent-queue" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4ca0197aee26d1ae37445ee532fefce43251d24cc7c166799f4d46817f1d3973" dependencies = [ - "crypto-common", - "inout", - "zeroize", + "crossbeam-utils", + "loom", ] [[package]] @@ -502,25 +284,9 @@ dependencies = [ "encode_unicode", "libc", "once_cell", - "windows-sys", -] - -[[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", + "windows-sys 0.59.0", ] -[[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" @@ -533,37 +299,30 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "773648b94d0e5d620f64f280777445740e61fe701025087ec8b57f45c791888b" [[package]] -name = "cpufeatures" -version = "0.2.17" +name = "core-models" +version = "0.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +checksum = "657f625ff361906f779745d08375ae3cc9fef87a35fba5f22874cf773010daf4" dependencies = [ - "libc", + "hax-lib", + "pastey", + "rand", ] [[package]] -name = "crc" -version = "3.3.0" +name = "cpufeatures" +version = "0.2.17" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9710d3b3739c2e349eb44fe848ad0b7c8cb1e42bd87ee49371df2f7acaf3e675" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" dependencies = [ - "crc-catalog", + "libc", ] [[package]] -name = "crc-catalog" -version = "2.4.0" +name = "crossbeam-utils" +version = "0.8.21" 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", -] +checksum = "d0a5c400df2834b80a4c3327b3aad3a4c4cd4de0629063962b03235697506a28" [[package]] name = "crunchy" @@ -571,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" @@ -590,201 +337,37 @@ 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" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "195516311f81c243495a4d907f05c9d9a78e0b2b53cc384a83b19f1cf9f04cf2" dependencies = [ - "chrono", - "half", - "hex", - "paste", - "thiserror", - "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 = "digest" -version = "0.10.7" -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", + "chrono", + "half", + "hex", + "paste", + "thiserror", + "unicode-normalization", ] [[package]] -name = "either" -version = "1.15.0" +name = "diatomic-waker" +version = "0.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "48c757948c5ede0e46177b7add2e67155f70e33c07fea8284df6576da70b3719" +checksum = "ab03c107fafeb3ee9f5925686dbb7a73bc76e3932abb0d2b365cb64b169cf04c" [[package]] -name = "elliptic-curve" -version = "0.13.8" +name = "digest" +version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b5e6043086bf7973472e0c7dff2142ea0b680d30e18d9cc40f267efbf222bd47" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ - "base16ct", - "crypto-bigint", - "digest", - "ff", - "generic-array", - "group", - "pkcs8", - "rand_core 0.6.4", - "sec1", - "subtle", - "zeroize", + "block-buffer", + "crypto-common", ] [[package]] @@ -795,128 +378,76 @@ checksum = "34aa73646ffb006b8f5147f3dc182bd4bcb190227ce861fc4a4844bf8e3cb2c0" [[package]] name = "env_filter" -version = "0.1.3" +version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "186e05a59d4c50738528153b83b0b0194d3a29507dfec16eccd4b342903397d0" +checksum = "32e90c2accc4b07a8456ea0debdc2e7587bdd890680d71173a15d4ae604f6eef" dependencies = [ "log", "regex", ] [[package]] -name = "equivalent" -version = "1.0.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" - -[[package]] -name = "ff" -version = "0.13.1" +name = "env_logger" +version = "0.11.10" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c0b50bfb653653f9ca9095b427bed08ab8d75a137839d9ad64eb11810d5b6393" +checksum = "0621c04f2196ac3f488dd583365b9c09be011a4ab8b9f37248ffcc8f6198b56a" dependencies = [ - "rand_core 0.6.4", - "subtle", + "anstream", + "anstyle", + "env_filter", + "jiff", + "log", ] [[package]] -name = "fiat-crypto" -version = "0.2.9" +name = "equivalent" +version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" [[package]] -name = "flutter_rust_bridge" -version = "2.11.1" +name = "errno" +version = "0.3.14" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dde126295b2acc5f0a712e265e91b6fdc0ed38767496483e592ae7134db83725" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" 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", + "libc", + "windows-sys 0.61.2", ] [[package]] -name = "flutter_rust_bridge_macros" -version = "2.11.1" +name = "event-listener" +version = "5.4.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d5f0420326b13675321b194928bb7830043b68cf8b810e1c651285c747abb080" +checksum = "e13b66accf52311f30a0db42147dadea9850cb48cd070028831ae5f5d4b856ab" dependencies = [ - "hex", - "md-5", - "proc-macro2", - "quote", - "syn 2.0.106", + "concurrent-queue", + "loom", + "parking", + "pin-project-lite", ] [[package]] -name = "form_urlencoded" -version = "1.2.2" +name = "event-listener-strategy" +version = "0.5.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" -dependencies = [ - "percent-encoding", -] - -[[package]] -name = "foundation-api" -version = "2.0.0" +checksum = "8be9f3dfaaffdae2972880079a491a1a8bb7cbed0b8dd7a347f668b4150a3b93" dependencies = [ - "bc-components", - "bc-envelope", - "bc-xid", - "chrono", - "dcbor", - "flutter_rust_bridge", - "gstp", - "insta", - "quantum-link-macros", - "rkyv", - "thiserror", + "event-listener", + "pin-project-lite", ] [[package]] -name = "futures" -version = "0.3.31" +name = "fastrand" +version = "2.3.0" 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", -] +checksum = "37909eebbb50d72f9059c3b6d82c0463f2ff062c9e95845c43a6c9c0355411be" [[package]] -name = "futures-channel" -version = "0.3.31" +name = "fnv" +version = "1.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2dff15bf788c671c1934e366d07e30c1814a8ef514e1af724a602e8a2fbe1b10" -dependencies = [ - "futures-core", - "futures-sink", -] +checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" [[package]] name = "futures-core" @@ -924,17 +455,6 @@ 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" @@ -942,44 +462,31 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9e5c1b78ca4aae1ac06c48a526a655760685149f0d465d21f37abfe57ce075c6" [[package]] -name = "futures-macro" -version = "0.3.31" +name = "futures-lite" +version = "2.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "162ee34ebcb7c64a8abebc059ce0fee27c2262618d7b60ed8faf72fef13c3650" +checksum = "f78e10609fe0e0b3f4157ffab1876319b5b0db102a2c60dc4626306dc46b44ad" dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", + "fastrand", + "futures-core", + "futures-io", + "parking", + "pin-project-lite", ] [[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" +name = "generator" +version = "0.8.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9fa08315bb612088cc391249efdc3bc77536f16c91f6cf495e6fbe85b20a4a81" +checksum = "52f04ae4152da20c76fe800fa48659201d5cf627c5149ca0b707b69d7eef6cf9" dependencies = [ - "futures-channel", - "futures-core", - "futures-io", - "futures-macro", - "futures-sink", - "futures-task", - "memchr", - "pin-project-lite", - "pin-utils", - "slab", + "cc", + "cfg-if", + "libc", + "log", + "rustversion", + "windows-link 0.2.1", + "windows-result", ] [[package]] @@ -990,7 +497,6 @@ checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" dependencies = [ "typenum", "version_check", - "zeroize", ] [[package]] @@ -1022,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" @@ -1063,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" @@ -1082,46 +551,47 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "841d1cc9bed7f9236f321df977030373f4a4163ae1a7dbfe1a51a2c1a51d9100" [[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" +name = "hax-lib" +version = "0.3.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" +checksum = "543f93241d32b3f00569201bfce9d7a93c92c6421b23c77864ac929dc947b9fc" dependencies = [ - "serde", + "hax-lib-macros", + "num-bigint", + "num-traits", ] [[package]] -name = "hex-conservative" -version = "0.2.1" +name = "hax-lib-macros" +version = "0.3.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5313b072ce3c597065a808dbf612c4c8e8590bdbf8b579508bf7a762c5eae6cd" +checksum = "f8755751e760b11021765bb04cb4a6c4e24742688d9f3aa14c2079638f537b0f" dependencies = [ - "arrayvec", + "hax-lib-macros-types", + "proc-macro-error2", + "proc-macro2", + "quote", + "syn", ] [[package]] -name = "hkdf" -version = "0.12.4" +name = "hax-lib-macros-types" +version = "0.3.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7b5f8eb2ad728638ea2c7d47a21db23b7b58a72ed6a38256b8a1849f15fbbdf7" +checksum = "f177c9ae8ea456e2f71ff3c1ea47bf4464f772a05133fcbba56cd5ba169035a2" dependencies = [ - "hmac", + "proc-macro2", + "quote", + "serde", + "serde_json", + "uuid", ] [[package]] -name = "hmac" -version = "0.12.1" +name = "hex" +version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" -dependencies = [ - "digest", -] +checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" [[package]] name = "iana-time-zone" @@ -1147,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" @@ -1264,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" @@ -1296,13 +650,10 @@ dependencies = [ ] [[package]] -name = "itertools" -version = "0.11.0" +name = "is_terminal_polyfill" +version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b1c173a5686ce8bfa551b3563d0c2170bf24ca44da99c7ca4bfdab5418c3fe57" -dependencies = [ - "either", -] +checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" [[package]] name = "itoa" @@ -1311,34 +662,37 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4a5f13b858c8d314ee3e8f639011f7ccefe71f97f96e50151fb991f267928e2c" [[package]] -name = "jobserver" -version = "0.1.33" +name = "jiff" +version = "0.2.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "38f262f097c174adebe41eb73d66ae9c06b2844fb0da69969647bbddd9b0538a" +checksum = "1a3546dc96b6d42c5f24902af9e2538e82e39ad350b0c766eb3fbf2d8f3d8359" dependencies = [ - "getrandom 0.3.3", - "libc", + "jiff-static", + "log", + "portable-atomic", + "portable-atomic-util", + "serde_core", ] [[package]] -name = "js-sys" -version = "0.3.77" +name = "jiff-static" +version = "0.2.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1cfaf33c695fc6e08064efbc1f72ec937429614f25eef83af942d0e227c3a28f" +checksum = "2a8c8b344124222efd714b73bb41f8b5120b27a7cc1c75593a6ff768d9d05aa4" dependencies = [ - "once_cell", - "wasm-bindgen", + "proc-macro2", + "quote", + "syn", ] [[package]] -name = "known-values" -version = "0.11.0" +name = "js-sys" +version = "0.3.77" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "efadaa833480ac053954ea1bf019eee3b3ed0123417434158d1d9ca46162dfed" +checksum = "1cfaf33c695fc6e08064efbc1f72ec937429614f25eef83af942d0e227c3a28f" dependencies = [ - "bc-components", - "dcbor", - "paste", + "once_cell", + "wasm-bindgen", ] [[package]] @@ -1346,9 +700,6 @@ name = "lazy_static" version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" -dependencies = [ - "spin", -] [[package]] name = "libc" @@ -1357,77 +708,122 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6a82ae493e598baaea5209805c49bbf2ea7de956d50d7da0da1164f9c6d28543" [[package]] -name = "libm" -version = "0.2.15" +name = "libcrux-aesgcm" +version = "0.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f9fbbcab51052fe104eb5e5d351cf728d30a5be1fe14d9be8a3b097481fb97de" +checksum = "99f2a019dab4097585a7d4f5b9deebe46cd1e628b16a5bc4cb0ce35e1da334e6" +dependencies = [ + "libcrux-intrinsics", + "libcrux-platform", + "libcrux-secrets", + "libcrux-traits", +] [[package]] -name = "litemap" -version = "0.8.0" +name = "libcrux-intrinsics" +version = "0.0.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "241eaef5fd12c88705a01fc1066c48c4b36e0dd4377dcdc7ec3942cea7a69956" +checksum = "b1b5db005ff8001e026b73a6842ee81bbef8ec5ff0e1915a67ae65fd2a9fafa5" +dependencies = [ + "core-models", + "hax-lib", +] [[package]] -name = "lock_api" -version = "0.4.13" +name = "libcrux-ml-kem" +version = "0.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "96936507f153605bddfcda068dd804796c84324ed2510809e5b2a624c81da765" +checksum = "aca7de713c6dddcf7aaf76e8ef9dc0097c8d7ce23a8eadf04c8761734714e184" dependencies = [ - "autocfg", - "scopeguard", + "hax-lib", + "libcrux-intrinsics", + "libcrux-platform", + "libcrux-secrets", + "libcrux-sha3", + "libcrux-traits", + "rand", + "tls_codec", ] [[package]] -name = "log" -version = "0.4.27" +name = "libcrux-platform" +version = "0.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "13dc2df351e3202783a1fe0d44375f7295ffb4049267b0f3018346dc122a1d94" +checksum = "1d9e21d7ed31a92ac539bd69a8c970b183ee883872d2d19ce27036e24cb8ecc4" +dependencies = [ + "libc", +] [[package]] -name = "md-5" -version = "0.10.6" +name = "libcrux-secrets" +version = "0.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d89e7ee0cfbedfc4da3340218492196241d89eefb6dab27de5df917a6d2e78cf" +checksum = "1ce650f3041b44ba40d4263852347d007cd2cd9d1cc856a6f6c8b2e10c3fd40b" dependencies = [ - "cfg-if", - "digest", + "hax-lib", ] [[package]] -name = "memchr" -version = "2.7.5" +name = "libcrux-sha3" +version = "0.0.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "32a282da65faaf38286cf3be983213fcf1d2e2a58700e808f83f4ea9a4804bc0" +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", +] [[package]] -name = "minicbor" -version = "0.19.1" +name = "linux-raw-sys" +version = "0.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df1d3c3b53da64cf5760482273a98e575c651a67eec7f77df96b5b642de8f039" + +[[package]] +name = "log" +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 = "d7005aaf257a59ff4de471a9d5538ec868a21586534fff7f85dd97d4043a6139" +checksum = "419e0dc8046cb947daa77eb95ae174acfbddb7673b4151f56d1eed8e93fbfaca" dependencies = [ - "minicbor-derive", + "cfg-if", + "generator", + "scoped-tls", + "tracing", + "tracing-subscriber", ] [[package]] -name = "minicbor-derive" -version = "0.13.0" +name = "matchers" +version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1154809406efdb7982841adb6311b3d095b46f78342dd646736122fe6b19e267" +checksum = "d1525a2a28c7f4fa0fc98bb91ae755d1e2d1505079e05539e35bc876b5d65ae9" dependencies = [ - "proc-macro2", - "quote", - "syn 1.0.109", + "regex-automata", ] [[package]] -name = "miniz_oxide" -version = "0.7.4" +name = "memchr" +version = "2.7.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b8a240ddb74feaf34a79a7add65a741f3167852fba007066dcac1ca548d89c08" -dependencies = [ - "adler", -] +checksum = "32a282da65faaf38286cf3be983213fcf1d2e2a58700e808f83f4ea9a4804bc0" [[package]] name = "miniz_oxide" @@ -1446,7 +842,7 @@ checksum = "78bed444cc8a2160f01cbcf811ef18cac863ad68ae8ca62092e8db51d51c761c" dependencies = [ "libc", "wasi 0.11.1+wasi-snapshot-preview1", - "windows-sys", + "windows-sys 0.59.0", ] [[package]] @@ -1466,42 +862,34 @@ checksum = "4568f25ccbd45ab5d5603dc34318c1ec56b117531781260002151b8530a9f931" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] -name = "num-bigint-dig" -version = "0.8.6" +name = "nu-ansi-term" +version = "0.50.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e661dda6640fad38e827a6d4a310ff4763082116fe217f279885c97f511bb0b7" +checksum = "7957b9740744892f114936ab4a57b3f487491bbeafaf8083688b16841a4240e5" dependencies = [ - "lazy_static", - "libm", - "num-integer", - "num-iter", - "num-traits", - "rand 0.8.5", - "smallvec", - "zeroize", + "windows-sys 0.61.2", ] [[package]] -name = "num-integer" -version = "0.1.46" +name = "num-bigint" +version = "0.4.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7969661fd2958a5cb096e56c8e1ad0444ac2bbcd0061bd28660485a44879858f" +checksum = "a5e44f723f1133c9deac646763579fdb3ac745e418f2a7af9cd0c431da1f20b9" dependencies = [ + "num-integer", "num-traits", ] [[package]] -name = "num-iter" -version = "0.1.45" +name = "num-integer" +version = "0.1.46" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1429034a0490724d0075ebb2bc9e875d6503c3cf69e235a8941aa757d83ef5bf" +checksum = "7969661fd2958a5cb096e56c8e1ad0444ac2bbcd0061bd28660485a44879858f" dependencies = [ - "autocfg", - "num-integer", "num-traits", ] @@ -1512,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]] @@ -1541,82 +918,24 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42f5e15c9953c5e4ccceeb2e7382a716482c34515315f7b03532b8b4e8393d2d" [[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" +name = "once_cell_polyfill" +version = "1.70.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0fc9e2161f1f215afdfce23677034ae137bbd45016a880c2eb3ba8eb95f085b2" -dependencies = [ - "base16ct", - "ecdsa", - "elliptic-curve", - "primeorder", - "rand_core 0.6.4", - "sha2", -] +checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" [[package]] -name = "parking_lot_core" -version = "0.9.11" +name = "oneshot" +version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bc838d2a56b5b1a6c25f55575dfc605fabb63bb2365f6c2353ef9159aa69e4a5" -dependencies = [ - "cfg-if", - "libc", - "redox_syscall", - "smallvec", - "windows-targets", -] +checksum = "b4ce411919553d3f9fa53a0880544cda985a112117a0444d5ff1e870a893d6ea" [[package]] -name = "password-hash" -version = "0.5.0" +name = "parking" +version = "2.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "346f04948ba92c43e8469c1ee6736c7563d71012b17d40745260fe106aac2166" +checksum = "f38d5652c16fde515bb1ecef450ab0f6a219d619a7274976324d5e377f7dceba" dependencies = [ - "base64ct", - "rand_core 0.6.4", - "subtle", + "loom", ] [[package]] @@ -1626,251 +945,186 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" [[package]] -name = "pbkdf2" -version = "0.12.2" +name = "pastey" +version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f8ed6a7761f76e3b9f92dfb0a60a6a6477c61024b775147ff0973a02653abaf2" -dependencies = [ - "digest", - "hmac", -] +checksum = "b867cad97c0791bbd3aaa6472142568c6c9e8f71937e98379f584cfb0cf35bec" [[package]] -name = "pem-rfc7468" -version = "0.7.0" +name = "pin-project-lite" +version = "0.2.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "88b39c9bfcfc231068454382784bb460aae594343fb030d46e9f50a645418412" -dependencies = [ - "base64ct", -] +checksum = "3b3cff922bd51709b605d9ead9aa71031d81447142d828eb4a6eba76fe619f9b" [[package]] -name = "percent-encoding" -version = "2.3.2" +name = "portable-atomic" +version = "1.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" +checksum = "f84267b20a16ea918e43c6a88433c2d54fa145c92a811b5b047ccbe153674483" [[package]] -name = "phf" -version = "0.11.3" +name = "portable-atomic-util" +version = "0.2.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1fd6780a80ae0c52cc120a26a1a42c1ae51b247a253e4e06113d23d2c2edd078" +checksum = "c2a106d1259c23fac8e543272398ae0e3c0b8d33c88ed73d0cc71b0f1d902618" dependencies = [ - "phf_macros", - "phf_shared", + "portable-atomic", ] [[package]] -name = "phf_generator" -version = "0.11.3" +name = "ppv-lite86" +version = "0.2.21" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c80231409c20246a13fddb31776fb942c38553c51e871f8cbd687a4cfb5843d" +checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" dependencies = [ - "phf_shared", - "rand 0.8.5", + "zerocopy", ] [[package]] -name = "phf_macros" -version = "0.11.3" +name = "proc-macro-error-attr2" +version = "2.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f84ac04429c13a7ff43785d75ad27569f2951ce0ffd30a3321230db2fc727216" +checksum = "96de42df36bb9bba5542fe9f1a054b8cc87e172759a1868aa05c1f3acc89dfc5" 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" +name = "proc-macro-error2" +version = "2.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c8ffb9f10fa047879315e6625af03c164b16962a5368d724ed16323b68ace47f" +checksum = "11ec05c52be0a07b08061f7dd003e7d7092e0472bc731b4af7bb1ef876109802" dependencies = [ - "der", - "pkcs8", - "spki", + "proc-macro-error-attr2", + "proc-macro2", + "quote", + "syn", ] [[package]] -name = "pkcs8" -version = "0.10.2" +name = "proc-macro2" +version = "1.0.101" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f950b2377845cebe5cf8b5165cb3cc1a5e0fa5cfa3e1f7f55707d8fd82e0a7b7" +checksum = "89ae43fd86e4158d6db51ad8e2b80f313af9cc74f5c0e03ccb87de09998732de" dependencies = [ - "der", - "spki", + "unicode-ident", ] [[package]] -name = "poly1305" -version = "0.8.0" +name = "proptest" +version = "1.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8159bd90725d2df49889a078b54f4f79e87f1f8a8444194cdca81d38f5393abf" +checksum = "4b45fcc2344c680f5025fe57779faef368840d0bd1f42f216291f0dc4ace4744" dependencies = [ - "cpufeatures", - "opaque-debug", - "universal-hash", + "bit-set", + "bit-vec", + "bitflags", + "num-traits", + "rand", + "rand_chacha", + "rand_xorshift", + "regex-syntax", + "rusty-fork", + "tempfile", + "unarray", ] [[package]] -name = "portable-atomic" -version = "1.11.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f84267b20a16ea918e43c6a88433c2d54fa145c92a811b5b047ccbe153674483" - -[[package]] -name = "potential_utf" -version = "0.1.2" +name = "ptr_meta" +version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e5a7c30837279ca13e7c867e9e40053bc68740f988cb07f7ca6df43cc734b585" +checksum = "0b9a0cf95a1196af61d4f1cbdab967179516d9a4a4312af1f31948f8f6224a79" dependencies = [ - "zerovec", + "ptr_meta_derive", ] [[package]] -name = "ppv-lite86" -version = "0.2.21" +name = "ptr_meta_derive" +version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" +checksum = "7347867d0a7e1208d93b46767be83e2b8f978c3dad35f775ac8d8847551d6fe1" dependencies = [ - "zerocopy", + "proc-macro2", + "quote", + "syn", ] [[package]] -name = "pqcrypto-internals" -version = "0.2.10" -source = "git+https://github.com/Foundation-Devices/pqcrypto?rev=ebadf71214f67cb970242fa1053b4acb65767737#ebadf71214f67cb970242fa1053b4acb65767737" +name = "ql-codec" +version = "0.1.0" dependencies = [ - "cc", - "dunce", - "getrandom 0.2.16", - "libc", + "bytes", ] [[package]] -name = "pqcrypto-mldsa" -version = "0.1.1" -source = "git+https://github.com/Foundation-Devices/pqcrypto?rev=ebadf71214f67cb970242fa1053b4acb65767737#ebadf71214f67cb970242fa1053b4acb65767737" +name = "ql-common" +version = "0.1.0" dependencies = [ - "cc", - "glob", - "libc", - "paste", - "pqcrypto-internals", - "pqcrypto-traits", + "bytes", + "ql-codec", ] [[package]] -name = "pqcrypto-mlkem" +name = "ql-fsm" version = "0.1.0" -source = "git+https://github.com/Foundation-Devices/pqcrypto?rev=ebadf71214f67cb970242fa1053b4acb65767737#ebadf71214f67cb970242fa1053b4acb65767737" dependencies = [ - "cc", - "glob", - "libc", - "pqcrypto-internals", - "pqcrypto-traits", + "bytes", + "indexmap", + "proptest", + "ql-codec", + "ql-common", + "ql-wire", ] [[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" +name = "ql-rpc" +version = "0.1.0" dependencies = [ - "elliptic-curve", + "bytes", + "ql-common", + "trait-variant", ] [[package]] -name = "proc-macro2" -version = "1.0.101" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "89ae43fd86e4158d6db51ad8e2b80f313af9cc74f5c0e03ccb87de09998732de" +name = "ql-runtime" +version = "0.1.0" dependencies = [ - "unicode-ident", + "async-channel", + "bytes", + "diatomic-waker", + "env_logger", + "event-listener", + "futures-lite", + "log", + "loom", + "oneshot", + "ql-codec", + "ql-common", + "ql-fsm", + "ql-rpc", + "ql-wire", + "tokio", ] [[package]] -name = "provenance-mark" -version = "0.16.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e8a2078ef3c515d873099557bdbcb4b1ad8e67af1d2b398b5a7aa224baca970d" +name = "ql-wire" +version = "0.1.0" dependencies = [ - "base64", - "bc-envelope", - "bc-rand", - "bc-tags", - "bc-ur", - "chacha20", - "chrono", - "dcbor", - "hex", - "hkdf", - "rand_core 0.6.4", - "serde", - "serde_json", + "bytes", + "getrandom 0.2.16", + "libcrux-aesgcm", + "libcrux-ml-kem", + "ql-codec", + "ql-common", "sha2", - "thiserror", - "url", ] [[package]] -name = "ptr_meta" -version = "0.3.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0b9a0cf95a1196af61d4f1cbdab967179516d9a4a4312af1f31948f8f6224a79" -dependencies = [ - "ptr_meta_derive", -] - -[[package]] -name = "ptr_meta_derive" -version = "0.3.1" +name = "quick-error" +version = "1.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7347867d0a7e1208d93b46767be83e2b8f978c3dad35f775ac8d8847551d6fe1" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] - -[[package]] -name = "quantum-link-macros" -version = "0.1.0" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.106", -] +checksum = "a1d01941d82fa2ab50be1e79e6714289dd7cde78eba4c074bc5a4374f650dfe0" [[package]] name = "quote" @@ -1896,35 +1150,14 @@ 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", + "rand_chacha", + "rand_core", ] [[package]] @@ -1934,16 +1167,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" dependencies = [ "ppv-lite86", - "rand_core 0.9.3", -] - -[[package]] -name = "rand_core" -version = "0.6.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c" -dependencies = [ - "getrandom 0.2.16", + "rand_core", ] [[package]] @@ -1956,28 +1180,19 @@ dependencies = [ ] [[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" +name = "rand_xorshift" +version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5407465600fb0548f1442edf71dd20683c6ed326200ace4b1ef0763521bb3b77" +checksum = "513962919efc330f829edb2535844d1b912b0fbe2ca165d613e4e8788bb05a5a" dependencies = [ - "bitflags", + "rand_core", ] [[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", @@ -1987,9 +1202,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", @@ -2011,16 +1226,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" @@ -2048,28 +1253,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]] @@ -2079,12 +1263,16 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "56f7d92ca342cea22a06f2121d944b4fd82af56988c270852495420f961d4ace" [[package]] -name = "rustc_version" -version = "0.4.1" +name = "rustix" +version = "1.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cfcb3a22ef46e85b45de6ee7e79d063319ebb6594faafcf1c225ea92ab6e9b92" +checksum = "cd15f8a2c5551a84d56efdc1cd049089e409ac19a3072d5037a17fd70719ff3e" dependencies = [ - "semver", + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", ] [[package]] @@ -2094,95 +1282,57 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" [[package]] -name = "ryu" -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 = "scopeguard" -version = "1.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" - -[[package]] -name = "scrypt" -version = "0.11.0" +name = "rusty-fork" +version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0516a385866c09368f0b5bcd1caff3366aace790fcd46e2bb032697bb172fd1f" +checksum = "cc6bf79ff24e648f6da1f8d1f011e9cac26491b619e6b9280f2b47f1774e6ee2" dependencies = [ - "pbkdf2", - "salsa20", - "sha2", + "fnv", + "quick-error", + "tempfile", + "wait-timeout", ] [[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", -] +name = "ryu" +version = "1.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "28d3b2b1366ec20994f1fd18c3c594f05c5dd4bc44d8bb0c1c632c8d6829481f" [[package]] -name = "secp256k1" -version = "0.30.0" +name = "scoped-tls" +version = "1.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b50c5943d326858130af85e049f2661ba3c78b26589b8ab98e65e80ae44a1252" -dependencies = [ - "bitcoin_hashes 0.14.0", - "rand 0.8.5", - "secp256k1-sys", -] +checksum = "e1cf6437eb19a8f4a6cc0f7dca544973b0b78843adbfeb3683d1a94a0024a294" [[package]] -name = "secp256k1-sys" -version = "0.10.1" +name = "serde" +version = "1.0.228" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d4387882333d3aa8cb20530a17c69a3752e97837832f34f6dccc760e715001d9" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" dependencies = [ - "cc", + "serde_core", + "serde_derive", ] [[package]] -name = "semver" -version = "1.0.26" +name = "serde_core" +version = "1.0.228" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "56e6fa9c48d24d85fb3de5ad847117517440f6beceb7798af16b4a87d616b8d0" - -[[package]] -name = "serde" -version = "1.0.219" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5f0e2c6ed6606019b4e29e69dbaba95b11854410e5347d525002456dbbb786b6" +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", - "syn 2.0.106", + "syn", ] [[package]] @@ -2198,10 +1348,10 @@ dependencies = [ ] [[package]] -name = "sha1" -version = "0.10.6" +name = "sha2" +version = "0.10.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" dependencies = [ "cfg-if", "cpufeatures", @@ -2209,14 +1359,12 @@ dependencies = [ ] [[package]] -name = "sha2" -version = "0.10.9" +name = "sharded-slab" +version = "0.1.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6" dependencies = [ - "cfg-if", - "cpufeatures", - "digest", + "lazy_static", ] [[package]] @@ -2225,16 +1373,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" @@ -2247,12 +1385,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" @@ -2266,188 +1398,178 @@ 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" +name = "syn" +version = "2.0.106" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d91ed6c858b01f942cd56b37a94b3e0a1798290327d1236e4d9cf4eaca44d29d" +checksum = "ede7c438028d4436d71104916910f5bb611972c5cfd7f89b8300a8186e6fada6" dependencies = [ - "base64ct", - "der", + "proc-macro2", + "quote", + "unicode-ident", ] [[package]] -name = "ssh-cipher" -version = "0.2.0" +name = "tempfile" +version = "3.23.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "caac132742f0d33c3af65bfcde7f6aa8f62f0e991d80db99149eb9d44708784f" +checksum = "2d31c77bdf42a745371d260a26ca7163f1e0924b64afa0b688e61b5a9fa02f16" dependencies = [ - "cipher", - "ssh-encoding", + "fastrand", + "getrandom 0.3.3", + "once_cell", + "rustix", + "windows-sys 0.61.2", ] [[package]] -name = "ssh-encoding" -version = "0.2.0" +name = "thiserror" +version = "2.0.17" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "eb9242b9ef4108a78e8cd1a2c98e193ef372437f8c22be363075233321dd4a15" +checksum = "f63587ca0f12b72a0600bcba1d40081f830876000bb46dd2337a3051618f4fc8" dependencies = [ - "base64ct", - "pem-rfc7468", - "sha2", + "thiserror-impl", ] [[package]] -name = "ssh-key" -version = "0.6.7" +name = "thiserror-impl" +version = "2.0.17" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3b86f5297f0f04d08cabaa0f6bff7cb6aec4d9c3b49d87990d63da9d9156a8c3" +checksum = "3ff15c8ecd7de3849db632e14d18d2571fa09dfc5ed93479bc4485c7a517c913" 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", + "proc-macro2", + "quote", + "syn", ] [[package]] -name = "sskr" -version = "0.11.0" +name = "thread_local" +version = "1.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7228e0234fae61785706c7f2b2bc5e47b6b34397e8c1052e1cfcba8030536234" +checksum = "f60246a4944f24f6e018aa17cdeffb7818b76356965d03b07d6a9886e8962185" dependencies = [ - "bc-rand", - "bc-shamir", - "thiserror", + "cfg-if", ] [[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" +name = "tinyvec" +version = "1.10.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" +checksum = "bfa5fdc3bce6191a1dbc8c02d5c8bffcf557bafa17c124c5264a458f1b0613fa" +dependencies = [ + "tinyvec_macros", +] [[package]] -name = "syn" -version = "1.0.109" +name = "tinyvec_macros" +version = "0.1.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "72b64191b275b66ffe2469e8af2c1cfe3bafa67b529ead792a6d0160888b4237" -dependencies = [ - "proc-macro2", - "quote", - "unicode-ident", -] +checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] -name = "syn" -version = "2.0.106" +name = "tls_codec" +version = "0.4.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ede7c438028d4436d71104916910f5bb611972c5cfd7f89b8300a8186e6fada6" +checksum = "0de2e01245e2bb89d6f05801c564fa27624dbd7b1846859876c7dad82e90bf6b" dependencies = [ - "proc-macro2", - "quote", - "unicode-ident", + "tls_codec_derive", + "zeroize", ] [[package]] -name = "synstructure" -version = "0.13.2" +name = "tls_codec_derive" +version = "0.4.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" +checksum = "2d2e76690929402faae40aebdda620a2c0e25dd6d3b9afe48867dfd95991f4bd" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] -name = "thiserror" -version = "2.0.17" +name = "tokio" +version = "1.47.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f63587ca0f12b72a0600bcba1d40081f830876000bb46dd2337a3051618f4fc8" +checksum = "89e49afdadebb872d3145a5638b59eb0691ea23e46ca484037cfab3b76b95038" dependencies = [ - "thiserror-impl", + "backtrace", + "io-uring", + "libc", + "mio", + "pin-project-lite", + "slab", + "tokio-macros", ] [[package]] -name = "thiserror-impl" -version = "2.0.17" +name = "tokio-macros" +version = "2.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3ff15c8ecd7de3849db632e14d18d2571fa09dfc5ed93479bc4485c7a517c913" +checksum = "6e06d43f1345a3bcd39f6a56dbb7dcab2ba47e68e8ac134855e7e2bdbaf8cab8" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] -name = "threadpool" -version = "1.8.1" +name = "tracing" +version = "0.1.44" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d050e60b33d41c19108b32cea32164033a9013fe3b46cbd4457559bfbf77afaa" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" dependencies = [ - "num_cpus", + "pin-project-lite", + "tracing-core", ] [[package]] -name = "tinystr" -version = "0.8.1" +name = "tracing-core" +version = "0.1.36" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5d4f6d1145dcb577acf783d4e601bc1d76a13337bb54e6233add580b07344c8b" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" dependencies = [ - "displaydoc", - "zerovec", + "once_cell", + "valuable", ] [[package]] -name = "tinyvec" -version = "1.10.0" +name = "tracing-log" +version = "0.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bfa5fdc3bce6191a1dbc8c02d5c8bffcf557bafa17c124c5264a458f1b0613fa" +checksum = "ee855f1f400bd0e5c02d150ae5de3840039a3f54b025156404e34c23c03f47c3" dependencies = [ - "tinyvec_macros", + "log", + "once_cell", + "tracing-core", ] [[package]] -name = "tinyvec_macros" -version = "0.1.1" +name = "tracing-subscriber" +version = "0.3.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" +checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319" +dependencies = [ + "matchers", + "nu-ansi-term", + "once_cell", + "regex-automata", + "sharded-slab", + "smallvec", + "thread_local", + "tracing", + "tracing-core", + "tracing-log", +] [[package]] -name = "tokio" -version = "1.47.1" +name = "trait-variant" +version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "89e49afdadebb872d3145a5638b59eb0691ea23e46ca484037cfab3b76b95038" +checksum = "70977707304198400eb4835a78f6a9f928bf41bba420deb8fdb175cd965d77a7" dependencies = [ - "backtrace", - "io-uring", - "libc", - "mio", - "pin-project-lite", - "slab", + "proc-macro2", + "quote", + "syn", ] [[package]] @@ -2456,6 +1578,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" @@ -2472,44 +1600,10 @@ dependencies = [ ] [[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" +name = "utf8parse" +version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" [[package]] name = "uuid" @@ -2517,16 +1611,32 @@ 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", ] +[[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" 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" @@ -2564,23 +1674,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" @@ -2599,7 +1696,7 @@ checksum = "8ae87ea40c9f689fc23f209965b6fb8a99ad69aeeb0231408be24920604395de" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", "wasm-bindgen-backend", "wasm-bindgen-shared", ] @@ -2613,16 +1710,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" @@ -2631,7 +1718,7 @@ checksum = "c0fdd3ddb90610c7638aa2b3a3ab2904fb9e5cdbecc643ddb3647212781c4ae3" dependencies = [ "windows-implement", "windows-interface", - "windows-link", + "windows-link 0.1.3", "windows-result", "windows-strings", ] @@ -2644,7 +1731,7 @@ checksum = "a47fddd13af08290e67f4acabf4b459f647552718f683a7b415d290ac744a836" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -2655,7 +1742,7 @@ checksum = "bd9211b69f8dcdfa817bfd14bf1c97c9188afa36f4750130fcdf3f400eca9fa8" dependencies = [ "proc-macro2", "quote", - "syn 2.0.106", + "syn", ] [[package]] @@ -2664,13 +1751,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]] @@ -2679,7 +1772,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]] @@ -2691,6 +1784,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" @@ -2764,48 +1866,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" @@ -2823,28 +1883,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]] @@ -2864,38 +1903,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 0fd0e755..056fd430 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,31 +1,31 @@ [workspace] resolver = "2" -members = ["api", "backup-shard", "btp", "quantum-link-macros"] +members = [ + "backup-shard", + "btp", + "ql-codec", + "ql-common", + "ql-fsm", + "ql-rpc", + "ql-runtime", + "ql-wire", +] [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" } - -[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" } +ql-codec = { path = "ql-codec" } +ql-common = { path = "ql-common" } +ql-fsm = { path = "ql-fsm" } +ql-rpc = { path = "ql-rpc" } +ql-wire = { path = "ql-wire" } 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/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..c6887364 --- /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` +- `StreamReset.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 | +| `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` | + +`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 + +`StreamReset` 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. 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/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..2a116083 --- /dev/null +++ b/ql-codec/src/codec.rs @@ -0,0 +1,244 @@ +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 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() + } +} + +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), + } + } +} + +/// 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") +} + +pub fn encoded_len_bytes(bytes: &B) -> usize { + let len = bytes.buf().remaining(); + varint::encoded_len(len_prefix(len)) + len +} + +pub fn encode_bytes(bytes: &B, out: &mut W) +where + B: BufView + ?Sized, + W: BufMut + ?Sized, +{ + varint::encode(len_prefix(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(); + while bytes.has_remaining() { + let chunk = bytes.chunk(); + out.put_slice(chunk); + bytes.advance(chunk.len()); + } +} 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..53e9966a --- /dev/null +++ b/ql-codec/src/lib.rs @@ -0,0 +1,60 @@ +//! Small binary codec primitives shared by QuantumLink crates. + +mod buf_view; +mod codec; +mod error; +mod macros; +mod reader; +mod slice; +pub mod varint; + +pub use buf_view::BufView; +pub use codec::{encode_bytes, encode_bytes_raw, 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; + + 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 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/macros.rs b/ql-codec/src/macros.rs new file mode 100644 index 00000000..50cee548 --- /dev/null +++ b/ql-codec/src/macros.rs @@ -0,0 +1,215 @@ +/// 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 + ( + $(#[$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()?)) + } + } + }; + + // struct with an optional byte container + ( + $(#[$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),*); + }; + + // 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! { + #[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))? = $value),* + } + }; + + // u8 discriminant enum + ( + $(#[$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() + } + } + }; + + // 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>()? { + $($kind::$variant => Self::$variant $((reader.decode::<$payload>()?))?,)* + }) + } + } + }; + + // 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>()? { + $($kind::$variant => Self::$variant $((reader.decode::<$payload>()?))?,)* + }) + } + } + }; +} diff --git a/ql-codec/src/reader.rs b/ql-codec/src/reader.rs new file mode 100644 index 00000000..35ed7033 --- /dev/null +++ b/ql-codec/src/reader.rs @@ -0,0 +1,60 @@ +use crate::{varint, ByteSlice, Decode, Error}; + +#[derive(Clone)] +pub struct Reader { + remaining: B, +} + +impl Reader { + #[inline] + pub fn new(bytes: B) -> Self { + Self { remaining: 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); + } + Ok(self.remaining.split_off_front(len)) + } + + #[inline] + pub fn take_u8(&mut self) -> Result { + self.remaining.take_u8().ok_or(Error::UnexpectedEof) + } + + pub fn take_all(&mut self) -> B { + self.remaining.split_off_front(self.remaining.len()) + } + + pub fn take_len_prefixed(&mut self) -> Result { + let len = self.decode_varint::()?; + self.take_n(len as usize) + } + + #[inline] + pub fn decode(&mut self) -> Result + where + T: Decode, + { + T::decode(self) + } + + #[inline] + pub fn decode_varint(&mut self) -> Result + where + T: varint::Primitive, + { + varint::decode(self) + } +} diff --git a/ql-codec/src/slice.rs b/ql-codec/src/slice.rs new file mode 100644 index 00000000..9fafbdbf --- /dev/null +++ b/ql-codec/src/slice.rs @@ -0,0 +1,60 @@ +use core::{mem, ops::Deref}; + +use bytes::{Buf, Bytes}; + +/// A byte slice owner used by the codec reader +pub trait ByteSlice: Deref + Sized { + /// splits `self[..mid]` off the front, leaving `self[mid..]` behind + /// + /// # Panics + /// + /// Panics if `mid` exceeds the slice length. + fn split_off_front(&mut self, mid: usize) -> Self; + + fn take_u8(&mut self) -> Option; +} + +impl ByteSlice for &[u8] { + #[inline] + fn split_off_front(&mut self, mid: usize) -> Self { + let (head, tail) = self.split_at(mid); + *self = tail; + head + } + + #[inline] + fn take_u8(&mut self) -> Option { + let (&byte, remaining) = self.split_first()?; + *self = remaining; + Some(byte) + } +} + +impl ByteSlice for &mut [u8] { + #[inline] + fn split_off_front(&mut self, mid: usize) -> Self { + let (head, tail) = mem::take(self).split_at_mut(mid); + *self = tail; + head + } + + #[inline] + fn take_u8(&mut self) -> Option { + let (byte, remaining) = mem::take(self).split_first_mut()?; + let byte = *byte; + *self = remaining; + Some(byte) + } +} + +impl ByteSlice for Bytes { + #[inline] + fn split_off_front(&mut self, mid: usize) -> Self { + self.split_to(mid) + } + + #[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..0e5e1cb9 --- /dev/null +++ b/ql-codec/src/varint.rs @@ -0,0 +1,221 @@ +use core::{fmt, ops::Deref}; + +use bytes::BufMut; + +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 Primitive: 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; +} + +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(); + len += 1; + } + len +} + +pub fn encode(mut value: T, out: &mut W) +where + T: Primitive, + 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: Primitive, + 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 Primitive 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: Primitive + 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: Primitive + Debug + PartialEq, + { + let mut reader = Reader::new(bytes); + assert_eq!(decode::(&mut reader), Err(error)); + } + + fn test_type(max: T) + where + T: Primitive + 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); +} diff --git a/ql-common/Cargo.toml b/ql-common/Cargo.toml new file mode 100644 index 00000000..be206ecc --- /dev/null +++ b/ql-common/Cargo.toml @@ -0,0 +1,10 @@ +[package] +name = "ql-common" +version = "0.1.0" +edition = "2021" +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 new file mode 100644 index 00000000..1f661aa0 --- /dev/null +++ b/ql-common/src/lib.rs @@ -0,0 +1,78 @@ +//! Shared QuantumLink primitive types. + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +#[repr(transparent)] +pub struct ResetCode(pub u64); + +impl ResetCode { + /// 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(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(7); + /// configured or encoded size limit was exceeded + pub const LIMIT: Self = Self(8); + /// route identifier was unknown + pub const UNKNOWN_ROUTE: Self = Self(9); +} + +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"), + 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}"), + } + } +} + +ql_codec::varint_wrapper!(ResetCode, u64); + +ql_codec::codec! { + #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] + #[repr(transparent)] + pub struct QID(pub [u8; Self::SIZE]); +} + +impl QID { + pub const SIZE: usize = 16; +} + +/// 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 { + pub qid: QID, + pub stream_id: StreamId, + pub header: Box<[u8]>, +} diff --git a/ql-fsm/Cargo.toml b/ql-fsm/Cargo.toml new file mode 100644 index 00000000..4b5a48cc --- /dev/null +++ b/ql-fsm/Cargo.toml @@ -0,0 +1,17 @@ +[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-common = { workspace = true } +ql-codec = { workspace = true } +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..8614a3c3 --- /dev/null +++ b/ql-fsm/src/error.rs @@ -0,0 +1,140 @@ +use std::{ + error::Error, + fmt::{Display, Formatter}, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum ReceiveError { + Wire { + stage: ReceiveStage, + source: ql_wire::Error, + }, + InvalidRecordVersion, + InvalidRemoteBundle, + InvalidQid, + NoPeer, + NoSession, + NotPairingMode, + 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::Wire { stage, source } => write!(f, "invalid {stage}: {source}"), + Self::InvalidRecordVersion => f.write_str("invalid record version"), + 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 => f.write_str("invalid pairing id"), + } + } +} + +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 { + 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, + NoSession, +} + +impl Display for StreamError { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + let message = match self { + Self::MissingStream => "missing stream", + 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..a532c9bd --- /dev/null +++ b/ql-fsm/src/fsm.rs @@ -0,0 +1,268 @@ +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}; + +use crate::{ + handshake, + session::{self, SessionEvent, TerminalFrame}, + state::LinkState, + Event, NoPeerError, NoSessionError, OutboundWrite, QlFsm, ReceiveError, ReceiveStage, + 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) => { + self.events.push_back(Event::Opened(stream_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::Reset(reset) => { + self.events.push_back(Event::Reset(reset)); + } + 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 = 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); + } + 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(|error| ReceiveError::wire(ReceiveStage::HandshakeRecord, error))?; + 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(|error| ReceiveError::wire(ReceiveStage::SessionRecord, error))?; + let payload = wire::decrypt_record( + crypto, + &header, + &record.header, + record.payload, + &conn.transport.rx_key, + ) + .map_err(|error| ReceiveError::wire(ReceiveStage::SessionPayload, error))?; + (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((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, + }); + } + + let QlFsm { state, .. } = fsm; + let conn = state.link.connected_mut()?; + let route = wire::RouteHeader { + sender: fsm.identity.qid, + recipient: conn.transport.remote_qid, + }; + + let (write_id, builder) = conn.session.take_next_write(state.now)?; + let record = builder.encrypt(crypto, route, &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, + header: Box<[u8]>, +) -> Result, NoSessionError> { + let QlFsm { state, events, .. } = fsm; + let conn = state.link.connected_mut_or_err()?; + let inner = conn.session.open_stream(header, 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..1b1c141e --- /dev/null +++ b/ql-fsm/src/handshake/ik.rs @@ -0,0 +1,189 @@ +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, +}; +use crate::{ + state::{InitiatorState, LinkState}, + QlFsm, ReceiveError, ReceiveStage, +}; + +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 = 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, + }; + + 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_1( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + route: RouteHeader, + message: &Ik1, + pattern: IkPattern, +) -> Result<(), ReceiveError> { + if should_ignore_inbound(fsm, route, message, pattern) { + return Ok(()); + } + + 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 = match (pattern, peer) { + (IkPattern::Ik, expected_remote) => wire::IkHandshake::new_ik_responder( + crypto, + &fsm.identity, + expected_remote, + super::local_transport_params(fsm), + ), + (IkPattern::Kk, Some(remote_bundle)) => wire::IkHandshake::new_kk_responder( + crypto, + &fsm.identity, + remote_bundle, + super::local_transport_params(fsm), + ), + (IkPattern::Kk, None) => unreachable!("KK peer was checked above"), + }; + handshake + .read_1(crypto, route, message) + .map_err(|source| wire_error(pattern, source))?; + let outbound = handshake + .write_2(crypto, message.handshake_id) + .map_err(|source| wire_error(pattern, source))?; + establish_session( + fsm, + 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, + }, + record, + ); + Ok(()) +} + +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) + { + 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 handshake initiator was checked above"); + }; + establish_session( + fsm, + message.handshake_id, + state + .handshake + .finalize(crypto) + .map_err(|source| wire_error(pattern, source))?, + ) +} + +fn should_ignore_inbound( + fsm: &QlFsm, + route: RouteHeader, + message: &Ik1, + pattern: IkPattern, +) -> bool { + match &fsm.state.link { + LinkState::Idle | LinkState::XxInitiator(_) | LinkState::XxResponder(_) => false, + LinkState::Connected(_) => { + super::is_connected_replay(fsm, message.handshake_id, route.sender) + } + LinkState::IkInitiator(state) => { + 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 + } + } + } +} + +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/mod.rs b/ql-fsm/src/handshake/mod.rs new file mode 100644 index 00000000..e70c2936 --- /dev/null +++ b/ql-fsm/src/handshake/mod.rs @@ -0,0 +1,172 @@ +mod ik; +mod xx; + +use ql_common::QID; +use ql_wire::{ + self as wire, EphemeralPublicKey, HandshakeId, IkPattern, MlKemPublicKey, QlCrypto, + QlHandshakeRecord, RouteHeader, +}; + +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, 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); + ik::start_initiator(fsm, crypto, peer, IkPattern::Kk); + 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_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); + handshake_id +} + +pub fn enqueue_handshake(fsm: &mut QlFsm, route: RouteHeader, record: QlHandshakeRecord) { + debug_assert!(fsm.state.handshake.is_none()); + fsm.state.handshake = Some((route, 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, + route: RouteHeader, + record: &QlHandshakeRecord, +) -> Result<(), ReceiveError> { + match record { + 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), + QlHandshakeRecord::Xx4(message) => xx::handle_xx4(fsm, crypto, route, 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, + handshake_id: HandshakeId, + 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 { + handshake_id, + transport, + session, + }); + emit_peer_status(fsm, fsm.state.link.status()); + Ok(()) +} + +pub fn establish_session( + fsm: &mut QlFsm, + handshake_id: HandshakeId, + finalized: wire::FinalizedHandshake, +) -> Result<(), ReceiveError> { + let transport = SessionTransport { + remote_qid: finalized.remote_bundle.qid, + tx_key: finalized.tx_key, + rx_key: finalized.rx_key, + 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; + } +} + +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 { + 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) +} diff --git a/ql-fsm/src/handshake/xx.rs b/ql-fsm/src/handshake/xx.rs new file mode 100644 index 00000000..3c0cbc91 --- /dev/null +++ b/ql-fsm/src/handshake/xx.rs @@ -0,0 +1,238 @@ +use ql_common::QID; +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, +}; +use crate::{ + state::{InitiatorState, LinkState, XxResponderState}, + QlFsm, ReceiveError, ReceiveStage, +}; + +pub fn start_initiator( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + token: PairingToken, + remote_qid: QID, +) { + let handshake_id = super::next_handshake_id(fsm); + let route = RouteHeader { + sender: fsm.identity.qid, + recipient: remote_qid, + }; + 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, handshake_id).unwrap(); + + fsm.state.link = LinkState::XxInitiator(InitiatorState { + handshake, + deadline: fsm.state.now + fsm.config.handshake_timeout, + }); + 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, route, message) { + return Ok(()); + } + match fsm.state.armed_pairing_token { + Some(expected) if expected.id(crypto) != message.pairing_id => { + Err(ReceiveError::InvalidPairingId) + } + Some(token) => { + reset_connected_session_if_needed(fsm); + + let mut handshake = wire::XxHandshake::new_responder( + crypto, + fsm.identity.clone(), + route.sender, + token, + super::local_transport_params(fsm), + ); + handshake + .read_1(crypto, route, message) + .map_err(wire_error)?; + let outbound = handshake + .write_2(crypto, message.handshake_id) + .map_err(wire_error)?; + fsm.state.link = LinkState::XxResponder(XxResponderState { + handshake, + deadline: fsm.state.now + fsm.config.handshake_timeout, + }); + fsm.state.handshake = None; + enqueue_handshake( + fsm, + RouteHeader { + sender: fsm.identity.qid, + recipient: route.sender, + }, + QlHandshakeRecord::Xx2(outbound), + ); + Ok(()) + } + None => Err(ReceiveError::NotPairingMode), + } +} + +pub fn handle_xx2( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + route: RouteHeader, + message: &Xx2, +) -> Result<(), ReceiveError> { + { + let LinkState::XxInitiator(state) = &mut fsm.state.link else { + return Ok(()); + }; + + if state.handshake.handshake_id() != Some(message.handshake_id) { + return Ok(()); + } + + state + .handshake + .read_2(crypto, route, message) + .map_err(wire_error)?; + let outbound = state + .handshake + .write_3(crypto, message.handshake_id) + .map_err(wire_error)?; + fsm.state.handshake = None; + enqueue_handshake( + fsm, + RouteHeader { + sender: fsm.identity.qid, + recipient: route.sender, + }, + QlHandshakeRecord::Xx3(outbound), + ); + } + + Ok(()) +} + +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 { + return Ok(()); + }; + + if state.handshake.handshake_id() != Some(message.handshake_id) { + return Ok(()); + } + + state + .handshake + .read_3(crypto, route, message) + .map_err(wire_error)?; + let LinkState::XxResponder(mut state) = fsm.state.link.take() else { + unreachable!("active XX responder was checked above"); + }; + let outbound = state + .handshake + .write_4(crypto, message.handshake_id) + .map_err(wire_error)?; + fsm.state.handshake = None; + enqueue_handshake( + fsm, + RouteHeader { + sender: fsm.identity.qid, + recipient: route.sender, + }, + QlHandshakeRecord::Xx4(outbound), + ); + establish_session( + fsm, + message.handshake_id, + state.handshake.finalize(crypto).map_err(wire_error)?, + ) +} + +pub fn handle_xx4( + fsm: &mut QlFsm, + crypto: &impl QlCrypto, + route: RouteHeader, + message: &Xx4, +) -> Result<(), ReceiveError> { + { + let LinkState::XxInitiator(state) = &mut fsm.state.link else { + return Ok(()); + }; + + if state.handshake.handshake_id() != Some(message.handshake_id) { + return Ok(()); + } + + state + .handshake + .read_4(crypto, route, message) + .map_err(wire_error)?; + } + + let LinkState::XxInitiator(state) = fsm.state.link.take() else { + unreachable!("active XX initiator was checked above"); + }; + establish_session( + fsm, + message.handshake_id, + state.handshake.finalize(crypto).map_err(wire_error)?, + ) +} + +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, + route: RouteHeader, + message: &Xx1, +) -> bool { + match &fsm.state.link { + LinkState::Idle => false, + LinkState::Connected(_) => { + super::is_connected_replay(fsm, message.handshake_id, route.sender) + } + LinkState::IkInitiator(_) | LinkState::XxResponder(_) => true, + LinkState::XxInitiator(state) => { + if state.handshake.pairing_id(crypto) != message.pairing_id { + return false; + } + if route.sender != state.handshake.remote_qid() { + return false; + } + super::local_start_wins( + state + .handshake + .local_ephemeral() + .expect("initiator has sent message 1"), + &message.ephemeral, + ) + } + } +} + +fn wire_error(source: ql_wire::Error) -> ReceiveError { + ReceiveError::wire(ReceiveStage::XxHandshake, source) +} diff --git a/ql-fsm/src/lib.rs b/ql-fsm/src/lib.rs new file mode 100644 index 00000000..02964bdb --- /dev/null +++ b/ql-fsm/src/lib.rs @@ -0,0 +1,359 @@ +//! 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_common::{ResetCode, StreamId}; +use ql_wire::{PairingToken, PeerBundle, QlCrypto, QlIdentity, SessionClose, SessionCloseCode}; +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(StreamId), + /// 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), + /// 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 + /// one final write remains: a record containing only `SessionFrame::Close` + /// the FSM does not wait for an ack for that record + 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); + +/// 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() + } + + pub fn header(&self) -> &[u8] { + self.inner.header() + } + + /// 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() + } + + /// resets the local read side, write side, or both sides of the stream + pub fn reset(&mut self, target: StreamResetTarget, code: ResetCode) { + self.inner.reset(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, + /// 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, + /// 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, header: Box<[u8]>) -> Result, NoSessionError> { + fsm::open_stream(self, header) + } + + /// 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..12eec3e8 --- /dev/null +++ b/ql-fsm/src/pairing.rs @@ -0,0 +1,16 @@ +use ql_common::QID; +use ql_wire::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/session/ack_tracker.rs b/ql-fsm/src/session/ack_tracker.rs new file mode 100644 index 00000000..7f947d2d --- /dev/null +++ b/ql-fsm/src/session/ack_tracker.rs @@ -0,0 +1,262 @@ +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.0; + 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(range.start)..=RecordSeq(end) +} + +fn from_ack_range(range: RangeInclusive) -> std::ops::Range { + let start = range.start().0; + let end = range.end().0.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 ack_ranges(pending_ack: &PendingAck) -> Vec<(u64, u64)> { + pending_ack + .ack + .ranges() + .map(|range| (range.start().0, range.end().0)) + .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(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(); + 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(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(); + 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(RecordSeq(10)), ReceiveOutcome::New); + assert_eq!(ack_tracker.insert(RecordSeq(15)), ReceiveOutcome::New); + + assert_eq!(ack_tracker.insert(RecordSeq(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(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(); + 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(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(); + 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..a24da78f --- /dev/null +++ b/ql-fsm/src/session/mod.rs @@ -0,0 +1,1063 @@ +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_codec::Varint; +use ql_common::StreamId; +use ql_wire::{ + RecordAck, RecordSeq, ResetTarget, SessionClose, SessionCloseCode, SessionFrame, + SessionRecordBuilder, StreamData, StreamReset, StreamWindow, +}; + +use self::{ + ack_tracker::{AckTracker, PendingAck, ReceiveOutcome}, + remote_stream_history::RemoteStreamHistory, + state::{InboundState, OutboundState, SessionPhase, SessionState, StreamRole, StreamState}, + stream_tx::StreamTxRange, + tracked::{LossRecovery, TrackedFrame, TrackedRecord, TrackedStreamData}, +}; +use crate::{NoSessionError, StreamError, StreamResetEvent, StreamResetTarget}; + +#[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_secs(1), + 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(StreamId), + Readable(StreamId), + Writable(StreamId), + Finished(StreamId), + OutboundFinished(StreamId), + Reset(StreamResetEvent), + 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(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, + ), + 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, + header: Box<[u8]>, + 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(Bytes::from(header)), + 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) = (|| { + 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)) + } + + 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, ql_wire::Error>>, + { + 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(now, &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::StreamReset(frame) => { + if self.handle_stream_reset(&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 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 + rto)) + .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)> { + const TRACKED_RECORD_LIMIT: usize = 64; + + 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); + + // 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); + 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_reset(&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_reset( + &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(reset) = stream.pending_reset.as_ref() else { + continue; + }; + if !builder.push_stream_reset(reset) { + break; + } + + outbound.frames.push(TrackedFrame::StreamReset( + stream.pending_reset.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(stream.recv_limit()), + }; + if !builder.push_stream_window(&frame) { + break; + } + + stream.pending_window = false; + stream.advertised_max_offset = *frame.maximum_offset; + outbound + .window_updates + .push((stream_id, *frame.maximum_offset)); + } + } + + fn push_next_stream_data( + &mut self, + builder: &mut SessionRecordBuilder, + outbound: &mut TrackedRecord, + ) { + const OVERHEAD: usize = 1 + StreamData::>::MAX_WIRE_OVERHEAD; + + 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; + } + // 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; + }; + let frame = StreamData { + stream_id, + offset: Varint(candidate.offset), + header: if candidate.offset == 0 { header } 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, now: Instant, ack: &RecordAck, sink: &mut impl EventSink) { + let stream_send_buffer_size = self.config.stream_send_buffer_size; + 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 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(); + } + + 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 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 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( + &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; + 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 = match (stream.role, stream.header.as_ref(), header, frame_offset) { + (StreamRole::Responder, None, Some(header), 0) => { + stream.header = Some(header); + true + } + (StreamRole::Initiator, _, Some(_), _) + | (StreamRole::Responder, None, Some(_), _) + | (StreamRole::Responder, None, None, 0) => return Err(()), + _ => false, + }; + + match stream.inbound_state { + InboundState::Open => {} + 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 { + 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 opened { + sink.emit(SessionEvent::Opened(stream_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 opened { + sink.emit(SessionEvent::Opened(stream_id)); + } + + if stream.header.is_some() && readable_before == 0 && stream.readable_bytes() > 0 { + sink.emit(SessionEvent::Readable(stream_id)); + } + + if stream.header.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; + 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_reset( + &mut self, + frame: &StreamReset, + 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(()), + }, + }; + + 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(); + } + if outbound { + stream.outbound_state = OutboundState::Closed; + stream.tx.clear(); + stream.pending_reset = None; + } + 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(()) + } + + 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(); + } + if Self::target_affects_outbound(stream.role, target) { + stream.outbound_state = OutboundState::Closed; + stream.tx.clear(); + } + } + + fn target_affects_inbound(role: StreamRole, target: ResetTarget) -> bool { + matches!(target, ResetTarget::Both) || role.inbound_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 { + 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::StreamReset(frame) => frame.stream_id == stream_id, + }) + }); + if tracked_refs_stream { + return false; + } + + if !stream.tx.is_empty() + || stream.pending_reset.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::Reset(_) | 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.0 < local_parity.make_stream_id(next_stream_ordinal).0 +} + +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::StreamReset(reset) => restore_stream_reset(streams, reset.clone()), + TrackedFrame::StreamData(frame) => restore_stream_data(streams, *frame), + } +} + +fn restore_stream_reset(streams: &mut IndexMap, reset: StreamReset) { + if let Some(stream) = streams.get_mut(&reset.stream_id) { + stream.pending_reset = Some(reset); + } +} + +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::StreamReset(_) => {} + 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 + .0 + .checked_add(1) + .map(RecordSeq) + .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..1983680f --- /dev/null +++ b/ql-fsm/src/session/remote_stream_history.rs @@ -0,0 +1,60 @@ +use ql_common::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 + .0 + .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..7e2e8a07 --- /dev/null +++ b/ql-fsm/src/session/state.rs @@ -0,0 +1,146 @@ +use std::time::Instant; + +use bytes::Bytes; +use indexmap::IndexMap; +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::{LossRecovery, 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 loss_recovery: LossRecovery, + 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 header: Option, + pub rx: StreamRx, + pub tx: StreamTx, + pub pending_reset: 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, + 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, + tx: StreamTx::new(), + pending_reset: 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) -> ResetTarget { + match self { + Self::Initiator => ResetTarget::Origin, + Self::Responder => ResetTarget::Return, + } + } + + pub fn inbound_target(self) -> ResetTarget { + match self { + Self::Initiator => ResetTarget::Return, + Self::Responder => ResetTarget::Origin, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum OutboundState { + Open, + FinQueued, + Finished, + Closed, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum InboundState { + Open, + Finished, + Reset(StreamReset), + 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..5ab1c1b8 --- /dev/null +++ b/ql-fsm/src/session/stream_ops.rs @@ -0,0 +1,162 @@ +use ql_common::{ResetCode, StreamId}; +use ql_wire::StreamReset; + +use super::{ + state::{InboundState, StreamState}, + stream_rx::StreamReadIter, + EventSink, SessionEvent, SessionFsm, +}; +use crate::{CommitReadError, StreamResetTarget}; + +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 + #[inline] + pub fn stream_id(&self) -> StreamId { + self.stream_id + } + + /// returns the streams details + #[inline] + pub fn header(&self) -> &[u8] { + 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() + } + + /// 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.header.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)) + } + + /// 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(); + 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: wire_target, + code, + }); + 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] + } +} + +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..7ad63957 --- /dev/null +++ b/ql-fsm/src/session/stream_parity.rs @@ -0,0 +1,44 @@ +use ql_common::{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.0 % 2 == 0, + Self::Odd => stream_id.0 % 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(u64::from( + 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..ca29906e --- /dev/null +++ b/ql-fsm/src/session/stream_tx.rs @@ -0,0 +1,586 @@ +use std::{collections::VecDeque, ops::Range}; + +use bytes::{Buf, Bytes}; +use ql_codec::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 + } + + /// 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() + } + + 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..6599ce8d --- /dev/null +++ b/ql-fsm/src/session/tests.rs @@ -0,0 +1,1046 @@ +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, + SessionRecordBuilder, StreamData, StreamReset, +}; + +use super::{SessionConfig, SessionEvent, SessionFsm}; +use crate::{session::stream_parity::StreamParity, StreamResetEvent}; + +const REFUSED: ResetCode = ResetCode(1); +const TIMEOUT: ResetCode = ResetCode(2); + +fn open_stream_id(fsm: &mut SessionFsm) -> StreamId { + fsm.open_stream(Box::from([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, RecordSeq(0)); + assert_eq!(second_seq, RecordSeq(1)); +} + +#[test] +fn retransmit_uses_new_record_seq() { + 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"retry"), 5); + let (first_seq, first) = next_outbound(&mut fsm, now).unwrap(); + + let mut emit = |_| {}; + 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::>::MAX_WIRE_OVERHEAD + + 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; + + let now = Instant::now(); + let mut fsm = SessionFsm::new( + SessionConfig { + record_max_size: SessionRecordBuilder::MIN_CAPACITY + + 1 // discriminator byte + + StreamData::>::MAX_WIRE_OVERHEAD + + 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'; PAYLOAD_LEN]; + let payload_b = vec![b'b'; PAYLOAD_LEN]; + + 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(); + 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), + RecordSeq(9), + std::iter::once(Ok(SessionFrame::Ack( + RecordAck::from_ranges([record_seq..=record_seq]).unwrap(), + ))), + &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), + RecordSeq(9), + std::iter::once(Ok(SessionFrame::Ack( + RecordAck::from_ranges([record_seq..=record_seq]).unwrap(), + ))), + &mut emit, + ); + } + assert_eq!(events, vec![SessionEvent::OutboundFinished(stream_id)]); + + { + let mut emit = |event| events.push(event); + fsm.receive( + now + Duration::from_millis(2), + RecordSeq(10), + std::iter::once(Ok(SessionFrame::Ack( + RecordAck::from_ranges([record_seq..=record_seq]).unwrap(), + ))), + &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 = StreamId(1); + let data = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: Varint(0), + header: Some(vec![1_u8]), + fin: false, + bytes: b"hi".to_vec(), + })]; + let events = receive_events(&mut fsm, now, RecordSeq(7), &data); + assert_eq!( + events, + vec![ + SessionEvent::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 = StreamId(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: Varint(0), + header: Some(vec![1_u8]), + fin: false, + bytes: b"hi".to_vec(), + })]; + + 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(); + 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 = StreamId(1); + let record = vec![SessionFrame::StreamData(ql_wire::StreamData { + stream_id, + offset: Varint(0), + header: Some(vec![1_u8]), + fin: true, + bytes: b"hello".to_vec(), + })]; + + let events = receive_events(&mut fsm, now, RecordSeq(0), &record); + assert_eq!( + events, + vec![ + SessionEvent::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 = StreamId(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: Varint(0), + header: Some(vec![1_u8]), + fin: true, + bytes: Vec::new(), + })]; + + let events = receive_events(&mut fsm, now, RecordSeq(0), &record); + assert_eq!( + events, + vec![ + SessionEvent::Opened(stream_id), + SessionEvent::Finished(stream_id) + ] + ); +} + +#[test] +fn remote_stream_reset_is_reliable_and_retried() { + 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); + + fsm.stream(stream_id, |_| {}) + .unwrap() + .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); + let first = decode_session_frames(builder.bytes()).unwrap(); + assert!(matches!( + first.as_slice(), + [SessionFrame::StreamReset(StreamReset { stream_id: id, .. })] if *id == stream_id + )); + + let mut emit = |_| {}; + 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_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(vec![1_u8].into_boxed_slice(), |_| {}) + .unwrap() + .stream_id(); + let odd_id = SessionFsm::new( + SessionConfig { + local_parity: odd, + ..SessionConfig::default() + }, + now, + ) + .open_stream(vec![1_u8].into_boxed_slice(), |_| {}) + .unwrap() + .stream_id(); + + 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 = StreamId(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: Varint(0), + header: Some(vec![1_u8]), + fin: false, + bytes: b"hi".to_vec(), + })]; + 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()); +} + +#[test] +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: StreamId(1), + target: ResetTarget::Both, + code: ResetCode(9), + }; + let record = vec![SessionFrame::StreamReset(reset.clone())]; + + let first = receive_events(&mut fsm, now, RecordSeq(1), &record); + assert_eq!( + first, + 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), + RecordSeq(2), + &record, + ); + assert!(second.is_empty()); +} + +#[test] +fn late_remote_stream_data_after_reset_is_ignored() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + let stream_id = StreamId(1); + let reset = vec![SessionFrame::StreamReset(StreamReset { + stream_id, + target: ResetTarget::Both, + code: ResetCode(9), + })]; + let data = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: Varint(0), + header: Some(vec![1_u8]), + fin: false, + bytes: b"hello".to_vec(), + })]; + + let first = receive_events(&mut fsm, now, RecordSeq(1), &reset); + assert_eq!( + first, + vec![SessionEvent::Reset(StreamResetEvent { + stream_id, + code: ResetCode(9), + target: crate::StreamResetTarget::Both, + })] + ); + + let second = receive_events( + &mut fsm, + now + Duration::from_millis(1), + RecordSeq(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 = StreamId(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: Varint(0), + header: Some(vec![1_u8]), + fin: true, + bytes: b"hello".to_vec(), + })]; + + let first = receive_events(&mut fsm, now, RecordSeq(1), &record); + assert_eq!( + first, + vec![ + SessionEvent::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), + RecordSeq(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 = StreamId(1); + let record = vec![SessionFrame::StreamData(StreamData { + stream_id, + offset: Varint(0), + header: Some(vec![1_u8]), + fin: true, + bytes: b"hello".to_vec(), + })]; + + let first = receive_events(&mut fsm, now, RecordSeq(1), &record); + assert_eq!( + first, + vec![ + SessionEvent::Opened(stream_id), + SessionEvent::Readable(stream_id) + ] + ); + + 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!( + 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 reset3 = vec![SessionFrame::StreamReset(StreamReset { + stream_id: StreamId(3), + target: ResetTarget::Both, + code: REFUSED, + })]; + let reset1 = vec![SessionFrame::StreamReset(StreamReset { + stream_id: StreamId(1), + target: ResetTarget::Both, + code: TIMEOUT, + })]; + + let first = receive_events(&mut fsm, now, RecordSeq(1), &reset3); + assert_eq!( + first, + vec![SessionEvent::Reset(StreamResetEvent { + stream_id: StreamId(3), + code: REFUSED, + target: crate::StreamResetTarget::Both, + })] + ); + + let second = receive_events( + &mut fsm, + now + Duration::from_millis(1), + RecordSeq(2), + &reset1, + ); + assert_eq!( + second, + vec![SessionEvent::Reset(StreamResetEvent { + stream_id: StreamId(1), + code: TIMEOUT, + target: crate::StreamResetTarget::Both, + })] + ); + + let third = receive_events( + &mut fsm, + now + Duration::from_millis(2), + RecordSeq(3), + &reset3, + ); + assert!(third.is_empty()); +} + +#[test] +fn invalid_remote_stream_reset_closes_session() { + let now = Instant::now(); + let mut fsm = SessionFsm::new(SessionConfig::default(), now); + + let invalid = vec![SessionFrame::StreamReset(StreamReset { + stream_id: StreamId(0), + target: ResetTarget::Both, + code: ResetCode(9), + })]; + let events = receive_events(&mut fsm, now, RecordSeq(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: StreamId(0), + offset: Varint(0), + header: Some(vec![1_u8]), + fin: false, + bytes: b"bad".to_vec(), + })]; + let events = receive_events(&mut fsm, now, RecordSeq(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), + RecordSeq(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, RecordSeq(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), + RecordSeq(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), + RecordSeq(9), + &[SessionFrame::StreamWindow(ql_wire::StreamWindow { + stream_id, + maximum_offset: Varint(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 == 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 + + 1 // discriminator byte + + 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), + 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'; 1200]; + assert_eq!( + write_stream_bytes(&mut sender, stream_id, &payload), + payload.len() + ); + + let originals = drain_outbound(&mut sender, now, 4096); + 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); + } + + 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()); +} + +#[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()); +} diff --git a/ql-fsm/src/session/tracked.rs b/ql-fsm/src/session/tracked.rs new file mode 100644 index 00000000..f8a234ca --- /dev/null +++ b/ql-fsm/src/session/tracked.rs @@ -0,0 +1,81 @@ +//! outbound record tracking state for ack and retransmit handling + +use std::time::{Duration, Instant}; + +use ql_common::StreamId; +use ql_wire::{RecordAck, RecordSeq, StreamReset}; + +#[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), + StreamReset(StreamReset), +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct TrackedStreamData { + pub stream_id: StreamId, + pub offset: u64, + 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); + } +} diff --git a/ql-fsm/src/state.rs b/ql-fsm/src/state.rs new file mode 100644 index 00000000..8ceeb711 --- /dev/null +++ b/ql-fsm/src/state.rs @@ -0,0 +1,102 @@ +use std::time::Instant; + +use ql_common::QID; +use ql_wire::{ + HandshakeId, IkHandshake, PairingToken, PeerBundle, QlHandshakeRecord, RouteHeader, 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<(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 remote_transport_params: TransportParams, +} + +#[allow(clippy::large_enum_variant)] +pub enum LinkState { + Idle, + IkInitiator(InitiatorState), + XxInitiator(InitiatorState), + XxResponder(XxResponderState), + Connected(ConnectedState), +} + +pub struct ConnectedState { + pub handshake_id: HandshakeId, + pub transport: SessionTransport, + pub session: SessionFsm, +} + +#[derive(Debug, Clone)] +pub struct InitiatorState { + pub handshake: H, + pub deadline: Instant, +} + +#[derive(Debug, Clone)] +pub struct XxResponderState { + pub handshake: XxHandshake, + 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::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::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..3ffe083a --- /dev/null +++ b/ql-fsm/src/tests/handshake.rs @@ -0,0 +1,383 @@ +use std::time::Duration; + +use ql_wire::{IkPattern, 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"); + 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 { + version: PairingInvite::VERSION, + qid: ql_common::QID([2; ql_common::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() { + 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)); +} + +#[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::IkInitiator(ref state) if state.handshake.pattern() == IkPattern::Kk + )); + + 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").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.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-fsm/src/tests/mod.rs b/ql-fsm/src/tests/mod.rs new file mode 100644 index 00000000..60986ae3 --- /dev/null +++ b/ql-fsm/src/tests/mod.rs @@ -0,0 +1,352 @@ +mod handshake; +mod proptest; +mod session; + +use std::time::{Duration, Instant}; + +use ql_common::QID; +use ql_wire::{ + self, generate_identity, test_identities, HandshakeId, PairingToken, QlCrypto, SessionKey, + SoftwareCrypto, TransportParams, +}; + +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([7; SessionKey::SIZE]); + let b_to_a_key = SessionKey([9; SessionKey::SIZE]); + + 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(), + 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 { + 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, + 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 { + version: PairingInvite::VERSION, + 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, + &header, + &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..3e397b2d --- /dev/null +++ b/ql-fsm/src/tests/proptest.rs @@ -0,0 +1,1004 @@ +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_common::{ResetCode, StreamId}; + +use super::*; +use crate::{ + state::LinkState, Event, PeerStatus, ReceiveError, ReceiveStage, StreamResetTarget, 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, + }, + Reset { + 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 reset(side: Side, slot: usize) -> Self { + Self::Reset { side, slot } + } +} + +#[derive(Clone, Debug)] +struct TakenWrite { + record: Vec, + write_id: Option, +} + +#[derive(Default)] +struct SideEventState { + opened: BTreeSet, + finished: BTreeSet, + outbound_finished: BTreeSet, + writable_reset: BTreeSet, + reset: 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], + reset_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()], + reset_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()], + reset_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(Box::from([1])) + .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::Reset { side, slot } => { + if let Some(stream_id) = self.slots[side.idx()][*slot] { + let reset = self + .harness + .node_mut(*side) + .fsm + .stream(stream_id) + .is_ok_and(|mut stream| { + stream.reset(StreamResetTarget::Both, ResetCode::CANCELLED); + true + }); + if reset { + self.reset_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()].reset.contains(&stream_id), + "side {side:?} emitted Finished after Reset 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::Reset(reset) => { + prop_assert!( + self.known_streams.contains(&reset.stream_id), + "side {side:?} emitted Reset for unknown stream {:?}", + 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()]; + 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::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:?}" + ); + } + + 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()].reset.contains(stream_id) + || !connected[opposite(side).idx()], + "side {side:?} finished {stream_id:?} but side {:?} saw neither Finished nor Reset", + opposite(side) + ); + } + + for stream_id in &self.reset_by[side.idx()] { + prop_assert!( + self.events[opposite(side).idx()].reset.contains(stream_id) + || !connected[opposite(side).idx()], + "side {side:?} reset {stream_id:?} but side {:?} saw no Reset 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.reset.is_empty() + && events.writable_reset.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()].reset.contains(&stream_id) + || self.reset_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::reset), + ] +} + +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::reset), + 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: ql_wire::SessionRecordBuilder::MIN_CAPACITY + 94, + 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..2f7386a6 --- /dev/null +++ b/ql-fsm/src/tests/session.rs @@ -0,0 +1,522 @@ +use std::time::Duration; + +use bytes::Bytes; +use ql_common::StreamId; +use ql_wire::SessionClose; + +use super::*; +use crate::{state::LinkState, CommitReadError, Event, NoSessionError, PeerStatus, StreamError}; + +fn open_stream_id(fsm: &mut QlFsm) -> StreamId { + fsm.open_stream(Box::from([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 mut writer = stream.writer().expect("stream is not writable"); + 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(Event::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(Event::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(Event::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(Event::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 = StreamId(0); + + assert!(matches!( + harness.a.fsm.open_stream(Box::from([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.reset( + crate::StreamResetTarget::Both, + ql_common::ResetCode::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 = StreamId(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(Event::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(Box::from([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(Box::from([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(Box::from([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)) + ); +} diff --git a/ql-rpc/Cargo.toml b/ql-rpc/Cargo.toml new file mode 100644 index 00000000..40af473c --- /dev/null +++ b/ql-rpc/Cargo.toml @@ -0,0 +1,11 @@ +[package] +name = "ql-rpc" +version = "0.1.0" +edition = "2021" +description = "QuantumLink RPC protocol" +license = "Proprietary" + +[dependencies] +bytes = { workspace = true } +ql-common = { 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..3a6f1559 --- /dev/null +++ b/ql-rpc/src/chunk_queue.rs @@ -0,0 +1,259 @@ +use std::collections::VecDeque; + +use bytes::{Buf, Bytes}; + +use crate::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<(), Error> { + if self.remaining > 0 { + Err(Error::TrailingBytes) + } else { + Ok(()) + } + } + + 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() { + 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<(), Error> { + if self.remaining > 0 { + Err(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..b2718c01 --- /dev/null +++ b/ql-rpc/src/codec.rs @@ -0,0 +1,90 @@ +use std::{convert::Infallible, str::Utf8Error}; + +use bytes::{Buf, BufMut, Bytes}; + +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); +} + +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(); + 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..6c176a7f --- /dev/null +++ b/ql-rpc/src/error.rs @@ -0,0 +1,87 @@ +use ql_common::ResetCode; + +#[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 {} + +impl Error { + pub const fn reset_code(self) -> ResetCode { + match self { + Self::LengthOverflow => ResetCode::LIMIT, + Self::Truncated + | Self::UnexpectedFrameKind(_) + | Self::MissingResponse + | Self::TrailingBytes => ResetCode::PROTOCOL, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum RpcError { + Protocol(Error), + Codec(C), + Transport(T), +} + +impl RpcError { + pub const fn reset_code(&self) -> Option { + match self { + Self::Protocol(error) => Some(error.reset_code()), + Self::Codec(_) => Some(ResetCode::CODEC), + Self::Transport(_) => None, + } + } +} + +impl std::fmt::Display for RpcError +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 RpcError +where + C: std::error::Error + 'static, + T: 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), + Self::Transport(error) => Some(error), + } + } +} + +impl From for RpcError { + fn from(error: Error) -> Self { + Self::Protocol(error) + } +} diff --git a/ql-rpc/src/lib.rs b/ql-rpc/src/lib.rs new file mode 100644 index 00000000..57fc4ba0 --- /dev/null +++ b/ql-rpc/src/lib.rs @@ -0,0 +1,17 @@ +#![allow(clippy::type_complexity)] + +//! QuantumLink RPC protocol + +mod chunk_queue; +mod codec; +mod error; +mod router; +mod rpc; +mod stream; + +pub use chunk_queue::ChunkQueue; +pub use codec::RpcCodec; +pub use error::*; +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 new file mode 100644 index 00000000..0c1f9dd7 --- /dev/null +++ b/ql-rpc/src/router/builder.rs @@ -0,0 +1,232 @@ +use std::marker::PhantomData; + +use super::*; +use crate::{ + download::*, duplex::*, notification::*, progress::*, request::*, subscription::*, upload::*, + Route, +}; + +pub struct LocalRoutes; +pub struct SendRoutes; + +pub struct RouterBuilder +where + K: RpcRouteKey, + Sp: Spawner, +{ + config: RouterConfig, + spawner: Sp, + routes: Vec>, + marker: PhantomData Mode>, +} + +impl RouterBuilder +where + K: RpcRouteKey, + 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.key.clone()); + self.routes.shrink_to_fit(); + Router { + config: self.config, + state, + spawner: self.spawner, + routes: self.routes, + } + } + + fn add_route(mut self, key: K, route: RouteFn) -> Self { + if self.routes.iter().any(|entry| entry.key == key) { + panic!("duplicate rpc route {key:?}"); + } + self.routes.push(RouteEntry::new(key, route)); + self + } +} + +impl RouterBuilder +where + K: RpcRouteKey, + Sp: LocalSpawner, + St: RpcStream + 'static, +{ + pub fn request(self) -> Self + where + M: Request + 'static, + S: RequestHandlerLocal + 'static, + { + add_route!(self, M, handle_request, S::handle, S::handle_error) + } + + pub fn notification(self) -> Self + where + M: Notification + 'static, + S: NotificationHandlerLocal + 'static, + { + add_route!(self, M, handle_notification, S::handle, S::handle_error) + } + + pub fn duplex(self) -> Self + where + M: Duplex + 'static, + S: DuplexHandlerLocal + 'static, + { + add_route!(self, M, handle_duplex, S::handle) + } + + pub fn download(self) -> Self + where + M: Download + 'static, + S: DownloadHandlerLocal + 'static, + { + add_route!(self, M, handle_download, S::handle, S::handle_error) + } + + pub fn subscription(self) -> Self + where + M: Subscription + 'static, + S: SubscriptionHandlerLocal + 'static, + { + add_route!(self, M, handle_subscription, S::handle, S::handle_error) + } + + pub fn progress(self) -> Self + where + M: Progress + 'static, + S: ProgressHandlerLocal + 'static, + { + add_route!(self, M, handle_progress, S::handle, S::handle_error) + } + + pub fn upload(self) -> Self + where + M: Upload + 'static, + S: UploadHandlerLocal + 'static, + { + add_route!(self, M, handle_upload, S::handle, S::handle_error) + } +} + +impl RouterBuilder +where + K: RpcRouteKey, + Sp: SendSpawner + Send, + St: RpcStream + 'static, +{ + pub fn request(self) -> Self + where + M: Request + 'static, + M::Request: Send + 'static, + S: RequestHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + add_route!(self, M, handle_request, S::handle, S::handle_error) + } + + pub fn notification(self) -> Self + where + M: Notification + 'static, + M::Payload: Send + 'static, + S: NotificationHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + add_route!(self, M, handle_notification, S::handle, S::handle_error) + } + + pub fn duplex(self) -> Self + where + M: Duplex + 'static, + M::InitiatorEvent: Send + 'static, + M::ResponderEvent: Send + 'static, + S: DuplexHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + add_route!(self, M, handle_duplex, S::handle) + } + + pub fn download(self) -> Self + where + M: Download + 'static, + M::Request: Send + 'static, + S: DownloadHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + add_route!(self, M, handle_download, S::handle, S::handle_error) + } + + pub fn subscription(self) -> Self + where + M: Subscription + 'static, + M::Request: Send + 'static, + S: SubscriptionHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + add_route!(self, M, handle_subscription, S::handle, S::handle_error) + } + + pub fn progress(self) -> Self + where + M: Progress + 'static, + M::Request: Send + 'static, + S: ProgressHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + add_route!(self, M, handle_progress, S::handle, S::handle_error) + } + + pub fn upload(self) -> Self + where + M: Upload + 'static, + M::Request: Send + 'static, + S: UploadHandler + Send + 'static, + St::Reader: Send + 'static, + St::Writer: Send + 'static, + { + 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( + <$rpc as Route>::key(), + |spawner, state, context, config, stream| { + spawner.spawn($handler( + state, + context, + config, + stream, + $($arg),+ + )) + }, + ) + }; +} + +use add_route; 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..0b867c6d --- /dev/null +++ b/ql-rpc/src/router/mod.rs @@ -0,0 +1,102 @@ +use ql_common::{ResetCode, StreamId, StreamInfo, QID}; + +mod builder; +mod config; +mod mode; + +pub use self::{builder::*, config::*, mode::*}; +use crate::{RpcRead, RpcRouteKey, RpcStream, RpcWrite}; + +pub struct Router +where + K: RpcRouteKey, + Sp: Spawner, +{ + config: RouterConfig, + state: S, + spawner: Sp, + routes: Vec>, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub struct Context { + pub qid: QID, + pub stream_id: StreamId, +} + +struct RouteEntry +where + K: RpcRouteKey, + Sp: Spawner, +{ + key: K, + route: RouteFn, +} + +impl RouteEntry +where + K: RpcRouteKey, + Sp: Spawner, +{ + fn new(key: K, route: RouteFn) -> Self { + Self { key, route } + } +} + +impl Router +where + K: RpcRouteKey, + 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, info: StreamInfo, stream: St) -> Option { + let StreamInfo { + qid, + stream_id, + header, + } = info; + let context = Context { qid, stream_id }; + 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.clone()) + else { + 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( + &self.spawner, + self.state.clone(), + context, + self.config, + stream, + )) + } + + pub fn route_keys(&self) -> impl ExactSizeIterator + '_ { + self.routes.iter().map(|entry| &entry.key) + } +} diff --git a/ql-rpc/src/router/mode.rs b/ql-rpc/src/router/mode.rs new file mode 100644 index 00000000..3a61ae4b --- /dev/null +++ b/ql-rpc/src/router/mode.rs @@ -0,0 +1,22 @@ +use std::future::Future; + +use super::Context; +use crate::RouterConfig; + +pub type RouteFn = fn(&Sp, S, Context, 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..81bbd615 --- /dev/null +++ b/ql-rpc/src/rpc/download/client.rs @@ -0,0 +1,214 @@ +use std::{future::poll_fn, marker::PhantomData}; + +use bytes::Bytes; +use ql_common::ResetCode; + +use crate::{ + download::Download, + parts::{PartFrameReader, PartReadStep}, + rpc::{parts::FrameKind, read_framed_prefix, write_eof_value}, + DropResetRead, RpcError, RpcRead, RpcStream, +}; + +pub async fn start( + stream: St, + request: &M::Request, +) -> Result, RpcError> +where + M: Download, + St: RpcStream, +{ + let (reader, mut writer) = stream.split(); + write_eof_value(&mut writer, request) + .await + .map_err(RpcError::Transport)?; + Ok(DownloadCall::new(reader)) +} + +pub struct DownloadCall +where + M: Download, + R: RpcRead, +{ + stream: DropResetRead, + marker: PhantomData M>, +} + +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: DropResetRead, + reader: PartFrameReader, +} + +impl DownloadCall +where + M: Download, + R: RpcRead, +{ + pub fn new(stream: R) -> Self { + Self { + stream: DropResetRead::new(stream), + marker: PhantomData, + } + } + + pub async fn start( + mut self, + ) -> Result<(M::ResponseHeader, DownloadReader), RpcError> { + 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) { + self.reset_inner(code); + } + + fn reset_inner(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.stream, code); + } +} + +impl DownloadReader +where + M: Download, + R: RpcRead, +{ + pub async fn next_part( + &mut self, + ) -> Result)>, RpcError> { + if !self.stream.is_some() { + return Ok(None); + } + + match self.read_frame().await? { + PartReadStep::PartHeader(value) => Ok(Some(( + value, + DownloadPart { + parent: self, + finished: false, + }, + ))), + PartReadStep::Finish => { + self.stream.disarm(); + 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<(), RpcError> { + match self.read_frame().await? { + PartReadStep::Finish => { + self.stream.disarm(); + 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 reset(mut self, code: ResetCode) { + self.reset_inner(code); + } + + async fn read_frame( + &mut self, + ) -> Result, RpcError> { + loop { + match self.reader.advance() { + Ok(PartReadStep::NeedMore) => {} + Ok(step) => return Ok(step), + Err(error) => return Err(error), + } + + match poll_fn(|cx| self.stream.poll_read(cx)).await { + Ok(Some(chunk)) => { + self.reader.push(chunk); + } + Ok(None) => return Err(crate::Error::Truncated.into()), + Err(error) => return Err(RpcError::Transport(error)), + } + } + } + + fn reset_inner(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.stream, code); + } +} + +impl DownloadPart<'_, M, R> +where + M: Download, + R: RpcRead, +{ + pub async fn read_chunk(&mut self) -> Result, RpcError> { + 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 reset(mut self, code: ResetCode) { + self.parent.reset_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.reset_inner(ResetCode::DROPPED); + } + } +} diff --git a/ql-rpc/src/rpc/download/mod.rs b/ql-rpc/src/rpc/download/mod.rs new file mode 100644 index 00000000..3e9a613e --- /dev/null +++ b/ql-rpc/src/rpc/download/mod.rs @@ -0,0 +1,23 @@ +use super::Route; +use crate::RpcCodec; + +mod client; +mod server; + +pub use self::{client::*, server::*}; + +/// 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..bfd95ae7 --- /dev/null +++ b/ql-rpc/src/rpc/download/server.rs @@ -0,0 +1,204 @@ +use std::{future::Future, marker::PhantomData}; + +use bytes::Bytes; +use ql_common::ResetCode; + +use crate::{ + codec, + download::Download, + finish_bytes, + rpc::{ + parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, + read_eof_request, + }, + write_bytes, Context, DropResetWrite, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, +}; + +#[trait_variant::make(DownloadHandler: Send)] +pub trait DownloadHandlerLocal +where + M: Download, + St: RpcStream, +{ + async fn handle( + self, + context: Context, + message: M::Request, + download: DownloadStart, + ); + + fn handle_error(&self, error: &RpcError) { + let _ = error; + } +} + +pub struct DownloadStart +where + M: Download, + W: RpcWrite, +{ + writer: DropResetWrite, + marker: PhantomData M>, +} + +pub struct DownloadWriter +where + M: Download, + W: RpcWrite, +{ + writer: DropResetWrite, + marker: PhantomData M>, +} + +pub struct DownloadPartWriter<'a, M, W> +where + M: Download, + W: RpcWrite, +{ + parent: &'a mut DownloadWriter, + finished: bool, +} + +impl DownloadStart +where + M: Download, + W: RpcWrite, +{ + pub(crate) fn new(writer: W) -> Self { + Self { + writer: DropResetWrite::new(writer), + marker: PhantomData, + } + } + + /// send the response header and begin streaming parts + pub async fn start( + 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); + write_bytes(&mut writer, Bytes::from(encoded)).await?; + Ok(DownloadWriter { + writer, + marker: PhantomData, + }) + } + + /// send a header-only response and finish the stream + 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); + write_bytes(&mut writer, Bytes::from(encoded)).await?; + finish_bytes(&mut writer).await + } + + /// reset the stream with a transport code + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); + } +} + +impl DownloadWriter +where + M: Download, + W: RpcWrite, +{ + pub async fn start_part( + &mut self, + part_header: M::PartHeader, + ) -> Result, W::Error> { + let writer = &mut self.writer; + 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(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?; + finish_bytes(&mut writer).await + } + + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); + } +} + +impl DownloadPartWriter<'_, M, W> +where + M: Download, + W: RpcWrite, +{ + pub async fn send(&mut self, bytes: Bytes) -> Result<(), W::Error> { + 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 = &mut self.parent.writer; + 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: Download, + W: RpcWrite, +{ + fn drop(&mut self) { + if !self.finished { + DropResetWrite::reset(&mut self.parent.writer, ResetCode::DROPPED); + } + } +} + +pub(crate) fn handle_download( + state: S, + context: Context, + config: RouterConfig, + stream: St, + handle: H, + handle_error: E, +) -> impl Future +where + M: Download + 'static, + St: RpcStream + 'static, + H: FnOnce(S, Context, M::Request, DownloadStart) -> HF, + HF: Future, + E: FnOnce(&S, &RpcError), +{ + 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; + } + }; + + 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 new file mode 100644 index 00000000..aa501e7f --- /dev/null +++ b/ql-rpc/src/rpc/duplex/client.rs @@ -0,0 +1,195 @@ +use std::{ + future::poll_fn, + marker::PhantomData, + task::{Context, Poll}, +}; + +use bytes::Bytes; +use ql_common::ResetCode; + +use crate::{ + codec, duplex::Duplex, finish_bytes, write_bytes, ChunkQueue, DropResetRead, DropResetWrite, + RpcCodec, RpcError, RpcRead, RpcStream, RpcWrite, +}; + +pub fn start(stream: St) -> DuplexCall +where + M: Duplex, + St: RpcStream, +{ + let (reader, writer) = stream.split(); + DuplexCall { + sender: DuplexSender::new(writer), + receiver: DuplexReceiver::new(reader), + } +} + +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: DropResetWrite, + marker: PhantomData T>, +} + +pub struct DuplexReceiver +where + T: RpcCodec, + R: RpcRead, +{ + stream: DropResetRead, + reader: EventReader, +} + +impl DuplexSender +where + T: RpcCodec, + W: RpcWrite, +{ + pub 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 + } + + /// 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); + } +} + +impl DuplexReceiver +where + T: RpcCodec, + R: RpcRead, +{ + pub fn new(stream: R) -> Self { + Self { + stream: DropResetRead::new(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_some() { + 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.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); + } +} + +enum ReadStep { + NeedMore, + Event(T), +} + +struct EventReader { + bytes: ChunkQueue, + marker: PhantomData T>, +} + +impl Default for EventReader { + fn default() -> Self { + Self { + bytes: 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/mod.rs b/ql-rpc/src/rpc/duplex/mod.rs new file mode 100644 index 00000000..8eb385bd --- /dev/null +++ b/ql-rpc/src/rpc/duplex/mod.rs @@ -0,0 +1,21 @@ +use super::Route; +use crate::RpcCodec; + +mod client; +mod server; + +pub use self::{client::*, server::*}; + +/// 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..fd2f941d --- /dev/null +++ b/ql-rpc/src/rpc/duplex/server.rs @@ -0,0 +1,53 @@ +use std::future::Future; + +use crate::{ + duplex::{Duplex, DuplexReceiver, DuplexSender}, + Context, RpcRead, RpcStream, RpcWrite, +}; + +#[trait_variant::make(DuplexHandler: Send)] +pub trait DuplexHandlerLocal +where + M: Duplex, + St: RpcStream, +{ + async fn handle(self, context: Context, peer: DuplexPeer); +} + +pub struct DuplexPeer +where + M: Duplex, + W: RpcWrite, + R: RpcRead, +{ + pub sender: DuplexSender, + pub receiver: DuplexReceiver, +} + +pub(crate) fn handle_duplex( + state: S, + context: Context, + _config: crate::RouterConfig, + stream: St, + handle: H, +) -> impl Future +where + M: Duplex + 'static, + St: RpcStream + 'static, + H: FnOnce(S, Context, DuplexPeer) -> HF, + HF: Future, +{ + 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 new file mode 100644 index 00000000..1d726b92 --- /dev/null +++ b/ql-rpc/src/rpc/mod.rs @@ -0,0 +1,39 @@ +//! 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 caller-defined route keys and the submodules provide the +//! matching client and server helpers for encoding, decoding, and handler glue + +use bytes::BufMut; + +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 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 { + type Key: RpcRouteKey; + + fn key() -> Self::Key; +} + +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/client.rs b/ql-rpc/src/rpc/notification/client.rs new file mode 100644 index 00000000..9d5d6bf2 --- /dev/null +++ b/ql-rpc/src/rpc/notification/client.rs @@ -0,0 +1,13 @@ +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 + M: Notification, + 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/notification/mod.rs b/ql-rpc/src/rpc/notification/mod.rs new file mode 100644 index 00000000..2a2e9691 --- /dev/null +++ b/ql-rpc/src/rpc/notification/mod.rs @@ -0,0 +1,18 @@ +use super::Route; +use crate::RpcCodec; + +mod client; +mod server; + +pub use self::{client::*, server::*}; + +/// 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..fc18efa5 --- /dev/null +++ b/ql-rpc/src/rpc/notification/server.rs @@ -0,0 +1,57 @@ +use std::future::Future; + +use ql_common::ResetCode; + +use crate::{ + notification::Notification, rpc::read_eof_request, Context, RouterConfig, RpcCodec, RpcError, + RpcRead, RpcStream, RpcWrite, +}; + +#[trait_variant::make(NotificationHandler: Send)] +pub trait NotificationHandlerLocal +where + M: Notification, + St: RpcStream, +{ + async fn handle(self, context: Context, message: M::Payload); + + fn handle_error(&self, error: &RpcError) { + let _ = error; + } +} + +pub(crate) fn handle_notification( + state: S, + context: Context, + config: RouterConfig, + stream: St, + handle: H, + handle_error: E, +) -> impl Future +where + Payload: RpcCodec + 'static, + St: RpcStream + 'static, + H: FnOnce(S, Context, Payload) -> HF, + HF: Future, + E: FnOnce(&S, &RpcError), +{ + 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; + } + }; + + 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 new file mode 100644 index 00000000..d206c9ca --- /dev/null +++ b/ql-rpc/src/rpc/parts.rs @@ -0,0 +1,272 @@ +use std::marker::PhantomData; + +use bytes::{BufMut, Bytes}; + +use crate::{codec, ChunkQueue, RpcCodec, RpcError}; + +pub enum PartReadStep { + NeedMore, + PartHeader(H), + BodyBytes(Bytes), + EndPart, + Finish, +} + +pub struct PartFrameReader { + bytes: 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, RpcError> { + 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(RpcError::Codec)?; + return Ok(PartReadStep::PartHeader(value)); + } + FrameKind::BodyChunk => unreachable!("body chunk is not a control frame"), + FrameKind::EndPart => { + body.expect_empty().map_err(RpcError::Protocol)?; + return Ok(PartReadStep::EndPart); + } + FrameKind::Finish => { + body.expect_empty().map_err(RpcError::Protocol)?; + drop(body); + self.bytes.expect_empty().map_err(RpcError::Protocol)?; + return Ok(PartReadStep::Finish); + } + } + } + PendingFrame::None => { + let Some((kind, len)) = self + .bytes + .try_take_tagged_part_header() + .map_err(RpcError::Protocol)? + else { + return Ok(PartReadStep::NeedMore); + }; + + let kind = FrameKind::try_from(kind).map_err(RpcError::Protocol)?; + 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]>)) { + 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); +} + +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_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..0998ae7b --- /dev/null +++ b/ql-rpc/src/rpc/progress/client.rs @@ -0,0 +1,166 @@ +use std::{ + future::{poll_fn, Future}, + pin::Pin, + task::{Context, Poll}, +}; + +use bytes::Bytes; +use ql_common::ResetCode; + +use crate::{ + codec, finish_bytes, + progress::Progress, + rpc::progress::codec::{ReadStep, ResponseReader}, + write_bytes, DropResetRead, Error, RpcError, RpcRead, RpcStream, +}; + +pub async fn start( + stream: St, + request: &M::Request, +) -> Result, RpcError> +where + M: Progress, + 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)) + .await + .map_err(RpcError::Transport)?; + finish_bytes(&mut writer) + .await + .map_err(RpcError::Transport)?; + Ok(ProgressCall::new(reader)) +} + +pub struct ProgressCall +where + M: Progress, + R: RpcRead, +{ + stream: DropResetRead, + 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: DropResetRead::new(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.stream.disarm(); + self.state = State::Terminal(Ok(response)); + return Poll::Ready(None); + } + Ok(ReadStep::NeedMore) => {} + Err(error) => { + self.stream.disarm(); + self.state = State::Terminal(Err(error)); + return Poll::Ready(None); + } + } + + match self.stream.poll_read(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.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); + } + Poll::Pending => return Poll::Pending, + } + } + } + + pub fn poll_next_progress(&mut self, cx: &mut Context<'_>) -> Poll> { + self.poll_step(cx) + } + + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); + } + + fn reset_inner(&mut self, code: ResetCode) { + self.state = State::Done; + DropResetRead::reset(&mut self.stream, code); + } +} + +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..a31c98c3 --- /dev/null +++ b/ql-rpc/src/rpc/progress/codec.rs @@ -0,0 +1,69 @@ +use std::marker::PhantomData; + +use bytes::Bytes; + +use crate::{progress::Progress, ChunkQueue, Error, RpcCodec, RpcError}; + +pub enum ReadStep { + NeedMore, + Progress(M::Progress), + Response(M::Response), +} + +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 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); + }; + + match kind { + x if x == FrameKind::Progress as u8 => { + let value = { + 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(RpcError::Codec)?; + drop(body); + if self.bytes.remaining() > 0 { + Err(RpcError::Protocol(Error::TrailingBytes)) + } else { + Ok(ReadStep::Response(response)) + } + } + other => Err(RpcError::Protocol(Error::UnexpectedFrameKind(other))), + } + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +#[repr(u8)] +pub enum FrameKind { + Progress = 1, + Response = 2, +} diff --git a/ql-rpc/src/rpc/progress/mod.rs b/ql-rpc/src/rpc/progress/mod.rs new file mode 100644 index 00000000..327a35fb --- /dev/null +++ b/ql-rpc/src/rpc/progress/mod.rs @@ -0,0 +1,25 @@ +use super::Route; +use crate::RpcCodec; + +mod client; +pub(crate) mod codec; +mod server; + +pub use self::{client::*, server::*}; + +/// 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..476cb950 --- /dev/null +++ b/ql-rpc/src/rpc/progress/server.rs @@ -0,0 +1,106 @@ +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, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, +}; + +#[trait_variant::make(ProgressHandler: Send)] +pub trait ProgressHandlerLocal +where + M: Progress, + St: RpcStream, +{ + async fn handle( + self, + context: Context, + request: M::Request, + responder: ProgressResponder, + ); + + fn handle_error(&self, error: &RpcError) { + let _ = error; + } +} + +pub struct ProgressResponder +where + M: Progress, + W: RpcWrite, +{ + writer: DropResetWrite, + marker: PhantomData M>, +} + +impl ProgressResponder +where + M: Progress, + W: RpcWrite, +{ + pub async fn send(&mut self, progress: M::Progress) -> Result<(), W::Error> { + let writer = &mut self.writer; + let mut encoded = Vec::new(); + 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(); + 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 + } + + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); + } +} + +pub(crate) fn handle_progress( + state: S, + context: Context, + config: RouterConfig, + stream: St, + handle: H, + handle_error: E, +) -> impl Future +where + M: Progress + 'static, + St: RpcStream + 'static, + H: FnOnce(S, Context, M::Request, ProgressResponder) -> HF, + HF: Future, + E: FnOnce(&S, &RpcError), +{ + 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; + } + }; + + 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 new file mode 100644 index 00000000..2bd08ef6 --- /dev/null +++ b/ql-rpc/src/rpc/request/client.rs @@ -0,0 +1,20 @@ +use crate::{ + request::Request, + rpc::{read_eof_value, write_eof_value}, + RpcError, RpcStream, +}; + +pub async fn call( + stream: St, + request: &M::Request, +) -> Result> +where + M: Request, + St: RpcStream, +{ + let (mut reader, mut writer) = stream.split(); + 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 new file mode 100644 index 00000000..3ad45db7 --- /dev/null +++ b/ql-rpc/src/rpc/request/mod.rs @@ -0,0 +1,22 @@ +use super::Route; +use crate::RpcCodec; + +mod client; +mod server; + +pub use self::{client::*, server::*}; + +/// 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..c8677043 --- /dev/null +++ b/ql-rpc/src/rpc/request/server.rs @@ -0,0 +1,97 @@ +use std::{future::Future, marker::PhantomData}; + +use bytes::Bytes; +use ql_common::ResetCode; + +use crate::{ + 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: Request, + St: RpcStream, +{ + async fn handle( + self, + context: Context, + message: M::Request, + responder: Response, + ); + + fn handle_error(&self, error: &RpcError) { + let _ = error; + } +} + +pub struct Response +where + W: RpcWrite, +{ + writer: DropResetWrite, + marker: PhantomData T>, +} + +impl Response +where + T: RpcCodec, + W: RpcWrite, +{ + pub(crate) fn new(writer: W) -> Self { + Self { + writer: DropResetWrite::new(writer), + marker: PhantomData, + } + } + + pub async fn respond(mut self, response: T) -> Result<(), W::Error> { + let writer = &mut self.writer; + let mut encoded = Vec::new(); + response.encode_value(&mut encoded); + write_bytes(writer, Bytes::from(encoded)).await?; + finish_bytes(writer).await?; + Ok(()) + } + + pub fn reset(mut self, code: ResetCode) { + DropResetWrite::reset(&mut self.writer, code); + } +} + +pub(crate) fn handle_request( + state: S, + context: Context, + config: RouterConfig, + stream: St, + handle: H, + handle_error: E, +) -> impl Future +where + Req: RpcCodec + 'static, + Res: RpcCodec + 'static, + St: RpcStream + 'static, + H: FnOnce(S, Context, Req, Response) -> HF, + HF: Future, + E: FnOnce(&S, &RpcError), +{ + 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; + } + }; + + handle(state, context, 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..37d51334 --- /dev/null +++ b/ql-rpc/src/rpc/subscription/client.rs @@ -0,0 +1,20 @@ +use crate::{ + 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, +) -> Result, RpcError> +where + M: Subscription, + St: RpcStream, +{ + let (reader, mut writer) = stream.split(); + write_eof_value(&mut writer, request) + .await + .map_err(RpcError::Transport)?; + Ok(DuplexReceiver::new(reader)) +} diff --git a/ql-rpc/src/rpc/subscription/mod.rs b/ql-rpc/src/rpc/subscription/mod.rs new file mode 100644 index 00000000..657f96dc --- /dev/null +++ b/ql-rpc/src/rpc/subscription/mod.rs @@ -0,0 +1,20 @@ +use super::Route; +use crate::RpcCodec; + +mod client; +mod server; + +pub use self::{client::*, server::*}; + +/// 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..97bcc8d5 --- /dev/null +++ b/ql-rpc/src/rpc/subscription/server.rs @@ -0,0 +1,62 @@ +use std::future::Future; + +use crate::{ + 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 + M: Subscription, + St: RpcStream, +{ + async fn handle( + self, + context: Context, + message: M::Request, + responder: DuplexSender, + ); + + fn handle_error(&self, error: &RpcError) { + let _ = error; + } +} + +pub(crate) fn handle_subscription( + state: S, + context: Context, + config: RouterConfig, + stream: St, + handle: H, + handle_error: E, +) -> impl Future +where + Req: RpcCodec + 'static, + Event: RpcCodec + 'static, + St: RpcStream + 'static, + H: FnOnce(S, Context, Req, DuplexSender) -> HF, + HF: Future, + E: FnOnce(&S, &RpcError), +{ + 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; + } + }; + + handle(state, context, request, DuplexSender::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..b5120599 --- /dev/null +++ b/ql-rpc/src/rpc/upload/client.rs @@ -0,0 +1,129 @@ +use bytes::Bytes; +use ql_common::ResetCode; + +use crate::{ + rpc::{ + parts::{encode_body_chunk, encode_end_part, encode_finish, encode_part_header}, + read_eof_value, + }, + upload::Upload, + write_bytes, DropResetRead, DropResetWrite, RpcError, RpcRead, RpcStream, RpcWrite, +}; + +pub async fn start( + stream: St, + request: &M::Request, +) -> Result, St::Error> +where + M: Upload, + 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?; + Ok(UploadCall::new(writer, reader)) +} + +pub struct UploadCall +where + M: Upload, + W: RpcWrite, + R: RpcRead, +{ + writer: DropResetWrite, + reader: DropResetRead, + 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: DropResetWrite::new(writer), + reader: DropResetRead::new(reader), + marker: std::marker::PhantomData, + } + } + + pub async fn start_part( + &mut self, + part_header: M::PartHeader, + ) -> Result, W::Error> { + let writer = &mut self.writer; + 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 writer = &mut self.writer; + let mut encoded = Vec::new(); + encode_finish(&mut encoded); + write_bytes(writer, Bytes::from(encoded)) + .await + .map_err(RpcError::Transport)?; + writer.queue_finish(); + + read_eof_value::(&mut self.reader).await + } + + fn reset(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.reader, code); + DropResetWrite::reset(&mut self.writer, code); + } +} + +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 = &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 = &mut self.parent.writer; + 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.reset(ResetCode::DROPPED); + } + } +} diff --git a/ql-rpc/src/rpc/upload/mod.rs b/ql-rpc/src/rpc/upload/mod.rs new file mode 100644 index 00000000..de73630b --- /dev/null +++ b/ql-rpc/src/rpc/upload/mod.rs @@ -0,0 +1,25 @@ +use super::*; +use crate::RpcCodec; + +mod client; +mod server; + +pub use self::{client::*, server::*}; + +/// 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..d74145b4 --- /dev/null +++ b/ql-rpc/src/rpc/upload/server.rs @@ -0,0 +1,216 @@ +use std::future::{poll_fn, Future}; + +use bytes::Bytes; +use ql_common::ResetCode; + +use crate::{ + request::Response, + rpc::{ + parts::{FrameKind, PartFrameReader, PartReadStep}, + read_framed_prefix, + }, + Context, DropResetRead, RouterConfig, RpcError, RpcRead, RpcStream, RpcWrite, Upload, +}; + +#[trait_variant::make(UploadHandler: Send)] +pub trait UploadHandlerLocal +where + M: Upload, + St: RpcStream, +{ + async fn handle( + self, + context: Context, + request: M::Request, + upload: UploadReader, + responder: UploadResponder, + ); + + fn handle_error(&self, error: &RpcError) { + let _ = error; + } +} + +pub struct UploadReader +where + M: Upload, + R: RpcRead, +{ + stream: DropResetRead, + reader: PartFrameReader, +} + +pub struct UploadPart<'a, M, R> +where + M: Upload, + R: RpcRead, +{ + parent: &'a mut UploadReader, + finished: bool, +} + +pub type UploadResponder = Response; + +impl UploadReader +where + M: Upload, + R: RpcRead, +{ + pub async fn next_part( + &mut self, + ) -> Result)>, crate::RpcError> + { + if !self.stream.is_some() { + return Ok(None); + } + + match self.read_frame().await? { + PartReadStep::PartHeader(value) => Ok(Some(( + value, + UploadPart { + parent: self, + finished: false, + }, + ))), + PartReadStep::Finish => { + self.stream.disarm(); + 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::RpcError> { + loop { + match self.reader.advance() { + Ok(PartReadStep::NeedMore) => {} + Ok(step) => return Ok(step), + Err(error) => return Err(error), + } + + match poll_fn(|cx| self.stream.poll_read(cx)).await { + Ok(Some(chunk)) => { + self.reader.push(chunk); + } + Ok(None) => return Err(crate::Error::Truncated.into()), + Err(error) => return Err(crate::RpcError::Transport(error)), + } + } + } + + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); + } + + fn reset_inner(&mut self, code: ResetCode) { + DropResetRead::reset(&mut self.stream, code); + } +} + +impl UploadPart<'_, M, R> +where + M: Upload, + R: RpcRead, +{ + pub async fn read_chunk( + &mut self, + ) -> Result, crate::RpcError> { + 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 reset(mut self, code: ResetCode) { + self.parent.reset_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.reset_inner(ResetCode::DROPPED); + } + } +} + +pub(crate) fn handle_upload( + state: S, + context: Context, + config: RouterConfig, + stream: St, + handle: H, + handle_error: E, +) -> impl Future +where + M: Upload + 'static, + St: RpcStream + 'static, + H: FnOnce( + S, + Context, + M::Request, + UploadReader, + UploadResponder, + ) -> HF, + HF: Future, + E: FnOnce(&S, &RpcError), +{ + let (mut reader, writer) = stream.split(); + + async move { + let (request, buffered) = + match read_framed_prefix::(&mut reader, Some(config.max_request_bytes)) + .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; + } + }; + + handle( + state, + context, + request, + UploadReader { + stream: DropResetRead::new(reader), + reader: PartFrameReader::new(buffered), + }, + Response::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..81208c00 --- /dev/null +++ b/ql-rpc/src/rpc/utils.rs @@ -0,0 +1,140 @@ +use bytes::Bytes; + +use crate::{ + 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> +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 async fn read_framed_request( + reader: &mut R, + config: RouterConfig, +) -> Result> +where + T: RpcCodec, + R: RpcRead, +{ + 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), + Ok(Some(_)) => Err(RpcError::Protocol(Error::TrailingBytes)), + Err(error) => Err(RpcError::Transport(error)), + } +} + +pub async fn read_framed_prefix( + reader: &mut R, + max_len: Option, +) -> Result<(T, ChunkQueue), RpcError> +where + T: RpcCodec, + R: RpcRead, +{ + let mut bytes = ChunkQueue::default(); + + loop { + if let Some(value) = try_take_framed_value(&mut bytes)? { + return Ok((value, bytes)); + } + + match read_bytes(reader).await { + Ok(Some(chunk)) => { + 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)), + } + } +} + +/// reads one eof-delimited value up to the configured request limit +pub 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); + match read_bytes(reader).await { + Ok(Some(chunk)) => { + if chunk.len() > remaining { + return Err(RpcError::Protocol(Error::LengthOverflow)); + } + total_read += chunk.len(); + bytes.push(chunk); + } + Ok(None) => break, + Err(error) => return Err(RpcError::Transport(error)), + } + } + + let value = T::decode_value(&mut bytes).map_err(RpcError::Codec)?; + if bytes.remaining() > 0 { + return Err(RpcError::Protocol(Error::TrailingBytes)); + } + Ok(value) +} + +fn try_take_framed_value(bytes: &mut ChunkQueue) -> Result, RpcError> +where + T: RpcCodec, +{ + 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 new file mode 100644 index 00000000..d5da6951 --- /dev/null +++ b/ql-rpc/src/stream.rs @@ -0,0 +1,173 @@ +use std::{ + future::{poll_fn, Future}, + task::{Context, Poll}, +}; + +use bytes::Bytes; +use ql_common::ResetCode; + +pub trait RpcStream { + type Error; + type Reader: RpcRead; + type Writer: RpcWrite; + + fn split(self) -> (Self::Reader, Self::Writer); +} + +pub trait RpcRead { + type Error; + + /// reads inbound bytes until eof or error + fn poll_read(&mut self, cx: &mut Context<'_>) -> Poll, Self::Error>>; + + /// aborts the read side + fn reset(self, code: ResetCode); +} + +pub trait RpcWrite { + type Error; + + /// writes outbound bytes before finish or reset + fn poll_write( + &mut self, + bytes: &mut Bytes, + cx: &mut Context<'_>, + ) -> Poll>; + + /// 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; must not replace a queued finish + fn reset(self, code: ResetCode); +} + +pub async fn read_bytes(reader: &mut R) -> Result, R::Error> +where + R: RpcRead, +{ + poll_fn(|cx| reader.poll_read(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 fn finish_bytes(writer: &mut W) -> impl Future> + '_ +where + W: RpcWrite, +{ + writer.queue_finish(); + poll_fn(|cx| writer.poll_finish(cx)) +} + +pub(crate) use drop::*; +mod drop { + use super::*; + + pub struct DropResetRead { + inner: Option, + } + + impl DropResetRead { + 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 reset(&mut self, code: ResetCode) { + if let Some(reader) = self.inner.take() { + reader.reset(code); + } + } + } + + impl RpcRead for DropResetRead { + type Error = R::Error; + + #[track_caller] + fn poll_read(&mut self, cx: &mut Context<'_>) -> Poll, Self::Error>> { + self.inner.as_mut().unwrap().poll_read(cx) + } + + fn reset(mut self, code: ResetCode) { + Self::reset(&mut self, code); + } + } + + impl Drop for DropResetRead { + fn drop(&mut self) { + self.reset(ResetCode::DROPPED); + } + } + + pub struct DropResetWrite { + inner: Option, + } + + impl DropResetWrite { + pub fn new(writer: W) -> Self { + Self { + inner: Some(writer), + } + } + + #[inline] + pub fn reset(&mut self, code: ResetCode) { + if let Some(writer) = self.inner.take() { + writer.reset(code); + } + } + } + + impl RpcWrite for DropResetWrite { + 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 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) + } + + fn reset(mut self, code: ResetCode) { + Self::reset(&mut self, code); + } + } + + impl Drop for DropResetWrite { + fn drop(&mut self) { + self.reset(ResetCode::DROPPED); + } + } +} diff --git a/ql-runtime/Cargo.toml b/ql-runtime/Cargo.toml new file mode 100644 index 00000000..fef1396c --- /dev/null +++ b/ql-runtime/Cargo.toml @@ -0,0 +1,37 @@ +[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-common = { workspace = true } +ql-rpc = { workspace = true, optional = true } +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"] } + +[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..22c26897 --- /dev/null +++ b/ql-runtime/src/command.rs @@ -0,0 +1,56 @@ +use ql_common::{ResetCode, StreamId}; +use ql_fsm::{NoSessionError, PairingInvite, StreamResetTarget}; +use ql_wire::{PairingToken, PeerBundle, SessionCloseCode}; + +use crate::{StreamReader, StreamWriter}; + +pub enum Command { + BindPeer { + peer: PeerBundle, + }, + Connect, + ArmPairing { + token: PairingToken, + }, + DisarmPairing, + StartPairing { + invite: PairingInvite, + }, + OpenStream { + header: Box<[u8]>, + start: oneshot::Sender>, + }, + PollInbound { + stream_id: StreamId, + }, + PollStream { + stream_id: StreamId, + }, + CloseSession { + code: SessionCloseCode, + }, + Unpair, + ResetStream { + stream_id: StreamId, + target: StreamResetTarget, + code: ResetCode, + }, +} + +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::ResetStream { .. } => "ResetStream", + } + } +} diff --git a/ql-runtime/src/driver/mod.rs b/ql-runtime/src/driver/mod.rs new file mode 100644 index 00000000..a00289e1 --- /dev/null +++ b/ql-runtime/src/driver/mod.rs @@ -0,0 +1,564 @@ +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_common::{ResetCode, StreamId, StreamInfo}; +use ql_fsm::{Event, QlFsm, StreamResetEvent, StreamResetTarget, WriteId}; + +use self::state::{DriverState, DriverStreamIo, InboundIo, InboundWriteResult, OutboundIo}; +use crate::{ + command::Command, + io, log, + platform::{QlInbound, QlPlatform, QlTimer}, + QlStreamError, ResetOrigin, Runtime, +}; + +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.crypto()) { + 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.crypto()).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.crypto()); + } + Command::CloseSession { code } => { + log::info!("closing session: code={code:?}"); + fsm.close_session(code); + } + Command::Unpair => { + log::info!("unpairing peer"); + fsm.unpair(); + } + Command::OpenStream { header, start } => { + log::info!("open stream requested"); + + let mut stream_ops = match fsm.open_stream(header) { + Ok(stream_ops) => stream_ops, + Err(error) => { + log::warn!("open stream failed"); + let _ = start.send(Err(error)); + return; + } + }; + let stream_id = stream_ops.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( + stream_id, + DriverStreamIo::new( + Some(OutboundIo::new(writer_io)), + Some(InboundIo::new(reader_io)), + ), + ); + 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(StreamResetTarget::Both, ResetCode::DROPPED); + 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::ResetStream { + stream_id, + target, + code, + } => { + log::debug!( + "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.reader() { + stream.inbound_close(); + } + if target.writer() { + stream.outbound_close(); + } + Self::try_reap_stream(entry); + } + if let Ok(mut stream) = fsm.stream(stream_id) { + stream.reset(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) => { + 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}"); + 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::Reset(reset) => { + self.handle_stream_reset(reset); + } + 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, + ) { + let (reader, writer, reader_io, writer_io) = + io::new_stream(stream_id, self.runtime_tx.clone()); + + self.streams.insert( + stream_id, + DriverStreamIo::new( + Some(OutboundIo::new(writer_io)), + Some(InboundIo::new(reader_io)), + ), + ); + + let qid = fsm.peer().unwrap().qid; + let stream = fsm.stream(stream_id).unwrap(); + let header = Box::<[u8]>::from(stream.header()); + + log::info!("delivering inbound stream to platform: stream_id={stream_id}",); + + platform.handle_inbound( + StreamInfo { + qid, + stream_id, + header, + }, + crate::QlStream { 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 Some(stream) = self.streams.get_mut(&stream_id) else { + return; + }; + 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}" + ); + 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.reset(StreamResetTarget::Reader, ResetCode::DROPPED); + 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_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 reset.target.reader() { + stream.inbound_fail(QlStreamError::StreamReset { + code: reset.code, + origin: ResetOrigin::Peer, + }); + } + if reset.target.writer() { + stream.outbound_fail(QlStreamError::StreamReset { + code: reset.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 { + 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.crypto()) 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..df725838 --- /dev/null +++ b/ql-runtime/src/driver/state.rs @@ -0,0 +1,140 @@ +use std::collections::HashMap; + +use bytes::Bytes; +use ql_common::StreamId; + +use crate::{ + command::Command, + io::{PushError, Rx, Tx}, + QlStreamError, +}; + +pub struct DriverState { + pub streams: HashMap, + pub runtime_tx: async_channel::Sender, + pub max_concurrent_message_writes: usize, +} + +pub struct DriverStreamIo { + outbound: Option, + inbound: Option, +} + +impl DriverStreamIo { + pub fn new(outbound: Option, inbound: Option) -> Self { + Self { outbound, inbound } + } + + 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..c30c0813 --- /dev/null +++ b/ql-runtime/src/driver/test.rs @@ -0,0 +1,187 @@ +use ql_common::{StreamInfo, QID}; +use ql_fsm::StreamResetEvent; +use ql_wire::{generate_identity, NoopCrypto, PeerBundle, SoftwareCrypto}; + +use super::*; +use crate::{ + driver::state::{InboundIo, OutboundIo}, + io, +}; + +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 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) + } + + 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, _info: StreamInfo, _stream: crate::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, + max_concurrent_message_writes: 1, + }, + QlFsm::new( + ql_fsm::QlFsmConfig::default(), + generate_identity(&SoftwareCrypto, "driver"), + 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()), 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()), runtime_tx); + let (_, _, _, writer_io) = stream; + OutboundIo::new(writer_io) +} + +#[test] +fn handle_inbound_finished_reaps_reset_initiator_stream() { + let (mut state, _fsm) = new_driver_state(); + let stream_id = StreamId(1u32.into()); + + state.streams.insert( + stream_id, + DriverStreamIo::new(None, Some(new_inbound_io(1))), + ); + + state.handle_inbound_finished(stream_id); + + assert!(!state.streams.contains_key(&stream_id)); +} + +#[test] +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(Some(new_outbound_io()), Some(new_inbound_io(1))), + ); + + state.handle_stream_reset(StreamResetEvent { + stream_id, + code: ResetCode::CANCELLED, + target: StreamResetTarget::Both, + }); + + 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, runtime_tx); + writer.queue_finish(); + state.streams.insert( + stream_id, + DriverStreamIo::new(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_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, runtime_tx); + + state.streams.insert( + stream_id, + DriverStreamIo::new(Some(OutboundIo::new(writer_io)), None), + ); + + state.drive_command( + &mut fsm, + Command::ResetStream { + stream_id, + target: StreamResetTarget::Writer, + code: ResetCode::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").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); + + state.streams.insert( + stream_id, + DriverStreamIo::new( + 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..329ca521 --- /dev/null +++ b/ql-runtime/src/error.rs @@ -0,0 +1,37 @@ +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 { + code: ResetCode, + origin: ResetOrigin, + }, + NoSession, +} + +impl std::fmt::Display for QlStreamError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match self { + Self::StreamReset { code, origin } => write!(f, "stream reset {code:?} ({origin:?})"), + Self::NoSession => f.write_str("no session"), + } + } +} + +impl std::error::Error for QlStreamError {} + +impl From for QlStreamError { + fn from(_: NoSessionError) -> Self { + Self::NoSession + } +} diff --git a/ql-runtime/src/handle/mod.rs b/ql-runtime/src/handle/mod.rs new file mode 100644 index 00000000..7256e2e8 --- /dev/null +++ b/ql-runtime/src/handle/mod.rs @@ -0,0 +1,98 @@ +use std::sync::Arc; + +use ql_fsm::{NoSessionError, PairingInvite}; +use ql_wire::{PairingToken, PeerBundle, SessionCloseCode}; + +use crate::command::Command; +pub use crate::io::{StreamReader, StreamWriter}; + +#[derive(Debug)] +pub struct QlStream { + pub writer: StreamWriter, + pub reader: StreamReader, +} + +#[derive(Clone)] +pub struct RuntimeHandle { + inner: Arc, +} + +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, header: Box<[u8]>) -> Result { + let (start_tx, start_rx) = oneshot::channel(); + self.send(Command::OpenStream { + header, + start: start_tx, + }); + + // runtime cannot be shutdown while we have a handle + let (reader, writer) = start_rx.await.unwrap()?; + + Ok(QlStream { 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 { + inner: Arc::new(Inner { tx }), + } + } + + #[inline] + #[track_caller] + pub(crate) fn send(&self, cmd: Command) { + self.inner.tx.try_send(cmd).expect("runtime is alive"); + } +} + +struct Inner { + tx: async_channel::Sender, +} + +impl Drop for Inner { + fn drop(&mut self) { + self.tx.close(); + } +} diff --git a/ql-runtime/src/io/inner.rs b/ql-runtime/src/io/inner.rs new file mode 100644 index 00000000..7a03cde7 --- /dev/null +++ b/ql-runtime/src/io/inner.rs @@ -0,0 +1,647 @@ +//! 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_common::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_common::ResetCode; + + use super::*; + use crate::{ + io::{sync::loom::*, Tx}, + QlStreamError, ResetOrigin, + }; + + #[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::StreamReset { + code: ResetCode::CANCELLED, + origin: ResetOrigin::Local, + }) + }) + }; + + 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::StreamReset { code, .. })) => { + assert_eq!(code, ResetCode::CANCELLED); + } + _ => panic!("expected terminal reader error"), + } + } + (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))); + } + _ => 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::StreamReset { + code: ResetCode::CANCELLED, + origin: ResetOrigin::Local, + }); + 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::StreamReset { code, .. })) => { + assert_eq!(code, ResetCode::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::StreamReset { + code: ResetCode::CANCELLED, + origin: ResetOrigin::Local, + }) + }) + }; + + 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::StreamReset { code, .. })) => { + assert_eq!(code, ResetCode::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::StreamReset { + code: ResetCode::CANCELLED, + origin: ResetOrigin::Local, + }) + }) + }; + + 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::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 new file mode 100644 index 00000000..7575cbea --- /dev/null +++ b/ql-runtime/src/io/mod.rs @@ -0,0 +1,56 @@ +mod inner; +mod reader; +mod slot; +mod sync; +mod writer; + +use std::ops::Deref; + +use ql_common::StreamId; + +pub use self::{reader::StreamReader, slot::PushError, writer::StreamWriter}; + +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, + runtime_tx: async_channel::Sender, +) -> (StreamReader, StreamWriter, Rx, Tx) { + let shared = inner::new(stream_id); + ( + 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 new file mode 100644 index 00000000..f8e16cea --- /dev/null +++ b/ql-runtime/src/io/reader.rs @@ -0,0 +1,197 @@ +use std::{ + future::poll_fn, + task::{Context, Poll}, +}; + +use bytes::Bytes; +use ql_common::ResetCode; +use ql_fsm::StreamResetTarget; + +use super::{ + inner::{Item, RxInner}, + slot::PopError, + Rx, +}; +use crate::{command::Command, log, QlStreamError}; + +pub struct StreamReader { + rx: Rx, + terminal: ReaderTerminalState, + runtime_tx: async_channel::Sender, +} + +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( + "terminal", + &matches!(self.terminal, ReaderTerminalState::Delivered), + ) + .finish_non_exhaustive() + } +} + +impl StreamReader { + pub(crate) fn new(shared: Rx, runtime_tx: async_channel::Sender) -> Self { + Self { + rx: shared, + terminal: ReaderTerminalState::Open, + runtime_tx, + } + } + + pub fn poll_read( + &mut self, + cx: &mut Context<'_>, + ) -> Poll, QlStreamError>> { + if matches!(self.terminal, ReaderTerminalState::Delivered) { + return Poll::Ready(Ok(None)); + } + + 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() { + Poll::Ready(result) => { + self.rx.unregister_waiter(); + Poll::Ready(result) + } + Poll::Pending => Poll::Pending, + } + } + + fn try_read_ready(&mut self) -> Poll, QlStreamError>> { + match self.rx.pop() { + Ok(Item::Chunk(bytes)) => { + log::trace!( + "byte reader received chunk: stream_id={} len={}", + self.rx.stream_id(), + bytes.len() + ); + let _ = self.runtime_tx.try_send(Command::PollInbound { + stream_id: self.rx.stream_id(), + }); + Poll::Ready(Ok(Some(bytes))) + } + Ok(Item::Error(error)) => { + log::debug!( + "byte reader delivered terminal error: stream_id={} error={:?}", + self.rx.stream_id(), + 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={}", + self.rx.stream_id() + ); + self.terminal = ReaderTerminalState::Delivered; + return Poll::Ready(Ok(None)); + } + Poll::Pending + } + } + } + + pub async fn read(&mut self) -> Result, QlStreamError> { + poll_fn(|cx| self.poll_read(cx)).await + } + + pub fn reset(mut self, code: ResetCode) { + self.reset_inner(code); + } + + fn reset_inner(&mut self, code: ResetCode) { + if matches!(self.terminal, ReaderTerminalState::Delivered) { + return; + } + log::debug!( + "byte reader explicit reset: stream_id={:?} code={:?}", + self.rx.stream_id(), + code + ); + self.terminal = ReaderTerminalState::Delivered; + let _ = self.runtime_tx.try_send(Command::ResetStream { + stream_id: self.rx.stream_id(), + target: StreamResetTarget::Reader, + code, + }); + } +} + +impl Drop for StreamReader { + fn drop(&mut self) { + if matches!(self.terminal, ReaderTerminalState::Delivered) { + return; + } + log::debug!( + "byte reader drop reset: stream_id={:?} code={:?}", + self.rx.stream_id(), + ResetCode::DROPPED + ); + let _ = self.runtime_tx.try_send(Command::ResetStream { + stream_id: self.rx.stream_id(), + target: StreamResetTarget::Reader, + code: ResetCode::DROPPED, + }); + } +} + +#[cfg(all(test, loom))] +mod loom_tests { + use std::task::{Context, Poll, Waker}; + + use bytes::Bytes; + use loom::thread; + use ql_fsm::StreamResetTarget; + + 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()), 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(&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(&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..4863478d --- /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_common::StreamId; + + use super::Arc; + use crate::{command::Command, io::inner::Inner}; + + 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() -> async_channel::Sender { + let (tx, _rx) = async_channel::unbounded(); + tx + } +} diff --git a/ql-runtime/src/io/writer.rs b/ql-runtime/src/io/writer.rs new file mode 100644 index 00000000..be354c03 --- /dev/null +++ b/ql-runtime/src/io/writer.rs @@ -0,0 +1,285 @@ +use std::{ + future::poll_fn, + task::{Context, Poll}, +}; + +use bytes::Bytes; +use ql_common::ResetCode; +use ql_fsm::StreamResetTarget; + +use super::{ + inner::{Item, TxInner}, + slot::PopError, + PushError, Tx, +}; +use crate::{command::Command, log, QlStreamError}; + +pub struct StreamWriter { + tx: Tx, + open: bool, + terminal: WriterTerminalState, + runtime_tx: async_channel::Sender, +} + +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("closed", &!self.open) + .finish_non_exhaustive() + } +} + +impl StreamWriter { + pub(crate) fn new(shared: Tx, runtime_tx: async_channel::Sender) -> Self { + Self { + tx: shared, + open: true, + terminal: WriterTerminalState::Pending, + runtime_tx, + } + } + + 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={}", + self.tx.stream_id() + ); + 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={}", + self.tx.stream_id() + ); + 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={}", self.tx.stream_id()); + 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 reset(mut self, code: ResetCode) { + self.reset_inner(code); + } + + fn poll_runtime(&self) { + let _ = self.runtime_tx.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 reset_inner(&mut self, code: ResetCode) { + if !self.open { + return; + } + self.open = false; + log::debug!( + "byte writer reset: stream_id={:?} code={:?}", + self.tx.stream_id(), + code + ); + let _ = self.runtime_tx.try_send(Command::ResetStream { + stream_id: self.tx.stream_id(), + target: StreamResetTarget::Writer, + code, + }); + } +} + +impl Drop for StreamWriter { + fn drop(&mut self) { + self.reset_inner(ResetCode::DROPPED); + } +} + +#[cfg(all(test, loom))] +mod loom_tests { + use std::task::{Context, Poll, Waker}; + + use bytes::Bytes; + use loom::thread; + use ql_fsm::StreamResetTarget; + + 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()), 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()), 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..a24a590d --- /dev/null +++ b/ql-runtime/src/lib.rs @@ -0,0 +1,68 @@ +pub use ql_fsm::{NoSessionError, PairingInvite}; + +pub use self::{ + error::{QlStreamError, ResetOrigin}, + 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::Sender, +} + +pub fn new_runtime

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

, RuntimeHandle) +where + P: QlPlatform, +{ + let (tx, rx) = async_channel::unbounded(); + let handle = RuntimeHandle::new(tx.clone()); + ( + Runtime { + identity, + platform, + config, + rx, + tx, + }, + handle, + ) +} 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..2e50bc65 --- /dev/null +++ b/ql-runtime/src/platform.rs @@ -0,0 +1,44 @@ +use std::{ + future::Future, + pin::Pin, + task::{Context, Poll}, + time::Instant, +}; + +use ql_common::{StreamInfo, QID}; +use ql_fsm::{PeerStatus, ReceiveError}; +use ql_wire::{PeerBundle, QlCrypto}; + +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 { + 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. + /// + /// 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, 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 new file mode 100644 index 00000000..293caa3d --- /dev/null +++ b/ql-runtime/src/rpc/adapter.rs @@ -0,0 +1,56 @@ +use std::task::{Context as TaskContext, Poll}; + +use bytes::Bytes; +use ql_common::ResetCode; +use ql_rpc::{RpcRead, RpcStream, RpcWrite}; + +use crate::{QlStream, QlStreamError, StreamReader, StreamWriter}; + +impl RpcStream for QlStream { + type Error = QlStreamError; + type Reader = StreamReader; + type Writer = StreamWriter; + + fn split(self) -> (Self::Reader, Self::Writer) { + (self.reader, self.writer) + } +} + +impl RpcRead for StreamReader { + type Error = QlStreamError; + + fn poll_read( + &mut self, + cx: &mut TaskContext<'_>, + ) -> Poll, QlStreamError>> { + StreamReader::poll_read(self, cx) + } + + fn reset(self, code: ResetCode) { + StreamReader::reset(self, code); + } +} + +impl RpcWrite for StreamWriter { + type Error = QlStreamError; + + fn poll_write( + &mut self, + bytes: &mut Bytes, + cx: &mut TaskContext<'_>, + ) -> Poll> { + StreamWriter::poll_write(self, bytes, cx) + } + + fn queue_finish(&mut self) { + StreamWriter::queue_finish(self); + } + + fn poll_finish(&mut self, cx: &mut TaskContext<'_>) -> Poll> { + StreamWriter::poll_finish(self, cx) + } + + fn reset(self, code: ResetCode) { + StreamWriter::reset(self, code); + } +} diff --git a/ql-runtime/src/rpc/mod.rs b/ql-runtime/src/rpc/mod.rs new file mode 100644 index 00000000..e75df6a2 --- /dev/null +++ b/ql-runtime/src/rpc/mod.rs @@ -0,0 +1,111 @@ +mod adapter; + +use ql_rpc::{ + download, duplex, notification, progress, request, subscription, upload, Route, RpcRouteKey, +}; + +use crate::{QlStream, QlStreamError, RuntimeHandle, StreamReader, StreamWriter}; + +type RpcResult = Result>; + +#[derive(Clone)] +pub struct RpcHandle { + inner: RuntimeHandle, +} + +impl RpcHandle { + pub async fn notification(&self, event: &M::Payload) -> RpcResult<(), M::Error> + where + M: notification::Notification, + { + let stream = self.open_rpc_stream::().await?; + notification::send::(stream, event) + .await + .map_err(ql_rpc::RpcError::Transport) + } + + pub async fn request(&self, request: &M::Request) -> RpcResult + where + M: request::Request, + { + let stream = self.open_rpc_stream::().await?; + request::call::(stream, request).await + } + + pub async fn subscribe( + &self, + request: &M::Request, + ) -> RpcResult, M::Error> + where + M: subscription::Subscription, + { + let stream = self.open_rpc_stream::().await?; + subscription::start::(stream, request).await + } + + pub async fn download( + &self, + request: &M::Request, + ) -> RpcResult, M::Error> + where + M: download::Download, + { + let stream = self.open_rpc_stream::().await?; + download::start::(stream, request).await + } + + pub async fn progress( + &self, + request: &M::Request, + ) -> RpcResult, M::Error> + where + M: progress::Progress, + { + let stream = self.open_rpc_stream::().await?; + progress::start::(stream, request).await + } + + pub async fn upload( + &self, + request: &M::Request, + ) -> RpcResult, M::Error> + where + M: upload::Upload, + { + let stream = self.open_rpc_stream::().await?; + upload::start::(stream, request) + .await + .map_err(ql_rpc::RpcError::Transport) + } + + pub async fn duplex( + &self, + ) -> RpcResult, M::Error> + where + M: duplex::Duplex, + { + let stream = self.open_rpc_stream::().await?; + Ok(duplex::start::(stream)) + } +} + +impl RpcHandle { + pub(super) fn new(inner: RuntimeHandle) -> Self { + Self { inner } + } + + 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(header.into_boxed_slice()) + .await + .map_err(QlStreamError::from) + .map_err(ql_rpc::RpcError::Transport) + } +} diff --git a/ql-runtime/src/tests/handshake.rs b/ql-runtime/src/tests/handshake.rs new file mode 100644 index 00000000..805371e0 --- /dev/null +++ b/ql-runtime/src/tests/handshake.rs @@ -0,0 +1,186 @@ +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_open_stream_params()) + .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_open_stream_params()) + .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 { + version: PairingInvite::VERSION, + 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 { + version: PairingInvite::VERSION, + 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..e356d7f9 --- /dev/null +++ b/ql-runtime/src/tests/mod.rs @@ -0,0 +1,661 @@ +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_codec::Decode; +use ql_common::{StreamInfo, QID}; +use ql_fsm::PeerStatus; +use ql_wire::{ + generate_identity, test_identities, PairingToken, PeerBundle, QlIdentity, RecordHeader, + RecordType, SoftwareCrypto, +}; +use tokio::{task::LocalSet, time::Sleep}; + +use crate::{ + new_runtime, platform::QlTimer, NoSessionError, PairingInvite, QlFsmConfig, QlStream, + QlStreamError, RuntimeConfig, RuntimeHandle, +}; + +type InboundStream = (StreamInfo, QlStream); + +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_open_stream_params() -> Box<[u8]> { + Box::from([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 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; + 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, info: StreamInfo, stream: QlStream) { + if let Some(tx) = &self.inbound { + let _ = tx.try_send((info, stream)); + } + } +} + +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(stream: &mut crate::StreamReader) -> Result>, QlStreamError> { + stream + .read() + .await + .map(|chunk| chunk.map(|bytes| bytes.to_vec())) +} + +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"); + 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"); + 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..32599389 --- /dev/null +++ b/ql-runtime/src/tests/rpc.rs @@ -0,0 +1,733 @@ +use std::{ + cell::RefCell, + future::Future, + rc::Rc, + str::Utf8Error, + sync::{Arc, Mutex}, + time::Duration, +}; + +use bytes::{BufMut, Bytes}; +use ql_codec::Encode; +use ql_common::ResetCode; +use ql_rpc::{ + 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::*; +use crate::{QlStream, QlStreamError, ResetOrigin, StreamWriter}; + +#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord)] +struct TestRouteKey(u64); + +impl ql_rpc::RpcRouteKey for TestRouteKey { + fn encoded_len(&self) -> usize { + ql_codec::Varint(self.0).encoded_len() + } + + fn encode(&self, out: &mut W) { + 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::>().ok()?; + Some(Self(*route_id)) + } +} + +macro_rules! test_route { + ($route:ty, $id:expr) => { + impl ql_rpc::Route for $route { + type Key = TestRouteKey; + + fn key() -> Self::Key { + TestRouteKey($id) + } + } + }; +} + +#[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; + +test_route!(Echo, 51); + +impl ql_rpc::request::Request for Echo { + type Error = Utf8Error; + + type Request = String; + type Response = String; +} + +struct Feed; + +test_route!(Feed, 52); + +impl ql_rpc::subscription::Subscription for Feed { + type Error = core::convert::Infallible; + type Request = Vec; + type Event = Vec; +} + +struct Notice; + +test_route!(Notice, 521); + +impl ql_rpc::notification::Notification for Notice { + type Error = core::convert::Infallible; + type Payload = Vec; +} + +struct Download; + +test_route!(Download, 53); + +impl ql_rpc::progress::Progress for Download { + type Error = core::convert::Infallible; + type Request = Vec; + type Progress = Vec; + type Response = Vec; +} + +struct BlobDownload; + +test_route!(BlobDownload, 54); + +impl ql_rpc::download::Download for BlobDownload { + type Error = core::convert::Infallible; + type Request = Vec; + type ResponseHeader = Vec; + type PartHeader = Vec; +} + +struct BlobUpload; + +test_route!(BlobUpload, 55); + +impl ql_rpc::upload::Upload for BlobUpload { + type Error = core::convert::Infallible; + type Request = Vec; + type PartHeader = Vec; + type Response = Vec; +} + +struct Chat; + +test_route!(Chat, 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, + _context: Context, + 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::::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(); + if let Some(fut) = router.handle(info, stream) { + 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, _context: Context, 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::::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(); + if let Some(fut) = router.handle(info, stream) { + 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, + _context: Context, + 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_wait().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::::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(); + if let Some(fut) = router.handle(info, stream) { + 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_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) + .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, + _context: Context, + 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::::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(); + if let Some(fut) = router.handle(info, stream) { + fut.await.unwrap(); + } + }); + + let rpc = pair.side_mut(Side::A).handle.rpc(); + let response = rpc.request::(&"hello".to_string()).await; + assert!(matches!( + response, + Err(ql_rpc::RpcError::Transport(QlStreamError::StreamReset { code, origin })) + if code == ResetCode::LIMIT && origin == ResetOrigin::Peer + )); + + 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, + _context: Context, + 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::::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(); + if let Some(fut) = router.handle(info, stream) { + 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_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()]); + + 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, + _context: Context, + 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::::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(); + if let Some(fut) = router.handle(info, stream) { + 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, + _context: Context, + 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::::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(); + if let Some(fut) = router.handle(info, stream) { + 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, + _context: Context, + 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::::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(); + if let Some(fut) = router.handle(info, stream) { + 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, + _context: Context, + 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(); + } + } + + 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::::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(); + if let Some(fut) = router.handle(info, stream) { + 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(); + 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..89066e00 --- /dev/null +++ b/ql-runtime/src/tests/session.rs @@ -0,0 +1,222 @@ +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_open_stream_params()) + .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_open_stream_params()) + .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_open_stream_params()) + .await, + Err(NoSessionError) + )); + assert!(matches!( + pair.side(Side::B) + .handle + .open_stream(test_open_stream_params()) + .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_open_stream_params()) + .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..af94b641 --- /dev/null +++ b/ql-runtime/src/tests/stream.rs @@ -0,0 +1,629 @@ +use std::time::Duration; + +use bytes::Bytes; +use ql_common::ResetCode; + +use super::*; +use crate::{QlStreamError, ResetOrigin}; + +#[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_open_stream_params()) + .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 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_open_stream_params()) + .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_open_stream_params()) + .await + .unwrap(); + let err = stream.writer.finish().await.unwrap_err(); + assert!(matches!( + err, + 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::StreamReset { code, origin } + if code == ResetCode::DROPPED && origin == ResetOrigin::Peer + )); + + 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::StreamReset { code, origin } + if code == ResetCode::DROPPED && origin == ResetOrigin::Peer + )); + }); + + let mut stream = pair + .side(Side::A) + .handle + .open_stream(test_open_stream_params()) + .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_open_stream_params()) + .await + .unwrap(); + let mut writer = stream.writer; + stream.reader.reset(ResetCode::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_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); + })); + } + + 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_open_stream_params()) + .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_open_stream_params()) + .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_open_stream_params()) + .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_open_stream_params()) + .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; +} diff --git a/ql-wire/Cargo.toml b/ql-wire/Cargo.toml new file mode 100644 index 00000000..f18f8a22 --- /dev/null +++ b/ql-wire/Cargo.toml @@ -0,0 +1,29 @@ +[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 } +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 } +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/crypto.rs b/ql-wire/src/crypto.rs new file mode 100644 index 00000000..380e98de --- /dev/null +++ b/ql-wire/src/crypto.rs @@ -0,0 +1,63 @@ +use crate::{ + MlKemCiphertext, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, SessionKey, + ENCRYPTED_MESSAGE_AUTH_SIZE, +}; + +ql_codec::codec! { + #[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) + } +} + +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..8be99b93 --- /dev/null +++ b/ql-wire/src/encrypted/ack.rs @@ -0,0 +1,427 @@ +use std::{fmt, ops::RangeInclusive}; + +use ql_codec::{ByteSlice, Encode, Error, Varint}; + +use crate::RecordSeq; + +#[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.0, + first_range_len: Some(*self.first_range_len), + previous_start: None, + blocks: self.blocks.iter(), + } + } + + pub fn contains(&self, seq: u64) -> bool { + let seq = RecordSeq(seq); + self.ranges().any(|range| range.contains(&seq)) + } + + /// 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() + } +} + +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; + self.previous_start = Some(start); + return Some(RecordSeq(start)..=RecordSeq(end)); + } + + 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 - 2; + let start = end - *block.range_len; + self.previous_start = Some(start); + Some(RecordSeq(start)..=RecordSeq(end)) + } +} + +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() + + self + .blocks + .iter() + .map(RecordAckBlock::encoded_len) + .sum::() + } + + fn encode(&self, out: &mut W) { + self.largest_acked.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); + block.range_len.encode(out); + } + } +} + +impl ql_codec::Decode for RecordAck { + fn decode(reader: &mut ql_codec::Reader) -> Result { + let largest_acked = 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 { + 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 + .0 + .checked_sub(*ack.first_range_len) + .ok_or(Error::InvalidRange)?; + + for block in &ack.blocks { + let end = previous_start + .checked_sub(block.gap.checked_add(2).ok_or(Error::InvalidRange)?) + .ok_or(Error::InvalidRange)?; + previous_start = end + .checked_sub(*block.range_len) + .ok_or(Error::InvalidRange)?; + } + } + 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().0; + let end = range.end().0; + if start > end { + return Err(RecordAckRangeError::InvertedRange); + } + + let range_len = end - start; + 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(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 + + (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(end); + let wire_len = largest_acked.encoded_len() + + RecordAck::block_count_len(0) + + 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(Varint(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 ql_codec::{Decode, Encode, Error, Varint}; + + use super::{RecordAck, RecordAckBlock, RecordAckBuilder, RecordAckRangeError}; + use crate::RecordSeq; + + #[test] + fn encode_decode_round_trip() { + 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_bytes(encoded.as_slice()).unwrap(), ack); + } + + #[test] + fn wire_fields_match_gap_encoding() { + let ack = RecordAck::from_ranges([ + RecordSeq(95)..=RecordSeq(100), + RecordSeq(90)..=RecordSeq(92), + RecordSeq(80)..=RecordSeq(80), + ]) + .unwrap(); + + assert_eq!(ack.largest_acked, RecordSeq(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_stops_when_budget_is_exhausted() { + let first_only = RecordAck::from_ranges([RecordSeq(95)..=RecordSeq(100)]).unwrap(); + let mut builder = RecordAckBuilder::new(); + + assert!(builder + .try_push_range(RecordSeq(95)..=RecordSeq(100), first_only.encoded_len()) + .unwrap()); + assert!(!builder + .try_push_range(RecordSeq(90)..=RecordSeq(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(RecordSeq(95)..=RecordSeq(100), usize::MAX) + .unwrap()); + assert_eq!( + builder.try_push_range(RecordSeq(90)..=RecordSeq(95), usize::MAX), + Err(RecordAckRangeError::NotCanonical) + ); + } + + #[test] + fn rejects_unsorted_ranges() { + assert_eq!( + RecordAck::from_ranges([ + RecordSeq(90)..=RecordSeq(92), + RecordSeq(95)..=RecordSeq(100) + ]), + Err(RecordAckRangeError::NotCanonical) + ); + } + + #[test] + fn rejects_touching_ranges() { + assert_eq!( + RecordAck::from_ranges([RecordSeq(10)..=RecordSeq(12), RecordSeq(7)..=RecordSeq(9)]), + Err(RecordAckRangeError::NotCanonical) + ); + } + + #[test] + fn rejects_overlapping_ranges() { + assert_eq!( + RecordAck::from_ranges([RecordSeq(10)..=RecordSeq(12), RecordSeq(8)..=RecordSeq(11)]), + Err(RecordAckRangeError::NotCanonical) + ); + } + + #[test] + fn contains_matches_range_membership() { + let ack = RecordAck::from_ranges([ + RecordSeq(150)..=RecordSeq(163), + RecordSeq(105)..=RecordSeq(110), + RecordSeq(100)..=RecordSeq(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([RecordSeq(5)..=RecordSeq(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_bytes(encoded.as_slice()), + Err(Error::InvalidRange) + ); + } + + #[test] + fn decode_rejects_truncated_payload() { + assert_eq!(RecordAck::decode_bytes(&[][..]), Err(Error::UnexpectedEof)); + + let encoded = RecordAck::from_ranges([RecordSeq(42)..=RecordSeq(42)]) + .unwrap() + .encode_vec(); + assert_eq!( + 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 new file mode 100644 index 00000000..82ccef0b --- /dev/null +++ b/ql-wire/src/encrypted/builder.rs @@ -0,0 +1,166 @@ +use bytes::BufMut; +use ql_codec::{BufView, Encode, Varint}; + +use super::{RecordAck, SessionClose, SessionFrame, StreamData, StreamReset, StreamWindow}; +use crate::{ + Nonce, QlCrypto, RecordHeader, RecordSeq, RecordType, RouteHeader, SessionHeader, SessionKey, +}; + +#[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 = RecordHeader::WIRE_SIZE + + Varint::::MAX_ENCODED_LEN + + crate::ENCRYPTED_MESSAGE_AUTH_SIZE; + + pub fn new(seq: RecordSeq, max_capacity: usize) -> Self { + let prefix_len = + RecordHeader::WIRE_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_reset(&mut self, frame: &StreamReset) -> bool { + self.push_frame_payload(super::SessionFrameKind::StreamReset, 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::StreamReset(frame) => self.push_stream_reset(frame), + SessionFrame::Close(close) => self.push_close(close), + } + } + + pub fn encrypt( + mut self, + crypto: &impl QlCrypto, + route: RouteHeader, + session_key: &SessionKey, + ) -> Vec { + self.ensure_prefix_capacity(0); + 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); + 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]; + record_header.encode(&mut prefix); + 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..1b76f3a9 --- /dev/null +++ b/ql-wire/src/encrypted/close.rs @@ -0,0 +1,19 @@ +ql_codec::codec! { + /// 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)] +#[repr(transparent)] +pub struct SessionCloseCode(pub u64); + +impl SessionCloseCode { + pub const CANCELLED: Self = Self(0); + pub const PROTOCOL: Self = Self(1); + pub const TIMEOUT: Self = Self(2); +} + +ql_codec::varint_wrapper!(SessionCloseCode, u64); diff --git a/ql-wire/src/encrypted/mod.rs b/ql-wire/src/encrypted/mod.rs new file mode 100644 index 00000000..efd1c918 --- /dev/null +++ b/ql-wire/src/encrypted/mod.rs @@ -0,0 +1,99 @@ +use ql_codec::{ByteSlice, Reader}; +use ql_common::StreamId; + +use crate::{ + encrypted_message::EncryptedMessage, Error, Nonce, QlCrypto, SessionHeader, SessionKey, +}; + +mod ack; +mod builder; +mod close; +mod stream_data; +mod stream_reset; +mod stream_window; + +pub use ack::*; +pub use builder::*; +pub use close::*; +pub use stream_data::*; +pub use stream_reset::*; +pub use stream_window::*; + +ql_codec::codec! { + #[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, + } +} + +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::StreamReset(frame) => SessionFrame::StreamReset(frame), + Self::Close(frame) => SessionFrame::Close(frame), + } + } +} + +pub fn parse_session_frames(bytes: B) -> SessionFrameIter { + SessionFrameIter { + reader: Reader::new(bytes), + } +} + +pub fn decode_session_frames(bytes: &[u8]) -> Result>>, Error> { + 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, Error>; + + fn next(&mut self) -> Option { + if self.reader.is_empty() { + None + } else { + Some(self.reader.decode::>().map_err(Into::into)) + } + } +} + +pub fn decrypt_record>( + crypto: &impl QlCrypto, + record_header: &crate::RecordHeader, + header: &SessionHeader, + encrypted: EncryptedMessage, + session_key: &SessionKey, +) -> Result { + let aad = header.aad(record_header.route); + let nonce = Nonce::from_counter(header.seq.0); + let mut ciphertext = encrypted.ciphertext; + if !crypto.aes256_gcm_decrypt( + session_key, + &nonce, + &aad, + ciphertext.as_mut(), + &encrypted.auth, + ) { + return Err(Error::DecryptFailed); + } + Ok(ciphertext) +} diff --git a/ql-wire/src/encrypted/stream_data.rs b/ql-wire/src/encrypted/stream_data.rs new file mode 100644 index 00000000..1f211795 --- /dev/null +++ b/ql-wire/src/encrypted/stream_data.rs @@ -0,0 +1,105 @@ +use ql_codec::{ + encode_bytes, encoded_len_bytes, BufView, ByteSlice, Decode, Encode, Error, Varint, +}; +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 fin: bool, + pub bytes: B, +} + +impl StreamData { + /// 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 { + fn decode(reader: &mut ql_codec::Reader) -> Result { + let stream_id = reader.decode()?; + let offset = 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.take_len_prefixed()?) + } else { + None + }; + let bytes = reader.take_len_prefixed()?; + + Ok(Self { + stream_id, + offset, + header, + fin, + bytes, + }) + } +} + +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.map(|header| header.to_vec()), + fin: self.fin, + bytes: self.bytes.to_vec(), + } + } +} + +impl Encode for StreamData { + fn encoded_len(&self) -> usize { + self.stream_id.encoded_len() + + self.offset.encoded_len() + + size_of::() + + 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 == 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 { + encode_bytes(header, out); + } + encode_bytes(&self.bytes, out); + } +} + +mod flag { + pub const FIN: u8 = 0x01; + pub const HEADER: u8 = 0x02; +} diff --git a/ql-wire/src/encrypted/stream_reset.rs b/ql-wire/src/encrypted/stream_reset.rs new file mode 100644 index 00000000..711d786b --- /dev/null +++ b/ql-wire/src/encrypted/stream_reset.rs @@ -0,0 +1,30 @@ +use ql_common::ResetCode; + +use super::StreamId; + +ql_codec::codec! { + /// 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, + } +} + +ql_codec::codec! { + /// 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 new file mode 100644 index 00000000..49b61342 --- /dev/null +++ b/ql-wire/src/encrypted/stream_window.rs @@ -0,0 +1,12 @@ +use ql_codec::Varint; + +use super::StreamId; + +ql_codec::codec! { + /// 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/encrypted_message.rs b/ql-wire/src/encrypted_message.rs new file mode 100644 index 00000000..0b2ddfe0 --- /dev/null +++ b/ql-wire/src/encrypted_message.rs @@ -0,0 +1,42 @@ +use bytes::Buf; +use ql_codec::{encode_bytes_raw, BufView, ByteSlice, Decode, Encode}; + +use crate::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 fn into_owned(self) -> EncryptedMessage> + where + B: ByteSlice, + { + EncryptedMessage { + auth: self.auth, + ciphertext: self.ciphertext.to_vec(), + } + } +} + +impl Decode for EncryptedMessage { + fn decode(reader: &mut ql_codec::Reader) -> Result { + Ok(Self { + auth: reader.decode()?, + ciphertext: reader.take_all(), + }) + } +} + +impl Encode for EncryptedMessage { + fn encoded_len(&self) -> usize { + self.auth.encoded_len() + self.ciphertext.buf().remaining() + } + + fn encode(&self, out: &mut W) { + self.auth.encode(out); + encode_bytes_raw(&self.ciphertext, out); + } +} diff --git a/ql-wire/src/error.rs b/ql-wire/src/error.rs new file mode 100644 index 00000000..71359231 --- /dev/null +++ b/ql-wire/src/error.rs @@ -0,0 +1,65 @@ +use core::fmt; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub enum Error { + // codec errors + UnexpectedEof, + InvalidData, + InvalidDiscriminant, + InvalidRange, + LengthOverflow, + InvalidVarint, + InvalidUtf8, + InvalidPayload, + + // protocol validation + InvalidRouteHeader, + InvalidHandshakeId, + InvalidPairingId, + InvalidRemoteBundle, + + // cryptographic/session + DecryptFailed, + Expired, + + InvalidState, +} + +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::InvalidHandshakeId => "invalid handshake id", + Self::InvalidPairingId => "invalid pairing id", + Self::InvalidRemoteBundle => "invalid remote bundle", + Self::DecryptFailed => "decryption failed", + Self::Expired => "expired", + Self::InvalidState => "invalid state", + }; + f.write_str(message) + } +} + +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/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/id.rs b/ql-wire/src/handshake/id.rs new file mode 100644 index 00000000..4f34602f --- /dev/null +++ b/ql-wire/src/handshake/id.rs @@ -0,0 +1,5 @@ +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 new file mode 100644 index 00000000..ec8da42b --- /dev/null +++ b/ql-wire/src/handshake/ik.rs @@ -0,0 +1,459 @@ +use std::borrow::Borrow; + +use ql_codec::{ByteSlice, Encode}; + +use super::{ + decrypt_mlkem_ciphertext, decrypt_peer_bundle, encrypt_mlkem_ciphertext, encrypt_peer_bundle, + 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, HandshakeId, HandshakeKind, MlKemCiphertext, MlKemPublicKey, PeerBundle, QlCrypto, + QlIdentity, +}; + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Ik1 { + pub handshake_id: HandshakeId, + pub transport_params: TransportParams, + pub skem_ciphertext: MlKemCiphertext, + pub ephemeral: EphemeralPublicKey, + pub static_bundle: Option, +} + +impl Encode for Ik1 { + fn encoded_len(&self) -> usize { + self.handshake_id.encoded_len() + + self.transport_params.encoded_len() + + self.skem_ciphertext.encoded_len() + + self.ephemeral.encoded_len() + + self + .static_bundle + .as_ref() + .map_or(0, EncryptedPeerBundle::encoded_len) + } + + fn encode(&self, out: &mut W) { + self.handshake_id.encode(out); + self.transport_params.encode(out); + self.skem_ciphertext.encode(out); + self.ephemeral.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, + }) + } +} + +ql_codec::codec! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Ik2 { + pub handshake_id: HandshakeId, + pub transport_params: TransportParams, + pub ekem_ciphertext: MlKemCiphertext, + pub skem_ciphertext: EncryptedMlKemCiphertext, + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Step { + Send1, + Recv1, + Send2, + Recv2, + Done, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum IkPattern { + Ik, + Kk, +} + +#[derive(Debug, Clone)] +pub struct IkHandshake { + pattern: IkPattern, + role: Role, + step: Step, + symmetric: SymmetricState, + local: I, + remote_bundle: Option, + local_ephemeral: Option, + remote_ephemeral: Option, + handshake_id: Option, + local_transport_params: TransportParams, + remote_transport_params: Option, +} + +impl> IkHandshake { + 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: I, + remote_bundle: PeerBundle, + local_transport_params: TransportParams, + ) -> Self { + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_IK); + symmetric.mix_hash(crypto, &remote_bundle.encode_vec()); + Self::new( + IkPattern::Ik, + Role::Initiator, + symmetric, + local, + Some(remote_bundle), + local_transport_params, + ) + } + + pub fn new_ik_responder( + crypto: &impl QlCrypto, + local: I, + expected_remote: Option, + local_transport_params: TransportParams, + ) -> Self { + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_IK); + symmetric.mix_hash(crypto, &local.borrow().bundle().encode_vec()); + Self::new( + IkPattern::Ik, + Role::Responder, + symmetric, + local, + expected_remote, + local_transport_params, + ) + } + + pub fn new_kk_initiator( + crypto: &impl QlCrypto, + local: I, + remote_bundle: PeerBundle, + local_transport_params: TransportParams, + ) -> Self { + let mut symmetric = SymmetricState::new(crypto, PROTOCOL_KK); + symmetric.mix_hash(crypto, &local.borrow().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, + ) + } + + pub fn new_kk_responder( + crypto: &impl QlCrypto, + 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.borrow().bundle().encode_vec()); + Self::new( + IkPattern::Kk, + Role::Responder, + symmetric, + local, + Some(remote_bundle), + local_transport_params, + ) + } + + fn new( + pattern: IkPattern, + role: Role, + symmetric: SymmetricState, + local: I, + 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, + } + } + pub fn is_finished(&self) -> bool { + self.step == Step::Done + } + + pub fn write_1( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { + if self.step != Step::Send1 { + return Err(Error::InvalidState); + } + 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: local.qid, + recipient: remote_bundle.qid, + }; + mix_hash_routed_handshake( + &mut self.symmetric, + crypto, + header, + 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()); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + let local_ephemeral = generate_ephemeral_keypair(crypto); + let ephemeral = local_ephemeral.public(); + self.symmetric.mix_hash_ephemeral(crypto, &ephemeral); + + let static_bundle = match self.pattern { + IkPattern::Ik => Some(encrypt_peer_bundle( + crypto, + &mut self.symmetric, + &local.bundle(), + )?), + IkPattern::Kk => None, + }; + + self.local_ephemeral = Some(local_ephemeral); + self.step = Step::Recv2; + Ok(Ik1 { + handshake_id, + transport_params: self.local_transport_params, + skem_ciphertext, + ephemeral, + static_bundle, + }) + } + + pub fn read_1( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Ik1, + ) -> Result<(), Error> { + if self.step != Step::Recv1 { + return Err(Error::InvalidState); + } + 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, + header, + match self.pattern { + IkPattern::Ik => HandshakeKind::Ik1, + IkPattern::Kk => HandshakeKind::Kk1, + }, + message.handshake_id, + message.transport_params, + ); + self.symmetric + .mix_hash(crypto, message.skem_ciphertext.as_bytes()); + let skem_secret = + crypto.mlkem_decapsulate(&local.mlkem_private_key, &message.skem_ciphertext); + self.symmetric + .mix_key_and_hash(crypto, skem_secret.as_bytes()); + + self.symmetric + .mix_hash_ephemeral(crypto, &message.ephemeral); + self.remote_ephemeral = Some(message.ephemeral.clone()); + + 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), + } + } + (IkPattern::Kk, None) => {} + _ => return Err(Error::InvalidState), + } + + self.remote_transport_params = Some(message.transport_params); + 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.borrow().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 != Step::Recv2 { + return Err(Error::InvalidState); + } + 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, + 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()); + 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.borrow().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 = Step::Done; + Ok(()) + } + + pub fn finalize(self, crypto: &impl QlCrypto) -> Result { + if !self.is_finished() { + return Err(Error::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, + self.role, + remote_bundle, + remote_transport_params, + )) + } + + fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), Error> { + if header.recipient != self.local.borrow().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/mod.rs b/ql-wire/src/handshake/mod.rs new file mode 100644 index 00000000..8e3c7c9e --- /dev/null +++ b/ql-wire/src/handshake/mod.rs @@ -0,0 +1,446 @@ +use ql_codec::{ByteSlice, Decode, Encode}; + +use crate::{ + 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}; +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 HANDSHAKE_PREAMBLE_DOMAIN: &[u8] = b"ql-wire:handshake-preamble:v1"; + +ql_codec::codec! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct EphemeralPublicKey { + pub mlkem_public_key: MlKemPublicKey, + } +} + +ql_codec::codec! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct EncryptedMlKemCiphertext(pub Box<[u8; Self::SIZE]>); +} + +impl EncryptedMlKemCiphertext { + pub const SIZE: usize = MlKemCiphertext::SIZE + ENCRYPTED_MESSAGE_AUTH_SIZE; + + pub fn new(data: Box<[u8; Self::SIZE]>) -> Self { + Self(data) + } + + pub fn as_bytes(&self) -> &[u8; Self::SIZE] { + self.0.as_ref() + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct EncryptedPeerBundle(pub Box<[u8]>); + +impl EncryptedPeerBundle { + pub fn as_bytes(&self) -> &[u8] { + self.0.as_ref() + } +} + +impl Encode for EncryptedPeerBundle { + fn encoded_len(&self) -> usize { + self.0.len() + } + + fn encode(&self, out: &mut W) { + out.put_slice(self.as_bytes()); + } +} + +impl ql_codec::Decode for EncryptedPeerBundle { + fn decode(reader: &mut ql_codec::Reader) -> Result { + let data = reader.take_all(); + Ok(Self(Box::from(&*data))) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct FinalizedHandshake { + pub tx_key: SessionKey, + pub rx_key: SessionKey, + 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, 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); + 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, Error> { + if ciphertext.len() < ENCRYPTED_MESSAGE_AUTH_SIZE { + 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(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(Error::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_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; + 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 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, + plaintext: &[u8], + ) -> Result, Error> { + 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, Error> { + 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(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), + } + } +} + +fn generate_ephemeral_keypair(crypto: &impl QlCrypto) -> EphemeralKeyPair { + EphemeralKeyPair { + mlkem: crypto.mlkem_generate_keypair(), + } +} + +fn mix_hash_routed_handshake( + symmetric: &mut SymmetricState, + crypto: &impl QlCrypto, + header: RouteHeader, + kind: HandshakeKind, + handshake_id: HandshakeId, + transport_params: TransportParams, +) { + mix_hash_handshake_preamble( + symmetric, + crypto, + &header.encode_vec(), + kind, + handshake_id, + transport_params, + ); +} + +fn mix_hash_pairing_handshake( + symmetric: &mut SymmetricState, + crypto: &impl QlCrypto, + header: RouteHeader, + kind: HandshakeKind, + 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, + handshake_id, + transport_params, + ); +} + +fn mix_hash_handshake_preamble( + symmetric: &mut SymmetricState, + crypto: &impl QlCrypto, + header: &[u8], + kind: HandshakeKind, + 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, &handshake_id.encode_vec()); + symmetric.mix_hash(crypto, &transport_params.encode_vec()); +} + +fn initialize_handshake_id( + expected: &mut Option, + handshake_id: HandshakeId, +) -> Result<(), Error> { + match expected { + Some(stored) if *stored != handshake_id => Err(Error::InvalidHandshakeId), + Some(_) => Ok(()), + None => { + *expected = Some(handshake_id); + Ok(()) + } + } +} + +fn require_handshake_id( + expected: Option<&HandshakeId>, + handshake_id: HandshakeId, +) -> Result<(), Error> { + match expected { + Some(stored) if *stored == handshake_id => Ok(()), + _ => Err(Error::InvalidHandshakeId), + } +} + +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_bytes(plaintext.as_slice())?; + bundle.validate(crypto)?; + 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::SIZE]> = + encrypted.try_into().map_err(|_| Error::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(|_| Error::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); + FinalizedHandshake { + tx_key, + rx_key, + handshake_hash, + remote_bundle, + remote_transport_params, + } +} + +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(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(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..54e22efa --- /dev/null +++ b/ql-wire/src/handshake/pairing.rs @@ -0,0 +1,55 @@ +use std::fmt::{self, Display, Formatter}; + +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! { + #[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(()) + } +} + +ql_codec::codec! { + #[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(()) + } +} diff --git a/ql-wire/src/handshake/transport_params.rs b/ql-wire/src/handshake/transport_params.rs new file mode 100644 index 00000000..b0c6cb5e --- /dev/null +++ b/ql-wire/src/handshake/transport_params.rs @@ -0,0 +1,16 @@ +ql_codec::codec! { + /// 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 Default for TransportParams { + fn default() -> Self { + Self { + initial_stream_receive_window: 16 * 1024, + } + } +} diff --git a/ql-wire/src/handshake/xx.rs b/ql-wire/src/handshake/xx.rs new file mode 100644 index 00000000..34c256e1 --- /dev/null +++ b/ql-wire/src/handshake/xx.rs @@ -0,0 +1,482 @@ +use ql_common::QID; + +use super::{ + decrypt_mlkem_ciphertext, decrypt_peer_bundle, encrypt_mlkem_ciphertext, encrypt_peer_bundle, + 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, HandshakeId, HandshakeKind, MlKemCiphertext, MlKemPublicKey, PairingId, PairingToken, + PeerBundle, QlCrypto, QlIdentity, +}; + +const PROTOCOL_XX: &[u8] = b"ql-wire:pq-xx:v1"; + +ql_codec::codec! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Xx1 { + pub handshake_id: HandshakeId, + pub pairing_id: PairingId, + pub transport_params: TransportParams, + pub ephemeral: EphemeralPublicKey, + } +} + +ql_codec::codec! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Xx2 { + pub handshake_id: HandshakeId, + pub transport_params: TransportParams, + pub ekem_ciphertext: MlKemCiphertext, + pub static_bundle: EncryptedPeerBundle, + } +} + +ql_codec::codec! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Xx3 { + pub handshake_id: HandshakeId, + pub skem_ciphertext: EncryptedMlKemCiphertext, + pub static_bundle: EncryptedPeerBundle, + } +} + +ql_codec::codec! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct Xx4 { + pub handshake_id: HandshakeId, + pub skem_ciphertext: EncryptedMlKemCiphertext, + } +} + +#[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_id: 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: SymmetricState::new(crypto, PROTOCOL_XX), + local, + remote_qid, + pairing_token, + remote_bundle: None, + local_ephemeral: None, + remote_ephemeral: None, + handshake_id: 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: SymmetricState::new(crypto, PROTOCOL_XX), + local, + remote_qid, + pairing_token, + remote_bundle: None, + local_ephemeral: None, + remote_ephemeral: None, + handshake_id: 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 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 + } + + pub fn remote_bundle(&self) -> Option<&PeerBundle> { + self.remote_bundle.as_ref() + } + + fn header(&self) -> RouteHeader { + RouteHeader { + sender: self.local.qid, + recipient: self.remote_qid, + } + } + + fn ensure_inbound_header(&self, header: RouteHeader) -> Result<(), Error> { + if header.sender != self.remote_qid || header.recipient != self.local.qid { + return Err(Error::InvalidRouteHeader); + } + Ok(()) + } + + fn ensure_remote_bundle(&self, bundle: &PeerBundle) -> Result<(), Error> { + if bundle.qid == self.remote_qid { + Ok(()) + } else { + Err(Error::InvalidRemoteBundle) + } + } + + pub fn write_1( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { + if self.step != XxStep::Send1 { + return Err(Error::InvalidState); + } + 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( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Xx1, + handshake_id, + pairing_id, + self.local_transport_params, + ); + self.symmetric + .mix_psk_pairing_token(crypto, self.pairing_token); + + let local_ephemeral = generate_ephemeral_keypair(crypto); + let ephemeral = local_ephemeral.public(); + self.symmetric.mix_hash_ephemeral(crypto, &ephemeral); + + self.local_ephemeral = Some(local_ephemeral); + self.step = XxStep::Recv2; + Ok(Xx1 { + handshake_id, + pairing_id, + transport_params: self.local_transport_params, + ephemeral, + }) + } + + pub fn read_1( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Xx1, + ) -> Result<(), Error> { + if self.step != XxStep::Recv1 { + return Err(Error::InvalidState); + } + 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.handshake_id, + message.pairing_id, + message.transport_params, + ); + 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()); + self.remote_transport_params = Some(message.transport_params); + self.step = XxStep::Send2; + Ok(()) + } + + pub fn write_2( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { + if self.step != XxStep::Send2 { + return Err(Error::InvalidState); + } + require_handshake_id(self.handshake_id.as_ref(), handshake_id)?; + let header = self.header(); + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Xx2, + handshake_id, + self.pairing_token.id(crypto), + 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 static_bundle = encrypt_peer_bundle(crypto, &mut self.symmetric, &self.local.bundle())?; + + self.step = XxStep::Recv3; + Ok(Xx2 { + handshake_id, + transport_params: self.local_transport_params, + ekem_ciphertext, + static_bundle, + }) + } + + pub fn read_2( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Xx2, + ) -> Result<(), Error> { + if self.step != XxStep::Recv2 { + return Err(Error::InvalidState); + } + 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.handshake_id, + self.pairing_token.id(crypto), + 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 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.remote_transport_params = Some(message.transport_params); + self.step = XxStep::Send3; + Ok(()) + } + + pub fn write_3( + &mut self, + crypto: &impl QlCrypto, + handshake_id: HandshakeId, + ) -> Result { + if self.step != XxStep::Send3 { + return Err(Error::InvalidState); + } + require_handshake_id(self.handshake_id.as_ref(), handshake_id)?; + let header = self.header(); + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Xx3, + handshake_id, + self.pairing_token.id(crypto), + self.local_transport_params, + ); + + 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()); + + let static_bundle = encrypt_peer_bundle(crypto, &mut self.symmetric, &self.local.bundle())?; + + self.step = XxStep::Recv4; + Ok(Xx3 { + handshake_id, + skem_ciphertext, + static_bundle, + }) + } + + pub fn read_3( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Xx3, + ) -> Result<(), Error> { + if self.step != XxStep::Recv3 { + return Err(Error::InvalidState); + } + 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.handshake_id, + self.pairing_token.id(crypto), + remote_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, + handshake_id: HandshakeId, + ) -> Result { + if self.step != XxStep::Send4 { + return Err(Error::InvalidState); + } + require_handshake_id(self.handshake_id.as_ref(), handshake_id)?; + let header = self.header(); + mix_hash_pairing_handshake( + &mut self.symmetric, + crypto, + header, + HandshakeKind::Xx4, + handshake_id, + self.pairing_token.id(crypto), + self.local_transport_params, + ); + + 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 = XxStep::Done; + Ok(Xx4 { + handshake_id, + skem_ciphertext, + }) + } + + pub fn read_4( + &mut self, + crypto: &impl QlCrypto, + header: RouteHeader, + message: &Xx4, + ) -> Result<(), Error> { + if self.step != XxStep::Recv4 { + return Err(Error::InvalidState); + } + 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.handshake_id, + self.pairing_token.id(crypto), + remote_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(Error::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, + 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..a40d96f5 --- /dev/null +++ b/ql-wire/src/header.rs @@ -0,0 +1,53 @@ +use ::bytes::BufMut; +use ql_codec::Encode; +use ql_common::QID; + +use crate::QL_WIRE_VERSION; + +ql_codec::codec! { + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub struct RouteHeader { + pub sender: QID, + pub recipient: QID, + } +} + +ql_codec::codec! { + #[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); + +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 { + const AAD_DOMAIN: &[u8] = b"ql-wire:session-aad:v1"; + const AAD_RECORD_KIND_SESSION: u8 = 1; + + pub fn aad(&self, route: RouteHeader) -> Vec { + let aad_len = Self::AAD_DOMAIN.len() + + size_of::() + + size_of::() + + route.encoded_len() + + 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.seq.encode(&mut aad); + debug_assert_eq!(aad.len(), aad_len); + aad + } +} diff --git a/ql-wire/src/identity.rs b/ql-wire/src/identity.rs new file mode 100644 index 00000000..ee39e371 --- /dev/null +++ b/ql-wire/src/identity.rs @@ -0,0 +1,69 @@ +use ql_common::QID; + +use crate::{derive_qid, Error, MlKemKeyPair, MlKemPrivateKey, MlKemPublicKey, QlCrypto, QlHash}; + +ql_codec::codec! { + #[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 { + 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(()) + } +} + +ql_codec::codec! { + #[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, + mlkem_private_key: MlKemPrivateKey, + mlkem_public_key: MlKemPublicKey, + name: impl Into, + ) -> Self { + let qid = derive_qid(crypto, &mlkem_public_key); + Self { + qid, + mlkem_private_key, + mlkem_public_key, + capabilities: 0, + name: name.into(), + } + } + + 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(), + } + } +} + +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 new file mode 100644 index 00000000..4b415562 --- /dev/null +++ b/ql-wire/src/lib.rs @@ -0,0 +1,37 @@ +//! +//! QuantumLink protocol wire format +//! + +#![allow(clippy::too_many_arguments)] + +mod crypto; +mod encrypted; +mod encrypted_message; +mod error; +mod handshake; +mod header; +mod identity; +mod pq; +mod qid; +mod record; +#[cfg(any(feature = "test-utils", test))] +mod testing; + +pub use crypto::*; +pub use encrypted::*; +pub use encrypted_message::*; +pub use error::*; +pub use handshake::*; +pub use header::*; +pub use identity::*; +pub use pq::*; +pub use qid::*; +pub use record::*; +#[cfg(any(feature = "test-utils", test))] +pub use testing::*; + +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/pq.rs b/ql-wire/src/pq.rs new file mode 100644 index 00000000..afd43f91 --- /dev/null +++ b/ql-wire/src/pq.rs @@ -0,0 +1,109 @@ +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; + +ql_codec::codec! { + #[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; + + 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); + } +} + +ql_codec::codec! { + #[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); + } +} + +ql_codec::codec! { + #[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); + } +} + +ql_codec::codec! { + #[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); + } +} + +#[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..7d32a1f8 --- /dev/null +++ b/ql-wire/src/qid.rs @@ -0,0 +1,14 @@ +use ql_common::QID; + +use crate::{MlKemPublicKey, QlHash, ML_KEM_SUITE_TAG}; + +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/record.rs b/ql-wire/src/record.rs new file mode 100644 index 00000000..6484b7d5 --- /dev/null +++ b/ql-wire/src/record.rs @@ -0,0 +1,92 @@ +use ql_codec::{ByteSlice, Decode, Encode}; +use ql_common::QID; + +use crate::{ + encrypted_message::EncryptedMessage, + handshake::{Ik1, Ik2, Xx1, Xx2, Xx3, Xx4}, + Error, RouteHeader, SessionHeader, QL_WIRE_VERSION, +}; + +pub fn encode_record(out: &mut W, header: RecordHeader, body: &T) +where + W: bytes::BufMut + ?Sized, + T: Encode + ?Sized, +{ + header.encode(out); + body.encode(out); +} + +pub fn encode_record_vec(header: RecordHeader, body: &T) -> Vec { + let mut out = Vec::with_capacity(header.encoded_len() + body.encoded_len()); + encode_record(&mut out, header, body); + out +} + +pub fn decode_record(bytes: B) -> Result<(RecordHeader, T), Error> +where + T: Decode, + B: ByteSlice, +{ + let mut reader = ql_codec::Reader::new(bytes); + Ok((reader.decode()?, reader.decode()?)) +} + +ql_codec::codec! { + #[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::() + QID::SIZE * 2 + size_of::(); + + pub fn new(route: RouteHeader, record_type: RecordType) -> Self { + Self { + version: QL_WIRE_VERSION, + route, + record_type, + } + } +} + +ql_codec::codec! { + #[derive(Debug, Clone, Copy, PartialEq, Eq)] + pub enum RecordType { + Handshake = 1, + Session = 2, + } +} + +ql_codec::codec! { + #[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, + } +} + +ql_codec::codec! { + #[derive(Debug, Clone, PartialEq, Eq)] + pub struct QlSessionRecord { + pub header: SessionHeader, + pub payload: EncryptedMessage, + } +} + +impl QlSessionRecord { + pub fn into_owned(self) -> QlSessionRecord> { + QlSessionRecord { + header: self.header, + payload: self.payload.into_owned(), + } + } +} diff --git a/ql-wire/src/testing.rs b/ql-wire/src/testing.rs new file mode 100644 index 00000000..1ce88b55 --- /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"), + crate::generate_identity(crypto, "bob"), + ) +} + +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.as_bytes()).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.as_bytes()).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(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(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([0; SessionKey::SIZE]), + ) + } + + fn mlkem_decapsulate( + &self, + _private_key: &MlKemPrivateKey, + _ciphertext: &MlKemCiphertext, + ) -> SessionKey { + SessionKey([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..680dc141 --- /dev/null +++ b/ql-wire/src/tests.rs @@ -0,0 +1,1065 @@ +use ql_codec::{Decode, Encode, Varint}; +use ql_common::{ResetCode, StreamId, QID}; + +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 handshake_id(id: u32) -> HandshakeId { + HandshakeId(id) +} + +fn handshake_transport_params(window: u32) -> TransportParams { + TransportParams { + initial_stream_receive_window: window, + } +} + +fn route(sender: u8, recipient: u8) -> RouteHeader { + RouteHeader { + sender: QID([sender; QID::SIZE]), + recipient: QID([recipient; QID::SIZE]), + } +} + +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>], +) -> 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, route, session_key).as_slice()) +} + +#[test] +fn peer_bundle_round_trip() { + let crypto = SoftwareCrypto; + let mut identity = generate_identity(&crypto, "alice"); + identity.capabilities = 1231; + 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"); +} + +#[test] +fn handshake_record_round_trip_supports_ik_kk_and_xx() { + let ik = QlHandshakeRecord::Ik1(Ik1 { + 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: 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); + 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(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); + 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 { + handshake_id: handshake_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_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, + } + ); + assert_eq!(decode_handshake_record(xx_encoded.as_slice()), xx); +} + +#[test] +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_ik_initiator( + &crypto, + initiator, + responder.bundle(), + TransportParams::default(), + ); + let mut responder_state = + IkHandshake::new_ik_responder(&crypto, responder, None, TransportParams::default()); + + 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_id(77)).unwrap(); + m2.handshake_id = HandshakeId(78); + + assert_eq!( + initiator_state.read_2(&crypto, responder_to_initiator, &m2), + Err(Error::InvalidHandshakeId) + ); +} + +#[test] +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 = IkHandshake::new_kk_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + TransportParams::default(), + ); + let mut responder_state = IkHandshake::new_kk_responder( + &crypto, + responder, + initiator.bundle(), + TransportParams::default(), + ); + + 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_id(88)).unwrap(); + let tampered_route = route(9, 1); + + assert_eq!( + initiator_state.read_2(&crypto, tampered_route, &m2), + Err(Error::InvalidRouteHeader) + ); +} + +#[test] +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_ik_initiator( + &crypto, + initiator, + responder.bundle(), + handshake_transport_params(4096), + ); + let mut responder_state = + IkHandshake::new_ik_responder(&crypto, responder, None, handshake_transport_params(8192)); + + 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_id(89)).unwrap(); + m2.transport_params.initial_stream_receive_window += 1; + + assert_eq!( + initiator_state.read_2(&crypto, responder_to_initiator, &m2), + Err(Error::DecryptFailed) + ); +} + +#[test] +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_ik_initiator( + &crypto, + initiator, + responder.bundle(), + TransportParams::default(), + ); + let mut responder_state = + IkHandshake::new_ik_responder(&crypto, responder, None, TransportParams::default()); + + let m1 = initiator_state.write_1(&crypto, handshake_id(90)).unwrap(); + initiator_to_responder.sender = QID([9; QID::SIZE]); + + assert_eq!( + responder_state.read_1(&crypto, initiator_to_responder, &m1), + Err(Error::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"); + let (initiator_to_responder, _) = identity_routes(&initiator, &responder); + + let mut initiator_state = IkHandshake::new_ik_initiator( + &crypto, + initiator, + responder.bundle(), + TransportParams::default(), + ); + let mut responder_state = IkHandshake::new_ik_responder( + &crypto, + responder, + Some(bogus.bundle()), + TransportParams::default(), + ); + + let m1 = initiator_state.write_1(&crypto, handshake_id(91)).unwrap(); + + assert_eq!( + responder_state.read_1(&crypto, initiator_to_responder, &m1), + Err(Error::InvalidRouteHeader) + ); +} + +#[test] +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); + let mut initiator_state = IkHandshake::new_ik_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + initiator_params, + ); + let mut responder_state = + IkHandshake::new_ik_responder(&crypto, responder.clone(), None, responder_params); + + 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_id(11)).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(); + + 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.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_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); + let mut initiator_state = IkHandshake::new_ik_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + initiator_params, + ); + let mut responder_state = IkHandshake::new_ik_responder( + &crypto, + responder.clone(), + Some(initiator.bundle()), + responder_params, + ); + + 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_id(12)).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(); + + 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.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_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); + let mut initiator_state = IkHandshake::new_kk_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + initiator_params, + ); + let mut responder_state = IkHandshake::new_kk_responder( + &crypto, + responder.clone(), + initiator.bundle(), + responder_params, + ); + + 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_id(21)).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(); + + 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.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 (initiator_to_responder, responder_to_initiator) = identity_routes(&initiator, &responder); + + let mut initiator_state = IkHandshake::new_kk_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + handshake_transport_params(12288), + ); + let mut responder_state = IkHandshake::new_kk_responder( + &crypto, + responder, + initiator.bundle(), + handshake_transport_params(24576), + ); + + 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_id(22)).unwrap(); + m2.transport_params.initial_stream_receive_window += 1; + + assert_eq!( + initiator_state.read_2(&crypto, responder_to_initiator, &m2), + Err(Error::DecryptFailed) + ); +} + +#[test] +fn xx_handshake_rejects_tampered_pairing_id() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + let token = PairingToken([7; PairingToken::SIZE]); + let (initiator_to_responder, _) = identity_routes(&initiator, &responder); + + 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_id(31)).unwrap(); + m1.pairing_id = PairingId([8; PairingId::SIZE]); + + assert_eq!( + responder_state.read_1(&crypto, initiator_to_responder, &m1), + Err(Error::InvalidPairingId) + ); +} + +#[test] +fn xx_handshake_rejects_tampered_sender_or_recipient() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + let token = PairingToken([7; PairingToken::SIZE]); + + 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 m1 = initiator_state.write_1(&crypto, handshake_id(31)).unwrap(); + let (mut route, _) = identity_routes(&initiator, &responder); + route.sender = responder.qid; + + assert_eq!( + responder_state.read_1(&crypto, route, &m1), + Err(Error::InvalidRouteHeader) + ); + + 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 m1 = initiator_state.write_1(&crypto, handshake_id(31)).unwrap(); + let (mut route, _) = identity_routes(&initiator, &responder); + route.recipient = initiator.qid; + + assert_eq!( + responder_state.read_1(&crypto, route, &m1), + Err(Error::InvalidRouteHeader) + ); +} + +#[test] +fn xx_handshake_rejects_tampered_transport_params() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + 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, + 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_id(32)).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .unwrap(); + + let mut m2 = responder_state.write_2(&crypto, handshake_id(32)).unwrap(); + m2.transport_params.initial_stream_receive_window += 1; + + assert_eq!( + initiator_state.read_2(&crypto, responder_to_initiator, &m2), + Err(Error::DecryptFailed) + ); +} + +#[test] +fn xx_handshake_round_trip_derives_matching_transport_and_learns_remote() { + let crypto = SoftwareCrypto; + let (initiator, responder) = test_identities(&crypto); + 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); + 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_id(33)).unwrap(); + responder_state + .read_1(&crypto, initiator_to_responder, &m1) + .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_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_id(33)).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(); + + 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.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_authenticates_header() { + let crypto = SoftwareCrypto; + let header = SessionHeader { seq: RecordSeq(11) }; + let session_route = route(1, 2); + let body = vec![ + SessionFrame::Ping, + SessionFrame::Unpair, + SessionFrame::Ack( + RecordAck::from_ranges([RecordSeq(20)..=RecordSeq(23), RecordSeq(12)..=RecordSeq(13)]) + .unwrap(), + ), + SessionFrame::StreamWindow(StreamWindow { + stream_id: StreamId(9), + maximum_offset: Varint(65_536), + }), + SessionFrame::StreamData(StreamData { + stream_id: StreamId(9), + offset: Varint(1024), + header: None, + bytes: b"hello".to_vec(), + fin: true, + }), + SessionFrame::StreamReset(StreamReset { + stream_id: StreamId(9), + target: ResetTarget::Both, + code: ResetCode::CANCELLED, + }), + SessionFrame::Close(SessionClose { + code: SessionCloseCode::TIMEOUT, + }), + ]; + let session_key = SessionKey([7; SessionKey::SIZE]); + let record = encrypt_record(&crypto, session_route, header, &session_key, &body); + + 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, + } + ); + let decoded = decode_session_record(bytes.as_slice()); + assert_eq!(decoded.header, header); + let encrypted = decoded.payload; + + let decrypted = encrypted::decrypt_record( + &crypto, + &record_header, + &header, + encrypted.clone(), + &session_key, + ) + .unwrap(); + assert_eq!(decode_session_frames(&decrypted).unwrap(), body); + + 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(Error::DecryptFailed) + ); + + let wrong_seq_header = SessionHeader { + seq: RecordSeq(header.seq.0 + 1), + }; + assert_eq!( + encrypted::decrypt_record( + &crypto, + &record_header, + &wrong_seq_header, + encrypted, + &session_key, + ), + Err(Error::DecryptFailed) + ); +} + +#[test] +fn protocol_record_size_breakdown() { + fn print_size(label: &str, size: usize) { + 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); + + let mut ik_initiator = IkHandshake::new_ik_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + TransportParams::default(), + ); + let mut ik_responder = + IkHandshake::new_ik_responder(&crypto, responder.clone(), None, TransportParams::default()); + + 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_id(101)).unwrap(); + ik_initiator + .read_2(&crypto, responder_to_initiator, &ik2) + .unwrap(); + + let ik1 = QlHandshakeRecord::Ik1(ik1); + let ik2 = QlHandshakeRecord::Ik2(ik2); + + let mut kk_initiator = IkHandshake::new_kk_initiator( + &crypto, + initiator.clone(), + responder.bundle(), + TransportParams::default(), + ); + let mut kk_responder = IkHandshake::new_kk_responder( + &crypto, + responder.clone(), + initiator.bundle(), + TransportParams::default(), + ); + + 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_id(201)).unwrap(); + kk_initiator + .read_2(&crypto, responder_to_initiator, &kk2) + .unwrap(); + + let kk1 = QlHandshakeRecord::Kk1(kk1); + let kk2 = QlHandshakeRecord::Kk2(kk2); + + let token = PairingToken([0x42; PairingToken::SIZE]); + 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_id(301)).unwrap(); + xx_responder + .read_1(&crypto, initiator_to_responder, &xx1) + .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_id(301)).unwrap(); + xx_responder + .read_3(&crypto, initiator_to_responder, &xx3) + .unwrap(); + + let xx4 = xx_responder.write_4(&crypto, handshake_id(301)).unwrap(); + xx_initiator + .read_4(&crypto, responder_to_initiator, &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_route = RouteHeader { + sender: initiator.qid, + recipient: responder.qid, + }; + let session_ping = encrypt_record( + &crypto, + session_route, + SessionHeader { seq: RecordSeq(1) }, + &session.tx_key, + &[SessionFrame::Ping], + ); + let session_ack = encrypt_record( + &crypto, + session_route, + SessionHeader { seq: RecordSeq(2) }, + &session.tx_key, + &[SessionFrame::Ack( + RecordAck::from_ranges([RecordSeq(6)..=RecordSeq(6), RecordSeq(1)..=RecordSeq(2)]) + .unwrap(), + )], + ); + let session_unpair = encrypt_record( + &crypto, + session_route, + SessionHeader { seq: RecordSeq(3) }, + &session.tx_key, + &[SessionFrame::Unpair], + ); + let session_stream_empty = encrypt_record( + &crypto, + session_route, + SessionHeader { seq: RecordSeq(4) }, + &session.tx_key, + &[SessionFrame::StreamData(StreamData { + stream_id: StreamId(1), + offset: Varint(0), + header: None, + fin: false, + bytes: Vec::new(), + })], + ); + let session_close = encrypt_record( + &crypto, + session_route, + SessionHeader { seq: RecordSeq(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", 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", + record_size(&session_stream_empty), + ); + 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::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(), + }, + ); +} 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) -}