diff options
| author | zirkonya <zirkonya@iridium.lan> | 2026-09-01 09:51:18 +0200 |
|---|---|---|
| committer | zirkonya <zirkonya@iridium.lan> | 2026-09-01 09:51:18 +0200 |
| commit | be62f065a62048b798e8bb2b8e3699bdd5b6d517 (patch) | |
| tree | 8b3f4b05838b5fa34c3ae24e35004343baadf2af | |
| parent | fdc02f07cbd1994c1efb057f24a37a96faaa51fa (diff) | |
add proc macros ; benchmark ; example
| -rw-r--r-- | Cargo.lock | 491 | ||||
| -rw-r--r-- | Cargo.toml | 9 | ||||
| -rw-r--r-- | benches/codec_bench.rs | 738 | ||||
| -rw-r--r-- | examples/codec.rs | 401 | ||||
| -rw-r--r-- | macros/Cargo.toml | 3 | ||||
| -rw-r--r-- | macros/src/decode.rs | 508 | ||||
| -rw-r--r-- | macros/src/encode.rs | 516 | ||||
| -rw-r--r-- | macros/src/lib.rs | 33 | ||||
| -rw-r--r-- | src/codec/decode.rs | 251 | ||||
| -rw-r--r-- | src/codec/encode.rs | 336 | ||||
| -rw-r--r-- | src/lib.rs | 11 | ||||
| -rw-r--r-- | src/transport/receiver.rs | 13 | ||||
| -rw-r--r-- | src/transport/sender.rs | 12 | ||||
| -rw-r--r-- | src/transport/tcp.rs | 32 | ||||
| -rw-r--r-- | src/transport/udp.rs | 26 | ||||
| -rw-r--r-- | src/types/prefix.rs | 3 | ||||
| -rw-r--r-- | src/types/prefix/count.rs | 23 | ||||
| -rw-r--r-- | src/types/prefix/length.rs | 42 | ||||
| -rw-r--r-- | tests/count_prefix.rs | 305 | ||||
| -rw-r--r-- | tests/encode_decode.rs | 589 | ||||
| -rw-r--r-- | tests/len_prefix.rs | 310 | ||||
| -rw-r--r-- | tests/tmp/todo.md | 356 |
22 files changed, 3114 insertions, 1894 deletions
@@ -3,6 +3,27 @@ version = 4 [[package]] +name = "aho-corasick" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c982642fa9e8606056828ee9a8505737230110bb1099153c79efe865c59d12ba" +dependencies = [ + "memchr", +] + +[[package]] +name = "anes" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4b46cbb362ab8752921c97e041f5e366ee6297bd428a31275b9fcf1e380f7299" + +[[package]] +name = "anstyle" +version = "1.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" + +[[package]] name = "async-channel" version = "2.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -153,12 +174,76 @@ dependencies = [ ] [[package]] +name = "bumpalo" +version = "3.20.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" + +[[package]] +name = "cast" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "37b2a672a2cb129a2e41c10b1224bb368f9f37a2b16b612598138befd7b37eb5" + +[[package]] name = "cfg-if" version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" [[package]] +name = "ciborium" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "42e69ffd6f0917f5c029256a24d0161db17cea3997d185db0d35926308770f0e" +dependencies = [ + "ciborium-io", + "ciborium-ll", + "serde", +] + +[[package]] +name = "ciborium-io" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "05afea1e0a06c9be33d539b876f1ce3692f4afea2cb41f740e7743225ed1c757" + +[[package]] +name = "ciborium-ll" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "57663b653d948a338bfb3eeba9bb2fd5fcfaecb9e199e87e1eda4d9e8b240fd9" +dependencies = [ + "ciborium-io", + "half", +] + +[[package]] +name = "clap" +version = "4.6.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "473c7e07f409a8d772161724aa8db6a765a2532a70f9667eeb7b49d3d02fbdca" +dependencies = [ + "clap_builder", +] + +[[package]] +name = "clap_builder" +version = "4.6.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b48fea5a88e9ae728a2dcbedbfc0e730f7d60da42e1cb049a83c9fb8b789889" +dependencies = [ + "anstyle", + "clap_lex", +] + +[[package]] +name = "clap_lex" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" + +[[package]] name = "concurrent-queue" version = "2.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -168,12 +253,79 @@ dependencies = [ ] [[package]] +name = "criterion" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2b12d017a929603d80db1831cd3a24082f8137ce19c69e6447f54f5fc8d692f" +dependencies = [ + "anes", + "cast", + "ciborium", + "clap", + "criterion-plot", + "is-terminal", + "itertools", + "num-traits", + "once_cell", + "oorandom", + "plotters", + "rayon", + "regex", + "serde", + "serde_derive", + "serde_json", + "tinytemplate", + "walkdir", +] + +[[package]] +name = "criterion-plot" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b50826342786a51a89e2da3a28f1c32b06e387201bc2d19791f622c673706b1" +dependencies = [ + "cast", + "itertools", +] + +[[package]] +name = "crossbeam-deque" +version = "0.8.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5181e0de7b61eb03a81e347d6dd8797bae9da5146707b51077e2d71a54ec0ceb" +dependencies = [ + "crossbeam-epoch", + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-epoch" +version = "0.9.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d6914041f254d6e9176c01941b21115dcfb7089e55135a35411081bd106ef3f" +dependencies = [ + "crossbeam-utils", +] + +[[package]] name = "crossbeam-utils" version = "0.8.22" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17" [[package]] +name = "crunchy" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" + +[[package]] +name = "either" +version = "1.18.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "252afb9ae5eaa683babdc6a068b3f5726eb19e05070c731f9b2a23a7c3e8ed34" + +[[package]] name = "errno" version = "0.3.14" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -235,6 +387,24 @@ dependencies = [ ] [[package]] +name = "futures-task" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cd417de3d1d015fc3bfd2b1ea46dfc7bab72ef86f1cc7cc9c78e728b34a6d1fd" + +[[package]] +name = "futures-util" +version = "0.3.34" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0d50a92467f8ba5dd6e3ee5d4bd04d73ab2e4e1c44474a0674821dfce14b79bc" +dependencies = [ + "futures-core", + "futures-task", + "pin-project-lite", + "slab", +] + +[[package]] name = "getset" version = "0.1.7" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -246,12 +416,60 @@ dependencies = [ ] [[package]] +name = "half" +version = "2.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" +dependencies = [ + "cfg-if", + "crunchy", + "zerocopy", +] + +[[package]] name = "hermit-abi" version = "0.5.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc0fef456e4baa96da950455cd02c081ca953b141298e41db3fc7e36b1da849c" [[package]] +name = "is-terminal" +version = "0.4.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3640c1c38b8e4e43584d8df18be5fc6b0aa314ce6ebf51b53313d4306cca8e46" +dependencies = [ + "hermit-abi", + "libc", + "windows-sys", +] + +[[package]] +name = "itertools" +version = "0.10.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b0fd2260e829bddf4cb6ea802289de2f86d6a7a690192fbe91b3f46e0f2c8473" +dependencies = [ + "either", +] + +[[package]] +name = "itoa" +version = "1.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" + +[[package]] +name = "js-sys" +version = "0.3.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0e0c1080212aad755ea003d18543e8768dd432c48819efd73a7bf1e39b7a5a3a" +dependencies = [ + "cfg-if", + "futures-util", + "wasm-bindgen", +] + +[[package]] name = "libc" version = "0.2.189" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -264,6 +482,33 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" [[package]] +name = "memchr" +version = "2.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" + +[[package]] +name = "num-traits" +version = "0.2.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071dfc062690e90b734c0b2273ce72ad0ffa95f0c74596bc250dcfd960262841" +dependencies = [ + "autocfg", +] + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "oorandom" +version = "11.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6790f58c7ff633d8771f42965289203411a5e5c68388703c06e14f24770b41e" + +[[package]] name = "parking" version = "2.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -287,6 +532,34 @@ dependencies = [ ] [[package]] +name = "plotters" +version = "0.3.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5aeb6f403d7a4911efb1e33402027fc44f29b5bf6def3effcc22d7bb75f2b747" +dependencies = [ + "num-traits", + "plotters-backend", + "plotters-svg", + "wasm-bindgen", + "web-sys", +] + +[[package]] +name = "plotters-backend" +version = "0.3.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df42e13c12958a16b3f7f4386b9ab1f3e7933914ecea48da7139435263a4172a" + +[[package]] +name = "plotters-svg" +version = "0.3.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "51bae2ac328883f7acdfea3d66a7c35751187f870bc81f94563733a154d7a670" +dependencies = [ + "plotters-backend", +] + +[[package]] name = "polling" version = "3.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -319,6 +592,55 @@ dependencies = [ ] [[package]] +name = "rayon" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fb39b166781f92d482534ef4b4b1b2568f42613b53e5b6c160e24cfbfa30926d" +dependencies = [ + "either", + "rayon-core", +] + +[[package]] +name = "rayon-core" +version = "1.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "22e18b0f0062d30d4230b2e85ff77fdfe4326feb054b9783a3460d8435c8ab91" +dependencies = [ + "crossbeam-deque", + "crossbeam-utils", +] + +[[package]] +name = "regex" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f020237b6c8eed93db2e2cb53c00c60a8e1bc73da7d073199a1180401450218d" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ad8553b9b26413251cbf30e620595c7a41b3887f03da04579c0e6b0d6a06b4b2" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + +[[package]] name = "rustix" version = "1.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -332,6 +654,64 @@ dependencies = [ ] [[package]] +name = "rustversion" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" + +[[package]] +name = "same-file" +version = "1.0.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93fc1dc3aaa9bfed95e02e6eadabb4baf7e3078b0bd1b4d7b6b0b68378900502" +dependencies = [ + "winapi-util", +] + +[[package]] +name = "serde" +version = "1.0.229" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4148590afebada386688f18773da617792bf2ef03ffc1e4cbd2b1d45b023e0ba" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.229" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "67dca2c9c51e58a4791a4b1ed58308b39c64224d349a935ab5039aa360942a48" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.229" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.3", +] + +[[package]] +name = "serde_json" +version = "1.0.151" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c841b55ecdae098c80dcae9cf767f6f8a0c2cdb3416bbef72181df4d0fe73f14" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] name = "signal-hook-registry" version = "1.4.8" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -407,6 +787,16 @@ dependencies = [ ] [[package]] +name = "tinytemplate" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be4d6b5f19ff7664e8c98d03e2139cb510db9b0a60b55f8e8709b689d939b6bc" +dependencies = [ + "serde", + "serde_json", +] + +[[package]] name = "tokio" version = "1.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -422,6 +812,80 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" [[package]] +name = "walkdir" +version = "2.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29790946404f91d9c5d06f9874efddea1dc06c5efe94541a7d6863108e3a5e4b" +dependencies = [ + "same-file", + "winapi-util", +] + +[[package]] +name = "wasm-bindgen" +version = "0.2.127" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1b70935747edd64d89de3efa29d73789b806c15798f8e7dca4d8ac356b50ce70" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.127" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77775f8f3f7217702089053b94958f8f54061a3f663417df76e19cbdcca29bc1" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.127" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e11d33f857dc2fb11b8bc75aee111aa9cbeb12cd9f25efd3d4c2a3dd4e235284" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn 2.0.119", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.127" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7ef64dbcc55df09c7e5a46182d181c2cfa3e925f3da937ea764728b4bbb9dcbf" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "web-sys" +version = "0.3.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c435338968042f4f59a557f690a253676d47ce13ceb55d70100e7facf6620a30" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "winapi-util" +version = "0.1.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c2a7b1c03c876122aa43f3020e6c3c3ee5c05081c9a00739faf7503aeba10d22" +dependencies = [ + "windows-sys", +] + +[[package]] name = "windows-link" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -437,9 +901,36 @@ dependencies = [ ] [[package]] +name = "zerocopy" +version = "0.8.56" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "556764e583adb45a9f8d413c2a147fa7e8d821e48e12b14fd560b607998b75eb" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.56" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2ab42fc20575779bd240faa45f94a74256f755c0fa9e89f0ede20d91d0cdfc1" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "zmij" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" + +[[package]] name = "zr_protocol" version = "0.1.1" dependencies = [ + "criterion", "getset", "smol", "thiserror", @@ -12,7 +12,16 @@ tokio = { version = "1.53.1", optional = true } # Macros zr_protocol_macros = { version = "0.1.0", path = "macros", optional = true } +[dev-dependencies] +criterion = "0.5" + +[[bench]] +name = "codec_bench" +harness = false +required-features = ["macros"] + [features] +default = ["macros"] # Macros macros = ["dep:zr_protocol_macros"] # Runtime diff --git a/benches/codec_bench.rs b/benches/codec_bench.rs new file mode 100644 index 0000000..522a4cc --- /dev/null +++ b/benches/codec_bench.rs @@ -0,0 +1,738 @@ +use criterion::{BenchmarkId, Criterion, Throughput, black_box, criterion_group, criterion_main}; +use zr_protocol::codec::{decode::Decode, encode::Encode}; +use zr_protocol::context::Context; +use zr_protocol::types::prefix::{CountPrefix, LenPrefixed}; +use zr_protocol_macros::Codec; + +const DEFAULT_CTX: Context<()> = Context { data: () }; + +#[derive(Codec, Clone, Debug, PartialEq)] +struct BenchStruct { + a: u8, + b: u16, + c: u32, + d: u64, + e: i8, + f: i16, + g: i32, + h: i64, + i: f32, + j: f64, + k: bool, +} + +#[derive(Codec, Clone, Debug, PartialEq)] +struct BenchStructWithCollections { + #[codec(count(u16))] + items: Vec<u32>, + #[codec(count(u16))] + tags: Vec<String>, + #[codec(len(u16))] + payload: Vec<u8>, +} + +#[derive(Codec, Clone, Debug, PartialEq)] +enum BenchEnum { + A(u8), + B { x: u16, y: u16 }, + C(u32, u32), + D, +} + +#[derive(Codec, Clone, Debug, PartialEq)] +struct ComplexStruct { + id: u64, + #[codec(count(u16))] + values: Vec<i32>, + #[codec(len(u16))] + data: Vec<u8>, + #[codec(count(u16))] + names: Vec<String>, + #[codec(count(u8))] + flags: Vec<bool>, + metadata: HashMap<String, u32>, + tags: HashSet<String>, +} + +use std::collections::{HashMap, HashSet}; + +fn encode_to_vec<T: Encode<()> + ?Sized>(value: &T) -> Vec<u8> { + let mut buf = Vec::with_capacity(1024); + value.encode(&mut buf, &DEFAULT_CTX).unwrap(); + buf +} + +fn bench_primitives_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("primitives_encode"); + + macro_rules! bench_encode { + ($name:expr, $value:expr) => { + group.bench_function($name, |b| { + b.iter(|| { + let mut buf = Vec::with_capacity(64); + black_box($value).encode(&mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }); + }; + } + + bench_encode!("u8", &0xFFu8); + bench_encode!("u16", &0xFFFFu16); + bench_encode!("u32", &0xFFFFFFFFu32); + bench_encode!("u64", &0xFFFFFFFFFFFFFFFFu64); + bench_encode!("u128", &u128::MAX); + bench_encode!("usize", &usize::MAX); + bench_encode!("i8", &-128i8); + bench_encode!("i16", &-32768i16); + bench_encode!("i32", &-2147483648i32); + bench_encode!("i64", &-9223372036854775808i64); + bench_encode!("i128", &i128::MIN); + bench_encode!("isize", &isize::MIN); + bench_encode!("f32", &std::f32::consts::PI); + bench_encode!("f64", &std::f64::consts::PI); + bench_encode!("bool_true", &true); + bench_encode!("bool_false", &false); +} + +fn bench_primitives_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("primitives_decode"); + + macro_rules! bench_decode { + ($name:expr, $bytes:expr, $ty:ty) => { + group.bench_function($name, |b| { + b.iter(|| { + let mut reader = &$bytes[..]; + black_box(<$ty as Decode<()>>::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }); + }; + } + + bench_decode!("u8", encode_to_vec(&0xFFu8), u8); + bench_decode!("u16", encode_to_vec(&0xFFFFu16), u16); + bench_decode!("u32", encode_to_vec(&0xFFFFFFFFu32), u32); + bench_decode!("u64", encode_to_vec(&0xFFFFFFFFFFFFFFFFu64), u64); + bench_decode!("u128", encode_to_vec(&u128::MAX), u128); + bench_decode!("usize", encode_to_vec(&usize::MAX), usize); + bench_decode!("i8", encode_to_vec(&-128i8), i8); + bench_decode!("i16", encode_to_vec(&-32768i16), i16); + bench_decode!("i32", encode_to_vec(&-2147483648i32), i32); + bench_decode!("i64", encode_to_vec(&-9223372036854775808i64), i64); + bench_decode!("i128", encode_to_vec(&i128::MIN), i128); + bench_decode!("isize", encode_to_vec(&isize::MIN), isize); + bench_decode!("f32", encode_to_vec(&std::f32::consts::PI), f32); + bench_decode!("f64", encode_to_vec(&std::f64::consts::PI), f64); + bench_decode!("bool_true", encode_to_vec(&true), bool); + bench_decode!("bool_false", encode_to_vec(&false), bool); +} + +fn bench_struct_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("struct_encode"); + let val = BenchStruct { + a: 0xFF, + b: 0xFFFF, + c: 0xFFFFFFFF, + d: 0xFFFFFFFFFFFFFFFF, + e: -128, + f: -32768, + g: -2147483648, + h: -9223372036854775808, + i: std::f32::consts::PI, + j: std::f64::consts::PI, + k: true, + }; + + group.bench_function("encode", |b| { + b.iter(|| { + let mut buf = Vec::with_capacity(128); + black_box(&val).encode(&mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }); +} + +fn bench_struct_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("struct_decode"); + let val = BenchStruct { + a: 0xFF, + b: 0xFFFF, + c: 0xFFFFFFFF, + d: 0xFFFFFFFFFFFFFFFF, + e: -128, + f: -32768, + g: -2147483648, + h: -9223372036854775808, + i: std::f32::consts::PI, + j: std::f64::consts::PI, + k: true, + }; + let bytes = encode_to_vec(&val); + + group.bench_function("decode", |b| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box(BenchStruct::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }); +} + +fn bench_collections_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("collections_encode"); + + for size in [0, 10, 100, 1000, 10000].iter() { + let items: Vec<u32> = (0..*size).collect(); + let tags: Vec<String> = (0..*size).map(|i| format!("tag{}", i)).collect(); + let payload: Vec<u8> = (0..*size).map(|i| (i % 256) as u8).collect(); + + let val = BenchStructWithCollections { + items: items.clone(), + tags: tags.clone(), + payload: payload.clone(), + }; + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input( + BenchmarkId::new("Vec_u32_count", size), + &items, + |b, items| { + b.iter(|| { + let mut buf = Vec::with_capacity(1024); + CountPrefix::<u32, u16, _>::new(items.clone()) + .encode(&mut buf, &DEFAULT_CTX) + .unwrap(); + black_box(buf) + }) + }, + ); + + group.bench_with_input( + BenchmarkId::new("Vec_String_count", size), + &tags, + |b, tags| { + b.iter(|| { + let mut buf = Vec::with_capacity(1024); + CountPrefix::<String, u16, _>::new(tags.clone()) + .encode(&mut buf, &DEFAULT_CTX) + .unwrap(); + black_box(buf) + }) + }, + ); + + group.bench_with_input( + BenchmarkId::new("Vec_u8_len", size), + &payload, + |b, payload| { + b.iter(|| { + let mut buf = Vec::with_capacity(1024); + LenPrefixed::<u16, _>::new(payload.clone()) + .encode(&mut buf, &DEFAULT_CTX) + .unwrap(); + black_box(buf) + }) + }, + ); + + group.bench_with_input(BenchmarkId::new("full_struct", size), &val, |b, val| { + b.iter(|| { + let mut buf = Vec::with_capacity(4096); + black_box(val).encode(&mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }); + } +} + +fn bench_collections_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("collections_decode"); + + for size in [0, 10, 100, 1000, 10000].iter() { + let items: Vec<u32> = (0..*size).collect(); + let tags: Vec<String> = (0..*size).map(|i| format!("tag{}", i)).collect(); + let payload: Vec<u8> = (0..*size).map(|i| (i % 256) as u8).collect(); + + let val = BenchStructWithCollections { + items: items.clone(), + tags: tags.clone(), + payload: payload.clone(), + }; + let bytes = encode_to_vec(&val); + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input( + BenchmarkId::new("Vec_u32_count", size), + &bytes, + |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box( + CountPrefix::<u32, u16, Vec<u32>>::decode(&mut reader, &DEFAULT_CTX) + .unwrap(), + ) + }) + }, + ); + + group.bench_with_input( + BenchmarkId::new("Vec_String_count", size), + &bytes, + |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box( + CountPrefix::<String, u16, Vec<String>>::decode(&mut reader, &DEFAULT_CTX) + .unwrap(), + ) + }) + }, + ); + + group.bench_with_input(BenchmarkId::new("Vec_u8_len", size), &bytes, |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box(LenPrefixed::<u16, Vec<u8>>::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }); + + group.bench_with_input(BenchmarkId::new("full_struct", size), &bytes, |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box(BenchStructWithCollections::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }); + } +} + +fn bench_string_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("string_encode"); + + for size in [0, 10, 100, 1000, 10000].iter() { + let s = "x".repeat(*size); + let s_clone = s.clone(); + + group.throughput(Throughput::Bytes(*size as u64)); + + group.bench_with_input(BenchmarkId::new("String_raw", size), &s_clone, |b, s| { + b.iter(|| { + let mut buf = Vec::with_capacity(*size + 8); + black_box(s).encode(&mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }); + + group.bench_with_input(BenchmarkId::new("str_raw", size), &s, |b, s| { + b.iter(|| { + let mut buf = Vec::with_capacity(*size + 8); + black_box(s.as_str()) + .encode(&mut buf, &DEFAULT_CTX) + .unwrap(); + black_box(buf) + }) + }); + + group.bench_with_input( + BenchmarkId::new("LenPrefixed_String", size), + &s_clone, + |b, s| { + b.iter(|| { + let mut buf = Vec::with_capacity(*size + 8); + LenPrefixed::<u16, _>::new(s.clone()) + .encode(&mut buf, &DEFAULT_CTX) + .unwrap(); + black_box(buf) + }) + }, + ); + } +} + +fn bench_string_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("string_decode"); + + for size in [0, 10, 100, 1000, 10000].iter() { + let s = "x".repeat(*size); + + let raw_bytes = encode_to_vec(&s); + let len_prefixed_bytes = encode_to_vec(&LenPrefixed::<u16, _>::new(s.clone())); + + group.throughput(Throughput::Bytes(*size as u64)); + + group.bench_with_input( + BenchmarkId::new("String_raw", size), + &raw_bytes, + |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box(String::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }, + ); + + group.bench_with_input( + BenchmarkId::new("LenPrefixed_String", size), + &len_prefixed_bytes, + |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box( + LenPrefixed::<u16, String>::decode(&mut reader, &DEFAULT_CTX).unwrap(), + ) + }) + }, + ); + } +} + +fn bench_enum_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("enum_encode"); + + let variants = [ + ("A", BenchEnum::A(0xFF)), + ( + "B", + BenchEnum::B { + x: 0xFFFF, + y: 0xFFFF, + }, + ), + ("C", BenchEnum::C(0xFFFFFFFF, 0xFFFFFFFF)), + ("D", BenchEnum::D), + ]; + + for (name, val) in variants { + group.bench_function(name, |b| { + b.iter(|| { + let mut buf = Vec::with_capacity(32); + black_box(&val).encode(&mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }); + } +} + +fn bench_enum_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("enum_decode"); + + let variants = [ + ("A", encode_to_vec(&BenchEnum::A(0xFF))), + ( + "B", + encode_to_vec(&BenchEnum::B { + x: 0xFFFF, + y: 0xFFFF, + }), + ), + ("C", encode_to_vec(&BenchEnum::C(0xFFFFFFFF, 0xFFFFFFFF))), + ("D", encode_to_vec(&BenchEnum::D)), + ]; + + for (name, bytes) in variants { + group.bench_function(name, |b| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box(BenchEnum::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }); + } +} + +fn bench_roundtrip(c: &mut Criterion) { + let mut group = c.benchmark_group("roundtrip"); + + let struct_val = BenchStruct { + a: 0xFF, + b: 0xFFFF, + c: 0xFFFFFFFF, + d: 0xFFFFFFFFFFFFFFFF, + e: -128, + f: -32768, + g: -2147483648, + h: -9223372036854775808, + i: std::f32::consts::PI, + j: std::f64::consts::PI, + k: true, + }; + + group.bench_function("struct", |b| { + b.iter(|| { + let mut buf = Vec::with_capacity(128); + struct_val.encode(&mut buf, &DEFAULT_CTX).unwrap(); + let mut reader = &buf[..]; + black_box(BenchStruct::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }); + + for size in [0, 10, 100, 1000].iter() { + let items: Vec<u32> = (0..*size).collect(); + let val = CountPrefix::<u32, u16, _>::new(items); + + group.throughput(Throughput::Elements(*size as u64)); + group.bench_with_input( + BenchmarkId::new("CountPrefix_Vec_u32", size), + &val, + |b, val| { + b.iter(|| { + let mut buf = Vec::with_capacity(4096); + val.encode(&mut buf, &DEFAULT_CTX).unwrap(); + let mut reader = &buf[..]; + black_box( + CountPrefix::<u32, u16, Vec<u32>>::decode(&mut reader, &DEFAULT_CTX) + .unwrap(), + ) + }) + }, + ); + } +} + +fn bench_complex_struct_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("complex_struct_encode"); + + for size in [0, 10, 100, 1000].iter() { + let mut values = Vec::with_capacity(*size); + let mut data = Vec::with_capacity(*size * 4); + let mut names = Vec::with_capacity(*size); + let mut flags = Vec::with_capacity(*size); + let mut metadata = HashMap::new(); + let mut tags = HashSet::new(); + + for i in 0..*size { + values.push(i as i32); + data.extend_from_slice(&(i as u32).to_be_bytes()); + names.push(format!("name_{}", i)); + flags.push(i % 2 == 0); + metadata.insert(format!("key_{}", i), i as u32); + tags.insert(format!("tag_{}", i)); + } + + let val = ComplexStruct { + id: 0xDEADBEEF, + values, + data, + names, + flags, + metadata, + tags, + }; + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input(BenchmarkId::new("ComplexStruct", size), &val, |b, val| { + b.iter(|| { + let mut buf = Vec::with_capacity(16384); + black_box(val).encode(&mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }); + } +} + +fn bench_complex_struct_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("complex_struct_decode"); + + for size in [0, 10, 100, 1000].iter() { + let mut values = Vec::with_capacity(*size); + let mut data = Vec::with_capacity(*size * 4); + let mut names = Vec::with_capacity(*size); + let mut flags = Vec::with_capacity(*size); + let mut metadata = HashMap::new(); + let mut tags = HashSet::new(); + + for i in 0..*size { + values.push(i as i32); + data.extend_from_slice(&(i as u32).to_be_bytes()); + names.push(format!("name_{}", i)); + flags.push(i % 2 == 0); + metadata.insert(format!("key_{}", i), i as u32); + tags.insert(format!("tag_{}", i)); + } + + let val = ComplexStruct { + id: 0xDEADBEEF, + values, + data, + names, + flags, + metadata, + tags, + }; + let bytes = encode_to_vec(&val); + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input( + BenchmarkId::new("ComplexStruct", size), + &bytes, + |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box(ComplexStruct::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }, + ); + } +} + +fn bench_slice_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("slice_encode"); + + for size in [10, 100, 1000, 10000].iter() { + let data: Vec<u32> = (0..*size).collect(); + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input(BenchmarkId::new("slice_u32", size), &data, |b, data| { + b.iter(|| { + let mut buf = Vec::with_capacity((*size * 4 + 8) as usize); + <_ as Encode<()>>::encode(data, &mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }); + } +} + +fn bench_slice_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("slice_decode"); + + for size in [10, 100, 1000, 10000].iter() { + let data: Vec<u32> = (0..*size).collect(); + let bytes = encode_to_vec(&data); + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input(BenchmarkId::new("Vec_u32", size), &bytes, |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box(<Vec<u32> as Decode<()>>::decode(&mut reader, &DEFAULT_CTX).unwrap()) + }) + }); + } +} + +fn bench_hashmap_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("hashmap_encode"); + + for size in [0, 10, 100, 1000].iter() { + let mut map = HashMap::with_capacity(*size); + for i in 0..*size { + map.insert(format!("key_{}", i), i as u32); + } + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input( + BenchmarkId::new("HashMap_String_u32", size), + &map, + |b, map| { + b.iter(|| { + let mut buf = Vec::with_capacity(4096); + black_box(map).encode(&mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }, + ); + } +} + +fn bench_hashmap_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("hashmap_decode"); + + for size in [0, 10, 100, 1000].iter() { + let mut map = HashMap::with_capacity(*size); + for i in 0..*size { + map.insert(format!("key_{}", i), i as u32); + } + let bytes = encode_to_vec(&map); + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input( + BenchmarkId::new("HashMap_String_u32", size), + &bytes, + |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box( + <HashMap<String, u32> as Decode<()>>::decode(&mut reader, &DEFAULT_CTX) + .unwrap(), + ) + }) + }, + ); + } +} + +fn bench_hashset_encode(c: &mut Criterion) { + let mut group = c.benchmark_group("hashset_encode"); + + for size in [0, 10, 100, 1000].iter() { + let mut set = HashSet::with_capacity(*size); + for i in 0..*size { + set.insert(format!("item_{}", i)); + } + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input(BenchmarkId::new("HashSet_String", size), &set, |b, set| { + b.iter(|| { + let mut buf = Vec::with_capacity(4096); + black_box(set).encode(&mut buf, &DEFAULT_CTX).unwrap(); + black_box(buf) + }) + }); + } +} + +fn bench_hashset_decode(c: &mut Criterion) { + let mut group = c.benchmark_group("hashset_decode"); + + for size in [0, 10, 100, 1000].iter() { + let mut set = HashSet::with_capacity(*size); + for i in 0..*size { + set.insert(format!("item_{}", i)); + } + let bytes = encode_to_vec(&set); + + group.throughput(Throughput::Elements(*size as u64)); + + group.bench_with_input( + BenchmarkId::new("HashSet_String", size), + &bytes, + |b, bytes| { + b.iter(|| { + let mut reader = &bytes[..]; + black_box( + <HashSet<String> as Decode<()>>::decode(&mut reader, &DEFAULT_CTX).unwrap(), + ) + }) + }, + ); + } +} + +criterion_group!( + benches, + bench_primitives_encode, + bench_primitives_decode, + bench_struct_encode, + bench_struct_decode, + bench_collections_encode, + bench_collections_decode, + bench_string_encode, + bench_string_decode, + bench_enum_encode, + bench_enum_decode, + bench_roundtrip, + bench_complex_struct_encode, + bench_complex_struct_decode, + bench_slice_encode, + bench_slice_decode, + bench_hashmap_encode, + bench_hashmap_decode, + bench_hashset_encode, + bench_hashset_decode, +); +criterion_main!(benches); diff --git a/examples/codec.rs b/examples/codec.rs new file mode 100644 index 0000000..dd20d57 --- /dev/null +++ b/examples/codec.rs @@ -0,0 +1,401 @@ +use zr_protocol::codec::decode::Decode as DecodeTrait; +use zr_protocol::codec::encode::Encode as EncodeTrait; +use zr_protocol::context::Context; + +// ── Test Codec derive (combines Encode + Decode) ─────────────────────── + +#[derive(zr_protocol::Codec)] +#[context(Ctx)] +pub struct CodecStruct { + id: u8, + #[codec(count(u8))] + items: Vec<u16>, +} + +// ── State type for context-aware enums ─────────────────────────────── + +#[derive(Debug, Clone, Copy, PartialEq)] +pub enum State { + Handshake, + Status, + Play, +} + +pub struct Ctx { + pub state: State, +} + +// ── Basic struct (no helpers) ──────────────────────────────────────── + +#[derive(zr_protocol::Encode, zr_protocol::Decode)] +pub struct BasicStruct { + a: u8, + b: u16, + c: u32, +} + +// ── Skip ───────────────────────────────────────────────────────────── + +#[derive(zr_protocol::Encode)] +pub struct WithSkip { + id: u8, + #[codec(skip)] + _cache: Vec<u8>, + name: u16, +} + +// ── Count prefix ───────────────────────────────────────────────────── + +#[derive(zr_protocol::Encode, zr_protocol::Decode)] +pub struct WithCount { + #[codec(count(u16))] + items: Vec<u32>, +} + +// ── Condition on struct field ──────────────────────────────────────── + +#[derive(zr_protocol::Encode)] +pub struct WithCondition { + version: u8, + #[codec(if(self.version >= 2))] + extra: Option<u16>, +} + +// ── Combined struct helpers ────────────────────────────────────────── + +#[derive(zr_protocol::Encode)] +pub struct Combined { + version: u8, + #[codec(count(u16))] + items: Vec<u32>, + #[codec(if(self.items.is_empty()))] + fallback: Option<u8>, +} + +// ── Enum with auto IDs (basic) ────────────────────────────────────── + +#[derive(zr_protocol::Encode, zr_protocol::Decode)] +pub enum BasicEnum { + A(u8), + B { x: u16, y: u16 }, + C, +} + +// ── Enum with custom IDs ──────────────────────────────────────────── + +#[derive(zr_protocol::Encode, zr_protocol::Decode)] +pub enum CustomIdEnum { + #[id(0x10)] + Alpha(u8), + #[id(0x20)] + Beta { val: u16 }, + #[id(0x30)] + Gamma, +} + +// ── Context-aware enum (same ID, different states) ────────────────── + +#[derive(zr_protocol::Encode, zr_protocol::Decode)] +#[context(Ctx)] +pub enum Packet { + #[id(0x00)] + #[codec(if(ctx.data.state == State::Handshake))] + Handshake(u16), + + #[id(0x00)] + #[codec(if(ctx.data.state == State::Status))] + Status(u32), + + #[id(0x01)] + #[codec(if(ctx.data.state == State::Play))] + PlayData(u8, u8), + + #[id(0x02)] + Ping, +} + +// ── Context enum with skip ────────────────────────────────────────── + +#[derive(zr_protocol::Encode, zr_protocol::Decode)] +#[context(Ctx)] +pub enum WithSkipVariant { + #[id(0x00)] + Visible(u8), + #[codec(skip)] + Internal(u16), + #[id(0x01)] + AlsoVisible, +} + +// ── Test runner ────────────────────────────────────────────────────── + +fn enc<T: EncodeTrait<Ctx>>(value: &T, state: State) -> Vec<u8> { + let mut buf = Vec::new(); + let ctx = Context::new(Ctx { state }); + value.encode(&mut buf, &ctx).unwrap(); + buf +} + +fn dec<T: DecodeTrait<Ctx>>(buf: &[u8], state: State) -> T { + let ctx = Context::new(Ctx { state }); + T::decode(&mut &buf[..], &ctx).unwrap() +} + +fn main() { + // ── BasicStruct encode ── + let basic = BasicStruct { + a: 1, + b: 0x203, + c: 0x4050607, + }; + let bytes = enc(&basic, State::Handshake); + assert_eq!(bytes, vec![0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07]); + println!("BasicStruct encode: OK"); + + // ── BasicStruct decode ── + let decoded: BasicStruct = dec(&bytes, State::Handshake); + assert_eq!(decoded.a, 1); + assert_eq!(decoded.b, 0x203); + assert_eq!(decoded.c, 0x4050607); + println!("BasicStruct decode: OK"); + + // ── WithSkip ── + let skip = WithSkip { + id: 0xAA, + _cache: vec![1, 2, 3], + name: 0xBBCC, + }; + let bytes = enc(&skip, State::Handshake); + assert_eq!(bytes, vec![0xAA, 0xBB, 0xCC]); + println!("WithSkip: OK"); + + // ── WithCount encode ── + let count_empty = WithCount { items: vec![] }; + let bytes = enc(&count_empty, State::Handshake); + assert_eq!(bytes, vec![0x00, 0x00]); + println!("WithCount (empty): OK"); + + // ── WithCount decode ── + let decoded: WithCount = dec(&bytes, State::Handshake); + assert_eq!(decoded.items, Vec::<u32>::new()); + println!("WithCount (empty) decode: OK"); + + // ── WithCount with items encode ── + let count_some = WithCount { + items: vec![0x0A, 0x0B], + }; + let bytes = enc(&count_some, State::Handshake); + assert_eq!( + bytes, + vec![0x00, 0x02, 0x00, 0x00, 0x00, 0x0A, 0x00, 0x00, 0x00, 0x0B] + ); + println!("WithCount (2 items): OK"); + + // ── WithCount with items decode ── + let decoded: WithCount = dec(&bytes, State::Handshake); + assert_eq!(decoded.items, vec![0x0A, 0x0B]); + println!("WithCount (2 items) decode: OK"); + + // ── WithCondition encode ── + let cond_false = WithCondition { + version: 1, + extra: None, + }; + let bytes = enc(&cond_false, State::Handshake); + assert_eq!(bytes, vec![0x01]); + println!("WithCondition (v1, None): OK"); + + let cond_true = WithCondition { + version: 2, + extra: Some(0xDEAD), + }; + let bytes = enc(&cond_true, State::Handshake); + assert_eq!(bytes, vec![0x02, 0xDE, 0xAD]); + println!("WithCondition (v2, Some): OK"); + + // ── Combined encode ── + let comb1 = Combined { + version: 1, + items: vec![42], + fallback: Some(99), + }; + let bytes = enc(&comb1, State::Handshake); + assert_eq!(bytes, vec![0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x2A]); + println!("Combined (items=[42]): OK"); + + let comb2 = Combined { + version: 1, + items: vec![], + fallback: Some(99), + }; + let bytes = enc(&comb2, State::Handshake); + assert_eq!(bytes, vec![0x01, 0x00, 0x00, 0x63]); + println!("Combined (items=[], fallback): OK"); + + // ── BasicEnum encode ── + let e_a = BasicEnum::A(0xFF); + let bytes = enc(&e_a, State::Handshake); + assert_eq!(bytes, vec![0x00, 0xFF]); + println!("Enum::A encode: OK"); + + let e_b = BasicEnum::B { x: 1, y: 2 }; + let bytes = enc(&e_b, State::Handshake); + assert_eq!(bytes, vec![0x01, 0x00, 0x01, 0x00, 0x02]); + println!("Enum::B encode: OK"); + + let e_c = BasicEnum::C; + let bytes = enc(&e_c, State::Handshake); + assert_eq!(bytes, vec![0x02]); + println!("Enum::C encode: OK"); + + // ── BasicEnum decode ── + let decoded: BasicEnum = dec(&[0x00, 0xFF], State::Handshake); + assert!(matches!(decoded, BasicEnum::A(0xFF))); + println!("Enum::A decode: OK"); + + let decoded: BasicEnum = dec(&[0x01, 0x00, 0x01, 0x00, 0x02], State::Handshake); + assert!(matches!(decoded, BasicEnum::B { x: 1, y: 2 })); + println!("Enum::B decode: OK"); + + let decoded: BasicEnum = dec(&[0x02], State::Handshake); + assert!(matches!(decoded, BasicEnum::C)); + println!("Enum::C decode: OK"); + + // ── CustomIdEnum encode ── + let bytes = enc(&CustomIdEnum::Alpha(42), State::Handshake); + assert_eq!(bytes, vec![0x10, 42]); + println!("CustomIdEnum::Alpha: OK"); + + let bytes = enc(&CustomIdEnum::Beta { val: 1234 }, State::Handshake); + assert_eq!(bytes, vec![0x20, 0x04, 0xD2]); + println!("CustomIdEnum::Beta: OK"); + + let bytes = enc(&CustomIdEnum::Gamma, State::Handshake); + assert_eq!(bytes, vec![0x30]); + println!("CustomIdEnum::Gamma: OK"); + + // ── CustomIdEnum decode ── + let decoded: CustomIdEnum = dec(&[0x10, 42], State::Handshake); + assert!(matches!(decoded, CustomIdEnum::Alpha(42))); + println!("CustomIdEnum decode Alpha: OK"); + + let decoded: CustomIdEnum = dec(&[0x20, 0x04, 0xD2], State::Handshake); + assert!(matches!(decoded, CustomIdEnum::Beta { val: 1234 })); + println!("CustomIdEnum decode Beta: OK"); + + let decoded: CustomIdEnum = dec(&[0x30], State::Handshake); + assert!(matches!(decoded, CustomIdEnum::Gamma)); + println!("CustomIdEnum decode Gamma: OK"); + + // ── Context-aware enum encode ── + + // Handshake state → ID 0x00 encodes Handshake + let bytes = enc(&Packet::Handshake(42), State::Handshake); + assert_eq!(bytes, vec![0x00, 0x00, 0x2A]); + println!("Packet Handshake encode: OK ({:02x?}", bytes); + + // Status state → ID 0x00 encodes Status + let bytes = enc(&Packet::Status(1000), State::Status); + assert_eq!(bytes, vec![0x00, 0x00, 0x00, 0x03, 0xE8]); + println!("Packet Status encode: OK ({:02x?}", bytes); + + // Play state → ID 0x01 encodes PlayData + let bytes = enc(&Packet::PlayData(1, 2), State::Play); + assert_eq!(bytes, vec![0x01, 0x01, 0x02]); + println!("Packet PlayData encode: OK ({:02x?}", bytes); + + // Ping → always valid (no condition) + let bytes = enc(&Packet::Ping, State::Handshake); + assert_eq!(bytes, vec![0x02]); + println!("Packet Ping encode: OK"); + + // ── Context-aware: wrong state → error ── + let result = std::panic::catch_unwind(|| enc(&Packet::Handshake(42), State::Status)); + assert!(result.is_err(), "should panic: Handshake in Status state"); + println!("Packet Handshake in Status state: correctly errors"); + + let result = std::panic::catch_unwind(|| enc(&Packet::Status(1000), State::Play)); + assert!(result.is_err(), "should panic: Status in Play state"); + println!("Packet Status in Play state: correctly errors"); + + // ── Context-aware enum decode ── + + // ID 0x00 in Handshake state → Handshake + let decoded: Packet = dec(&[0x00, 0x00, 0x2A], State::Handshake); + assert!(matches!(decoded, Packet::Handshake(42))); + println!("Packet decode Handshake: OK"); + + // ID 0x00 in Status state → Status + let decoded: Packet = dec(&[0x00, 0x00, 0x00, 0x03, 0xE8], State::Status); + assert!(matches!(decoded, Packet::Status(1000))); + println!("Packet decode Status: OK"); + + // ID 0x01 in Play state → PlayData + let decoded: Packet = dec(&[0x01, 0x01, 0x02], State::Play); + assert!(matches!(decoded, Packet::PlayData(1, 2))); + println!("Packet decode PlayData: OK"); + + // ID 0x02 → Ping (always valid) + let decoded: Packet = dec(&[0x02], State::Handshake); + assert!(matches!(decoded, Packet::Ping)); + println!("Packet decode Ping: OK"); + + // ── Context-aware decode: wrong state → error ── + let result: Result<Packet, _> = { + let ctx = Context::new(Ctx { + state: State::Status, + }); + let mut reader = &[0x00, 0x00, 0x2A][..]; + Packet::decode(&mut reader, &ctx) + }; + assert!( + result.is_err(), + "should error: ID 0x00 in Status decodes as Status, not Handshake" + ); + println!("Packet decode wrong context: correctly errors"); + + // ── WithSkipVariant encode ── + let bytes = enc(&WithSkipVariant::Visible(42), State::Handshake); + assert_eq!(bytes, vec![0x00, 42]); + println!("WithSkipVariant::Visible encode: OK"); + + let bytes = enc(&WithSkipVariant::AlsoVisible, State::Handshake); + assert_eq!(bytes, vec![0x01]); + println!("WithSkipVariant::AlsoVisible encode: OK"); + + // ── WithSkipVariant decode ── + let decoded: WithSkipVariant = dec(&[0x00, 42], State::Handshake); + assert!(matches!(decoded, WithSkipVariant::Visible(42))); + println!("WithSkipVariant decode Visible: OK"); + + let decoded: WithSkipVariant = dec(&[0x01], State::Handshake); + assert!(matches!(decoded, WithSkipVariant::AlsoVisible)); + println!("WithSkipVariant decode AlsoVisible: OK"); + + // Unknown discriminant → error + let result: Result<WithSkipVariant, _> = { + let ctx = Context::new(Ctx { + state: State::Handshake, + }); + let mut reader = &[0xFF][..]; + WithSkipVariant::decode(&mut reader, &ctx) + }; + assert!(result.is_err(), "should error on unknown discriminant"); + println!("WithSkipVariant unknown discriminant: correctly errors"); + + // ── Codec derive test ── + let codec_val = CodecStruct { + id: 0x42, + items: vec![1, 2, 3], + }; + let bytes = enc(&codec_val, State::Handshake); + assert_eq!(bytes, vec![0x42, 0x03, 0x00, 0x01, 0x00, 0x02, 0x00, 0x03]); + println!("CodecStruct encode: OK"); + + let decoded: CodecStruct = dec(&bytes, State::Handshake); + assert_eq!(decoded.id, 0x42); + assert_eq!(decoded.items, vec![1, 2, 3]); + println!("CodecStruct decode: OK"); + + println!("\nAll tests passed!"); +} diff --git a/macros/Cargo.toml b/macros/Cargo.toml index 3820f50..9a3fcaf 100644 --- a/macros/Cargo.toml +++ b/macros/Cargo.toml @@ -7,3 +7,6 @@ edition = "2024" proc-macro = true [dependencies] +proc-macro2 = "1.0.107" +quote = "1.0.47" +syn = "3.0.3" diff --git a/macros/src/decode.rs b/macros/src/decode.rs new file mode 100644 index 0000000..c289cce --- /dev/null +++ b/macros/src/decode.rs @@ -0,0 +1,508 @@ +use proc_macro::TokenStream; +use proc_macro2::TokenStream as TokenStream2; +use quote::quote; +use syn::{Data, DataEnum, DataStruct, DeriveInput, Generics, Ident, PathArguments, Type}; + +use crate::encode::{next_group, next_ident, parse_variant_attrs, skip_commas}; + +// ── Helper: extract element type from collection ──────────────────── + +fn extract_element_type(ty: &Type) -> Option<Type> { + if let Type::Path(type_path) = ty + && let Some(segment) = type_path.path.segments.last() + { + match segment.ident.to_string().as_str() { + "Vec" | "Option" | "HashSet" | "BTreeSet" | "BinaryHeap" | "LinkedList" + | "VecDeque" => { + if let PathArguments::AngleBracketed(args) = &segment.arguments + && let Some(syn::GenericArgument::Type(inner_ty)) = args.args.first() + { + return Some(inner_ty.clone()); + } + } + "HashMap" | "BTreeMap" => { + if let PathArguments::AngleBracketed(args) = &segment.arguments { + let mut types = args.args.iter().filter_map(|arg| { + if let syn::GenericArgument::Type(t) = arg { + Some(t.clone()) + } else { + None + } + }); + if let (Some(k), Some(v)) = (types.next(), types.next()) { + return Some(syn::parse_quote! { (#k, #v) }); + } + } + } + _ => {} + } + } + None +} + +// ── Field attribute parsing ───────────────────────────────────────── + +enum DecodeAttr { + Normal, + Count(syn::Type), + Len(syn::Type), + Custom(syn::Path), +} + +struct FieldAttrs { + decode: DecodeAttr, + element_type: Option<Type>, + condition: Option<syn::Expr>, +} + +fn parse_field_attrs(field: &syn::Field) -> FieldAttrs { + let mut decode: Option<DecodeAttr> = None; + let mut condition: Option<syn::Expr> = None; + + for attr in &field.attrs { + if !attr.path().is_ident("codec") { + continue; + } + + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + + let mut iter = list.tokens.clone().into_iter(); + + loop { + skip_commas(&mut iter); + + let ident = match next_ident(&mut iter) { + Some(i) => i, + None => break, + }; + + match ident.to_string().as_str() { + "skip" => { + assert!(decode.is_none(), "multiple codec attributes on one field"); + decode = Some(DecodeAttr::Normal); + } + "count" => { + assert!(decode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected count(Type)"); + let len_ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in count(...)"); + let element_type = extract_element_type(&field.ty); + decode = Some(DecodeAttr::Count(len_ty)); + return FieldAttrs { + decode: decode.unwrap(), + element_type, + condition, + }; + } + "len" => { + assert!(decode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected len(Type)"); + let ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in len(...)"); + decode = Some(DecodeAttr::Len(ty)); + } + "with" => { + assert!(decode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected with(path)"); + let path: syn::Path = + syn::parse2(group.stream()).expect("expected path in with(...)"); + decode = Some(DecodeAttr::Custom(path)); + } + "if" => { + assert!(condition.is_none(), "multiple codec(if(...)) on one field"); + let group = next_group(&mut iter).expect("expected if(expr)"); + let expr: syn::Expr = + syn::parse2(group.stream()).expect("expected expression in if(...)"); + condition = Some(expr); + } + other => panic!("unknown codec helper on decode: `{other}`"), + } + } + } + + let element_type = if matches!(decode, Some(DecodeAttr::Count(_))) { + extract_element_type(&field.ty) + } else { + None + }; + + FieldAttrs { + decode: decode.unwrap_or(DecodeAttr::Normal), + element_type, + condition, + } +} + +// ── Context type extraction ───────────────────────────────────────── + +fn extract_context_ty(attrs: &[syn::Attribute]) -> Option<Type> { + attrs.iter().find_map(|attr| { + if attr.path().is_ident("context") { + let syn::Meta::List(list) = &attr.meta else { + return None; + }; + syn::parse2::<Type>(list.tokens.clone()).ok() + } else { + None + } + }) +} + +// ── Build impl generics ───────────────────────────────────────────── + +fn build_decode_impl_generics( + params: &syn::punctuated::Punctuated<syn::GenericParam, syn::Token![,]>, + where_clause: &Option<syn::WhereClause>, + context_ty: &Option<Type>, +) -> (TokenStream2, TokenStream2, TokenStream2, TokenStream2) { + let params = params.iter().cloned().collect::<Vec<_>>(); + let has_params = !params.is_empty(); + + if let Some(ctx_ty) = context_ty { + let impl_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { #ctx_ty }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } else { + let mut all_params = vec![syn::parse_quote! { Data }]; + all_params.extend(params.clone()); + let impl_generics = quote! { <#(#all_params),*> }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { Data }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } +} + +// ── Decode code generation ────────────────────────────────────────── + +pub fn derive_decode( + DeriveInput { + ident, + generics, + data, + attrs, + .. + }: DeriveInput, +) -> TokenStream { + let context_ty = extract_context_ty(&attrs); + match data { + Data::Struct(data_struct) => impl_decode_struct(ident, generics, data_struct, context_ty), + Data::Enum(data_enum) => impl_decode_enum(ident, generics, data_enum, context_ty), + Data::Union(_) => panic!("Not implemented for Union"), + } +} + +fn gen_decode_field( + field_type: &Type, + attrs: &FieldAttrs, + data_param: &TokenStream2, +) -> TokenStream2 { + let inner = match &attrs.decode { + DecodeAttr::Normal => quote! { + <#field_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)? + }, + DecodeAttr::Count(len_type) => { + let element_type = attrs + .element_type + .clone() + .unwrap_or_else(|| field_type.clone()); + quote! { + { + let count: #len_type = <#len_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)?; + let mut items = Vec::with_capacity(count as usize); + for _ in 0..count { + items.push( + <#element_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)? + ); + } + items + } + } + } + DecodeAttr::Len(len_type) => quote! { + { + let byte_len: #len_type = <#len_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)?; + let byte_len = byte_len as usize; + if buf.len() < byte_len { + return Err(zr_protocol::codec::error::CodecError::IoError(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, + "insufficient bytes for len-prefixed field", + ))); + } + let (mut data, rest) = buf.split_at(byte_len); + *buf = rest; + <#field_type as zr_protocol::codec::decode::Decode<#data_param>>::decode(&mut data, ctx)? + } + }, + DecodeAttr::Custom(fn_path) => quote! { + #fn_path(buf, ctx)? + }, + }; + + match &attrs.condition { + Some(expr) => quote! { + if #expr { + #inner + } else { + Default::default() + } + }, + None => inner, + } +} + +fn impl_decode_struct( + ident: Ident, + generics: Generics, + data_struct: DataStruct, + context_ty: Option<Type>, +) -> TokenStream { + let (impl_generics, ty_generics, where_clause, data_param) = + build_decode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + let decode_fields = match &data_struct.fields { + syn::Fields::Named(fields) => { + let field_decodes: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let ty = &field.ty; + let attrs = parse_field_attrs(field); + let decoded = gen_decode_field(ty, &attrs, &data_param); + quote! { #name: #decoded } + }) + .collect(); + quote! { Ok(Self { #(#field_decodes),* }) } + } + syn::Fields::Unnamed(fields) => { + let field_decodes: Vec<_> = fields + .unnamed + .iter() + .map(|field| { + let ty = &field.ty; + let attrs = parse_field_attrs(field); + gen_decode_field(ty, &attrs, &data_param) + }) + .collect(); + quote! { Ok(Self(#(#field_decodes),*)) } + } + syn::Fields::Unit => quote! { Ok(Self) }, + }; + + quote! { + impl #impl_generics zr_protocol::codec::decode::Decode<#data_param> for #ident #ty_generics #where_clause { + fn decode(buf: &mut &[u8], ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<Self> { + #decode_fields + } + } + } + .into() +} + +fn impl_decode_enum( + ident: Ident, + generics: Generics, + data_enum: DataEnum, + context_ty: Option<Type>, +) -> TokenStream { + let (_, _, _, data_param) = + build_decode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + // Parse all variant attributes + let parsed: Vec<_> = data_enum.variants.iter().map(parse_variant_attrs).collect(); + + // Build explicit ID set and auto-ID counter + let mut explicit_ids = std::collections::HashSet::new(); + for attrs in &parsed { + if let Some(id) = attrs.id { + explicit_ids.insert(id); + } + } + + // Assign discriminants (same logic as encode) + let mut auto_counter: u8 = 0; + let mut discriminants: Vec<u8> = Vec::with_capacity(data_enum.variants.len()); + for attrs in &parsed { + if let Some(id) = attrs.id { + discriminants.push(id); + } else { + while explicit_ids.contains(&auto_counter) { + auto_counter = auto_counter.wrapping_add(1); + } + let d = auto_counter; + auto_counter = auto_counter.wrapping_add(1); + discriminants.push(d); + } + } + + // Group variants by discriminant ID + use std::collections::BTreeMap; + type VariantEntry<'a> = (usize, &'a Ident, &'a syn::Fields, &'a Option<syn::Expr>); + let mut groups: BTreeMap<u8, Vec<VariantEntry>> = BTreeMap::new(); + + for (i, (variant, attrs)) in data_enum.variants.iter().zip(parsed.iter()).enumerate() { + let disc = discriminants[i]; + if attrs.skip { + continue; + } + groups.entry(disc).or_default().push(( + i, + &variant.ident, + &variant.fields, + &attrs.condition, + )); + } + + // Generate decode match arms + let match_arms: Vec<_> = groups + .iter() + .map(|(disc, entries)| { + if entries.len() == 1 { + let (_, variant_ident, fields, condition) = &entries[0]; + let field_decode = gen_variant_decode(&ident, variant_ident, fields, &data_param); + + match condition { + Some(expr) => quote! { + #disc => { + if #expr { + #field_decode + } else { + Err(zr_protocol::codec::error::CodecError::Custom( + concat!("variant ", stringify!(#variant_ident), " not valid in current context").into() + )) + } + } + }, + None => quote! { + #disc => { #field_decode } + }, + } + } else { + let mut arms: Vec<TokenStream2> = Vec::new(); + let mut has_unconditional = false; + + for (_, variant_ident, fields, condition) in entries { + let field_decode = gen_variant_decode(&ident, variant_ident, fields, &data_param); + + match condition { + Some(expr) => { + arms.push(quote! { + if #expr { + #field_decode + } + }); + } + None => { + has_unconditional = true; + arms.push(field_decode); + } + } + } + + if has_unconditional { + let last = arms.pop().unwrap(); + let chain = arms.into_iter().rev().fold(last, |acc, arm| { + quote! { #arm else { #acc } } + }); + quote! { #disc => { #chain } } + } else { + let chain = arms.into_iter().rev().fold( + quote! { + Err(zr_protocol::codec::error::CodecError::Custom( + format!("no variant for ID {:#04x} in current context", #disc).into() + )) + }, + |acc, arm| { + quote! { #arm else { #acc } } + }, + ); + quote! { #disc => { #chain } } + } + } + }) + .collect(); + + let (impl_generics, ty_generics, where_clause, _) = + build_decode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + quote! { + impl #impl_generics zr_protocol::codec::decode::Decode<#data_param> for #ident #ty_generics #where_clause { + fn decode(buf: &mut &[u8], ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<Self> { + let id = <u8 as zr_protocol::codec::decode::Decode<#data_param>>::decode(buf, ctx)?; + match id { + #(#match_arms),* + other => Err(zr_protocol::codec::error::CodecError::Custom( + format!("unknown discriminant: {other:#04x}").into() + )), + } + } + } + } + .into() +} + +fn gen_variant_decode( + enum_ident: &Ident, + variant_ident: &Ident, + fields: &syn::Fields, + data_param: &TokenStream2, +) -> TokenStream2 { + match fields { + syn::Fields::Named(fields) => { + let field_decodes: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let ty = &field.ty; + let attrs = parse_field_attrs(field); + let decoded = gen_decode_field(ty, &attrs, data_param); + quote! { #name: #decoded } + }) + .collect(); + quote! { + Ok(#enum_ident::#variant_ident { #(#field_decodes),* }) + } + } + syn::Fields::Unnamed(fields) => { + let field_decodes: Vec<_> = fields + .unnamed + .iter() + .map(|field| { + let ty = &field.ty; + let attrs = parse_field_attrs(field); + gen_decode_field(ty, &attrs, data_param) + }) + .collect(); + quote! { + Ok(#enum_ident::#variant_ident(#(#field_decodes),*)) + } + } + syn::Fields::Unit => quote! { + Ok(#enum_ident::#variant_ident) + }, + } +} diff --git a/macros/src/encode.rs b/macros/src/encode.rs new file mode 100644 index 0000000..d8df604 --- /dev/null +++ b/macros/src/encode.rs @@ -0,0 +1,516 @@ +use proc_macro::TokenStream; +use proc_macro2::TokenStream as TokenStream2; +use quote::quote; +use syn::{Data, DataEnum, DataStruct, DeriveInput, Generics, Ident, Type}; + +// ── Shared token helpers ──────────────────────────────────────────── + +pub(crate) fn next_ident(iter: &mut proc_macro2::token_stream::IntoIter) -> Option<Ident> { + if let Some(proc_macro2::TokenTree::Ident(ident)) = iter.clone().next() { + iter.next(); + return Some(ident); + } + None +} + +pub(crate) fn next_group( + iter: &mut proc_macro2::token_stream::IntoIter, +) -> Option<proc_macro2::Group> { + if let Some(proc_macro2::TokenTree::Group(group)) = iter.clone().next() { + iter.next(); + return Some(group); + } + None +} + +pub(crate) fn skip_commas(iter: &mut proc_macro2::token_stream::IntoIter) { + loop { + let peek = iter.clone().next(); + match peek { + Some(proc_macro2::TokenTree::Punct(p)) if p.as_char() == ',' => { + iter.next(); + } + _ => break, + } + } +} + +// ── Field attribute parsing ───────────────────────────────────────── + +enum EncodeAttr { + Normal, + Skip, + Count(syn::Type), + Len(syn::Type), + Custom(syn::Path), +} + +struct FieldAttrs { + encode: EncodeAttr, + condition: Option<syn::Expr>, +} + +fn parse_field_attrs(field: &syn::Field) -> FieldAttrs { + let mut encode: Option<EncodeAttr> = None; + let mut condition: Option<syn::Expr> = None; + + for attr in &field.attrs { + if !attr.path().is_ident("codec") { + continue; + } + + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + + let mut iter = list.tokens.clone().into_iter(); + + loop { + skip_commas(&mut iter); + + let ident = match next_ident(&mut iter) { + Some(i) => i, + None => break, + }; + + match ident.to_string().as_str() { + "skip" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + encode = Some(EncodeAttr::Skip); + } + "count" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected count(Type)"); + let ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in count(...)"); + encode = Some(EncodeAttr::Count(ty)); + } + "len" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected len(Type)"); + let ty: syn::Type = + syn::parse2(group.stream()).expect("expected type in len(...)"); + encode = Some(EncodeAttr::Len(ty)); + } + "with" => { + assert!(encode.is_none(), "multiple codec attributes on one field"); + let group = next_group(&mut iter).expect("expected with(path)"); + let path: syn::Path = + syn::parse2(group.stream()).expect("expected path in with(...)"); + encode = Some(EncodeAttr::Custom(path)); + } + "if" => { + assert!(condition.is_none(), "multiple codec(if(...)) on one field"); + let group = next_group(&mut iter).expect("expected if(expr)"); + let expr: syn::Expr = + syn::parse2(group.stream()).expect("expected expression in if(...)"); + condition = Some(expr); + } + other => panic!("unknown codec helper: `{other}`"), + } + } + } + + FieldAttrs { + encode: encode.unwrap_or(EncodeAttr::Normal), + condition, + } +} + +// ── Variant attribute parsing ─────────────────────────────────────── + +pub(crate) struct VariantAttrs { + pub id: Option<u8>, + pub condition: Option<syn::Expr>, + pub skip: bool, +} + +pub(crate) fn parse_variant_attrs(variant: &syn::Variant) -> VariantAttrs { + let mut id: Option<u8> = None; + let mut condition: Option<syn::Expr> = None; + let mut skip = false; + + for attr in &variant.attrs { + // #[id(0x00)] + if attr.path().is_ident("id") { + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + let lit: syn::LitInt = + syn::parse2(list.tokens.clone()).expect("expected u8 literal in #[id(...)]"); + let val: u8 = lit.base10_parse().expect("id must be a u8 value"); + assert!(id.is_none(), "multiple #[id(...)] on one variant"); + id = Some(val); + continue; + } + + // #[codec(if(...)), codec(skip)] + if !attr.path().is_ident("codec") { + continue; + } + + let syn::Meta::List(list) = &attr.meta else { + continue; + }; + + let mut iter = list.tokens.clone().into_iter(); + + loop { + skip_commas(&mut iter); + + let ident = match next_ident(&mut iter) { + Some(i) => i, + None => break, + }; + + match ident.to_string().as_str() { + "skip" => { + skip = true; + } + "if" => { + assert!( + condition.is_none(), + "multiple codec(if(...)) on one variant" + ); + let group = next_group(&mut iter).expect("expected if(expr)"); + let expr: syn::Expr = + syn::parse2(group.stream()).expect("expected expression in if(...)"); + condition = Some(expr); + } + other => panic!("unknown codec helper on variant: `{other}`"), + } + } + } + + VariantAttrs { + id, + condition, + skip, + } +} + +// ── Context type extraction ───────────────────────────────────────── + +fn extract_context_ty(attrs: &[syn::Attribute]) -> Option<Type> { + attrs.iter().find_map(|attr| { + if attr.path().is_ident("context") { + let syn::Meta::List(list) = &attr.meta else { + return None; + }; + syn::parse2::<Type>(list.tokens.clone()).ok() + } else { + None + } + }) +} + +// ── Build impl generics ───────────────────────────────────────────── + +fn build_encode_impl_generics( + params: &syn::punctuated::Punctuated<syn::GenericParam, syn::Token![,]>, + where_clause: &Option<syn::WhereClause>, + context_ty: &Option<Type>, +) -> (TokenStream2, TokenStream2, TokenStream2, TokenStream2) { + let params = params.iter().cloned().collect::<Vec<_>>(); + let has_params = !params.is_empty(); + + if let Some(ctx_ty) = context_ty { + // With context type: impl<#params> Encode<Ctx> for Ident + let impl_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { #ctx_ty }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } else { + // Without context: impl<Data, #params> Encode<Data> for Ident + let mut all_params = vec![syn::parse_quote! { Data }]; + all_params.extend(params.iter().cloned()); + let impl_generics = quote! { <#(#all_params),*> }; + let ty_generics = if has_params { + quote! { <#(#params),*> } + } else { + quote! {} + }; + let data_param = quote! { Data }; + let where_clause_tokens = where_clause + .as_ref() + .map(|wc| quote! { #wc }) + .unwrap_or_default(); + (impl_generics, ty_generics, where_clause_tokens, data_param) + } +} + +// ── Encode code generation ────────────────────────────────────────── + +pub fn derive_encode( + DeriveInput { + ident, + generics, + data, + attrs, + .. + }: DeriveInput, +) -> TokenStream { + let context_ty = extract_context_ty(&attrs); + match data { + Data::Struct(data_struct) => impl_encode_struct(ident, generics, data_struct, context_ty), + Data::Enum(data_enum) => impl_encode_enum(ident, generics, data_enum, context_ty), + Data::Union(_) => panic!("Not implemented for Union"), + } +} + +fn gen_encode_field( + field_access: TokenStream2, + attrs: &FieldAttrs, + _data_param: &TokenStream2, +) -> TokenStream2 { + let inner = match &attrs.encode { + EncodeAttr::Normal => quote! { + len += #field_access.encode(buffer, ctx)?; + }, + EncodeAttr::Skip => quote! {}, + EncodeAttr::Count(len_type) => quote! { + let count: #len_type = #field_access.len() as #len_type; + len += count.encode(buffer, ctx)?; + for item in &#field_access { + len += item.encode(buffer, ctx)?; + } + }, + EncodeAttr::Len(len_type) => quote! { + let mut __buf = Vec::with_capacity(2048); + #field_access.encode(&mut __buf, ctx)?; + let __len: #len_type = zr_protocol::types::size::Size::from_size(__buf.len()); + len += __len.encode(buffer, ctx)?; + buffer.extend_from_slice(&__buf); + len += __buf.len(); + }, + EncodeAttr::Custom(fn_path) => quote! { + len += #fn_path(&#field_access, buffer, ctx)?; + }, + }; + + match &attrs.condition { + Some(expr) => quote! { + if #expr { + #inner + } + }, + None => inner, + } +} + +fn gen_encode_body( + enum_ident: &Ident, + variant_ident: &Ident, + pattern: &TokenStream2, + encode_fields: &TokenStream2, + discriminant: u8, + condition: &Option<syn::Expr>, +) -> TokenStream2 { + let encode_disc_and_fields = quote! { + len += #discriminant.encode(buffer, ctx)?; + #encode_fields + }; + + match condition { + Some(expr) => quote! { + #enum_ident::#variant_ident #pattern => { + if #expr { + #encode_disc_and_fields + } else { + return Err(zr_protocol::codec::error::CodecError::Custom( + concat!("variant ", stringify!(#variant_ident), " not valid in current context").into() + )); + } + } + }, + None => quote! { + #enum_ident::#variant_ident #pattern => { + #encode_disc_and_fields + } + }, + } +} + +fn impl_encode_struct( + ident: Ident, + generics: Generics, + data_struct: DataStruct, + context_ty: Option<Type>, +) -> TokenStream { + let (impl_generics, ty_generics, where_clause, data_param) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + let encode_fields = match &data_struct.fields { + syn::Fields::Named(fields) => { + let bodies: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let attrs = parse_field_attrs(field); + gen_encode_field(quote! { self.#name }, &attrs, &data_param) + }) + .collect(); + quote! { #(#bodies)* } + } + syn::Fields::Unnamed(fields) => { + let bodies: Vec<_> = fields + .unnamed + .iter() + .enumerate() + .map(|(i, field)| { + let idx = syn::Index::from(i); + let attrs = parse_field_attrs(field); + gen_encode_field(quote! { self.#idx }, &attrs, &data_param) + }) + .collect(); + quote! { #(#bodies)* } + } + syn::Fields::Unit => quote! {}, + }; + + quote! { + impl #impl_generics zr_protocol::codec::encode::Encode<#data_param> for #ident #ty_generics #where_clause { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<usize> { + let mut len = 0; + #encode_fields + Ok(len) + } + } + } + .into() +} + +fn impl_encode_enum( + ident: Ident, + generics: Generics, + data_enum: DataEnum, + context_ty: Option<Type>, +) -> TokenStream { + let (_, _, _, data_param) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + // First pass: collect explicit IDs and count auto-IDs needed + let mut explicit_ids = std::collections::HashSet::new(); + + // Pre-parse all variant attributes + let parsed: Vec<_> = data_enum + .variants + .iter() + .map(|v| { + let attrs = parse_variant_attrs(v); + if let Some(id) = attrs.id { + explicit_ids.insert(id); + } + attrs + }) + .collect(); + + let match_arms: Vec<_> = { + let mut auto_counter: u8 = 0; + data_enum + .variants + .iter() + .zip(parsed.iter()) + .map(|(variant, attrs)| { + let variant_ident = &variant.ident; + + if attrs.skip { + return quote! { + #ident::#variant_ident { .. } => { + unreachable!("skipped variant") + } + }; + } + + // Resolve discriminant: explicit #[id(...)] or auto-increment + let discriminant = if let Some(id) = attrs.id { + id + } else { + while explicit_ids.contains(&auto_counter) { + auto_counter = auto_counter.wrapping_add(1); + } + let d = auto_counter; + auto_counter = auto_counter.wrapping_add(1); + d + }; + + let (pattern, encode_fields) = match &variant.fields { + syn::Fields::Named(fields) => { + let names: Vec<_> = fields + .named + .iter() + .map(|f| f.ident.as_ref().unwrap()) + .collect(); + let pattern = quote! { { #(#names),* } }; + let bodies: Vec<_> = fields + .named + .iter() + .map(|field| { + let name = field.ident.as_ref().unwrap(); + let fa = parse_field_attrs(field); + gen_encode_field(quote! { #name }, &fa, &data_param) + }) + .collect(); + (pattern, quote! { #(#bodies)* }) + } + syn::Fields::Unnamed(fields) => { + let names: Vec<_> = (0..fields.unnamed.len()) + .map(|i| Ident::new(&format!("_{}", i), proc_macro2::Span::call_site())) + .collect(); + let pattern = quote! { ( #(#names),* ) }; + let bodies: Vec<_> = fields + .unnamed + .iter() + .enumerate() + .map(|(i, field)| { + let name = &names[i]; + let fa = parse_field_attrs(field); + gen_encode_field(quote! { #name }, &fa, &data_param) + }) + .collect(); + (pattern, quote! { #(#bodies)* }) + } + syn::Fields::Unit => (quote! {}, quote! {}), + }; + + gen_encode_body( + &ident, + variant_ident, + &pattern, + &encode_fields, + discriminant, + &attrs.condition, + ) + }) + .collect() + }; + + let (impl_generics, ty_generics, where_clause, _) = + build_encode_impl_generics(&generics.params, &generics.where_clause, &context_ty); + + quote! { + impl #impl_generics zr_protocol::codec::encode::Encode<#data_param> for #ident #ty_generics #where_clause { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &zr_protocol::context::Context<#data_param>) -> zr_protocol::codec::error::Result<usize> { + let mut len = 0; + match self { + #(#match_arms),* + } + Ok(len) + } + } + } + .into() +} diff --git a/macros/src/lib.rs b/macros/src/lib.rs index 8b13789..f20edae 100644 --- a/macros/src/lib.rs +++ b/macros/src/lib.rs @@ -1 +1,34 @@ +use proc_macro::TokenStream; +use proc_macro2::TokenStream as TokenStream2; +use syn::DeriveInput; +mod decode; +mod encode; + +#[proc_macro_derive(Encode, attributes(codec, id, context))] +pub fn encode(input: TokenStream) -> TokenStream { + let input = syn::parse_macro_input!(input as DeriveInput); + encode::derive_encode(input) +} + +#[proc_macro_derive(Decode, attributes(codec, id, context))] +pub fn decode(input: TokenStream) -> TokenStream { + let input = syn::parse_macro_input!(input as DeriveInput); + decode::derive_decode(input) +} + +#[proc_macro_derive(Codec, attributes(codec, id, context))] +pub fn codec(input: TokenStream) -> TokenStream { + let input = syn::parse_macro_input!(input as DeriveInput); + let encode_impl = encode::derive_encode(input.clone()); + let decode_impl = decode::derive_decode(input.clone()); + + let encode_ts: TokenStream2 = encode_impl.into(); + let decode_ts: TokenStream2 = decode_impl.into(); + + quote::quote! { + #encode_ts + #decode_ts + } + .into() +} diff --git a/src/codec/decode.rs b/src/codec/decode.rs index 529c448..d068760 100644 --- a/src/codec/decode.rs +++ b/src/codec/decode.rs @@ -1,150 +1,133 @@ -//! Decoding: reading protocol values from a byte stream. use std::{ collections::{BTreeMap, BTreeSet, BinaryHeap, HashMap, HashSet}, hash::Hash, - io::Read, sync::Arc, }; -use crate::{DEFAULT_BUFFER_LEN, codec::error::Result}; -use crate::{codec::error::CodecError, context::Context}; +use crate::{codec::error::Result, codec::error::CodecError, context::Context}; macro_rules! impl_decode { - (u8) => { - impl<Data> Decode<Data> for u8 { - fn decode(reader: &mut dyn Read, _: &Context<Data>) -> Result<Self> + ($t:ty) => { + impl<Data> Decode<Data> for $t { + fn decode(buf: &mut &[u8], _: &Context<Data>) -> Result<Self> where Self: Sized, { - let mut byte = [0u8; 1]; - reader.read_exact(&mut byte)?; - Ok(byte[0]) + const BYTES: usize = std::mem::size_of::<$t>(); + if buf.len() < BYTES { + return Err(CodecError::IoError(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, + "insufficient bytes", + ))); + } + let (bytes, rest) = buf.split_at(BYTES); + *buf = rest; + Ok(<$t>::from_be_bytes(bytes.try_into().unwrap())) } - fn decode_slice(reader: &mut dyn Read, _: &Context<Data>) -> Result<Vec<Self>> + fn decode_slice(buf: &mut &[u8], _: &Context<Data>) -> Result<Vec<Self>> where Self: Sized, { - let mut buf = Vec::new(); - reader.read_to_end(&mut buf)?; - Ok(buf) + const BYTES: usize = std::mem::size_of::<$t>(); + let count = buf.len() / BYTES; + let mut vec = Vec::with_capacity(count); + for _ in 0..count { + let (bytes, rest) = buf.split_at(BYTES); + *buf = rest; + vec.push(<$t>::from_be_bytes(bytes.try_into().unwrap())); + } + Ok(vec) } } }; - ($type: ty) => { - impl<Data> Decode<Data> for $type { - // decode number using big endian - fn decode(reader: &mut dyn Read, _: &Context<Data>) -> Result<Self> - where - Self: Sized, - { - const BYTES: usize = (<$type>::BITS / 8) as usize; - let mut bytes = [0; BYTES]; - reader.read_exact(&mut bytes)?; - Ok(<$type>::from_be_bytes(bytes)) - } +} - fn decode_slice(reader: &mut dyn Read, _: &Context<Data>) -> Result<Vec<Self>> +macro_rules! impl_decode_tuples { + ($($generic: ident),+) => { + impl<Data, $($generic),+> Decode<Data> for ($($generic,)+) + where + $($generic: Decode<Data>),+ + { + #[allow(non_snake_case)] + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let mut buf = Vec::with_capacity(DEFAULT_BUFFER_LEN); - reader.read_to_end(&mut buf); - const BYTES: usize = <$type>::BITS as usize / 8; - Ok(buf - .chunks(BYTES) - .map(|slice| { - let Some(bytes): Option<&[u8; BYTES]> = slice.as_array() else { - unreachable!() - }; - <$type>::from_be_bytes(*bytes) - }) - .collect()) + $( + let $generic = $generic::decode(buf, ctx)?; + )+ + Ok(($($generic,)+)) } } }; } -macro_rules! impl_decode_tuples { - ($($generic: ident),+) => { - impl<Data, $($generic),+> Decode<Data> for ($($generic,)+) where $($generic: Decode<Data>,)+ { - #[allow(non_snake_case)] - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> - where - Self: Sized { - $( - let $generic = $generic::decode(reader, ctx)?; - )+ - Ok(($($generic,)+)) - } - } - }; -} - -/// Read a value from a byte stream +/// Zero-copy decoding: reads directly from a byte slice pub trait Decode<Data> { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized; - // TODO : change for anything else than Vec - /// assume reader contains only the slice to decode - fn decode_slice(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Vec<Self>> + fn decode_slice(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Vec<Self>> where Self: Sized, { - let mut buf = Vec::new(); - loop { - match Self::decode(reader, ctx) { - Ok(value) => buf.push(value), + let mut vec = Vec::new(); + while !buf.is_empty() { + match Self::decode(buf, ctx) { + Ok(value) => vec.push(value), Err(CodecError::IoError(err)) - if let std::io::ErrorKind::UnexpectedEof = err.kind() => + if err.kind() == std::io::ErrorKind::UnexpectedEof => { - return Ok(buf); + return Ok(vec); } Err(err) => return Err(err), } } + Ok(vec) } } -// ~ Decode arrays +// ── Decode arrays ─────────────────────────────────────────────────────── impl<Data, T: Decode<Data>> Decode<Data> for Vec<T> { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let len = usize::decode(reader, ctx)?; + let len = usize::decode(buf, ctx)?; let mut vec = Vec::with_capacity(len); for _ in 0..len { - vec.push(T::decode(reader, ctx)?); + vec.push(T::decode(buf, ctx)?); } Ok(vec) } } impl<Data, T: Decode<Data>> Decode<Data> for Option<T> { - /// Assume is always some - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - T::decode(reader, ctx).map(Some) + if buf.is_empty() { + return Ok(None); + } + T::decode(buf, ctx).map(Some) } } -// ~ Decode slices +// ── Decode slices ─────────────────────────────────────────────────────── impl<Data, T, const S: usize> Decode<Data> for [T; S] where T: Decode<Data> + Default + Copy, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let slice = T::decode_slice(reader, ctx)?; + let slice = T::decode_slice(buf, ctx)?; let got = slice.len(); slice .as_array() @@ -157,7 +140,7 @@ impl<Data, T> Decode<Data> for &[T] where T: Decode<Data> + Default + Copy, { - fn decode(_: &mut dyn Read, _: &Context<Data>) -> Result<Self> + fn decode(_: &mut &[u8], _: &Context<Data>) -> Result<Self> where Self: Sized, { @@ -165,18 +148,22 @@ where } } -// ~ Decode set +// ── Decode set ────────────────────────────────────────────────────────── impl<Data, T> Decode<Data> for HashSet<T> where T: Decode<Data> + Eq + Hash, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let slice = T::decode_slice(reader, ctx)?; - Ok(slice.into_iter().collect()) + let len = usize::decode(buf, ctx)?; + let mut set = HashSet::with_capacity(len); + for _ in 0..len { + set.insert(T::decode(buf, ctx)?); + } + Ok(set) } } @@ -184,43 +171,57 @@ impl<Data, T> Decode<Data> for BTreeSet<T> where T: Decode<Data> + Ord, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let slice = T::decode_slice(reader, ctx)?; - Ok(slice.into_iter().collect()) + let len = usize::decode(buf, ctx)?; + let mut set = BTreeSet::new(); + for _ in 0..len { + set.insert(T::decode(buf, ctx)?); + } + Ok(set) } } -// ~ Decode misc +// ── Decode misc ───────────────────────────────────────────────────────── impl<Data, T> Decode<Data> for BinaryHeap<T> where T: Decode<Data> + Ord, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let slice = T::decode_slice(reader, ctx)?; - Ok(BinaryHeap::from_iter(slice)) + let len = usize::decode(buf, ctx)?; + let mut heap = BinaryHeap::with_capacity(len); + for _ in 0..len { + heap.push(T::decode(buf, ctx)?); + } + Ok(heap) } } -// ~ Decode maps +// ── Decode maps ───────────────────────────────────────────────────────── impl<Data, K, V> Decode<Data> for HashMap<K, V> where K: Decode<Data> + Eq + Hash, V: Decode<Data>, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let slice = <(K, V) as Decode<Data>>::decode_slice(reader, ctx)?; - Ok(slice.into_iter().collect()) + let len = usize::decode(buf, ctx)?; + let mut map = HashMap::with_capacity(len); + for _ in 0..len { + let key = K::decode(buf, ctx)?; + let val = V::decode(buf, ctx)?; + map.insert(key, val); + } + Ok(map) } } @@ -229,19 +230,25 @@ where K: Decode<Data> + Ord, V: Decode<Data>, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let slice = <(K, V) as Decode<Data>>::decode_slice(reader, ctx)?; - Ok(slice.into_iter().collect()) + let len = usize::decode(buf, ctx)?; + let mut map = BTreeMap::new(); + for _ in 0..len { + let key = K::decode(buf, ctx)?; + let val = V::decode(buf, ctx)?; + map.insert(key, val); + } + Ok(map) } } -// ~ Decode string +// ── Decode string - requires length prefix ────────────────────────────── impl<Data> Decode<Data> for &str { - fn decode(_: &mut dyn Read, _: &Context<Data>) -> Result<Self> + fn decode(_: &mut &[u8], _: &Context<Data>) -> Result<Self> where Self: Sized, { @@ -250,17 +257,17 @@ impl<Data> Decode<Data> for &str { } impl<Data> Decode<Data> for String { - fn decode(reader: &mut dyn Read, _: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let mut bytes: Vec<u8> = Vec::new(); - reader.read_to_end(&mut bytes)?; - Ok(String::from_utf8_lossy(&bytes).to_string()) + // String must be length-prefixed - decode as LenPrefixed<u16, String> + crate::types::prefix::length::LenPrefixed::<u16, String>::decode(buf, ctx) + .map(|p| p.data().clone()) } } -// ~ Decode primitive +// ── Decode primitive ──────────────────────────────────────────────────── impl_decode!(u8); impl_decode!(u16); @@ -277,39 +284,43 @@ impl_decode!(i128); impl_decode!(isize); impl<Data> Decode<Data> for f64 { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let bits = u64::decode(reader, ctx)?; - let value = f64::from_bits(bits); - Ok(value) + let bits = u64::decode(buf, ctx)?; + Ok(f64::from_bits(bits)) } } impl<Data> Decode<Data> for f32 { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let bits = u32::decode(reader, ctx)?; - let value = f32::from_bits(bits); - Ok(value) + let bits = u32::decode(buf, ctx)?; + Ok(f32::from_bits(bits)) } } impl<Data> Decode<Data> for bool { - fn decode(reader: &mut dyn Read, _: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], _: &Context<Data>) -> Result<Self> where Self: Sized, { - let mut byte = [0_u8]; - reader.read_exact(&mut byte)?; - Ok(byte[0] != 0) + if buf.is_empty() { + return Err(CodecError::IoError(std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, + "insufficient bytes for bool", + ))); + } + let b = buf[0]; + *buf = &buf[1..]; + Ok(b != 0) } } -// ~ Decode tuple +// ── Decode tuple ──────────────────────────────────────────────────────── impl_decode_tuples!(A); impl_decode_tuples!(A, B); @@ -329,14 +340,14 @@ impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O, P); impl_decode_tuples!(A, B, C, D, E, F, G, H, I, J, K, L, M, N, O, P, Q); -// ~ Decode pointer +// ── Decode pointer ────────────────────────────────────────────────────── impl<Data, T> Decode<Data> for std::sync::Arc<T> where T: Decode<Data>, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> { - T::decode(reader, ctx).map(Arc::new) + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> { + T::decode(buf, ctx).map(Arc::new) } } @@ -344,11 +355,11 @@ impl<Data, T> Decode<Data> for Arc<[T]> where T: Decode<Data>, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> Result<Self> where Self: Sized, { - let slice = T::decode_slice(reader, ctx)?; + let slice = T::decode_slice(buf, ctx)?; Ok(slice.into()) } } diff --git a/src/codec/encode.rs b/src/codec/encode.rs index 7463995..735aa70 100644 --- a/src/codec/encode.rs +++ b/src/codec/encode.rs @@ -1,49 +1,36 @@ -//! Encoding: writing protocol values into a byte stream. use crate::codec::error::Result; use crate::context::Context; -use std::{ - collections::{BTreeMap, BTreeSet, BinaryHeap, HashMap, HashSet, LinkedList, VecDeque}, - io::Write, -}; +use std::collections::{BTreeMap, BTreeSet, BinaryHeap, HashMap, HashSet, LinkedList, VecDeque}; macro_rules! impl_encode { - (u8) => { - impl<Data> Encode<Data> for u8 { - /// write number using big endian - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { - let size = buffer.write(&self.to_be_bytes())?; - Ok(size) + ($t:ty) => { + impl<Data> Encode<Data> for $t { + fn encode(&self, buffer: &mut Vec<u8>, _: &Context<Data>) -> Result<usize> { + buffer.extend_from_slice(&self.to_be_bytes()); + Ok(std::mem::size_of::<$t>()) } - /// write raw bytes - fn encode_slice( - slice: &[Self], - buffer: &mut dyn Write, - _: &Context<Data>, - ) -> Result<usize> { - buffer.write(slice).map_err(Into::into) - } - } - }; - ($t: ty) => { - impl<Data> Encode<Data> for $t { - /// write number using big endian - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { - let size = buffer.write(&self.to_be_bytes())?; - Ok(size) + fn encode_slice(slice: &[Self], buffer: &mut Vec<u8>, _: &Context<Data>) -> Result<usize> + where + Self: Sized, + { + let len = slice.len() * std::mem::size_of::<$t>(); + buffer.reserve(len); + for item in slice { + buffer.extend_from_slice(&item.to_be_bytes()); + } + Ok(len) } - fn encode_slice( - slice: &[Self], - buffer: &mut dyn Write, - _: &Context<Data>, - ) -> Result<usize> { - let len = buffer.write( - &slice - .iter() - .flat_map(|n| n.to_be_bytes()) - .collect::<Box<[u8]>>(), - )?; + fn encode_iter<'a, I>(iter: I, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> + where + Self: 'a + Sized, + I: IntoIterator<Item = &'a Self>, + { + let mut len = 0; + for item in iter { + len += item.encode(buffer, ctx)?; + } Ok(len) } } @@ -57,7 +44,7 @@ macro_rules! impl_encode_tuples { $($generic: Encode<Data>),+ { #[allow(non_snake_case)] - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { let ($($generic,)+): &($($generic,)+) = self; let mut len = 0; $(len += $generic.encode(buffer, ctx)?;)+ @@ -67,10 +54,11 @@ macro_rules! impl_encode_tuples { }; } -/// Write a value into a byte stream +/// Zero-copy encoding: writes directly into a byte buffer pub trait Encode<Data> { - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize>; - fn encode_slice(slice: &[Self], buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize>; + + fn encode_slice(slice: &[Self], buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> where Self: Sized, { @@ -80,204 +68,257 @@ pub trait Encode<Data> { } Ok(len) } + + fn encode_iter<'a, I>(iter: I, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> + where + Self: 'a + Sized, + I: IntoIterator<Item = &'a Self>, + { + let mut len = 0; + for item in iter { + len += item.encode(buffer, ctx)?; + } + Ok(len) + } } -// ~ Encode arrays +// ── Encode arrays ─────────────────────────────────────────────────────── impl<Data, T: Encode<Data>> Encode<Data> for Vec<T> { - /// Encode each element of Vec - /// use `CountPrefix` to prefix the vector with number of elements - /// use `LenPrefix` to prefix the vector with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { T::encode_slice(self, buffer, ctx) } } impl<Data, T: Encode<Data>> Encode<Data> for VecDeque<T> { - /// Encode each element of VecDeque - /// use `CountPrefix` to prefix the VecDeque with number of elements - /// use `LenPrefix` to prefix the VecDeque with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { - let (front, _) = self.as_slices(); - T::encode_slice(front, buffer, ctx) + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + let (front, back) = self.as_slices(); + let mut len = 0; + len += T::encode_slice(front, buffer, ctx)?; + len += T::encode_slice(back, buffer, ctx)?; + Ok(len) } } -impl<Data, T: Encode<Data>> Encode<Data> for LinkedList<T> -where - T: Clone, -{ - /// Encode each element of LinkedList (use clone..) - /// use `CountPrefix` to prefix the LinkedList with number of elements - /// use `LenPrefix` to prefix the LinkedList with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { - let mut view = Vec::with_capacity(self.len()); - view.extend(self.iter().cloned()); - T::encode_slice(&view, buffer, ctx) +impl<Data, T: Encode<Data>> Encode<Data> for LinkedList<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + let mut len = 0; + for item in self { + len += item.encode(buffer, ctx)?; + } + Ok(len) } } -// ~ Encode slices +// ── Encode slices ─────────────────────────────────────────────────────── impl<Data, T: Encode<Data>> Encode<Data> for &[T] { - /// Encode each element of slice (use clone..) - /// use `CountPrefix` to prefix the slice with number of elements - /// use `LenPrefix` to prefix the slice with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { T::encode_slice(self, buffer, ctx) } } impl<Data, T: Encode<Data>, const S: usize> Encode<Data> for [T; S] { - /// Encode each element of slice (use clone..) - /// use `CountPrefix` to prefix the slice with number of elements - /// use `LenPrefix` to prefix the slice with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { T::encode_slice(self, buffer, ctx) } } impl<Data, T: Encode<Data>> Encode<Data> for [T] { - /// Encode each element of slice (use clone..) - /// use `CountPrefix` to prefix the slice with number of elements - /// use `LenPrefix` to prefix the slice with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + T::encode_slice(self, buffer, ctx) + } +} + +// ── Encode references to collections ──────────────────────────────────── + +impl<Data, T: Encode<Data>> Encode<Data> for &Vec<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { T::encode_slice(self, buffer, ctx) } } -// ~ Encode set +impl<Data, T: Encode<Data>> Encode<Data> for &VecDeque<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + let (front, back) = self.as_slices(); + let mut len = 0; + len += T::encode_slice(front, buffer, ctx)?; + len += T::encode_slice(back, buffer, ctx)?; + Ok(len) + } +} -impl<Data, T: Encode<Data>> Encode<Data> for HashSet<T> +impl<Data, T: Encode<Data>> Encode<Data> for &LinkedList<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + let mut len = 0; + for item in (*self).iter() { + len += item.encode(buffer, ctx)?; + } + Ok(len) + } +} + +impl<Data, T: Encode<Data>> Encode<Data> for &HashSet<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + T::encode_iter(self.iter(), buffer, ctx) + } +} + +impl<Data, T: Encode<Data>> Encode<Data> for &BTreeSet<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + T::encode_iter(self.iter(), buffer, ctx) + } +} + +impl<Data, T: Encode<Data>> Encode<Data> for &BinaryHeap<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + T::encode_slice(self.as_slice(), buffer, ctx) + } +} + +impl<Data, K, V> Encode<Data> for &HashMap<K, V> where - T: Clone, + K: Encode<Data>, + V: Encode<Data>, { - /// Encode each element of HashSet (use clone..) - /// use `CountPrefix` to prefix the HashSet with number of elements - /// use `LenPrefix` to prefix the HashSet with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { - let mut view = Vec::with_capacity(self.len()); - view.extend(self.iter().cloned()); - T::encode_slice(&view, buffer, ctx) + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + let mut len = 0; + for (k, v) in (*self).iter() { + len += k.encode(buffer, ctx)?; + len += v.encode(buffer, ctx)?; + } + Ok(len) } } -impl<Data, T: Encode<Data>> Encode<Data> for BTreeSet<T> +impl<Data, K, V> Encode<Data> for &BTreeMap<K, V> where - T: Clone, + K: Encode<Data>, + V: Encode<Data>, { - /// Encode each element of BTreeSet (use clone..) - /// use `CountPrefix` to prefix the BTreeSet with number of elements - /// use `LenPrefix` to prefix the BTreeSet with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { - let mut view = Vec::with_capacity(self.len()); - view.extend(self.iter().cloned()); - T::encode_slice(&view, buffer, ctx) + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + let mut len = 0; + for (k, v) in (*self).iter() { + len += k.encode(buffer, ctx)?; + len += v.encode(buffer, ctx)?; + } + Ok(len) } } -// ~ Encode misc +// ── Encode set (owned) ────────────────────────────────────────────────── + +impl<Data, T: Encode<Data>> Encode<Data> for HashSet<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + T::encode_iter(self.iter(), buffer, ctx) + } +} + +impl<Data, T: Encode<Data>> Encode<Data> for BTreeSet<T> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + T::encode_iter(self.iter(), buffer, ctx) + } +} + +// ── Encode misc ───────────────────────────────────────────────────────── impl<Data, T> Encode<Data> for BinaryHeap<T> where T: Encode<Data>, { - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { T::encode_slice(self.as_slice(), buffer, ctx) } } -// ~ Encode maps +// ── Encode maps (owned) ───────────────────────────────────────────────── impl<Data, K, V> Encode<Data> for HashMap<K, V> where - K: Encode<Data> + Clone, - V: Encode<Data> + Clone, + K: Encode<Data>, + V: Encode<Data>, { - /// Encode each element of HashMap (use clone..) - /// use `CountPrefix` to prefix the HashMap with number of pairs - /// use `LenPrefix` to prefix the HashMap with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { - let slice: Vec<(K, V)> = self.clone().into_iter().collect(); - <(K, V) as Encode<Data>>::encode_slice(&slice, buffer, ctx) + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + let mut len = 0; + for (k, v) in self { + len += k.encode(buffer, ctx)?; + len += v.encode(buffer, ctx)?; + } + Ok(len) } } impl<Data, K, V> Encode<Data> for BTreeMap<K, V> where - K: Encode<Data> + Clone, - V: Encode<Data> + Clone, + K: Encode<Data>, + V: Encode<Data>, { - /// Encode each element of BTreeMap (use clone..) - /// use `CountPrefix` to prefix the BTreeMap with number of pairs - /// use `LenPrefix` to prefix the BTreeMap with encoded byte size - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { - let slice: Vec<(K, V)> = self.clone().into_iter().collect(); - <(K, V) as Encode<Data>>::encode_slice(&slice, buffer, ctx) + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + let mut len = 0; + for (k, v) in self { + len += k.encode(buffer, ctx)?; + len += v.encode(buffer, ctx)?; + } + Ok(len) } } -// ~ Encode string +// ── Encode string ─────────────────────────────────────────────────────── + impl<Data> Encode<Data> for String { - /// Encode the string using utf8 - /// use `CountPrefix` to prefix the string with char length - /// use `LenPrefix` to prefix the string with utf8 bytes length - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, _: &Context<Data>) -> Result<usize> { let utf8 = self.as_bytes(); - buffer.write_all(utf8)?; + buffer.extend_from_slice(utf8); Ok(utf8.len()) } } impl<Data> Encode<Data> for &str { - /// Encode the string using utf8 - /// use `CountPrefix` to prefix the string with char length - /// use `LenPrefix` to prefix the string with utf8 bytes length - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, _: &Context<Data>) -> Result<usize> { let utf8 = self.as_bytes(); - buffer.write_all(utf8)?; + buffer.extend_from_slice(utf8); Ok(utf8.len()) } } impl<Data, T: Encode<Data>> Encode<Data> for Option<T> { - /// Encode `T` if Some(T) or do nothing if None - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { match self { - Some(val) => { - let size = val.encode(buffer, ctx)?; - Ok(size) - } + Some(val) => val.encode(buffer, ctx), None => Ok(0), } } } -// ~ Encode primitive +// ── Encode primitive ──────────────────────────────────────────────────── + impl<Data> Encode<Data> for bool { - fn encode(&self, buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> { - let size = buffer.write(&[*self as u8])?; - Ok(size) + fn encode(&self, buffer: &mut Vec<u8>, _: &Context<Data>) -> Result<usize> { + buffer.push(*self as u8); + Ok(1) } - fn encode_slice(slice: &[Self], buffer: &mut dyn Write, _: &Context<Data>) -> Result<usize> + fn encode_slice(slice: &[Self], buffer: &mut Vec<u8>, _: &Context<Data>) -> Result<usize> where Self: Sized, { - let len = buffer.write(&slice.iter().map(|n| *n as u8).collect::<Box<[u8]>>())?; + let len = slice.len(); + buffer.reserve(len); + for &b in slice { + buffer.push(b as u8); + } Ok(len) } } -// ~ Encode pointer +// ── Encode pointer ────────────────────────────────────────────────────── impl<Data, T> Encode<Data> for std::sync::Arc<T> where T: Encode<Data>, { - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { - let data = self.as_ref(); - data.encode(buffer, ctx) + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + self.as_ref().encode(buffer, ctx) } } @@ -285,9 +326,8 @@ impl<Data, T> Encode<Data> for std::sync::Arc<[T]> where T: Encode<Data>, { - fn encode(&self, buffer: &mut dyn Write, ctx: &Context<Data>) -> Result<usize> { - let data = self.as_ref(); - data.encode(buffer, ctx) + fn encode(&self, buffer: &mut Vec<u8>, ctx: &Context<Data>) -> Result<usize> { + self.as_ref().encode(buffer, ctx) } } @@ -312,14 +352,14 @@ impl_encode!(f64); #[cfg(feature = "f128")] impl_encode!(f128); -// ~ Encode tuple +// ── Encode tuple ──────────────────────────────────────────────────────── impl<Data> Encode<Data> for () { - fn encode(&self, _: &mut dyn Write, _: &Context<Data>) -> Result<usize> { + fn encode(&self, _: &mut Vec<u8>, _: &Context<Data>) -> Result<usize> { Ok(0) } - fn encode_slice(_: &[Self], _: &mut dyn Write, _: &Context<Data>) -> Result<usize> + fn encode_slice(_: &[Self], _: &mut Vec<u8>, _: &Context<Data>) -> Result<usize> where Self: Sized, { @@ -8,11 +8,8 @@ pub mod event; pub mod transport; pub mod types; -#[cfg(feature = "macros")] -pub use zr_protocol_macros as macros; +pub use zr_protocol_macros::Codec; +pub use zr_protocol_macros::Decode; +pub use zr_protocol_macros::Encode; -/// The default buffer length used to allocate chunks -pub(crate) const DEFAULT_BUFFER_LEN: usize = 2048; - -// TODO : rewrite tests -// TODO : benchmark with some client / server example +pub use types::prefix::{CountPrefix, LenPrefixed}; diff --git a/src/transport/receiver.rs b/src/transport/receiver.rs index 3a7c948..a165110 100644 --- a/src/transport/receiver.rs +++ b/src/transport/receiver.rs @@ -3,3 +3,16 @@ use crate::{codec::Codec, context::Context, transport::Result}; pub trait PacketReceiver<Data, Uid: PartialEq>: Send + Sync { fn recv<P: Codec<Data>>(&self, ctx: &Context<Data>) -> Result<(Uid, P)>; } + +pub trait PacketReceiverBuf<Data, Uid: PartialEq>: Send + Sync { + fn recv_buf(&self, ctx: &Context<Data>) -> Result<(Uid, Vec<u8>)>; +} + +impl<Data, Uid: PartialEq, R: PacketReceiverBuf<Data, Uid>> PacketReceiver<Data, Uid> for R { + fn recv<P: Codec<Data>>(&self, ctx: &Context<Data>) -> Result<(Uid, P)> { + let (uid, bytes) = self.recv_buf(ctx)?; + let mut reader = bytes.as_slice(); + let packet = P::decode(&mut reader, ctx)?; + Ok((uid, packet)) + } +}
\ No newline at end of file diff --git a/src/transport/sender.rs b/src/transport/sender.rs index 804914f..3c83c6e 100644 --- a/src/transport/sender.rs +++ b/src/transport/sender.rs @@ -3,3 +3,15 @@ use crate::{codec::Codec, context::Context, transport::Result}; pub trait PacketSender<Data>: Send + Sync { fn send<P: Codec<Data>>(&self, packet: P, ctx: &Context<Data>) -> Result<()>; } + +pub trait PacketSenderBuf<Data>: Send + Sync { + fn send_buf(&self, buf: &[u8]) -> Result<()>; +} + +impl<Data, S: PacketSenderBuf<Data>> PacketSender<Data> for S { + fn send<P: Codec<Data>>(&self, packet: P, ctx: &Context<Data>) -> Result<()> { + let mut buf = Vec::new(); + packet.encode(&mut buf, ctx)?; + self.send_buf(&buf) + } +}
\ No newline at end of file diff --git a/src/transport/tcp.rs b/src/transport/tcp.rs index 19d02d6..12ecaa0 100644 --- a/src/transport/tcp.rs +++ b/src/transport/tcp.rs @@ -1,35 +1,30 @@ use std::hash::Hash; -use std::io::Write; +use std::io::{Read, Write}; use std::marker::PhantomData; use std::net::{SocketAddr, TcpListener as StdTcpListener, TcpStream}; use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::{Arc, Mutex}; -use crate::codec::Codec; use crate::context::Context; use crate::transport::connection::Connection; use crate::transport::error::TransportError; use crate::transport::listener::Listener; -use crate::transport::receiver::PacketReceiver; -use crate::transport::sender::PacketSender; +use crate::transport::receiver::PacketReceiverBuf; +use crate::transport::sender::PacketSenderBuf; use crate::transport::{Result, Transport}; -// TODO : check Tcp Transport - #[derive(Clone)] pub struct TcpSender { stream: Arc<Mutex<TcpStream>>, } -impl<Data> PacketSender<Data> for TcpSender { - fn send<P: Codec<Data>>(&self, packet: P, ctx: &Context<Data>) -> Result<()> { - let mut buf = Vec::new(); - packet.encode(&mut buf, ctx)?; +impl<Data> PacketSenderBuf<Data> for TcpSender { + fn send_buf(&self, buf: &[u8]) -> Result<()> { let mut guard = self .stream .lock() .map_err(|_| TransportError::LockPoisoned)?; - guard.write_all(&buf)?; + guard.write_all(buf)?; guard.flush()?; Ok(()) } @@ -40,17 +35,24 @@ pub struct TcpReceiver<Uid> { uid: Uid, } -impl<Data, Uid> PacketReceiver<Data, Uid> for TcpReceiver<Uid> +impl<Data, Uid> PacketReceiverBuf<Data, Uid> for TcpReceiver<Uid> where Uid: PartialEq + Clone + Send + Sync, { - fn recv<P: Codec<Data>>(&self, ctx: &Context<Data>) -> Result<(Uid, P)> { + fn recv_buf(&self, _ctx: &Context<Data>) -> Result<(Uid, Vec<u8>)> { let mut guard = self .stream .lock() .map_err(|_| TransportError::LockPoisoned)?; - let packet = P::decode(&mut *guard, ctx)?; - Ok((self.uid.clone(), packet)) + + // Read 4-byte big-endian length prefix + let mut len_buf = [0u8; 4]; + guard.read_exact(&mut len_buf)?; + let len = u32::from_be_bytes(len_buf) as usize; + + let mut buf = vec![0u8; len]; + guard.read_exact(&mut buf)?; + Ok((self.uid.clone(), buf)) } } diff --git a/src/transport/udp.rs b/src/transport/udp.rs index acb36f2..9b312cf 100644 --- a/src/transport/udp.rs +++ b/src/transport/udp.rs @@ -10,11 +10,11 @@ use crate::context::Context; use crate::transport::Result; use crate::transport::connection::Connection; use crate::transport::listener::Listener; -use crate::transport::receiver::PacketReceiver; -use crate::transport::sender::PacketSender; +use crate::transport::receiver::PacketReceiverBuf; +use crate::transport::sender::PacketSenderBuf; use crate::transport::{Transport, TransportError}; -// TODO : check Udp Transport +const UDP_BUF_SIZE: usize = 1500; #[derive(Clone)] pub struct UdpPeerSender { @@ -22,11 +22,9 @@ pub struct UdpPeerSender { peer: SocketAddr, } -impl<Data> PacketSender<Data> for UdpPeerSender { - fn send<P: crate::codec::Codec<Data>>(&self, packet: P, ctx: &Context<Data>) -> Result<()> { - let mut buf = Vec::new(); - packet.encode(&mut buf, ctx)?; - self.sock.send_to(&buf, self.peer)?; +impl<Data> PacketSenderBuf<Data> for UdpPeerSender { + fn send_buf(&self, buf: &[u8]) -> Result<()> { + self.sock.send_to(buf, self.peer)?; Ok(()) } } @@ -57,11 +55,11 @@ impl<Uid> UdpReceiver<Uid> { } } -impl<Data, Uid> PacketReceiver<Data, Uid> for UdpReceiver<Uid> +impl<Data, Uid> PacketReceiverBuf<Data, Uid> for UdpReceiver<Uid> where Uid: PartialEq + Clone + Send + Sync, { - fn recv<P: crate::codec::Codec<Data>>(&self, ctx: &Context<Data>) -> Result<(Uid, P)> { + fn recv_buf(&self, _ctx: &Context<Data>) -> Result<(Uid, Vec<u8>)> { let bytes = match &self.inner { UdpReceiverInner::Channel(rx) => rx .lock() @@ -69,15 +67,13 @@ where .recv() .map_err(|_| TransportError::ChannelClosed)?, UdpReceiverInner::Socket(sock) => { - let mut buf = vec![0u8; 65535]; + let mut buf = vec![0u8; UDP_BUF_SIZE]; let (n, _) = sock.recv_from(&mut buf)?; buf.truncate(n); buf } }; - let mut reader = bytes.as_slice(); - let packet = P::decode(&mut reader, ctx)?; - Ok((self.uid.clone(), packet)) + Ok((self.uid.clone(), bytes)) } } @@ -100,7 +96,7 @@ where fn accept(&mut self) -> Result<Connection<Uid, Data, UdpPeerSender, UdpReceiver<Uid>>> { loop { - let mut buf = vec![0u8; 65535]; + let mut buf = vec![0u8; UDP_BUF_SIZE]; let (n, src) = self.sock.recv_from(&mut buf)?; buf.truncate(n); diff --git a/src/types/prefix.rs b/src/types/prefix.rs index 67cf047..1f007cf 100644 --- a/src/types/prefix.rs +++ b/src/types/prefix.rs @@ -1,2 +1,5 @@ pub mod count; pub mod length; + +pub use count::CountPrefix; +pub use length::LenPrefixed; diff --git a/src/types/prefix/count.rs b/src/types/prefix/count.rs index 7099cec..04bc7c6 100644 --- a/src/types/prefix/count.rs +++ b/src/types/prefix/count.rs @@ -1,4 +1,4 @@ -use std::{io::Read, marker::PhantomData}; +use std::marker::PhantomData; use getset::Getters; @@ -68,20 +68,18 @@ where { fn encode( &self, - buffer: &mut dyn std::io::prelude::Write, + buffer: &mut Vec<u8>, ctx: &Context<Data>, - ) -> Result<usize, crate::codec::error::CodecError> + ) -> crate::codec::error::Result<usize> where Self: Sized, { - let vec = self.data.clone().into_iter().collect::<Vec<I>>(); + let vec: Vec<I> = self.data.clone().into_iter().collect(); let len: L = vec.len().try_into().map_err(|_| { crate::codec::error::CodecError::Custom("count exceeds prefix capacity".into()) })?; let mut l = len.encode(buffer, ctx)?; - for item in vec { - l += item.encode(buffer, ctx)?; - } + l += I::encode_slice(&vec, buffer, ctx)?; Ok(l) } } @@ -92,14 +90,17 @@ where L: Codec<Data> + Size, D: IntoIterator<Item = I> + FromIterator<I>, { - fn decode(reader: &mut dyn Read, ctx: &Context<Data>) -> crate::codec::error::Result<Self> + fn decode( + buf: &mut &[u8], + ctx: &Context<Data>, + ) -> crate::codec::error::Result<Self> where Self: Sized, { - let len = L::decode(reader, ctx)?.into_size(); + let len = L::decode(buf, ctx)?.into_size(); let mut data = Vec::with_capacity(len); for _ in 0..len { - let item = I::decode(reader, ctx)?; + let item = I::decode(buf, ctx)?; data.push(item); } Ok(Self { @@ -107,4 +108,4 @@ where _len: PhantomData, }) } -} +}
\ No newline at end of file diff --git a/src/types/prefix/length.rs b/src/types/prefix/length.rs index f30eea0..fe839d5 100644 --- a/src/types/prefix/length.rs +++ b/src/types/prefix/length.rs @@ -1,10 +1,9 @@ -use std::{io::Write, marker::PhantomData}; +use std::marker::PhantomData; use getset::Getters; use crate::{ - DEFAULT_BUFFER_LEN, - codec::{self, Codec, decode::Decode, encode::Encode}, + codec::{Codec, decode::Decode, encode::Encode}, context::Context, types::size::Size, }; @@ -56,14 +55,16 @@ where { fn encode( &self, - writer: &mut dyn Write, + buffer: &mut Vec<u8>, ctx: &Context<Data>, - ) -> Result<usize, codec::error::CodecError> { - let mut buf = Vec::with_capacity(DEFAULT_BUFFER_LEN); - self.data.encode(&mut buf, ctx)?; - let len = L::from_size(buf.len()); - let mut len = len.encode(writer, ctx)?; - len += writer.write(&buf)?; + ) -> crate::codec::error::Result<usize> { + let mut len = 0; + let mut inner_buf = Vec::with_capacity(64); + self.data.encode(&mut inner_buf, ctx)?; + let len_val = L::from_size(inner_buf.len()); + len += len_val.encode(buffer, ctx)?; + buffer.extend_from_slice(&inner_buf); + len += inner_buf.len(); Ok(len) } } @@ -73,17 +74,22 @@ where L: Codec<Data> + Size, D: Codec<Data>, { - fn decode( - reader: &mut dyn std::io::prelude::Read, - ctx: &Context<Data>, - ) -> crate::codec::error::Result<Self> + fn decode(buf: &mut &[u8], ctx: &Context<Data>) -> crate::codec::error::Result<Self> where Self: Sized, { - let len = L::decode(reader, ctx)?.into_size(); - let mut limited = vec![0_u8; len]; - reader.read_exact(&mut limited)?; - let data = D::decode(&mut &limited[..], ctx)?; + let len = L::decode(buf, ctx)?.into_size(); + if buf.len() < len { + return Err(crate::codec::error::CodecError::IoError( + std::io::Error::new( + std::io::ErrorKind::UnexpectedEof, + "insufficient bytes for len-prefixed data", + ), + )); + } + let (mut data, rest) = buf.split_at(len); + *buf = rest; + let data = D::decode(&mut data, ctx)?; Ok(Self { data, _len: PhantomData, diff --git a/tests/count_prefix.rs b/tests/count_prefix.rs deleted file mode 100644 index ab25e30..0000000 --- a/tests/count_prefix.rs +++ /dev/null @@ -1,305 +0,0 @@ -use zr_protocol::{ - codec::{decode::Decode, encode::Encode}, - types::prefix::count::CountPrefix, -}; - -fn encode_to_bytes<T: Encode>(value: T) -> Vec<u8> { - let mut buf = Vec::new(); - value.encode(&mut buf).unwrap(); - buf -} - -// ────────────────────── Vecteur vide ──────────── - -#[test] -fn empty_vec_u8_with_u8_count() { - let cp = CountPrefix::<u8, u8, Vec<u8>>::new(vec![]); - let buf = encode_to_bytes(cp); - assert_eq!(buf, vec![0x00]); -} - -#[test] -fn empty_vec_u32_with_u16_count() { - let cp = CountPrefix::<u32, u16, Vec<u32>>::new(vec![]); - let buf = encode_to_bytes(cp); - assert_eq!(buf, vec![0x00, 0x00]); -} - -#[test] -fn empty_vec_u8_with_u32_count() { - let cp = CountPrefix::<u8, u32, Vec<u8>>::new(vec![]); - let buf = encode_to_bytes(cp); - assert_eq!(buf, vec![0x00, 0x00, 0x00, 0x00]); -} - -// ────────────────────── Un seul élément ───────── - -#[test] -fn single_u8_u8_count() { - let cp = CountPrefix::<u8, u8, Vec<u8>>::new(vec![42]); - let buf = encode_to_bytes(cp); - assert_eq!(buf, vec![0x01, 0x2A]); -} - -#[test] -fn single_u32_u8_count() { - let cp = CountPrefix::<u32, u8, Vec<u32>>::new(vec![0x01020304]); - let buf = encode_to_bytes(cp); - assert_eq!(buf, vec![0x01, 0x01, 0x02, 0x03, 0x04]); -} - -// ────────────────────── Plusieurs éléments ────── - -#[test] -fn multiple_u8_u8_count() { - let cp = CountPrefix::<u8, u8, Vec<u8>>::new(vec![1, 2, 3]); - let buf = encode_to_bytes(cp); - assert_eq!(buf, vec![0x03, 0x01, 0x02, 0x03]); -} - -#[test] -fn multiple_u16_u8_count() { - let cp = CountPrefix::<u16, u8, Vec<u16>>::new(vec![256, 512, 1024]); - let buf = encode_to_bytes(cp); - assert_eq!( - buf, - vec![ - 0x03, // count = 3 - 0x01, 0x00, // 256 - 0x02, 0x00, // 512 - 0x04, 0x00, // 1024 - ] - ); -} - -#[test] -fn multiple_u32_u16_count() { - let cp = CountPrefix::<u32, u16, Vec<u32>>::new(vec![100, 200, 300]); - let buf = encode_to_bytes(cp); - assert_eq!( - buf, - vec![ - 0x00, 0x03, // count = 3 (u16) - 0x00, 0x00, 0x00, 0x64, // 100 - 0x00, 0x00, 0x00, 0xC8, // 200 - 0x00, 0x00, 0x01, 0x2C, // 300 - ] - ); -} - -// ────────────────────── Valeurs extrêmes ──────── - -#[test] -fn u8_count_max_255() { - let data: Vec<u8> = (0..255).collect(); - let cp = CountPrefix::<u8, u8, Vec<u8>>::new(data.clone()); - let buf = encode_to_bytes(cp); - assert_eq!(buf[0], 0xFF); - assert_eq!(buf.len(), 1 + 255); -} - -#[test] -fn large_u32_items_u16_count() { - let data: Vec<u32> = vec![u32::MAX, 0, u32::MIN + 1]; - let cp = CountPrefix::<u32, u16, Vec<u32>>::new(data); - let buf = encode_to_bytes(cp); - assert_eq!(buf.len(), 2 + 3 * 4); - assert_eq!(&buf[0..2], &[0x00, 0x03]); // count = 3 -} - -// ────────────────────── Roundtrip ─────────────── - -#[test] -fn roundtrip_u8_u8() { - let data = vec![10, 20, 30, 40, 50]; - let cp = CountPrefix::<u8, u8, Vec<u8>>::new(data.clone()); - let buf = encode_to_bytes(cp); - - let mut reader = &buf[..]; - let decoded = CountPrefix::<u8, u8, Vec<u8>>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), data); -} - -#[test] -fn roundtrip_u32_u16() { - let data = vec![1, 2, 3, 4, 5]; - let cp = CountPrefix::<u32, u16, Vec<u32>>::new(data.clone()); - let buf = encode_to_bytes(cp); - - let mut reader = &buf[..]; - let decoded = CountPrefix::<u32, u16, Vec<u32>>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), data); -} - -#[test] -fn roundtrip_empty() { - let cp = CountPrefix::<u32, u8, Vec<u32>>::new(vec![]); - let buf = encode_to_bytes(cp); - - let mut reader = &buf[..]; - let decoded = CountPrefix::<u32, u8, Vec<u32>>::decode(&mut reader).unwrap(); - assert!(decoded.as_ref().is_empty()); -} - -#[test] -fn roundtrip_single_element() { - let cp = CountPrefix::<u64, u8, Vec<u64>>::new(vec![u64::MAX]); - let buf = encode_to_bytes(cp); - - let mut reader = &buf[..]; - let decoded = CountPrefix::<u64, u8, Vec<u64>>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), vec![u64::MAX]); -} - -#[test] -fn roundtrip_negative_i32() { - let data = vec![-1, -100, 0, 100, 1]; - let cp = CountPrefix::<i32, u16, Vec<i32>>::new(data.clone()); - let buf = encode_to_bytes(cp); - - let mut reader = &buf[..]; - let decoded = CountPrefix::<i32, u16, Vec<i32>>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), data); -} - -#[test] -fn roundtrip_f64() { - let data = vec![0.0, 3.14, -2.71, f64::INFINITY, f64::NEG_INFINITY]; - let cp = CountPrefix::<f64, u8, Vec<f64>>::new(data.clone()); - let buf = encode_to_bytes(cp); - - let mut reader = &buf[..]; - let decoded = CountPrefix::<f64, u8, Vec<f64>>::decode(&mut reader).unwrap(); - let decoded_data = decoded.as_ref().clone(); - assert_eq!(decoded_data.len(), data.len()); - for (a, b) in decoded_data.iter().zip(data.iter()) { - if a.is_nan() { - assert!(b.is_nan()); - } else { - assert_eq!(a, b); - } - } -} - -#[test] -fn roundtrip_bool() { - let data = vec![true, false, true, true, false]; - let cp = CountPrefix::<bool, u8, Vec<bool>>::new(data.clone()); - let buf = encode_to_bytes(cp); - - let mut reader = &buf[..]; - let decoded = CountPrefix::<bool, u8, Vec<bool>>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), data); -} - -#[test] -fn roundtrip_large_dataset() { - let data: Vec<u32> = (0..1000).collect(); - let cp = CountPrefix::<u32, u16, Vec<u32>>::new(data.clone()); - let buf = encode_to_bytes(cp); - - let mut reader = &buf[..]; - let decoded = CountPrefix::<u32, u16, Vec<u32>>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), data); -} - -// ────────────────────── Imbrication ───────────── - -#[test] -fn nested_count_prefix() { - let inner1 = CountPrefix::<u8, u8, Vec<u8>>::new(vec![1, 2, 3]); - let inner2 = CountPrefix::<u8, u8, Vec<u8>>::new(vec![4, 5]); - let outer = - CountPrefix::<CountPrefix<u8, u8, Vec<u8>>, u8, Vec<CountPrefix<u8, u8, Vec<u8>>>>::new( - vec![inner1, inner2], - ); - - let buf = encode_to_bytes(outer); - - let mut reader = &buf[..]; - let decoded_outer = - CountPrefix::<CountPrefix<u8, u8, Vec<u8>>, u8, Vec<CountPrefix<u8, u8, Vec<u8>>>>::decode( - &mut reader, - ) - .unwrap(); - - let outer_data: Vec<_> = decoded_outer - .as_ref() - .iter() - .map(|c| c.as_ref().clone()) - .collect(); - assert_eq!(outer_data.len(), 2); - assert_eq!(outer_data[0], vec![1, 2, 3]); - assert_eq!(outer_data[1], vec![4, 5]); -} - -// ────────────────────── Erreurs de décodage ───── - -#[test] -fn decode_empty_buffer_fails() { - let result = CountPrefix::<u8, u8, Vec<u8>>::decode(&mut &[][..]); - assert!(result.is_err()); -} - -#[test] -fn decode_truncated_data_fails() { - let cp = CountPrefix::<u32, u8, Vec<u32>>::new(vec![1, 2, 3]); - let buf = encode_to_bytes(cp); - let truncated = &buf[..buf.len() - 1]; - let result = CountPrefix::<u32, u8, Vec<u32>>::decode(&mut &truncated[..]); - assert!(result.is_err()); -} - -#[test] -fn decode_wrong_count_fails() { - let cp = CountPrefix::<u32, u8, Vec<u32>>::new(vec![1, 2]); - let mut buf = encode_to_bytes(cp); - buf[0] = 0x05; - let result = CountPrefix::<u32, u8, Vec<u32>>::decode(&mut &buf[..]); - assert!(result.is_err()); -} - -// ────────────────────── AsRef / AsMut ─────────── - -#[test] -fn as_ref_returns_inner_data() { - let data = vec![10u32, 20, 30]; - let cp = CountPrefix::<u32, u8, Vec<u32>>::new(data.clone()); - assert_eq!(*cp.as_ref(), data); -} - -#[test] -fn as_mut_allows_mutation() { - let mut cp = CountPrefix::<u32, u8, Vec<u32>>::new(vec![1, 2, 3]); - cp.as_mut().push(4); - assert_eq!(*cp.as_ref(), vec![1, 2, 3, 4]); -} - -// ────────────────────── Counts byte exactitude ── - -#[test] -fn u8_count_header() { - let cp = CountPrefix::<u8, u8, Vec<u8>>::new(vec![1; 10]); - let buf = encode_to_bytes(cp); - assert_eq!(buf[0], 10); - assert_eq!(buf.len(), 1 + 10); -} - -#[test] -fn u16_count_header() { - let cp = CountPrefix::<u8, u16, Vec<u8>>::new(vec![1; 300]); - let buf = encode_to_bytes(cp); - assert_eq!(buf[0], 0x01); - assert_eq!(buf[1], 0x2C); - assert_eq!(buf.len(), 2 + 300); -} - -#[test] -fn u32_count_header() { - let data: Vec<u8> = vec![0xAA; 300]; - let len = data.len(); - let cp = CountPrefix::<u8, u32, Vec<u8>>::new(data); - let buf = encode_to_bytes(cp); - assert_eq!(buf.len(), 4 + len); - assert_eq!(&buf[0..4], &[0x00, 0x00, 0x01, 0x2C]); -} diff --git a/tests/encode_decode.rs b/tests/encode_decode.rs deleted file mode 100644 index 92a1e71..0000000 --- a/tests/encode_decode.rs +++ /dev/null @@ -1,589 +0,0 @@ -use std::sync::Arc; -use zr_protocol::codec::{decode::Decode, encode::Encode}; - -fn roundtrip<T: Encode + Decode + PartialEq + Clone + std::fmt::Debug>(value: T) { - let mut buf = Vec::new(); - let written = value.clone().encode(&mut buf).expect("encode failed"); - let decoded = T::decode(&mut &buf[..]).expect("decode failed"); - assert_eq!(written, buf.len(), "encode returned wrong byte count"); - assert_eq!(decoded, value, "roundtrip value mismatch"); -} - -// ────────────────────── u8 ────────────────────── - -#[test] -fn u8_zero() { - roundtrip(0u8); -} - -#[test] -fn u8_one() { - roundtrip(1u8); -} - -#[test] -fn u8_max() { - roundtrip(u8::MAX); -} - -#[test] -fn u8_arbitrary() { - roundtrip(42u8); -} - -#[test] -fn u8_bytes() { - let mut buf = Vec::new(); - 255u8.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![255]); -} - -// ────────────────────── u16 ───────────────────── - -#[test] -fn u16_zero() { - roundtrip(0u16); -} - -#[test] -fn u16_one() { - roundtrip(1u16); -} - -#[test] -fn u16_max() { - roundtrip(u16::MAX); -} - -#[test] -fn u16_big_endian_bytes() { - let mut buf = Vec::new(); - 256u16.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0x01, 0x00]); -} - -#[test] -fn u16_one_big_endian_bytes() { - let mut buf = Vec::new(); - 1u16.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0x00, 0x01]); -} - -#[test] -fn u16_arbitrary() { - roundtrip(1337u16); -} - -// ────────────────────── u32 ───────────────────── - -#[test] -fn u32_zero() { - roundtrip(0u32); -} - -#[test] -fn u32_max() { - roundtrip(u32::MAX); -} - -#[test] -fn u32_big_endian_bytes() { - let mut buf = Vec::new(); - 0x01020304u32.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0x01, 0x02, 0x03, 0x04]); -} - -#[test] -fn u32_arbitrary() { - roundtrip(429496729u32); -} - -// ────────────────────── u64 ───────────────────── - -#[test] -fn u64_zero() { - roundtrip(0u64); -} - -#[test] -fn u64_max() { - roundtrip(u64::MAX); -} - -#[test] -fn u64_big_endian_bytes() { - let mut buf = Vec::new(); - 0x0102030405060708u64.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08]); -} - -#[test] -fn u64_arbitrary() { - roundtrip(999_999_999_999u64); -} - -// ────────────────────── u128 ──────────────────── - -#[test] -fn u128_zero() { - roundtrip(0u128); -} - -#[test] -fn u128_max() { - roundtrip(u128::MAX); -} - -#[test] -fn u128_arbitrary() { - roundtrip(170141183460469231731687303715884105727u128); -} - -// ────────────────────── usize ─────────────────── - -#[test] -fn usize_zero() { - roundtrip(0usize); -} - -#[test] -fn usize_max() { - roundtrip(usize::MAX); -} - -#[test] -fn usize_arbitrary() { - roundtrip(42usize); -} - -// ────────────────────── i8 ────────────────────── - -#[test] -fn i8_zero() { - roundtrip(0i8); -} - -#[test] -fn i8_positive() { - roundtrip(42i8); -} - -#[test] -fn i8_negative() { - roundtrip(-42i8); -} - -#[test] -fn i8_min() { - roundtrip(i8::MIN); -} - -#[test] -fn i8_max() { - roundtrip(i8::MAX); -} - -#[test] -fn i8_negative_bytes() { - let mut buf = Vec::new(); - (-1i8).encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0xFF]); -} - -// ────────────────────── i16 ───────────────────── - -#[test] -fn i16_zero() { - roundtrip(0i16); -} - -#[test] -fn i16_positive() { - roundtrip(1234i16); -} - -#[test] -fn i16_negative() { - roundtrip(-1234i16); -} - -#[test] -fn i16_min() { - roundtrip(i16::MIN); -} - -#[test] -fn i16_max() { - roundtrip(i16::MAX); -} - -#[test] -fn i16_big_endian_negative() { - let mut buf = Vec::new(); - (-1i16).encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0xFF, 0xFF]); -} - -#[test] -fn i16_negative_bytes() { - let mut buf = Vec::new(); - (-256i16).encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0xFF, 0x00]); -} - -// ────────────────────── i32 ───────────────────── - -#[test] -fn i32_zero() { - roundtrip(0i32); -} - -#[test] -fn i32_positive() { - roundtrip(100000i32); -} - -#[test] -fn i32_negative() { - roundtrip(-100000i32); -} - -#[test] -fn i32_min() { - roundtrip(i32::MIN); -} - -#[test] -fn i32_max() { - roundtrip(i32::MAX); -} - -// ────────────────────── i64 ───────────────────── - -#[test] -fn i64_zero() { - roundtrip(0i64); -} - -#[test] -fn i64_positive() { - roundtrip(9_999_999_999i64); -} - -#[test] -fn i64_negative() { - roundtrip(-9_999_999_999i64); -} - -#[test] -fn i64_min() { - roundtrip(i64::MIN); -} - -#[test] -fn i64_max() { - roundtrip(i64::MAX); -} - -// ────────────────────── i128 ──────────────────── - -#[test] -fn i128_zero() { - roundtrip(0i128); -} - -#[test] -fn i128_min() { - roundtrip(i128::MIN); -} - -#[test] -fn i128_max() { - roundtrip(i128::MAX); -} - -#[test] -fn i128_negative() { - roundtrip(-42i128); -} - -// ────────────────────── isize ─────────────────── - -#[test] -fn isize_zero() { - roundtrip(0isize); -} - -#[test] -fn isize_positive() { - roundtrip(42isize); -} - -#[test] -fn isize_negative() { - roundtrip(-42isize); -} - -// ────────────────────── f32 ───────────────────── - -#[test] -fn f32_zero() { - roundtrip(0.0f32); -} - -#[test] -fn f32_positive() { - roundtrip(3.14f32); -} - -#[test] -fn f32_negative() { - roundtrip(-2.71f32); -} - -#[test] -fn f32_infinity() { - roundtrip(f32::INFINITY); -} - -#[test] -fn f32_neg_infinity() { - roundtrip(f32::NEG_INFINITY); -} - -#[test] -fn f32_nan() { - let mut buf = Vec::new(); - f32::NAN.encode(&mut buf).unwrap(); - let decoded = f32::decode(&mut &buf[..]).unwrap(); - assert!(decoded.is_nan(), "NaN roundtrip failed"); -} - -#[test] -fn f32_big_endian_bytes() { - let mut buf = Vec::new(); - 1.0f32.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0x3F, 0x80, 0x00, 0x00]); -} - -// ────────────────────── f64 ───────────────────── - -#[test] -fn f64_zero() { - roundtrip(0.0f64); -} - -#[test] -fn f64_positive() { - roundtrip(3.141592653589793f64); -} - -#[test] -fn f64_negative() { - roundtrip(-2.718281828459045f64); -} - -#[test] -fn f64_infinity() { - roundtrip(f64::INFINITY); -} - -#[test] -fn f64_neg_infinity() { - roundtrip(f64::NEG_INFINITY); -} - -#[test] -fn f64_nan() { - let mut buf = Vec::new(); - f64::NAN.encode(&mut buf).unwrap(); - let decoded = f64::decode(&mut &buf[..]).unwrap(); - assert!(decoded.is_nan(), "NaN roundtrip failed"); -} - -#[test] -fn f64_big_endian_bytes() { - let mut buf = Vec::new(); - 1.0f64.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0x3F, 0xF0, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00]); -} - -// ────────────────────── bool ──────────────────── - -#[test] -fn bool_true() { - roundtrip(true); -} - -#[test] -fn bool_false() { - roundtrip(false); -} - -#[test] -fn bool_true_byte() { - let mut buf = Vec::new(); - true.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![1]); -} - -#[test] -fn bool_false_byte() { - let mut buf = Vec::new(); - false.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0]); -} - -#[test] -fn bool_non_zero_is_true() { - let decoded = bool::decode(&mut &[0x42u8][..]).unwrap(); - assert!(decoded, "any non-zero byte should decode as true"); -} - -// ────────────────────── [u8; S] ───────────────── - -#[test] -fn array_1_byte() { - roundtrip([0xABu8; 1]); -} - -#[test] -fn array_4_bytes() { - roundtrip([0x01u8, 0x02, 0x03, 0x04]); -} - -#[test] -fn array_32_bytes() { - let data: [u8; 32] = std::array::from_fn(|i| i as u8); - roundtrip(data); -} - -#[test] -fn array_zeroes() { - roundtrip([0u8; 16]); -} - -#[test] -fn array_256_bytes() { - let data: [u8; 256] = std::array::from_fn(|i| (i % 256) as u8); - roundtrip(data); -} - -// ────────────────────── Arc<[u8]> ─────────────── - -#[test] -fn arc_slice_simple() { - roundtrip(Arc::<[u8]>::from(vec![1, 2, 3])); -} - -#[test] -fn arc_slice_empty() { - roundtrip(Arc::<[u8]>::from(Vec::<u8>::new())); -} - -#[test] -fn arc_slice_large() { - let data: Vec<u8> = (0..1024).map(|i| (i % 256) as u8).collect(); - roundtrip(Arc::<[u8]>::from(data)); -} - -// ────────────────────── Retour de bytes written ─ - -#[test] -fn encode_returns_correct_byte_count_u32() { - let mut buf = Vec::new(); - let written = 42u32.encode(&mut buf).unwrap(); - assert_eq!(written, 4); -} - -#[test] -fn encode_returns_correct_byte_count_bool() { - let mut buf = Vec::new(); - let written = true.encode(&mut buf).unwrap(); - assert_eq!(written, 1); -} - -#[test] -fn encode_returns_correct_byte_count_u128() { - let mut buf = Vec::new(); - let written = 123u128.encode(&mut buf).unwrap(); - assert_eq!(written, 16); -} - -#[test] -fn encode_returns_correct_byte_count_array() { - let mut buf = Vec::new(); - let written = [1u8, 2, 3, 4, 5].encode(&mut buf).unwrap(); - assert_eq!(written, 5); -} - -// ────────────────────── Buffer trop court ─────── - -#[test] -fn decode_u16_from_empty_buffer_fails() { - let result = u16::decode(&mut &[][..]); - assert!(result.is_err(), "decoding u16 from empty buffer should fail"); -} - -#[test] -fn decode_u32_from_too_short_buffer_fails() { - let result = u32::decode(&mut &[0x01, 0x02][..]); - assert!(result.is_err(), "decoding u32 from 2-byte buffer should fail"); -} - -#[test] -fn decode_bool_from_empty_buffer_fails() { - let result = bool::decode(&mut &[][..]); - assert!(result.is_err(), "decoding bool from empty buffer should fail"); -} - -#[test] -fn decode_array_from_short_buffer_fails() { - let result = <[u8; 4]>::decode(&mut &[0x01, 0x02][..]); - assert!( - result.is_err(), - "decoding [u8;4] from 2-byte buffer should fail" - ); -} - -// ────────────────────── Sérialisation exacte ──── - -#[test] -fn u16_258_bytes() { - let mut buf = Vec::new(); - 258u16.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0x01, 0x02]); -} - -#[test] -fn u32_16909060_bytes() { - let mut buf = Vec::new(); - 0x01020304u32.encode(&mut buf).unwrap(); - assert_eq!(buf, vec![1, 2, 3, 4]); -} - -#[test] -fn i16_minus_256_bytes() { - let mut buf = Vec::new(); - (-256i16).encode(&mut buf).unwrap(); - assert_eq!(buf, vec![0xFF, 0x00]); -} - -// ────────────────────── Sérialisation multiple ── - -#[test] -fn multiple_values_sequential_encode() { - let mut buf = Vec::new(); - 1u8.encode(&mut buf).unwrap(); - 2u16.encode(&mut buf).unwrap(); - 3u32.encode(&mut buf).unwrap(); - true.encode(&mut buf).unwrap(); - - assert_eq!(buf.len(), 1 + 2 + 4 + 1); - - let mut reader = &buf[..]; - assert_eq!(u8::decode(&mut reader).unwrap(), 1); - assert_eq!(u16::decode(&mut reader).unwrap(), 2); - assert_eq!(u32::decode(&mut reader).unwrap(), 3); - assert!(bool::decode(&mut reader).unwrap()); -} diff --git a/tests/len_prefix.rs b/tests/len_prefix.rs deleted file mode 100644 index d9597eb..0000000 --- a/tests/len_prefix.rs +++ /dev/null @@ -1,310 +0,0 @@ -use zr_protocol::{ - codec::{decode::Decode, encode::Encode}, - types::prefix::length::LenPrefixed, -}; - -fn encode_to_bytes<T: Encode>(value: T) -> Vec<u8> { - let mut buf = Vec::new(); - value.encode(&mut buf).unwrap(); - buf -} - -// ────────────────────── Données simples ───────── - -#[test] -fn simple_u32_with_u8_length() { - let lp = LenPrefixed::<u8, u32>::new(42u32); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x04, 0x00, 0x00, 0x00, 0x2A]); -} - -#[test] -fn simple_u16_with_u16_length() { - let lp = LenPrefixed::<u16, u16>::new(256u16); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x00, 0x02, 0x01, 0x00]); -} - -#[test] -fn simple_u8_with_u8_length() { - let lp = LenPrefixed::<u8, u8>::new(0xABu8); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x01, 0xAB]); -} - -#[test] -fn simple_u64_with_u32_length() { - let lp = LenPrefixed::<u32, u64>::new(0x0102030405060708u64); - let buf = encode_to_bytes(lp); - assert_eq!( - buf, - vec![0x00, 0x00, 0x00, 0x08, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x08] - ); -} - -// ────────────────────── Booléen ───────────────── - -#[test] -fn bool_true_with_u8_length() { - let lp = LenPrefixed::<u8, bool>::new(true); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x01, 0x01]); -} - -#[test] -fn bool_false_with_u8_length() { - let lp = LenPrefixed::<u8, bool>::new(false); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x01, 0x00]); -} - -// ────────────────────── Tableaux ──────────────── - -#[test] -fn array_4_bytes_with_u8_length() { - let lp = LenPrefixed::<u8, [u8; 4]>::new([0x01, 0x02, 0x03, 0x04]); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x04, 0x01, 0x02, 0x03, 0x04]); -} - -#[test] -fn array_16_bytes_with_u16_length() { - let data: [u8; 16] = std::array::from_fn(|i| i as u8); - let lp = LenPrefixed::<u16, [u8; 16]>::new(data); - let buf = encode_to_bytes(lp); - assert_eq!(buf[0..2], [0x00, 0x10]); - assert_eq!(&buf[2..], &data[..]); -} - -#[test] -fn array_256_bytes_with_u16_length() { - let data: [u8; 256] = std::array::from_fn(|i| (i % 256) as u8); - let lp = LenPrefixed::<u16, [u8; 256]>::new(data); - let buf = encode_to_bytes(lp); - assert_eq!(buf[0..2], [0x01, 0x00]); - assert_eq!(&buf[2..], &data[..]); -} - -// ────────────────────── Roundtrip ─────────────── - -#[test] -fn roundtrip_u32_u8() { - let lp = LenPrefixed::<u8, u32>::new(42u32); - let buf = encode_to_bytes(lp); - - let mut reader = &buf[..]; - let decoded = LenPrefixed::<u8, u32>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), 42u32); -} - -#[test] -fn roundtrip_u64_u16() { - let lp = LenPrefixed::<u16, u64>::new(u64::MAX); - let buf = encode_to_bytes(lp); - - let mut reader = &buf[..]; - let decoded = LenPrefixed::<u16, u64>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), u64::MAX); -} - -#[test] -fn roundtrip_bool_u8() { - let lp = LenPrefixed::<u8, bool>::new(true); - let buf = encode_to_bytes(lp); - - let mut reader = &buf[..]; - let decoded = LenPrefixed::<u8, bool>::decode(&mut reader).unwrap(); - assert!(*decoded.as_ref()); -} - -#[test] -fn roundtrip_array_u8() { - let data = [0xDE, 0xAD, 0xBE, 0xEF]; - let lp = LenPrefixed::<u8, [u8; 4]>::new(data); - let buf = encode_to_bytes(lp); - - let mut reader = &buf[..]; - let decoded = LenPrefixed::<u8, [u8; 4]>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), data); -} - -#[test] -fn roundtrip_negative_i32() { - let lp = LenPrefixed::<u8, i32>::new(-12345); - let buf = encode_to_bytes(lp); - - let mut reader = &buf[..]; - let decoded = LenPrefixed::<u8, i32>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), -12345); -} - -#[test] -fn roundtrip_f64() { - let lp = LenPrefixed::<u8, f64>::new(3.141592653589793); - let buf = encode_to_bytes(lp); - - let mut reader = &buf[..]; - let decoded = LenPrefixed::<u8, f64>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), 3.141592653589793); -} - -#[test] -fn roundtrip_u128() { - let lp = LenPrefixed::<u16, u128>::new(u128::MAX); - let buf = encode_to_bytes(lp); - - let mut reader = &buf[..]; - let decoded = LenPrefixed::<u16, u128>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), u128::MAX); -} - -// ────────────────────── Imbrication ───────────── - -#[test] -fn nested_len_prefixed_u16_u8_u32() { - let inner = LenPrefixed::<u8, u32>::new(42u32); - let outer = LenPrefixed::<u16, LenPrefixed<u8, u32>>::new(inner); - let buf = encode_to_bytes(outer); - - // outer length = 1 (u8 prefix of inner) + 4 (u32) = 5 - assert_eq!(buf[0..2], [0x00, 0x05]); - // inner length = 4 - assert_eq!(buf[2], 0x04); - // inner data = 42 - assert_eq!(&buf[3..7], &[0x00, 0x00, 0x00, 0x2A]); -} - -#[test] -fn roundtrip_nested() { - let inner = LenPrefixed::<u8, u16>::new(1337u16); - let outer = LenPrefixed::<u16, LenPrefixed<u8, u16>>::new(inner); - let buf = encode_to_bytes(outer); - - let mut reader = &buf[..]; - let decoded_outer = LenPrefixed::<u16, LenPrefixed<u8, u16>>::decode(&mut reader).unwrap(); - let decoded_inner = decoded_outer.as_ref(); - assert_eq!(decoded_inner.as_ref(), &1337u16); -} - -#[test] -fn triple_nested() { - let l1 = LenPrefixed::<u8, u8>::new(0xAA); - let l2 = LenPrefixed::<u8, LenPrefixed<u8, u8>>::new(l1); - let l3 = LenPrefixed::<u16, LenPrefixed<u8, LenPrefixed<u8, u8>>>::new(l2); - let buf = encode_to_bytes(l3); - - let mut reader = &buf[..]; - let d3 = LenPrefixed::<u16, LenPrefixed<u8, LenPrefixed<u8, u8>>>::decode(&mut reader).unwrap(); - let d2 = d3.as_ref(); - let d1 = d2.as_ref(); - assert_eq!(*d1.as_ref(), 0xAAu8); -} - -// ────────────────────── Erreurs de décodage ───── - -#[test] -fn decode_empty_buffer_fails() { - let result = LenPrefixed::<u8, u32>::decode(&mut &[][..]); - assert!(result.is_err()); -} - -#[test] -fn decode_truncated_length_fails() { - let result = LenPrefixed::<u16, u32>::decode(&mut &[0x01][..]); - assert!(result.is_err()); -} - -#[test] -fn decode_length_exceeds_data_fails() { - let mut buf = Vec::new(); - 100u8.encode(&mut buf).unwrap(); - buf.extend_from_slice(&[0x01, 0x02]); - let result = LenPrefixed::<u8, u32>::decode(&mut &buf[..]); - assert!(result.is_err()); -} - -#[test] -fn decode_zero_length_with_zero_sized_type() { - let lp = LenPrefixed::<u8, [u8; 0]>::new([]); - let buf = encode_to_bytes(lp); - let mut reader = &buf[..]; - let decoded = LenPrefixed::<u8, [u8; 0]>::decode(&mut reader).unwrap(); - assert_eq!(*decoded.as_ref(), []); -} - -// ────────────────────── AsRef / AsMut ─────────── - -#[test] -fn as_ref_returns_inner() { - let lp = LenPrefixed::<u8, u32>::new(99u32); - assert_eq!(*lp.as_ref(), 99u32); -} - -#[test] -fn as_mut_allows_mutation() { - let mut lp = LenPrefixed::<u8, u32>::new(1u32); - *lp.as_mut() = 42; - assert_eq!(*lp.as_ref(), 42); -} - -// ────────────────────── Taille du buffer ──────── - -#[test] -fn u8_length_header_size() { - let lp = LenPrefixed::<u8, u8>::new(0); - let buf = encode_to_bytes(lp); - assert_eq!(buf.len(), 2); // 1 byte length + 1 byte data -} - -#[test] -fn u16_length_header_size() { - let lp = LenPrefixed::<u16, u8>::new(0); - let buf = encode_to_bytes(lp); - assert_eq!(buf.len(), 3); // 2 byte length + 1 byte data -} - -#[test] -fn u32_length_header_size() { - let lp = LenPrefixed::<u32, u8>::new(0); - let buf = encode_to_bytes(lp); - assert_eq!(buf.len(), 5); // 4 byte length + 1 byte data -} - -#[test] -fn u64_length_header_size() { - let lp = LenPrefixed::<u64, u8>::new(0); - let buf = encode_to_bytes(lp); - assert_eq!(buf.len(), 9); // 8 byte length + 1 byte data -} - -#[test] -fn large_data_length_correct() { - let data: [u8; 1000] = std::array::from_fn(|i| (i % 256) as u8); - let lp = LenPrefixed::<u16, [u8; 1000]>::new(data); - let buf = encode_to_bytes(lp); - assert_eq!(buf.len(), 2 + 1000); - assert_eq!(&buf[0..2], &[0x03, 0xE8]); // 1000 in u16 big-endian -} - -// ────────────────────── Sérialisation exacte ──── - -#[test] -fn exact_bytes_u8_prefix() { - let lp = LenPrefixed::<u8, u16>::new(256u16); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x02, 0x01, 0x00]); -} - -#[test] -fn exact_bytes_negative_i16() { - let lp = LenPrefixed::<u8, i16>::new(-1); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x02, 0xFF, 0xFF]); -} - -#[test] -fn exact_bytes_zero_u32() { - let lp = LenPrefixed::<u8, u32>::new(0u32); - let buf = encode_to_bytes(lp); - assert_eq!(buf, vec![0x04, 0x00, 0x00, 0x00, 0x00]); -} diff --git a/tests/tmp/todo.md b/tests/tmp/todo.md deleted file mode 100644 index a0c6a83..0000000 --- a/tests/tmp/todo.md +++ /dev/null @@ -1,356 +0,0 @@ -# zr_protocol - Todo V1 - -> Crate Rust de communication réseau : protocole custom, sérialisation binaire, proc-macro pour définir des packets via un format `.packet`. - ---- - -## Phase 0 — Workspace & Scaffold - -- [ ] **0.1** Créer le workspace virtuel `Cargo.toml` à la racine (members: `zr-protocol`, `zr-protocol-macros`) -- [ ] **0.2** Créer le crate `zr-protocol/` (library, edition 2024) - - [ ] `Cargo.toml` avec dépendances : `bytes = "1"`, `thiserror = "2"`, features optionnelles (`tokio`, `serde`, `derive`) - - [ ] `src/lib.rs` avec modules déclarés : `encode`, `decode`, `packet`, `codec`, `framing`, `builder`, `error`, `impls` -- [ ] **0.3** Créer le crate `zr-protocol-macros/` (proc-macro = true, edition 2024) - - [ ] `Cargo.toml` avec dépendances : `syn = { version = "2", features = ["full"] }`, `quote = "1"`, `proc-macro2 = "1"` - - [ ] `src/lib.rs` minimal (entry point vide) -- [ ] **0.4** Supprimer les anciens fichiers `src/bin/client.rs` et `src/bin/server.rs` ( remplacer par des exemples dans `examples/` ) -- [ ] **0.5** Vérifier que `cargo check --workspace` compile sans erreur -- [ ] **0.6** Créer la structure de dossiers dans les deux crates : - - `zr-protocol/src/encode.rs`, `decode.rs`, `packet.rs`, `codec.rs`, `framing.rs`, `builder.rs`, `error.rs` - - `zr-protocol/src/impls/mod.rs`, `primitives.rs`, `arrays.rs`, `option.rs` - - `zr-protocol-macros/src/parser/ast.rs`, `lexer.rs`, `grammar.rs`, `mod.rs` - - `zr-protocol-macros/src/codegen/struct_gen.rs`, `encode_gen.rs`, `decode_gen.rs`, `registry_gen.rs`, `mod.rs` - - `examples/login/packets/`, `examples/login/server.rs`, `examples/login/client.rs` - - `tests/` (racine workspace) - ---- - -## Phase 1 — Traits Core : `Encode` & `Decode` - -> Les traits fondamentaux de sérialisation/déserialization binaire. - -- [ ] **1.1** Définir `src/error.rs` — type d'erreur `ZrError` - ```rust - use thiserror::Error; - #[derive(Debug, Error)] - pub enum ZrError { - #[error("io error: {0}")] - Io(#[from] std::io::Error), - #[error("invalid packet id: 0x{0:02X}")] - InvalidPacketId(u32), - #[error("buffer underflow: expected {expected} bytes, got {got}")] - Underflow { expected: usize, got: usize }, - #[error("packet too large: {0} bytes (max {1})")] - PacketTooLarge(usize, usize), - #[error("unknown field type: {0}")] - UnknownType(String), - #[error("decode error: {0}")] - Decode(String), - } - ``` -- [ ] **1.2** Définir `src/encode.rs` — trait `Encode` - ```rust - pub trait Encode { - fn encode(&self, buf: &mut BytesMut) -> io::Result<()>; - fn encoded_size(&self) -> usize; // optionnel, pour pré-réserver - } - ``` -- [ ] **1.3** Définir `src/decode.rs` — trait `Decode` - ```rust - pub trait Decode: Sized { - fn decode(buf: &mut BytesMut) -> io::Result<Self>; - } - ``` -- [ ] **1.4** Impl `Encode` pour les types primitifs dans `src/impls/primitives.rs` - - `u8`, `u16`, `u32`, `u64`, `u128` - - `i8`, `i16`, `i32`, `i64`, `i128` - - `bool` (1 octet : 0x00 = false, 0x01 = true) - - `f32`, `f64` (via `to_be_bytes()` / `to_le_bytes()`) - - Tous en big-endian par défaut (network byte order) -- [ ] **1.5** Impl `Encode` pour `String` et `Vec<u8>` - - Format : `u32` (longueur en octets) + octets bruts -- [ ] **1.6** Impl `Encode` pour `[u8; N]` dans `src/impls/arrays.rs` (écriture directe, pas de longueur) -- [ ] **1.7** Impl `Encode` pour `Option<T: Encode>` dans `src/impls/option.rs` - - Format : `u8` tag (0x00 = None, 0x01 = Some) + données si Some -- [ ] **1.8** Impl `Decode` pour tous les mêmes types (mirror de 1.4 à 1.7) - - Avec gestion propre des erreurs (underflow → `ZrError::Underflow`) -- [ ] **1.9** Impl `Encode`/`Decode` pour `Vec<T: Encode/Decode>` - - Format : `u32` nombre d'éléments + sérialisation de chaque élément -- [ ] **1.10** Tests unitaires pour chaque type implémenté - - Roundtrip : encode → decode → assert_eq - - Cas limites : chaîne vide, vec vide, None, MAX values - ---- - -## Phase 2 — Trait `PacketMeta` & Types Runtime - -> Métadonnées des packets et typage dynamique. - -- [ ] **2.1** Définir `src/packet.rs` — trait `PacketMeta` - ```rust - pub trait PacketMeta { - const ID: u32; - const NAME: &'static str; - const SIZE_HINT: Option<usize>; - } - ``` -- [ ] **2.2** Définir un enum `PacketId` (ou type alias `u32`) pour les IDs réservés - - Réservé `0x00` = Reserved, `0xFF` = KeepAlive/Ping -- [ ] **2.3** Définir `pub struct ZrPacket` (wrapper type-érasé pour le codec) - ```rust - pub struct ZrPacket { - pub id: u32, - pub payload: BytesMut, - } - ``` -- [ ] **2.4** Définir `pub trait IntoZrPacket: PacketMeta + Encode` pour convertir un typed packet → `ZrPacket` -- [ ] **2.5** Tests pour `PacketMeta` et `ZrPacket` - ---- - -## Phase 3 — Format `.packet` : Lexer & Parser - -> Parsing du format texte de définition de protocole. - -- [ ] **3.1** Définir l'AST dans `zr-protocol-macros/src/parser/ast.rs` - ```rust - pub struct ProtocolFile { - pub packets: Vec<PacketDef>, - } - pub struct PacketDef { - pub attributes: Vec<Attribute>, - pub name: Ident, - pub id: u32, - pub fields: Vec<FieldDef>, - } - pub struct FieldDef { - pub attributes: Vec<Attribute>, - pub name: Ident, - pub ty: TypeRef, - } - pub enum TypeRef { - Primitive(String), // u32, String, bool, etc. - Option(Box<TypeRef>), // Option<T> - Vec(Box<TypeRef>), // Vec<T> - Array(String, usize), // [u8; 64] - Custom(String), // un autre type packet - } - pub struct Attribute { - pub name: String, - pub value: Option<String>, - } - ``` -- [ ] **3.2** Implémenter le lexer dans `zr-protocol-macros/src/parser/lexer.rs` - - Tokens : `Packet`, `Ident`, `Number` (hex 0x.. et décimal), `Colon`, `BraceOpen`, `BraceClose`, `BracketOpen`, `BracketClose`, `Eq`, `String`, `Comma`, `EndAttribute`, `Hash` - - Whitespace et comments (`//`) ignorés -- [ ] **3.3** Implémenter le parser dans `zr-protocol-macros/src/parser/grammar.rs` - - Parse un `ProtocolFile` à partir de tokens - - Validation : nom unique, ID unique, types connus - - Erreurs avec span (ligne/colonne) pour de bons messages d'erreur -- [ ] **3.4** Tests du parser : - - [ ] Format valide basique - - [ ] Attributs optionnels (`#[endian = "big"]`) - - [ ] Types `Option<T>`, `Vec<T>`, `[u8; N]` - - [ ] Erreur : type inconnu - - [ ] Erreur : syntaxe invalide - - [ ] Erreur : ID manquant - - [ ] Fichier vide (pas de packet) - ---- - -## Phase 4 — Proc-macro `packet!` - -> Génération de code Rust à partir du format `.packet`. - -- [ ] **4.1** Implémenter l'entry point `packet!` dans `zr-protocol-macros/src/lib.rs` - - Accepte `include!("path/to/file.packet")` OU du code inline - - Lit le fichier via `CARGO_MANIFEST_DIR` + chemin relatif -- [ ] **4.2** Codegen struct dans `zr-protocol-macros/src/codegen/struct_gen.rs` - - Génère `#[derive(Debug, Clone)] pub struct NomPacket { pub field: Type, ... }` - - Gère `Option<T>` → `Option<T>`, `Vec<T>` → `Vec<T>`, `[u8; N]` → `[u8; N]` -- [ ] **4.3** Codegen `impl PacketMeta` dans `zr-protocol-macros/src/codegen/registry_gen.rs` - - `const ID: u32 = ...; const NAME: &'static str = "..."; const SIZE_HINT: Option<usize> = None;` -- [ ] **4.4** Codegen `impl Encode` dans `zr-protocol-macros/src/codegen/encode_gen.rs` - - Pour chaque champ : `Encode::encode(&self.champ, buf)?;` -- [ ] **4.5** Codegen `impl Decode` dans `zr-protocol-macros/src/codegen/decode_gen.rs` - - Pour chaque champ : `champ: Decode::decode(buf)?` -- [ ] **4.6** Codegen registry globale (un seul `match` ID → nom) dans `registry_gen.rs` - - Fonction `pub fn packet_name_by_id(id: u32) -> Option<&'static str>` -- [ ] **4.7** Codegen `impl IntoZrPacket` pour chaque packet - - Sérialise → `ZrPacket { id, payload }` -- [ ] **4.8** Gestion des attributs dans le codegen - - `#[endian = "big"]` → big-endian (défaut), `#[endian = "little"]` → little-endian - - `#[skip_if_none]` → ne sérialise pas le champ si `None` (sans le tag) -- [ ] **4.9** Tests d'intégration : - - [ ] Un fichier `.packet` simple → compiles, encode/decode roundtrip - - [ ] Fichier avec `Option<T>` et `Vec<T>` - - [ ] Compilation échoue sur type inconnu (trybuild) - - [ ] Compilation échoue sur syntaxe invalide (trybuild) -- [ ] **4.10** Macro `packet!` inline (pour les petits définitions sans fichier) - ```rust - packet! { - packet Ping 0xFF { - sequence: u32, - } - } - ``` - ---- - -## Phase 5 — Builder Pattern (secondaire) - -> API alternative pour construire des packets sans struct literals. - -- [ ] **5.1** Définir le trait `PacketBuilder` dans `src/builder.rs` - ```rust - pub trait PacketBuilder: Sized { - type Packet: Encode + PacketMeta; - fn new() -> Self; - fn field<T: IntoFieldValue>(mut self, name: &str, value: T) -> Self; - fn build(self) -> io::Result<Self::Packet>; - } - ``` -- [ ] **5.2** Codegen du builder dans `zr-protocol-macros/src/codegen/builder_gen.rs` - - Génère une struct `NomPacketBuilder { username: Option<String>, ... }` avec chaque champ en `Option` - - `field()` matche sur le nom (string) et set la valeur - - `build()` vérifie que tous les champs sont présents, sinon erreur -- [ ] **5.3** Opt-out via attribute : `#[packet(no_builder)]` désactive la génération du builder -- [ ] **5.4** Tests : - - [ ] Builder crée un packet valide - - [ ] Builder erreur si champ manquant - - [ ] `#[packet(no_builder)]` ne génère pas le builder - ---- - -## Phase 6 — Feature Flags & Extensibilité - -> Support optionnel de crates externes. - -- [ ] **6.1** Feature `tokio` : active les dépendances `tokio` + `tokio-util` - - Active le module `codec.rs` et `framing.rs` -- [ ] **6.2** Feature `serde` : active `serde` + `bincode` - - Génère `#[derive(serde::Serialize, serde::Deserialize)]` sur les structs via le proc-macro - - Attribut `#[packet(serde)]` pour forcer ou `#[packet(no_serde)]` pour désactiver par packet -- [ ] **6.3** Feature `derive` : active `zr-protocol-macros` - - Les macros `packet!`, `include_packets!` ne sont disponibles que avec cette feature -- [ ] **6.4** Feature `std` (défaut) : support `String`, `Vec`, etc. - - Feature `no_std` future (pas v1) : uniquement `[u8; N]`, `u8`, etc. -- [ ] **6.5** Documentation des features dans le `Cargo.toml` et `lib.rs` - ---- - -## Phase 7 — Codec & Framing (tokio-util) - -> Intégration avec tokio pour la communication async. - -- [ ] **7.1** Implémenter `src/framing.rs` — length-prefix framing - - Header : `u32` big-endian = longueur du payload - - Configurable : taille du header (2 ou 4 octets), endianness - - `max_packet_size` avec défaut (ex: 16 Mo) -- [ ] **7.2** Implémenter `src/codec.rs` — `ZrCodec` - - `impl Encoder<Box<dyn Encode>> for ZrCodec` — écrit header + payload - - `impl Decoder for ZrCodec` — lit header, vérifie taille, lit payload, retourne `ZrPacket` - - Gestion propre de `BytesMut` (reserve, split_to, advance) -- [ ] **7.3** Type `FramedPacket` pour le dispatch dynamique - - Le codec lit l'ID depuis le payload, lookup dans la registry - - Retourne un `ZrPacket` (type-érasé) que l'utilisateur cast avec un match sur l'ID -- [ ] **7.4** Helper `fn framed_read(stream) -> impl Stream<Item = ZrPacket>` (optionnel, wrapper) -- [ ] **7.5** Tests : - - [ ] Roundtrip TCP : encode → send → receive → decode → assert_eq - - [ ] Rejet des packets trop gros - - [ ] Gestion des reads partiels (TCP peut diviser les données) - ---- - -## Phase 8 — Client & Server Helpers - -> Utilitaires pour simplifier l'usage réseau. - -- [ ] **8.1** Struct `ZrConnection` wrapper autour de `Framed<TcpStream, ZrCodec>` - ```rust - pub struct ZrConnection { - framed: Framed<TcpStream, ZrCodec>, - } - impl ZrConnection { - pub async fn connect(addr: &str) -> io::Result<Self>; - pub async fn send_packet<P: Encode + PacketMeta>(&mut self, packet: &P) -> io::Result<()>; - pub async fn next_packet(&mut self) -> io::Result<Option<ZrPacket>>; - } - ``` -- [ ] **8.2** Struct `ZrListener` wrapper autour de `TcpListener` - ```rust - pub struct ZrListener { listener: TcpListener } - impl ZrListener { - pub async fn bind(addr: &str) -> io::Result<Self>; - pub async fn accept(&self) -> io::Result<(ZrConnection, SocketAddr)>; - } - ``` -- [ ] **8.3** Macro `#[tokio::main]` compatible — les helpers utilisent tokio derrière -- [ ] **8.4** Tests d'intégration : client/server qui s'échangent des packets - ---- - -## Phase 9 — Exemples - -> Démonstrations complètes d'utilisation. - -- [ ] **9.1** Exemple `examples/simple.rs` — packet inline, encode/decode sans réseau -- [ ] **9.2** Exemple `examples/login/` — client/server complet - - Fichiers `.packet` de définition - - Server qui écoute, reçoit `LoginRequest`, répond `LoginResponse` - - Client qui se connecte, envoie `LoginRequest`, lit `LoginResponse` -- [ ] **9.3** Exemple avec `Option<T>` et `Vec<T>` pour montrer les types composés -- [ ] **9.4** README avec badges, description, et example usage - ---- - -## Phase 10 — Tests & Qualité - -> Fiabilité et robustesse. - -- [ ] **10.1** Tests unitaires : chaque type primitif encode/decode roundtrip -- [ ] **10.2** Tests du parser lexer/grammar (bonne et mauvaise syntaxe) -- [ ] **10.3** Tests trybuild : erreurs de compilation avec bons messages - - [ ] Type inconnu dans `.packet` - - [ ] Syntaxe invalide - - [ ] ID manquant - - [ ] Champ dupliqué -- [ ] **10.4** Tests d'intégration TCP : client ↔ server roundtrip -- [ ] **10.5** Benchmark (criterion) : throughput sérialisation vs bincode -- [ ] **10.6** `cargo clippy --workspace` sans warning -- [ ] **10.7** `cargo fmt --check` passe - ---- - -## Phase 11 — Documentation - -- [ ] **11.1** `README.md` avec description, quick start, features -- [ ] **11.2** Doc comments (`///`) sur tous les traits publics -- [ ] **11.3** Doc comments sur le format `.packet` (guide de syntaxe) -- [ ] **11.4** Exemples dans les doc comments (doc-tests) -- [ ] **11.5** CHANGELOG.md - ---- - -## Ordre d'exécution recommandé - -``` -Phase 0 → Phase 1 → Phase 2 → Phase 3 → Phase 4 → Phase 5 - ↓ - Phase 6 (features) - ↓ - Phase 7 (codec) - ↓ - Phase 8 → Phase 9 → Phase 10 → Phase 11 -``` - -**Dépendances critiques :** -- Phase 4 (proc-macro) dépend de Phase 1 (traits) et Phase 2 (AST) -- Phase 7 (codec) dépend de Phase 1 (Encode/Decode) et Phase 4 (PacketMeta) -- Phase 8 (helpers) dépend de Phase 7 (codec) -- Phase 9-11 dépendent de tout le reste - -**Peut être fait en parallèle :** -- Phase 5 (builder) peut démarrer après Phase 4 -- Phase 6 (features) peut démarrer après Phase 4 -- Les tests (Phase 10) peuvent être écrits au fur et à mesure de chaque phase |
